| 23 | |
| 24 | class 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 |