MCPcopy Create free account
hub / github.com/Royalvice/DocDiff / __init__

Method __init__

schedule/diffusionSample.py:25–47  ·  view source on GitHub ↗
(self, model, T, schedule)

Source from the content-addressed store, hash-verified

23
24class GaussianDiffusion(nn.Module):
25 def __init__(self, model, T, schedule):
26 super().__init__()
27 self.visual = False
28 if self.visual:
29 self.num = 0
30 self.model = model
31 self.T = T
32 self.schedule = schedule
33 betas = self.schedule.get_betas()
34 self.register_buffer('betas', betas.float())
35 alphas = 1. - self.betas
36 alphas_bar = torch.cumprod(alphas, dim=0)
37 alphas_bar_prev = F.pad(alphas_bar, [1, 0], value=1)[:T]
38 gammas = alphas_bar
39
40 self.register_buffer('coeff1', torch.sqrt(1. / alphas))
41 self.register_buffer('coeff2', self.coeff1 * (1. - alphas) / torch.sqrt(1. - alphas_bar))
42 self.register_buffer('posterior_var', self.betas * (1. - alphas_bar_prev) / (1. - alphas_bar))
43
44 # calculation for q(y_t|y_{t-1})
45 self.register_buffer('gammas', gammas)
46 self.register_buffer('sqrt_one_minus_gammas', np.sqrt(1 - gammas))
47 self.register_buffer('sqrt_gammas', np.sqrt(gammas))
48
49 def predict_xt_prev_mean_from_eps(self, x_t, t, eps):
50 assert x_t.shape == eps.shape

Callers

nothing calls this directly

Calls 1

get_betasMethod · 0.80

Tested by

no test coverage detected