From 689593a5792512f81c159f7f441afcd2fcc26492 Mon Sep 17 00:00:00 2001 From: Phil Wang Date: Wed, 10 Aug 2022 10:52:07 -0700 Subject: [PATCH] add the new self conditioning technique from hintons group from bit diffusion paper 0.27.0 --- README.md | 11 +++ .../continuous_time_gaussian_diffusion.py | 1 + .../denoising_diffusion_pytorch.py | 80 +++++++++++++------ .../elucidated_diffusion.py | 1 + .../learned_gaussian_diffusion.py | 2 + .../weighted_objective_gaussian_diffusion.py | 1 + setup.py | 2 +- 7 files changed, 73 insertions(+), 25 deletions(-) diff --git a/README.md b/README.md index 323819c..f910378 100644 --- a/README.md +++ b/README.md @@ -170,3 +170,14 @@ $ accelerate launch train.py volume = {abs/2010.02502} } ``` + +```bibtex +@misc{chen2022analog, + title = {Analog Bits: Generating Discrete Data using Diffusion Models with Self-Conditioning}, + author = {Ting Chen and Ruixiang Zhang and Geoffrey Hinton}, + year = {2022}, + eprint = {2208.04202}, + archivePrefix = {arXiv}, + primaryClass = {cs.CV} +} +``` diff --git a/denoising_diffusion_pytorch/continuous_time_gaussian_diffusion.py b/denoising_diffusion_pytorch/continuous_time_gaussian_diffusion.py index 63e4123..7822d8b 100644 --- a/denoising_diffusion_pytorch/continuous_time_gaussian_diffusion.py +++ b/denoising_diffusion_pytorch/continuous_time_gaussian_diffusion.py @@ -127,6 +127,7 @@ class ContinuousTimeGaussianDiffusion(nn.Module): ): super().__init__() assert model.learned_sinusoidal_cond + assert not model.self_condition, 'not supported yet' self.model = model diff --git a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py index d9a73bb..d577c5e 100644 --- a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py +++ b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py @@ -1,23 +1,23 @@ import math import copy +from pathlib import Path +from random import random +from functools import partial +from collections import namedtuple +from multiprocessing import cpu_count + import torch from torch import nn, einsum import torch.nn.functional as F -from inspect import isfunction -from collections import namedtuple -from functools import partial - from torch.utils.data import Dataset, DataLoader -from multiprocessing import cpu_count -from pathlib import Path from torch.optim import Adam from torchvision import transforms as T, utils -from PIL import Image from einops import rearrange, reduce from einops.layers.torch import Rearrange +from PIL import Image from tqdm.auto import tqdm from ema_pytorch import EMA @@ -35,7 +35,7 @@ def exists(x): def default(val, d): if exists(val): return val - return d() if isfunction(d) else d + return d() if callable(d) else d def cycle(dl): while True: @@ -251,6 +251,7 @@ class Unet(nn.Module): out_dim = None, dim_mults=(1, 2, 4, 8), channels = 3, + self_condition = False, resnet_block_groups = 8, learned_variance = False, learned_sinusoidal_cond = False, @@ -261,9 +262,11 @@ class Unet(nn.Module): # determine dimensions self.channels = channels + self.self_condition = self_condition + input_channels = channels * (2 if self_condition else 1) init_dim = default(init_dim, dim) - self.init_conv = nn.Conv2d(channels, init_dim, 7, padding = 3) + self.init_conv = nn.Conv2d(input_channels, init_dim, 7, padding = 3) dims = [init_dim, *map(lambda m: dim * m, dim_mults)] in_out = list(zip(dims[:-1], dims[1:])) @@ -327,7 +330,11 @@ class Unet(nn.Module): self.final_res_block = block_klass(dim * 2, dim, time_emb_dim = time_dim) self.final_conv = nn.Conv2d(dim, self.out_dim, 1) - def forward(self, x, time): + def forward(self, x, time, x_self_cond = None): + if self.self_condition: + x_self_cond = default(x_self_cond, lambda: torch.zeros_like(x)) + x = torch.cat((x_self_cond, x), dim = 1) + x = self.init_conv(x) r = x.clone() @@ -395,7 +402,6 @@ class GaussianDiffusion(nn.Module): model, *, image_size, - channels = 3, timesteps = 1000, sampling_timesteps = None, loss_type = 'l1', @@ -408,9 +414,12 @@ class GaussianDiffusion(nn.Module): super().__init__() assert not (type(self) == GaussianDiffusion and model.channels != model.out_dim) - self.channels = channels - self.image_size = image_size self.model = model + self.channels = self.model.channels + self.self_condition = self.model.self_condition + + self.image_size = image_size + self.objective = objective assert objective in {'pred_noise', 'pred_x0'}, 'objective must be either pred_noise (predict noise) or pred_x0 (predict image start)' @@ -493,8 +502,8 @@ 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): - model_output = self.model(x, t) + def model_predictions(self, x, t, x_self_cond = None): + model_output = self.model(x, t, x_self_cond) if self.objective == 'pred_noise': pred_noise = model_output @@ -506,23 +515,24 @@ class GaussianDiffusion(nn.Module): return ModelPrediction(pred_noise, x_start) - def p_mean_variance(self, x, t, clip_denoised: bool): - preds = self.model_predictions(x, t) + def p_mean_variance(self, x, t, x_self_cond = None, clip_denoised = True): + preds = self.model_predictions(x, t, x_self_cond) x_start = preds.pred_x_start if clip_denoised: x_start.clamp_(-1., 1.) model_mean, posterior_variance, posterior_log_variance = self.q_posterior(x_start = x_start, x_t = x, t = t) - return model_mean, posterior_variance, posterior_log_variance + return model_mean, posterior_variance, posterior_log_variance, x_start @torch.no_grad() - def p_sample(self, x, t: int, clip_denoised = True): + def p_sample(self, x, t: int, x_self_cond = None, clip_denoised = True): b, *_, device = *x.shape, x.device batched_times = torch.full((x.shape[0],), t, device = x.device, dtype = torch.long) - model_mean, _, model_log_variance = self.p_mean_variance(x = x, t = batched_times, clip_denoised = clip_denoised) + model_mean, _, model_log_variance, x_start = self.p_mean_variance(x = x, t = batched_times, x_self_cond = x_self_cond, clip_denoised = clip_denoised) noise = torch.randn_like(x) if t > 0 else 0. # no noise if t == 0 - return model_mean + (0.5 * model_log_variance).exp() * noise + pred_img = model_mean + (0.5 * model_log_variance).exp() * noise + return pred_img, x_start @torch.no_grad() def p_sample_loop(self, shape): @@ -530,8 +540,11 @@ class GaussianDiffusion(nn.Module): img = torch.randn(shape, device=device) + x_start = None + for t in tqdm(reversed(range(0, self.num_timesteps)), desc = 'sampling loop time step'): - img = self.p_sample(img, t) + self_cond = x_start if self.self_condition else None + img, x_start = self.p_sample(img, t, self_cond) img = unnormalize_to_zero_to_one(img) return img @@ -546,13 +559,17 @@ class GaussianDiffusion(nn.Module): img = torch.randn(shape, device = device) + x_start = None + for time, time_next in tqdm(time_pairs, desc = 'sampling loop time step'): alpha = self.alphas_cumprod_prev[time] alpha_next = self.alphas_cumprod_prev[time_next] time_cond = torch.full((batch,), time, device = device, dtype = torch.long) - pred_noise, x_start, *_ = self.model_predictions(img, time_cond) + 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.) @@ -612,8 +629,23 @@ class GaussianDiffusion(nn.Module): b, c, h, w = x_start.shape noise = default(noise, lambda: torch.randn_like(x_start)) + # noise sample + x = self.q_sample(x_start = x_start, t = t, noise = noise) - model_out = self.model(x, t) + + # if doing self-conditioning, 50% of the time, predict x_start from current set of times + # and condition with unet with that + # this technique will slow down training by 25%, but seems to lower FID significantly + + x_self_cond = None + if self.self_condition and random() < 0.5: + with torch.no_grad(): + x_self_cond = self.model_predictions(x, t).pred_x_start + x_self_cond.detach_() + + # predict and take gradient step + + model_out = self.model(x, t, x_self_cond) if self.objective == 'pred_noise': target = noise diff --git a/denoising_diffusion_pytorch/elucidated_diffusion.py b/denoising_diffusion_pytorch/elucidated_diffusion.py index d430516..2dc40b5 100644 --- a/denoising_diffusion_pytorch/elucidated_diffusion.py +++ b/denoising_diffusion_pytorch/elucidated_diffusion.py @@ -52,6 +52,7 @@ class ElucidatedDiffusion(nn.Module): ): super().__init__() assert net.learned_sinusoidal_cond + assert not net.self_condition, 'not supported yet' self.net = net diff --git a/denoising_diffusion_pytorch/learned_gaussian_diffusion.py b/denoising_diffusion_pytorch/learned_gaussian_diffusion.py index 7d970f1..d6a4f8f 100644 --- a/denoising_diffusion_pytorch/learned_gaussian_diffusion.py +++ b/denoising_diffusion_pytorch/learned_gaussian_diffusion.py @@ -77,6 +77,8 @@ class LearnedGaussianDiffusion(GaussianDiffusion): ): super().__init__(model, *args, **kwargs) assert model.out_dim == (model.channels * 2), 'dimension out of unet must be twice the number of channels for learned variance - you can also set the `learned_variance` keyword argument on the Unet to be `True`' + assert not model.self_condition, 'not supported yet' + self.vb_loss_weight = vb_loss_weight def model_predictions(self, x, t): diff --git a/denoising_diffusion_pytorch/weighted_objective_gaussian_diffusion.py b/denoising_diffusion_pytorch/weighted_objective_gaussian_diffusion.py index 7d6b437..b24eebb 100644 --- a/denoising_diffusion_pytorch/weighted_objective_gaussian_diffusion.py +++ b/denoising_diffusion_pytorch/weighted_objective_gaussian_diffusion.py @@ -31,6 +31,7 @@ class WeightedObjectiveGaussianDiffusion(GaussianDiffusion): super().__init__(model, *args, **kwargs) channels = model.channels assert model.out_dim == (channels * 2 + 2), 'dimension out (out_dim) of unet must be twice the number of channels + 2 (for the softmax weighted sum) - for channels of 3, this should be (3 * 2) + 2 = 8' + assert not model.self_condition, 'not supported yet' assert not self.is_ddim_sampling, 'ddim sampling cannot be used' self.split_dims = (channels, channels, 2) diff --git a/setup.py b/setup.py index eec0773..9a8fbaa 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.26.5', + version = '0.27.0', license='MIT', description = 'Denoising Diffusion Probabilistic Models - Pytorch', author = 'Phil Wang',