diff --git a/README.md b/README.md index 4caed09..9f5ac57 100644 --- a/README.md +++ b/README.md @@ -108,3 +108,14 @@ Samples and model checkpoints will be logged to `./results` periodically url = {https://proceedings.mlr.press/v139/nichol21a.html}, } ``` + +```bibtex +@inproceedings{kingma2021on, + title = {On Density Estimation with Diffusion Models}, + author = {Diederik P Kingma and Tim Salimans and Ben Poole and Jonathan Ho}, + booktitle = {Advances in Neural Information Processing Systems}, + editor = {A. Beygelzimer and Y. Dauphin and P. Liang and J. Wortman Vaughan}, + year = {2021}, + url = {https://openreview.net/forum?id=2LdBqxc1Yv} +} +``` diff --git a/denoising_diffusion_pytorch/__init__.py b/denoising_diffusion_pytorch/__init__.py index 5082fd7..75ab437 100644 --- a/denoising_diffusion_pytorch/__init__.py +++ b/denoising_diffusion_pytorch/__init__.py @@ -1,4 +1,5 @@ from denoising_diffusion_pytorch.denoising_diffusion_pytorch import GaussianDiffusion, Unet, Trainer from denoising_diffusion_pytorch.learned_gaussian_diffusion import LearnedGaussianDiffusion +from denoising_diffusion_pytorch.continuous_time_gaussian_diffusion import ContinuousTimeGaussianDiffusion from denoising_diffusion_pytorch.weighted_objective_gaussian_diffusion import WeightedObjectiveGaussianDiffusion diff --git a/denoising_diffusion_pytorch/continuous_time_gaussian_diffusion.py b/denoising_diffusion_pytorch/continuous_time_gaussian_diffusion.py new file mode 100644 index 0000000..3a4f91a --- /dev/null +++ b/denoising_diffusion_pytorch/continuous_time_gaussian_diffusion.py @@ -0,0 +1,187 @@ +import torch +from torch import sqrt +from torch import nn, einsum +import torch.nn.functional as F +from torch.special import expm1 + +from tqdm import tqdm +from einops import rearrange, repeat + +# helpers + +def exists(val): + return val is not None + +def default(val, d): + if exists(val): + return val + return d() if callable(d) else d + +# normalization functions + +def normalize_to_neg_one_to_one(img): + return img * 2 - 1 + +def unnormalize_to_zero_to_one(t): + return (t + 1) * 0.5 + +# diffusion helpers + +def right_pad_dims_to(x, t): + padding_dims = x.ndim - t.ndim + if padding_dims <= 0: + return t + return t.view(*t.shape, *((1,) * padding_dims)) + +# continuous schedules + +# equations are taken from https://openreview.net/attachment?id=2LdBqxc1Yv&name=supplementary_material +# @crowsonkb Katherine's repository also helped here https://github.com/crowsonkb/v-diffusion-jax/blob/master/diffusion/utils.py + +# log(snr) that approximates the original linear schedule + +def beta_linear_log_snr(t): + return -torch.log(expm1(1e-4 + 10 * (t ** 2))) + +def alpha_cosine_log_snr(t): + raise NotImplementedError + +class learned_noise_schedule(nn.Module): + def __init__(self): + super().__init__() + raise NotImplementedError + # learned noise schedule, using learned monotonic MLP (weights kept positive) in the paper + +class ContinuousTimeGaussianDiffusion(nn.Module): + def __init__( + self, + denoise_fn, + *, + image_size, + channels = 3, + cond_scale = 500, + loss_type = 'l1', + noise_schedule = 'linear', + num_sample_steps = 500 + ): + super().__init__() + + self.denoise_fn = denoise_fn + + # image dimensions + + self.channels = channels + self.image_size = image_size + + # continuous noise schedule related stuff + + self.cond_scale = cond_scale # the log(snr) will be scaled by this value + self.loss_type = loss_type + + if noise_schedule == 'linear': + self.log_snr = beta_linear_log_snr + else: + raise ValueError(f'unknown noise schedule {noise_schedule}') + + # sampling + + self.num_sample_steps = num_sample_steps + + @property + def device(self): + return next(self.denoise_fn.parameters()).device + + @property + def loss_fn(self): + if self.loss_type == 'l1': + return F.l1_loss + elif self.loss_type == 'l2': + return F.mse_loss + else: + raise ValueError(f'invalid loss type {self.loss_type}') + + def p_mean_variance(self, x, time, time_next): + # reviewer found an error in the equation in the paper (missing sigma) + # following - https://openreview.net/forum?id=2LdBqxc1Yv¬eId=rIQgH0zKsRt + + # todo - derive x_start from the posterior mean and do dynamic thresholding + # assumed that is what is going on in Imagen + + batch = x.shape[0] + batch_time = repeat(time, ' -> b', b = batch) + + pred_noise = self.denoise_fn(x, batch_time * self.cond_scale) + + log_snr = self.log_snr(time) + log_snr_next = self.log_snr(time_next) + c = -expm1(log_snr - log_snr_next) + + squared_alpha, squared_alpha_next = log_snr.sigmoid(), log_snr_next.sigmoid() + squared_sigma, squared_sigma_next = (-log_snr).sigmoid(), (-log_snr_next).sigmoid() + + model_mean = sqrt(squared_alpha_next / squared_alpha) * (x - c * sqrt(squared_sigma) * pred_noise) + posterior_variance = squared_sigma_next * c + + return model_mean, posterior_variance + + # sampling related functions + + @torch.no_grad() + def p_sample(self, x, time, time_next): + batch, *_, device = *x.shape, x.device + + model_mean, model_variance = self.p_mean_variance(x = x, time = time, time_next = time_next) + noise = torch.randn_like(x) + return model_mean + sqrt(model_variance) * noise + + @torch.no_grad() + def p_sample_loop(self, shape): + batch = shape[0] + + img = torch.randn(shape, device = self.device) + steps = torch.linspace(1., 0., self.num_sample_steps + 1, device = self.device) + + for i in tqdm(range(self.num_sample_steps), desc = 'sampling loop time step', total = self.num_sample_steps): + times = steps[i] + times_next = steps[i + 1] + img = self.p_sample(img, times, times_next) + + img = unnormalize_to_zero_to_one(img) + return img + + @torch.no_grad() + def sample(self, batch_size = 16): + return self.p_sample_loop((batch_size, self.channels, self.image_size, self.image_size)) + + # training related functions - noise prediction + + def q_sample(self, x_start, times, noise = None): + noise = default(noise, lambda: torch.randn_like(x_start)) + + log_snr = self.log_snr(times) + + log_snr_padded = right_pad_dims_to(x_start, log_snr) + alpha, sigma = sqrt(log_snr_padded.sigmoid()), sqrt((-log_snr_padded).sigmoid()) + x_noised = x_start * alpha + noise * sigma + + return x_noised, log_snr + + def random_times(self, batch_size): + # times are now uniform from 0 to 1 + return torch.zeros((batch_size,), device = self.device).float().uniform_(0, 1) + + def p_losses(self, x_start, times, noise = None): + noise = default(noise, lambda: torch.randn_like(x_start)) + + x, log_snr = self.q_sample(x_start = x_start, times = times, noise = noise) + + model_out = self.denoise_fn(x, log_snr * self.cond_scale) + return self.loss_fn(model_out, noise) + + def forward(self, img, *args, **kwargs): + b, c, h, w, device, img_size, = *img.shape, img.device, self.image_size + assert h == img_size and w == img_size, f'height and width of image must be {img_size}' + + times = self.random_times(b) + img = normalize_to_neg_one_to_one(img) + return self.p_losses(img, times, *args, **kwargs) diff --git a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py index 21202a3..5958d15 100644 --- a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py +++ b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py @@ -324,11 +324,6 @@ def extract(a, t, x_shape): out = a.gather(-1, t) return out.reshape(b, *((1,) * (len(x_shape) - 1))) -def noise_like(shape, device, repeat=False): - repeat_noise = lambda: torch.randn((1, *shape[1:]), device=device).repeat(shape[0], *((1,) * (len(shape) - 1))) - noise = lambda: torch.randn(shape, device=device) - return repeat_noise() if repeat else noise() - def linear_beta_schedule(timesteps): scale = 1000 / timesteps beta_start = scale * 0.0001 @@ -444,10 +439,10 @@ class GaussianDiffusion(nn.Module): return model_mean, posterior_variance, posterior_log_variance @torch.no_grad() - def p_sample(self, x, t, clip_denoised=True, repeat_noise=False): + def p_sample(self, x, t, clip_denoised=True): b, *_, device = *x.shape, x.device model_mean, _, model_log_variance = self.p_mean_variance(x=x, t=t, clip_denoised=clip_denoised) - noise = noise_like(x.shape, device, repeat_noise) + noise = torch.randn_like(x) # no noise when t == 0 nonzero_mask = (1 - (t == 0).float()).reshape(b, *((1,) * (len(x.shape) - 1))) return model_mean + nonzero_mask * (0.5 * model_log_variance).exp() * noise diff --git a/setup.py b/setup.py index 926fb66..f0690f4 100644 --- a/setup.py +++ b/setup.py @@ -3,12 +3,13 @@ from setuptools import setup, find_packages setup( name = 'denoising-diffusion-pytorch', packages = find_packages(), - version = '0.16.1', + version = '0.16.3', license='MIT', description = 'Denoising Diffusion Probabilistic Models - Pytorch', author = 'Phil Wang', author_email = 'lucidrains@gmail.com', url = 'https://github.com/lucidrains/denoising-diffusion-pytorch', + long_description_content_type = 'text/markdown', keywords = [ 'artificial intelligence', 'generative models'