From 6dda508ff62ad014fcf1f387e5107c7106d6c99d Mon Sep 17 00:00:00 2001 From: Phil Wang Date: Thu, 1 Sep 2022 09:34:50 -0700 Subject: [PATCH] in ddim, clip x0 before calculation of predicted noise, thanks to @lukovnikov again for pointing out this inconsistency with glides implementation --- .../denoising_diffusion_pytorch.py | 17 ++++++++++------- setup.py | 2 +- 2 files changed, 11 insertions(+), 8 deletions(-) diff --git a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py index cfae2e6..2e73fa0 100644 --- a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py +++ b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py @@ -37,6 +37,9 @@ def default(val, d): return val return d() if callable(d) else d +def identity(t, *args, **kwargs): + return t + def cycle(dl): while True: for data in dl: @@ -517,16 +520,19 @@ class GaussianDiffusion(nn.Module): posterior_log_variance_clipped = extract(self.posterior_log_variance_clipped, t, x_t.shape) return posterior_mean, posterior_variance, posterior_log_variance_clipped - def model_predictions(self, x, t, x_self_cond = None): + def model_predictions(self, x, t, x_self_cond = None, clip_x_start = False): model_output = self.model(x, t, x_self_cond) + maybe_clip = partial(torch.clamp, min = -1., max = 1.) if clip_x_start else identity if self.objective == 'pred_noise': pred_noise = model_output - x_start = self.predict_start_from_noise(x, t, model_output) + x_start = self.predict_start_from_noise(x, t, pred_noise) + x_start = maybe_clip(x_start) elif self.objective == 'pred_x0': - pred_noise = self.predict_noise_from_start(x, t, model_output) x_start = model_output + x_start = maybe_clip(x_start) + pred_noise = self.predict_noise_from_start(x, t, x_start) return ModelPrediction(pred_noise, x_start) @@ -579,10 +585,7 @@ class GaussianDiffusion(nn.Module): for time, time_next in tqdm(time_pairs, desc = 'sampling loop time step'): time_cond = torch.full((batch,), time, device=device, dtype=torch.long) self_cond = x_start if self.self_condition else None - pred_noise, x_start, *_ = self.model_predictions(img, time_cond, self_cond) - - if clip_denoised: - x_start.clamp_(-1., 1.) + pred_noise, x_start, *_ = self.model_predictions(img, time_cond, self_cond, clip_x_start = clip_denoised) if time_next < 0: img = x_start diff --git a/setup.py b/setup.py index ba5b1e8..1a5707e 100644 --- a/setup.py +++ b/setup.py @@ -3,7 +3,7 @@ from setuptools import setup, find_packages setup( name = 'denoising-diffusion-pytorch', packages = find_packages(), - version = '0.27.7', + version = '0.27.8', license='MIT', description = 'Denoising Diffusion Probabilistic Models - Pytorch', author = 'Phil Wang',