Compare commits

..
3 Commits
7 changed files with 80 additions and 33 deletions
+11
View File
@@ -170,3 +170,14 @@ $ accelerate launch train.py
volume = {abs/2010.02502} 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}
}
```
@@ -127,6 +127,7 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
): ):
super().__init__() super().__init__()
assert model.learned_sinusoidal_cond assert model.learned_sinusoidal_cond
assert not model.self_condition, 'not supported yet'
self.model = model self.model = model
@@ -1,23 +1,23 @@
import math import math
import copy 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 import torch
from torch import nn, einsum from torch import nn, einsum
import torch.nn.functional as F 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 torch.utils.data import Dataset, DataLoader
from multiprocessing import cpu_count
from pathlib import Path
from torch.optim import Adam from torch.optim import Adam
from torchvision import transforms as T, utils from torchvision import transforms as T, utils
from PIL import Image
from einops import rearrange, reduce from einops import rearrange, reduce
from einops.layers.torch import Rearrange from einops.layers.torch import Rearrange
from PIL import Image
from tqdm.auto import tqdm from tqdm.auto import tqdm
from ema_pytorch import EMA from ema_pytorch import EMA
@@ -35,7 +35,7 @@ def exists(x):
def default(val, d): def default(val, d):
if exists(val): if exists(val):
return val return val
return d() if isfunction(d) else d return d() if callable(d) else d
def cycle(dl): def cycle(dl):
while True: while True:
@@ -89,16 +89,15 @@ def Downsample(dim, dim_out = None):
return nn.Conv2d(dim, default(dim_out, dim), 4, 2, 1) return nn.Conv2d(dim, default(dim_out, dim), 4, 2, 1)
class LayerNorm(nn.Module): class LayerNorm(nn.Module):
def __init__(self, dim, eps = 1e-5): def __init__(self, dim):
super().__init__() super().__init__()
self.eps = eps
self.g = nn.Parameter(torch.ones(1, dim, 1, 1)) self.g = nn.Parameter(torch.ones(1, dim, 1, 1))
self.b = nn.Parameter(torch.zeros(1, dim, 1, 1))
def forward(self, x): def forward(self, x):
eps = 1e-5 if x.dtype == torch.float32 else 1e-3
var = torch.var(x, dim = 1, unbiased = False, keepdim = True) var = torch.var(x, dim = 1, unbiased = False, keepdim = True)
mean = torch.mean(x, dim = 1, keepdim = True) mean = torch.mean(x, dim = 1, keepdim = True)
return (x - mean) / (var + self.eps).sqrt() * self.g + self.b return (x - mean) * (var + eps).rsqrt() * self.g
class PreNorm(nn.Module): class PreNorm(nn.Module):
def __init__(self, dim, fn): def __init__(self, dim, fn):
@@ -252,6 +251,7 @@ class Unet(nn.Module):
out_dim = None, out_dim = None,
dim_mults=(1, 2, 4, 8), dim_mults=(1, 2, 4, 8),
channels = 3, channels = 3,
self_condition = False,
resnet_block_groups = 8, resnet_block_groups = 8,
learned_variance = False, learned_variance = False,
learned_sinusoidal_cond = False, learned_sinusoidal_cond = False,
@@ -262,9 +262,11 @@ class Unet(nn.Module):
# determine dimensions # determine dimensions
self.channels = channels self.channels = channels
self.self_condition = self_condition
input_channels = channels * (2 if self_condition else 1)
init_dim = default(init_dim, dim) 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)] dims = [init_dim, *map(lambda m: dim * m, dim_mults)]
in_out = list(zip(dims[:-1], dims[1:])) in_out = list(zip(dims[:-1], dims[1:]))
@@ -328,7 +330,11 @@ class Unet(nn.Module):
self.final_res_block = block_klass(dim * 2, dim, time_emb_dim = time_dim) self.final_res_block = block_klass(dim * 2, dim, time_emb_dim = time_dim)
self.final_conv = nn.Conv2d(dim, self.out_dim, 1) 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) x = self.init_conv(x)
r = x.clone() r = x.clone()
@@ -396,7 +402,6 @@ class GaussianDiffusion(nn.Module):
model, model,
*, *,
image_size, image_size,
channels = 3,
timesteps = 1000, timesteps = 1000,
sampling_timesteps = None, sampling_timesteps = None,
loss_type = 'l1', loss_type = 'l1',
@@ -409,9 +414,12 @@ class GaussianDiffusion(nn.Module):
super().__init__() super().__init__()
assert not (type(self) == GaussianDiffusion and model.channels != model.out_dim) assert not (type(self) == GaussianDiffusion and model.channels != model.out_dim)
self.channels = channels
self.image_size = image_size
self.model = model self.model = model
self.channels = self.model.channels
self.self_condition = self.model.self_condition
self.image_size = image_size
self.objective = objective self.objective = objective
assert objective in {'pred_noise', 'pred_x0'}, 'objective must be either pred_noise (predict noise) or pred_x0 (predict image start)' assert objective in {'pred_noise', 'pred_x0'}, 'objective must be either pred_noise (predict noise) or pred_x0 (predict image start)'
@@ -494,8 +502,8 @@ class GaussianDiffusion(nn.Module):
posterior_log_variance_clipped = extract(self.posterior_log_variance_clipped, t, x_t.shape) posterior_log_variance_clipped = extract(self.posterior_log_variance_clipped, t, x_t.shape)
return posterior_mean, posterior_variance, posterior_log_variance_clipped return posterior_mean, posterior_variance, posterior_log_variance_clipped
def model_predictions(self, x, t): def model_predictions(self, x, t, x_self_cond = None):
model_output = self.model(x, t) model_output = self.model(x, t, x_self_cond)
if self.objective == 'pred_noise': if self.objective == 'pred_noise':
pred_noise = model_output pred_noise = model_output
@@ -507,23 +515,24 @@ class GaussianDiffusion(nn.Module):
return ModelPrediction(pred_noise, x_start) return ModelPrediction(pred_noise, x_start)
def p_mean_variance(self, x, t, clip_denoised: bool): def p_mean_variance(self, x, t, x_self_cond = None, clip_denoised = True):
preds = self.model_predictions(x, t) preds = self.model_predictions(x, t, x_self_cond)
x_start = preds.pred_x_start x_start = preds.pred_x_start
if clip_denoised: if clip_denoised:
x_start.clamp_(-1., 1.) x_start.clamp_(-1., 1.)
model_mean, posterior_variance, posterior_log_variance = self.q_posterior(x_start = x_start, x_t = x, t = t) 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() @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 b, *_, device = *x.shape, x.device
batched_times = torch.full((x.shape[0],), t, device = x.device, dtype = torch.long) 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 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() @torch.no_grad()
def p_sample_loop(self, shape): def p_sample_loop(self, shape):
@@ -531,8 +540,11 @@ class GaussianDiffusion(nn.Module):
img = torch.randn(shape, device=device) img = torch.randn(shape, device=device)
for t in tqdm(reversed(range(0, self.num_timesteps)), desc = 'sampling loop time step'): x_start = None
img = self.p_sample(img, t)
for t in tqdm(reversed(range(0, self.num_timesteps)), desc = 'sampling loop time step', total = self.num_timesteps):
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) img = unnormalize_to_zero_to_one(img)
return img return img
@@ -547,13 +559,17 @@ class GaussianDiffusion(nn.Module):
img = torch.randn(shape, device = device) img = torch.randn(shape, device = device)
x_start = None
for time, time_next in tqdm(time_pairs, desc = 'sampling loop time step'): for time, time_next in tqdm(time_pairs, desc = 'sampling loop time step'):
alpha = self.alphas_cumprod_prev[time] alpha = self.alphas_cumprod_prev[time]
alpha_next = self.alphas_cumprod_prev[time_next] alpha_next = self.alphas_cumprod_prev[time_next]
time_cond = torch.full((batch,), time, device = device, dtype = torch.long) 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: if clip_denoised:
x_start.clamp_(-1., 1.) x_start.clamp_(-1., 1.)
@@ -583,11 +599,11 @@ class GaussianDiffusion(nn.Module):
assert x1.shape == x2.shape assert x1.shape == x2.shape
t_batched = torch.stack([torch.tensor(t, device=device)] * b) t_batched = torch.stack([torch.tensor(t, device = device)] * b)
xt1, xt2 = map(lambda x: self.q_sample(x, t=t_batched), (x1, x2)) xt1, xt2 = map(lambda x: self.q_sample(x, t = t_batched), (x1, x2))
img = (1 - lam) * xt1 + lam * xt2 img = (1 - lam) * xt1 + lam * xt2
for i in tqdm(reversed(range(0, t)), desc='interpolation sample time step', total=t): for i in tqdm(reversed(range(0, t)), desc = 'interpolation sample time step', total = t):
img = self.p_sample(img, torch.full((b,), i, device=device, dtype=torch.long)) img = self.p_sample(img, torch.full((b,), i, device=device, dtype=torch.long))
return img return img
@@ -613,8 +629,23 @@ class GaussianDiffusion(nn.Module):
b, c, h, w = x_start.shape b, c, h, w = x_start.shape
noise = default(noise, lambda: torch.randn_like(x_start)) noise = default(noise, lambda: torch.randn_like(x_start))
# noise sample
x = self.q_sample(x_start = x_start, t = t, noise = noise) 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': if self.objective == 'pred_noise':
target = noise target = noise
@@ -52,6 +52,7 @@ class ElucidatedDiffusion(nn.Module):
): ):
super().__init__() super().__init__()
assert net.learned_sinusoidal_cond assert net.learned_sinusoidal_cond
assert not net.self_condition, 'not supported yet'
self.net = net self.net = net
@@ -77,6 +77,8 @@ class LearnedGaussianDiffusion(GaussianDiffusion):
): ):
super().__init__(model, *args, **kwargs) 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 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 self.vb_loss_weight = vb_loss_weight
def model_predictions(self, x, t): def model_predictions(self, x, t):
@@ -31,6 +31,7 @@ class WeightedObjectiveGaussianDiffusion(GaussianDiffusion):
super().__init__(model, *args, **kwargs) super().__init__(model, *args, **kwargs)
channels = model.channels 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 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' assert not self.is_ddim_sampling, 'ddim sampling cannot be used'
self.split_dims = (channels, channels, 2) self.split_dims = (channels, channels, 2)
+1 -1
View File
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
setup( setup(
name = 'denoising-diffusion-pytorch', name = 'denoising-diffusion-pytorch',
packages = find_packages(), packages = find_packages(),
version = '0.26.4', version = '0.27.1',
license='MIT', license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch', description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang', author = 'Phil Wang',