Compare commits

..
1 Commits
Author SHA1 Message Date
Phil Wang b72235e71d allow for mixed precision training with fp16 flag 2020-09-08 17:25:26 -07:00
10 changed files with 258 additions and 1388 deletions
-3
View File
@@ -1,6 +1,3 @@
# Generation results
results/
# Byte-compiled / optimized / DLL files # Byte-compiled / optimized / DLL files
__pycache__/ __pycache__/
*.py[cod] *.py[cod]
+23 -106
View File
@@ -2,18 +2,10 @@
## Denoising Diffusion Probabilistic Model, in Pytorch ## Denoising Diffusion Probabilistic Model, in Pytorch
Implementation of <a href="https://arxiv.org/abs/2006.11239">Denoising Diffusion Probabilistic Model</a> in Pytorch. It is a new approach to generative modeling that may <a href="https://ajolicoeur.wordpress.com/the-new-contender-to-gans-score-matching-with-langevin-sampling/">have the potential</a> to rival GANs. It uses denoising score matching to estimate the gradient of the data distribution, followed by Langevin sampling to sample from the true distribution. Implementation of <a href="https://arxiv.org/abs/2006.11239">Denoising Diffusion Probabilistic Model</a> in Pytorch. It is a new approach to generative modeling that may <a href="https://ajolicoeur.wordpress.com/the-new-contender-to-gans-score-matching-with-langevin-sampling/">have the potential</a> to rival GANs. It uses denoising score matching to estimate the gradient of the data distribution, followed by Langevin sampling to sample from the true distribution. This implementation was transcribed from the official Tensorflow version <a href="https://github.com/hojonathanho/diffusion">here</a>.
This implementation was transcribed from the official Tensorflow version <a href="https://github.com/hojonathanho/diffusion">here</a>
Youtube AI Educators - <a href="https://www.youtube.com/watch?v=W-O7AZNzbzQ">Yannic Kilcher</a> | <a href="https://www.youtube.com/watch?v=344w5h24-h8">AI Coffeebreak with Letitia</a> | <a href="https://www.youtube.com/watch?v=HoKDTa5jHvg">Outlier</a>
<a href="https://huggingface.co/blog/annotated-diffusion">Annotated code</a> by Research Scientists / Engineers from <a href="https://huggingface.co/">🤗 Huggingface</a>
<img src="./sample.png" width="500px"><img> <img src="./sample.png" width="500px"><img>
[![PyPI version](https://badge.fury.io/py/denoising-diffusion-pytorch.svg)](https://badge.fury.io/py/denoising-diffusion-pytorch)
## Install ## Install
```bash ```bash
@@ -33,17 +25,18 @@ model = Unet(
diffusion = GaussianDiffusion( diffusion = GaussianDiffusion(
model, model,
image_size = 128, beta_start = 0.0001,
timesteps = 1000, # number of steps beta_end = 0.02,
loss_type = 'l1' # L1 or L2 num_diffusion_timesteps = 1000, # number of steps
loss_type = 'l1' # L1 or L2 (wavegrad paper claims l1 is better?)
) )
training_images = torch.randn(8, 3, 128, 128) # images are normalized from 0 to 1 training_images = torch.randn(8, 3, 128, 128)
loss = diffusion(training_images) loss = diffusion(training_images)
loss.backward() loss.backward()
# after a lot of training # after a lot of training
sampled_images = diffusion.sample(batch_size = 4) sampled_images = diffusion.sample(128, batch_size = 4)
sampled_images.shape # (4, 3, 128, 128) sampled_images.shape # (4, 3, 128, 128)
``` ```
@@ -59,114 +52,38 @@ model = Unet(
diffusion = GaussianDiffusion( diffusion = GaussianDiffusion(
model, model,
image_size = 128, beta_start = 0.0001,
timesteps = 1000, # number of steps beta_end = 0.02,
sampling_timesteps = 250, # number of sampling timesteps (using ddim for faster inference [see citation for ddim paper]) num_diffusion_timesteps = 1000, # number of steps
loss_type = 'l1' # L1 or L2 loss_type = 'l1' # L1 or L2
).cuda() ).cuda()
trainer = Trainer( trainer = Trainer(
diffusion, diffusion,
'path/to/your/images', 'path/to/your/images',
image_size = 128,
train_batch_size = 32, train_batch_size = 32,
train_lr = 1e-4, train_lr = 2e-5,
train_num_steps = 700000, # total training steps train_num_steps = 100000, # total training steps
gradient_accumulate_every = 2, # gradient accumulation steps gradient_accumulate_every = 2, # gradient accumulation steps
ema_decay = 0.995, # exponential moving average decay ema_decay = 0.995, # exponential moving average decay
amp = True # turn on mixed precision fp16 = True # turn on mixed precision training with apex
) )
trainer.train() trainer.train()
``` ```
Samples and model checkpoints will be logged to `./results` periodically Todo: Command line tool for one-line training
## Multi-GPU Training
The `Trainer` class is now equipped with <a href="https://huggingface.co/docs/accelerate/accelerator">🤗 Accelerator</a>. You can easily do multi-gpu training in two steps using their `accelerate` CLI
At the project root directory, where the training script is, run
```python
$ accelerate config
```
Then, in the same directory
```python
$ accelerate launch train.py
```
## Citations ## Citations
```bibtex ```bibtex
@inproceedings{NEURIPS2020_4c5bcfec, @misc{ho2020denoising,
author = {Ho, Jonathan and Jain, Ajay and Abbeel, Pieter}, title={Denoising Diffusion Probabilistic Models},
booktitle = {Advances in Neural Information Processing Systems}, author={Jonathan Ho and Ajay Jain and Pieter Abbeel},
editor = {H. Larochelle and M. Ranzato and R. Hadsell and M.F. Balcan and H. Lin}, year={2020},
pages = {6840--6851}, eprint={2006.11239},
publisher = {Curran Associates, Inc.}, archivePrefix={arXiv},
title = {Denoising Diffusion Probabilistic Models}, primaryClass={cs.LG}
url = {https://proceedings.neurips.cc/paper/2020/file/4c5bcfec8584af0d967f1ab10179ca4b-Paper.pdf},
volume = {33},
year = {2020}
}
```
```bibtex
@InProceedings{pmlr-v139-nichol21a,
title = {Improved Denoising Diffusion Probabilistic Models},
author = {Nichol, Alexander Quinn and Dhariwal, Prafulla},
booktitle = {Proceedings of the 38th International Conference on Machine Learning},
pages = {8162--8171},
year = {2021},
editor = {Meila, Marina and Zhang, Tong},
volume = {139},
series = {Proceedings of Machine Learning Research},
month = {18--24 Jul},
publisher = {PMLR},
pdf = {http://proceedings.mlr.press/v139/nichol21a/nichol21a.pdf},
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}
}
```
```bibtex
@article{Choi2022PerceptionPT,
title = {Perception Prioritized Training of Diffusion Models},
author = {Jooyoung Choi and Jungbeom Lee and Chaehun Shin and Sungwon Kim and Hyunwoo J. Kim and Sung-Hoon Yoon},
journal = {ArXiv},
year = {2022},
volume = {abs/2204.00227}
}
```
```bibtex
@article{Karras2022ElucidatingTD,
title = {Elucidating the Design Space of Diffusion-Based Generative Models},
author = {Tero Karras and Miika Aittala and Timo Aila and Samuli Laine},
journal = {ArXiv},
year = {2022},
volume = {abs/2206.00364}
}
```
```bibtex
@article{Song2021DenoisingDI,
title = {Denoising Diffusion Implicit Models},
author = {Jiaming Song and Chenlin Meng and Stefano Ermon},
journal = {ArXiv},
year = {2021},
volume = {abs/2010.02502}
} }
``` ```
-5
View File
@@ -1,6 +1 @@
from denoising_diffusion_pytorch.denoising_diffusion_pytorch import GaussianDiffusion, Unet, Trainer 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
from denoising_diffusion_pytorch.elucidated_diffusion import ElucidatedDiffusion
@@ -1,287 +0,0 @@
import math
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, reduce
from einops.layers.torch import Rearrange
# 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))
# neural net helpers
class Residual(nn.Module):
def __init__(self, fn):
super().__init__()
self.fn = fn
def forward(self, x):
return x + self.fn(x)
class MonotonicLinear(nn.Module):
def __init__(self, *args, **kwargs):
super().__init__()
self.net = nn.Linear(*args, **kwargs)
def forward(self, x):
return F.linear(x, self.net.weight.abs(), self.net.bias.abs())
# 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 log(t, eps = 1e-20):
return torch.log(t.clamp(min = eps))
def beta_linear_log_snr(t):
return -log(expm1(1e-4 + 10 * (t ** 2)))
def alpha_cosine_log_snr(t, s = 0.008):
return -log((torch.cos((t + s) / (1 + s) * math.pi * 0.5) ** -2) - 1, eps = 1e-5)
class learned_noise_schedule(nn.Module):
""" described in section H and then I.2 of the supplementary material for variational ddpm paper """
def __init__(
self,
*,
log_snr_max,
log_snr_min,
hidden_dim = 1024,
frac_gradient = 1.
):
super().__init__()
self.slope = log_snr_min - log_snr_max
self.intercept = log_snr_max
self.net = nn.Sequential(
Rearrange('... -> ... 1'),
MonotonicLinear(1, 1),
Residual(nn.Sequential(
MonotonicLinear(1, hidden_dim),
nn.Sigmoid(),
MonotonicLinear(hidden_dim, 1)
)),
Rearrange('... 1 -> ...'),
)
self.frac_gradient = frac_gradient
def forward(self, x):
frac_gradient = self.frac_gradient
device = x.device
out_zero = self.net(torch.zeros_like(x))
out_one = self.net(torch.ones_like(x))
x = self.net(x)
normed = self.slope * ((x - out_zero) / (out_one - out_zero)) + self.intercept
return normed * frac_gradient + normed.detach() * (1 - frac_gradient)
class ContinuousTimeGaussianDiffusion(nn.Module):
def __init__(
self,
model,
*,
image_size,
channels = 3,
loss_type = 'l1',
noise_schedule = 'linear',
num_sample_steps = 500,
clip_sample_denoised = True,
learned_schedule_net_hidden_dim = 1024,
learned_noise_schedule_frac_gradient = 1., # between 0 and 1, determines what percentage of gradients go back, so one can update the learned noise schedule more slowly
p2_loss_weight_gamma = 0., # p2 loss weight, from https://arxiv.org/abs/2204.00227 - 0 is equivalent to weight of 1 across time
p2_loss_weight_k = 1
):
super().__init__()
assert model.learned_sinusoidal_cond
self.model = model
# image dimensions
self.channels = channels
self.image_size = image_size
# continuous noise schedule related stuff
self.loss_type = loss_type
if noise_schedule == 'linear':
self.log_snr = beta_linear_log_snr
elif noise_schedule == 'cosine':
self.log_snr = alpha_cosine_log_snr
elif noise_schedule == 'learned':
log_snr_max, log_snr_min = [beta_linear_log_snr(torch.tensor([time])).item() for time in (0., 1.)]
self.log_snr = learned_noise_schedule(
log_snr_max = log_snr_max,
log_snr_min = log_snr_min,
hidden_dim = learned_schedule_net_hidden_dim,
frac_gradient = learned_noise_schedule_frac_gradient
)
else:
raise ValueError(f'unknown noise schedule {noise_schedule}')
# sampling
self.num_sample_steps = num_sample_steps
self.clip_sample_denoised = clip_sample_denoised
# p2 loss weight
# proposed https://arxiv.org/abs/2204.00227
assert p2_loss_weight_gamma <= 2, 'in paper, they noticed any gamma greater than 2 is harmful'
self.p2_loss_weight_gamma = p2_loss_weight_gamma # recommended to be 0.5 or 1
self.p2_loss_weight_k = p2_loss_weight_k
@property
def device(self):
return next(self.model.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&noteId=rIQgH0zKsRt
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()
alpha, sigma, alpha_next = map(sqrt, (squared_alpha, squared_sigma, squared_alpha_next))
batch_log_snr = repeat(log_snr, ' -> b', b = x.shape[0])
pred_noise = self.model(x, batch_log_snr)
if self.clip_sample_denoised:
x_start = (x - sigma * pred_noise) / alpha
# in Imagen, this was changed to dynamic thresholding
x_start.clamp_(-1., 1.)
model_mean = alpha_next * (x * (1 - c) / alpha + c * x_start)
else:
model_mean = alpha_next / alpha * (x - c * 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)
if time_next == 0:
return model_mean
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.clamp_(-1., 1.)
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.model(x, log_snr)
losses = self.loss_fn(model_out, noise, reduction = 'none')
losses = reduce(losses, 'b ... -> b', 'mean')
if self.p2_loss_weight_gamma >= 0:
# following eq 8. in https://arxiv.org/abs/2204.00227
loss_weight = (self.p2_loss_weight_k + log_snr.exp()) ** -self.p2_loss_weight_gamma
losses = losses * loss_weight
return losses.mean()
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)
@@ -4,28 +4,29 @@ 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 inspect import isfunction
from collections import namedtuple
from functools import partial from functools import partial
from torch.utils.data import Dataset, DataLoader from torch.utils import data
from multiprocessing import cpu_count
from pathlib import Path 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, utils
from PIL import Image from PIL import Image
from einops import rearrange, reduce import numpy as np
from einops.layers.torch import Rearrange from tqdm import tqdm
from einops import rearrange
from tqdm.auto import tqdm try:
from ema_pytorch import EMA from apex import amp
APEX_AVAILABLE = True
from accelerate import Accelerator except:
APEX_AVAILABLE = False
# constants # constants
ModelPrediction = namedtuple('ModelPrediction', ['pred_noise', 'pred_x_start']) SAVE_AND_SAMPLE_EVERY = 1000
UPDATE_EMA_EVERY = 10
EXTS = ['jpg', 'png']
# helpers functions # helpers functions
@@ -42,32 +43,30 @@ def cycle(dl):
for data in dl: for data in dl:
yield data yield data
def has_int_squareroot(num): def loss_backwards(fp16, loss, optimizer, **kwargs):
return (math.sqrt(num) ** 2) == num if fp16:
with amp.scale_loss(loss, optimizer) as scaled_loss:
def num_to_groups(num, divisor): scaled_loss.backward(**kwargs)
groups = num // divisor else:
remainder = num % divisor loss.backward(**kwargs)
arr = [divisor] * groups
if remainder > 0:
arr.append(remainder)
return arr
def convert_image_to(img_type, image):
if image.mode != img_type:
return image.convert(img_type)
return image
# 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
# small helper modules # small helper modules
class EMA():
def __init__(self, beta):
super().__init__()
self.beta = beta
def update_model_average(self, ma_model, current_model):
for current_params, ma_params in zip(current_model.parameters(), ma_model.parameters()):
old_weight, up_weight = ma_params.data, current_params.data
ma_params.data = self.update_average(old_weight, up_weight)
def update_average(self, old, new):
if old is None:
return new
return old * self.beta + (1 - self.beta) * new
class Residual(nn.Module): class Residual(nn.Module):
def __init__(self, fn): def __init__(self, fn):
super().__init__() super().__init__()
@@ -76,39 +75,6 @@ class Residual(nn.Module):
def forward(self, x, *args, **kwargs): def forward(self, x, *args, **kwargs):
return self.fn(x, *args, **kwargs) + x return self.fn(x, *args, **kwargs) + x
def Upsample(dim, dim_out = None):
return nn.Sequential(
nn.Upsample(scale_factor = 2, mode = 'nearest'),
nn.Conv2d(dim, default(dim_out, dim), 3, padding = 1)
)
def Downsample(dim, dim_out = None):
return nn.Conv2d(dim, default(dim_out, dim), 4, 2, 1)
class LayerNorm(nn.Module):
def __init__(self, dim, eps = 1e-5):
super().__init__()
self.eps = eps
self.g = nn.Parameter(torch.ones(1, dim, 1, 1))
self.b = nn.Parameter(torch.zeros(1, dim, 1, 1))
def forward(self, x):
var = torch.var(x, dim = 1, unbiased = False, keepdim = True)
mean = torch.mean(x, dim = 1, keepdim = True)
return (x - mean) / (var + self.eps).sqrt() * self.g + self.b
class PreNorm(nn.Module):
def __init__(self, dim, fn):
super().__init__()
self.fn = fn
self.norm = LayerNorm(dim)
def forward(self, x):
x = self.norm(x)
return self.fn(x)
# sinusoidal positional embeds
class SinusoidalPosEmb(nn.Module): class SinusoidalPosEmb(nn.Module):
def __init__(self, dim): def __init__(self, dim):
super().__init__() super().__init__()
@@ -123,171 +89,99 @@ class SinusoidalPosEmb(nn.Module):
emb = torch.cat((emb.sin(), emb.cos()), dim=-1) emb = torch.cat((emb.sin(), emb.cos()), dim=-1)
return emb return emb
class LearnedSinusoidalPosEmb(nn.Module): class Mish(nn.Module):
""" following @crowsonkb 's lead with learned sinusoidal pos emb """ def forward(self, x):
""" https://github.com/crowsonkb/v-diffusion-jax/blob/master/diffusion/models/danbooru_128.py#L8 """ return x * torch.tanh(F.softplus(x))
class Upsample(nn.Module):
def __init__(self, dim): def __init__(self, dim):
super().__init__() super().__init__()
assert (dim % 2) == 0 self.conv = nn.ConvTranspose2d(dim, dim, 4, 2, 1)
half_dim = dim // 2
self.weights = nn.Parameter(torch.randn(half_dim))
def forward(self, x): def forward(self, x):
x = rearrange(x, 'b -> b 1') return self.conv(x)
freqs = x * rearrange(self.weights, 'd -> 1 d') * 2 * math.pi
fouriered = torch.cat((freqs.sin(), freqs.cos()), dim = -1) class Downsample(nn.Module):
fouriered = torch.cat((x, fouriered), dim = -1) def __init__(self, dim):
return fouriered super().__init__()
self.conv = nn.Conv2d(dim, dim, 3, 2, 1)
def forward(self, x):
return self.conv(x)
class Rezero(nn.Module):
def __init__(self, dim):
super().__init__()
self.g = nn.Parameter(torch.zeros(1))
def forward(self, x):
return x * self.g
# building block modules # building block modules
class Block(nn.Module): class Block(nn.Module):
def __init__(self, dim, dim_out, groups = 8): def __init__(self, dim, dim_out, groups = 8):
super().__init__() super().__init__()
self.proj = nn.Conv2d(dim, dim_out, 3, padding = 1) self.block = nn.Sequential(
self.norm = nn.GroupNorm(groups, dim_out) nn.Conv2d(dim, dim_out, 3, padding=1),
self.act = nn.SiLU() nn.GroupNorm(groups, dim_out),
Mish()
def forward(self, x, scale_shift = None): )
x = self.proj(x) def forward(self, x):
x = self.norm(x) return self.block(x)
if exists(scale_shift):
scale, shift = scale_shift
x = x * (scale + 1) + shift
x = self.act(x)
return x
class ResnetBlock(nn.Module): class ResnetBlock(nn.Module):
def __init__(self, dim, dim_out, *, time_emb_dim = None, groups = 8): def __init__(self, dim, dim_out, *, time_emb_dim, groups = 8):
super().__init__() super().__init__()
self.mlp = nn.Sequential( self.mlp = nn.Sequential(
nn.SiLU(), Mish(),
nn.Linear(time_emb_dim, dim_out * 2) nn.Linear(time_emb_dim, dim_out)
) if exists(time_emb_dim) else None )
self.block1 = Block(dim, dim_out, groups = groups) self.block1 = Block(dim, dim_out)
self.block2 = Block(dim_out, dim_out, groups = groups) self.block2 = Block(dim_out, dim_out)
self.res_conv = nn.Conv2d(dim, dim_out, 1) if dim != dim_out else nn.Identity() self.res_conv = nn.Conv2d(dim, dim_out, 1) if dim != dim_out else nn.Identity()
def forward(self, x, time_emb = None): def forward(self, x, time_emb):
h = self.block1(x)
scale_shift = None h += self.mlp(time_emb)[:, :, None, None]
if exists(self.mlp) and exists(time_emb):
time_emb = self.mlp(time_emb)
time_emb = rearrange(time_emb, 'b c -> b c 1 1')
scale_shift = time_emb.chunk(2, dim = 1)
h = self.block1(x, scale_shift = scale_shift)
h = self.block2(h) h = self.block2(h)
return h + self.res_conv(x) return h + self.res_conv(x)
class LinearAttention(nn.Module): class LinearAttention(nn.Module):
def __init__(self, dim, heads = 4, dim_head = 32): def __init__(self, dim, heads = 8, dim_head = 32):
super().__init__() super().__init__()
self.scale = dim_head ** -0.5
self.heads = heads self.heads = heads
hidden_dim = dim_head * heads hidden_dim = dim_head * heads
self.to_qkv = nn.Conv2d(dim, hidden_dim * 3, 1, bias = False) self.to_qkv = nn.Conv2d(dim, hidden_dim, 1, bias = False)
self.to_out = nn.Sequential(
nn.Conv2d(hidden_dim, dim, 1),
LayerNorm(dim)
)
def forward(self, x):
b, c, h, w = x.shape
qkv = self.to_qkv(x).chunk(3, dim = 1)
q, k, v = map(lambda t: rearrange(t, 'b (h c) x y -> b h c (x y)', h = self.heads), qkv)
q = q.softmax(dim = -2)
k = k.softmax(dim = -1)
q = q * self.scale
context = torch.einsum('b h d n, b h e n -> b h d e', k, v)
out = torch.einsum('b h d e, b h d n -> b h e n', context, q)
out = rearrange(out, 'b h c (x y) -> b (h c) x y', h = self.heads, x = h, y = w)
return self.to_out(out)
class Attention(nn.Module):
def __init__(self, dim, heads = 4, dim_head = 32):
super().__init__()
self.scale = dim_head ** -0.5
self.heads = heads
hidden_dim = dim_head * heads
self.to_qkv = nn.Conv2d(dim, hidden_dim * 3, 1, bias = False)
self.to_out = nn.Conv2d(hidden_dim, dim, 1) self.to_out = nn.Conv2d(hidden_dim, dim, 1)
def forward(self, x): def forward(self, x):
b, c, h, w = x.shape b, c, h, w = x.shape
qkv = self.to_qkv(x).chunk(3, dim = 1) qkv = self.to_qkv(x)
q, k, v = map(lambda t: rearrange(t, 'b (h c) x y -> b h c (x y)', h = self.heads), qkv) q, k, v = rearrange(qkv, 'b (qkv heads c) h w -> qkv b heads c (h w)', heads = self.heads)
q = q * self.scale q = q.softmax(dim=-2)
k = k.softmax(dim=-1)
sim = einsum('b h d i, b h d j -> b h i j', q, k) context = torch.einsum('bhdn,bhen->bhde', k, v)
sim = sim - sim.amax(dim = -1, keepdim = True).detach() out = torch.einsum('bhde,bhdn->bhen', context, q)
attn = sim.softmax(dim = -1) out = rearrange(out, 'b heads c (h w) -> b (heads c) h w', heads=self.heads, h=h, w=w)
out = einsum('b h i j, b h d j -> b h i d', attn, v)
out = rearrange(out, 'b h (x y) d -> b (h d) x y', x = h, y = w)
return self.to_out(out) return self.to_out(out)
# model # model
class Unet(nn.Module): class Unet(nn.Module):
def __init__( def __init__(self, dim, out_dim = None, dim_mults=(1, 2, 4, 8), groups = 8):
self,
dim,
init_dim = None,
out_dim = None,
dim_mults=(1, 2, 4, 8),
channels = 3,
resnet_block_groups = 8,
learned_variance = False,
learned_sinusoidal_cond = False,
learned_sinusoidal_dim = 16
):
super().__init__() super().__init__()
dims = [3, *map(lambda m: dim * m, dim_mults)]
# determine dimensions
self.channels = channels
init_dim = default(init_dim, dim)
self.init_conv = nn.Conv2d(channels, init_dim, 7, padding = 3)
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:]))
block_klass = partial(ResnetBlock, groups = resnet_block_groups) self.time_pos_emb = SinusoidalPosEmb(dim)
self.mlp = nn.Sequential(
# time embeddings nn.Linear(dim, dim * 4),
Mish(),
time_dim = dim * 4 nn.Linear(dim * 4, dim)
self.learned_sinusoidal_cond = learned_sinusoidal_cond
if learned_sinusoidal_cond:
sinu_pos_emb = LearnedSinusoidalPosEmb(learned_sinusoidal_dim)
fourier_dim = learned_sinusoidal_dim + 1
else:
sinu_pos_emb = SinusoidalPosEmb(dim)
fourier_dim = dim
self.time_mlp = nn.Sequential(
sinu_pos_emb,
nn.Linear(fourier_dim, time_dim),
nn.GELU(),
nn.Linear(time_dim, time_dim)
) )
# layers
self.downs = nn.ModuleList([]) self.downs = nn.ModuleList([])
self.ups = nn.ModuleList([]) self.ups = nn.ModuleList([])
num_resolutions = len(in_out) num_resolutions = len(in_out)
@@ -296,68 +190,57 @@ class Unet(nn.Module):
is_last = ind >= (num_resolutions - 1) is_last = ind >= (num_resolutions - 1)
self.downs.append(nn.ModuleList([ self.downs.append(nn.ModuleList([
block_klass(dim_in, dim_in, time_emb_dim = time_dim), ResnetBlock(dim_in, dim_out, time_emb_dim = dim),
block_klass(dim_in, dim_in, time_emb_dim = time_dim), ResnetBlock(dim_out, dim_out, time_emb_dim = dim),
Residual(PreNorm(dim_in, LinearAttention(dim_in))), Residual(Rezero(LinearAttention(dim_out))),
Downsample(dim_in, dim_out) if not is_last else nn.Conv2d(dim_in, dim_out, 3, padding = 1) Downsample(dim_out) if not is_last else nn.Identity()
])) ]))
mid_dim = dims[-1] mid_dim = dims[-1]
self.mid_block1 = block_klass(mid_dim, mid_dim, time_emb_dim = time_dim) self.mid_block1 = ResnetBlock(mid_dim, mid_dim, time_emb_dim = dim)
self.mid_attn = Residual(PreNorm(mid_dim, Attention(mid_dim))) self.mid_attn = Residual(Rezero(LinearAttention(mid_dim)))
self.mid_block2 = block_klass(mid_dim, mid_dim, time_emb_dim = time_dim) self.mid_block2 = ResnetBlock(mid_dim, mid_dim, time_emb_dim = dim)
for ind, (dim_in, dim_out) in enumerate(reversed(in_out)): for ind, (dim_in, dim_out) in enumerate(reversed(in_out[1:])):
is_last = ind == (len(in_out) - 1) is_last = ind >= (num_resolutions - 1)
self.ups.append(nn.ModuleList([ self.ups.append(nn.ModuleList([
block_klass(dim_out + dim_in, dim_out, time_emb_dim = time_dim), ResnetBlock(dim_out * 2, dim_in, time_emb_dim = dim),
block_klass(dim_out + dim_in, dim_out, time_emb_dim = time_dim), ResnetBlock(dim_in, dim_in, time_emb_dim = dim),
Residual(PreNorm(dim_out, LinearAttention(dim_out))), Residual(Rezero(LinearAttention(dim_in))),
Upsample(dim_out, dim_in) if not is_last else nn.Conv2d(dim_out, dim_in, 3, padding = 1) Upsample(dim_in) if not is_last else nn.Identity()
])) ]))
default_out_dim = channels * (1 if not learned_variance else 2) out_dim = default(out_dim, 3)
self.out_dim = default(out_dim, default_out_dim) self.final_conv = nn.Sequential(
Block(dim, dim),
self.final_res_block = block_klass(dim * 2, dim, time_emb_dim = time_dim) nn.Conv2d(dim, 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.init_conv(x) t = self.time_pos_emb(time)
r = x.clone() t = self.mlp(t)
t = self.time_mlp(time)
h = [] h = []
for block1, block2, attn, downsample in self.downs: for resnet, resnet2, attn, downsample in self.downs:
x = block1(x, t) x = resnet(x, t)
h.append(x) x = resnet2(x, t)
x = block2(x, t)
x = attn(x) x = attn(x)
h.append(x) h.append(x)
x = downsample(x) x = downsample(x)
x = self.mid_block1(x, t) x = self.mid_block1(x, t)
x = self.mid_attn(x) x = self.mid_attn(x)
x = self.mid_block2(x, t) x = self.mid_block2(x, t)
for block1, block2, attn, upsample in self.ups: for resnet, resnet2, attn, upsample in self.ups:
x = torch.cat((x, h.pop()), dim = 1) x = torch.cat((x, h.pop()), dim=1)
x = block1(x, t) x = resnet(x, t)
x = resnet2(x, t)
x = torch.cat((x, h.pop()), dim = 1)
x = block2(x, t)
x = attn(x) x = attn(x)
x = upsample(x) x = upsample(x)
x = torch.cat((x, r), dim = 1)
x = self.final_res_block(x, t)
return self.final_conv(x) return self.final_conv(x)
# gaussian diffusion trainer class # gaussian diffusion trainer class
@@ -367,106 +250,58 @@ def extract(a, t, x_shape):
out = a.gather(-1, t) out = a.gather(-1, t)
return out.reshape(b, *((1,) * (len(x_shape) - 1))) return out.reshape(b, *((1,) * (len(x_shape) - 1)))
def linear_beta_schedule(timesteps): def noise_like(shape, device, repeat=False):
scale = 1000 / timesteps repeat_noise = lambda: torch.randn((1, *shape[1:]), device=device).repeat(shape[0], *((1,) * (len(shape) - 1)))
beta_start = scale * 0.0001 noise = lambda: torch.randn(shape, device=device)
beta_end = scale * 0.02 return repeat_noise() if repeat else noise()
return torch.linspace(beta_start, beta_end, timesteps, dtype = torch.float64)
def cosine_beta_schedule(timesteps, s = 0.008):
"""
cosine schedule
as proposed in https://openreview.net/forum?id=-NEXDKk8gZ
"""
steps = timesteps + 1
x = torch.linspace(0, timesteps, steps, dtype = torch.float64)
alphas_cumprod = torch.cos(((x / timesteps) + s) / (1 + s) * math.pi * 0.5) ** 2
alphas_cumprod = alphas_cumprod / alphas_cumprod[0]
betas = 1 - (alphas_cumprod[1:] / alphas_cumprod[:-1])
return torch.clip(betas, 0, 0.999)
class GaussianDiffusion(nn.Module): class GaussianDiffusion(nn.Module):
def __init__( def __init__(self, denoise_fn, beta_start=0.0001, beta_end=0.02, num_diffusion_timesteps=1000, loss_type='l1', betas = None):
self,
model,
*,
image_size,
channels = 3,
timesteps = 1000,
sampling_timesteps = None,
loss_type = 'l1',
objective = 'pred_noise',
beta_schedule = 'cosine',
p2_loss_weight_gamma = 0., # p2 loss weight, from https://arxiv.org/abs/2204.00227 - 0 is equivalent to weight of 1 across time - 1. is recommended
p2_loss_weight_k = 1,
ddim_sampling_eta = 1.
):
super().__init__() super().__init__()
assert not (type(self) == GaussianDiffusion and model.channels != model.out_dim) self.denoise_fn = denoise_fn
self.channels = channels if exists(betas):
self.image_size = image_size self.np_betas = betas.detach().cpu().numpy() if isinstance(betas, torch.Tensor) else betas
self.model = model
self.objective = objective
assert objective in {'pred_noise', 'pred_x0'}, 'objective must be either pred_noise (predict noise) or pred_x0 (predict image start)'
if beta_schedule == 'linear':
betas = linear_beta_schedule(timesteps)
elif beta_schedule == 'cosine':
betas = cosine_beta_schedule(timesteps)
else: else:
raise ValueError(f'unknown beta schedule {beta_schedule}') self.np_betas = betas = np.linspace(beta_start, beta_end, num_diffusion_timesteps).astype(np.float64)
alphas = 1. - betas
alphas_cumprod = torch.cumprod(alphas, axis=0)
alphas_cumprod_prev = F.pad(alphas_cumprod[:-1], (1, 0), value = 1.)
timesteps, = betas.shape timesteps, = betas.shape
self.num_timesteps = int(timesteps) self.num_timesteps = int(timesteps)
self.loss_type = loss_type self.loss_type = loss_type
# sampling related parameters alphas = 1. - betas
alphas_cumprod = np.cumprod(alphas, axis=0)
alphas_cumprod_prev = np.append(1., alphas_cumprod[:-1])
self.sampling_timesteps = default(sampling_timesteps, timesteps) # default num sampling timesteps to number of timesteps at training to_torch = partial(torch.tensor, dtype=torch.float32)
assert self.sampling_timesteps <= timesteps self.register_buffer('betas', to_torch(betas))
self.is_ddim_sampling = self.sampling_timesteps < timesteps self.register_buffer('alphas_cumprod', to_torch(alphas_cumprod))
self.ddim_sampling_eta = ddim_sampling_eta self.register_buffer('alphas_cumprod_prev', to_torch(alphas_cumprod_prev))
# helper function to register buffer from float64 to float32
register_buffer = lambda name, val: self.register_buffer(name, val.to(torch.float32))
register_buffer('betas', betas)
register_buffer('alphas_cumprod', alphas_cumprod)
register_buffer('alphas_cumprod_prev', alphas_cumprod_prev)
# calculations for diffusion q(x_t | x_{t-1}) and others # calculations for diffusion q(x_t | x_{t-1}) and others
self.register_buffer('sqrt_alphas_cumprod', to_torch(np.sqrt(alphas_cumprod)))
register_buffer('sqrt_alphas_cumprod', torch.sqrt(alphas_cumprod)) self.register_buffer('sqrt_one_minus_alphas_cumprod', to_torch(np.sqrt(1. - alphas_cumprod)))
register_buffer('sqrt_one_minus_alphas_cumprod', torch.sqrt(1. - alphas_cumprod)) self.register_buffer('log_one_minus_alphas_cumprod', to_torch(np.log(1. - alphas_cumprod)))
register_buffer('log_one_minus_alphas_cumprod', torch.log(1. - alphas_cumprod)) self.register_buffer('sqrt_recip_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod)))
register_buffer('sqrt_recip_alphas_cumprod', torch.sqrt(1. / alphas_cumprod)) self.register_buffer('sqrt_recipm1_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod - 1)))
register_buffer('sqrt_recipm1_alphas_cumprod', torch.sqrt(1. / alphas_cumprod - 1))
# calculations for posterior q(x_{t-1} | x_t, x_0) # calculations for posterior q(x_{t-1} | x_t, x_0)
posterior_variance = betas * (1. - alphas_cumprod_prev) / (1. - alphas_cumprod) posterior_variance = betas * (1. - alphas_cumprod_prev) / (1. - alphas_cumprod)
# above: equal to 1. / (1. / (1. - alpha_cumprod_tm1) + alpha_t / beta_t) # above: equal to 1. / (1. / (1. - alpha_cumprod_tm1) + alpha_t / beta_t)
self.register_buffer('posterior_variance', to_torch(posterior_variance))
register_buffer('posterior_variance', posterior_variance)
# below: log calculation clipped because the posterior variance is 0 at the beginning of the diffusion chain # below: log calculation clipped because the posterior variance is 0 at the beginning of the diffusion chain
self.register_buffer('posterior_log_variance_clipped', to_torch(np.log(np.maximum(posterior_variance, 1e-20))))
self.register_buffer('posterior_mean_coef1', to_torch(
betas * np.sqrt(alphas_cumprod_prev) / (1. - alphas_cumprod)))
self.register_buffer('posterior_mean_coef2', to_torch(
(1. - alphas_cumprod_prev) * np.sqrt(alphas) / (1. - alphas_cumprod)))
register_buffer('posterior_log_variance_clipped', torch.log(posterior_variance.clamp(min =1e-20))) def q_mean_variance(self, x_start, t):
register_buffer('posterior_mean_coef1', betas * torch.sqrt(alphas_cumprod_prev) / (1. - alphas_cumprod)) mean = extract(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start
register_buffer('posterior_mean_coef2', (1. - alphas_cumprod_prev) * torch.sqrt(alphas) / (1. - alphas_cumprod)) variance = extract(1. - self.alphas_cumprod, t, x_start.shape)
log_variance = extract(self.log_one_minus_alphas_cumprod, t, x_start.shape)
# calculate p2 reweighting return mean, variance, log_variance
register_buffer('p2_loss_weight', (p2_loss_weight_k + alphas_cumprod / (1 - alphas_cumprod)) ** -p2_loss_weight_gamma)
def predict_start_from_noise(self, x_t, t, noise): def predict_start_from_noise(self, x_t, t, noise):
return ( return (
@@ -474,12 +309,6 @@ class GaussianDiffusion(nn.Module):
extract(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape) * noise extract(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape) * noise
) )
def predict_noise_from_start(self, x_t, t, x0):
return (
(x0 - extract(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t) / \
extract(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape)
)
def q_posterior(self, x_start, x_t, t): def q_posterior(self, x_start, x_t, t):
posterior_mean = ( posterior_mean = (
extract(self.posterior_mean_coef1, t, x_t.shape) * x_start + extract(self.posterior_mean_coef1, t, x_t.shape) * x_start +
@@ -489,87 +318,38 @@ 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):
model_output = self.model(x, t)
if self.objective == 'pred_noise':
pred_noise = model_output
x_start = self.predict_start_from_noise(x, t, model_output)
elif self.objective == 'pred_x0':
pred_noise = self.predict_noise_from_start(x, t, model_output)
x_start = model_output
return ModelPrediction(pred_noise, x_start)
def p_mean_variance(self, x, t, clip_denoised: bool): def p_mean_variance(self, x, t, clip_denoised: bool):
preds = self.model_predictions(x, t) x_recon = self.predict_start_from_noise(x, t=t, noise=self.denoise_fn(x, t))
x_start = preds.pred_x_start
if clip_denoised: if clip_denoised:
x_start.clamp_(-1., 1.) x_recon.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_recon, x_t=x, t=t)
return model_mean, posterior_variance, posterior_log_variance return model_mean, posterior_variance, posterior_log_variance
@torch.no_grad() @torch.no_grad()
def p_sample(self, x, t: int, clip_denoised = True): def p_sample(self, x, t, clip_denoised=True, repeat_noise=False):
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) model_mean, _, model_log_variance = self.p_mean_variance(x=x, t=t, clip_denoised=clip_denoised)
model_mean, _, model_log_variance = self.p_mean_variance(x = x, t = batched_times, clip_denoised = clip_denoised) noise = noise_like(x.shape, device, repeat_noise)
noise = torch.randn_like(x) if t > 0 else 0. # no noise if t == 0 # no noise when t == 0
return model_mean + (0.5 * model_log_variance).exp() * noise 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
@torch.no_grad() @torch.no_grad()
def p_sample_loop(self, shape): def p_sample_loop(self, shape):
batch, device = shape[0], self.betas.device device = self.betas.device
b = shape[0]
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'): for i in tqdm(reversed(range(0, self.num_timesteps)), desc='sampling loop time step', total=self.num_timesteps):
img = self.p_sample(img, t) img = self.p_sample(img, torch.full((b,), i, device=device, dtype=torch.long))
img = unnormalize_to_zero_to_one(img)
return img return img
@torch.no_grad() @torch.no_grad()
def ddim_sample(self, shape, clip_denoised = True): def sample(self, image_size, batch_size = 16):
batch, device, total_timesteps, sampling_timesteps, eta, objective = shape[0], self.betas.device, self.num_timesteps, self.sampling_timesteps, self.ddim_sampling_eta, self.objective return self.p_sample_loop((16, 3, image_size, image_size))
times = torch.linspace(0., total_timesteps, steps = sampling_timesteps + 2)[:-1]
times = list(reversed(times.int().tolist()))
time_pairs = list(zip(times[:-1], times[1:]))
img = torch.randn(shape, device = device)
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)
if clip_denoised:
x_start.clamp_(-1., 1.)
sigma = eta * ((1 - alpha / alpha_next) * (1 - alpha_next) / (1 - alpha)).sqrt()
c = ((1 - alpha_next) - sigma ** 2).sqrt()
noise = torch.randn_like(img) if time_next > 0 else 0.
img = x_start * alpha_next.sqrt() + \
c * pred_noise + \
sigma * noise
img = unnormalize_to_zero_to_one(img)
return img
@torch.no_grad()
def sample(self, batch_size = 16):
image_size, channels = self.image_size, self.channels
sample_fn = self.p_sample_loop if not self.is_ddim_sampling else self.ddim_sample
return sample_fn((batch_size, channels, image_size, image_size))
@torch.no_grad() @torch.no_grad()
def interpolate(self, x1, x2, t = None, lam = 0.5): def interpolate(self, x1, x2, t = None, lam = 0.5):
@@ -595,67 +375,41 @@ class GaussianDiffusion(nn.Module):
extract(self.sqrt_one_minus_alphas_cumprod, t, x_start.shape) * noise extract(self.sqrt_one_minus_alphas_cumprod, t, x_start.shape) * noise
) )
@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_losses(self, x_start, t, noise = None): def p_losses(self, x_start, t, noise = None):
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))
x = self.q_sample(x_start = x_start, t = t, noise = noise) x_noisy = self.q_sample(x_start=x_start, t=t, noise=noise)
model_out = self.model(x, t) x_recon = self.denoise_fn(x_noisy, t)
if self.objective == 'pred_noise': if self.loss_type == 'l1':
target = noise loss = (noise - x_recon).abs().mean()
elif self.objective == 'pred_x0': elif self.loss_type == 'l2':
target = x_start loss = F.mse_loss(noise, x_recon)
else: else:
raise ValueError(f'unknown objective {self.objective}') raise NotImplementedError()
loss = self.loss_fn(model_out, target, reduction = 'none') return loss
loss = reduce(loss, 'b ... -> b (...)', 'mean')
loss = loss * extract(self.p2_loss_weight, t, loss.shape) def forward(self, x, *args, **kwargs):
return loss.mean() b, *_, device = *x.shape, x.device
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}'
t = torch.randint(0, self.num_timesteps, (b,), device=device).long() t = torch.randint(0, self.num_timesteps, (b,), device=device).long()
return self.p_losses(x, t, *args, **kwargs)
img = normalize_to_neg_one_to_one(img)
return self.p_losses(img, t, *args, **kwargs)
# dataset classes # dataset classes
class Dataset(Dataset): class Dataset(data.Dataset):
def __init__( def __init__(self, folder, image_size):
self,
folder,
image_size,
exts = ['jpg', 'jpeg', 'png', 'tiff'],
augment_horizontal_flip = False,
convert_image_to = None
):
super().__init__() super().__init__()
self.folder = folder self.folder = folder
self.image_size = image_size self.image_size = image_size
self.paths = [p for ext in exts for p in Path(f'{folder}').glob(f'**/*.{ext}')] self.paths = [p for ext in EXTS for p in Path(f'{folder}').glob(f'**/*.{ext}')]
maybe_convert_fn = partial(convert_image_to, convert_image_to) if exists(convert_image_to) else nn.Identity() self.transform = transforms.Compose([
transforms.Resize(image_size),
self.transform = T.Compose([ transforms.RandomHorizontalFlip(),
T.Lambda(maybe_convert_fn), transforms.CenterCrop(image_size),
T.Resize(image_size), transforms.ToTensor()
T.RandomHorizontalFlip() if augment_horizontal_flip else nn.Identity(),
T.CenterCrop(image_size),
T.ToTensor()
]) ])
def __len__(self): def __len__(self):
@@ -674,137 +428,83 @@ class Trainer(object):
diffusion_model, diffusion_model,
folder, folder,
*, *,
train_batch_size = 16,
gradient_accumulate_every = 1,
augment_horizontal_flip = True,
train_lr = 1e-4,
train_num_steps = 100000,
ema_update_every = 10,
ema_decay = 0.995, ema_decay = 0.995,
save_and_sample_every = 1000, image_size = 128,
num_samples = 25, train_batch_size = 32,
results_folder = './results', train_lr = 2e-5,
amp = False, train_num_steps = 100000,
fp16 = False, gradient_accumulate_every = 2,
split_batches = True, fp16 = False
convert_image_to = None
): ):
super().__init__() super().__init__()
self.accelerator = Accelerator(
split_batches = split_batches,
mixed_precision = 'fp16' if fp16 else 'no'
)
self.accelerator.native_amp = amp
self.model = diffusion_model self.model = diffusion_model
self.ema = EMA(ema_decay)
self.ema_model = copy.deepcopy(self.model)
assert has_int_squareroot(num_samples), 'number of samples must have an integer square root' self.image_size = image_size
self.num_samples = num_samples
self.save_and_sample_every = save_and_sample_every
self.batch_size = train_batch_size
self.gradient_accumulate_every = gradient_accumulate_every self.gradient_accumulate_every = gradient_accumulate_every
self.train_num_steps = train_num_steps self.train_num_steps = train_num_steps
self.image_size = diffusion_model.image_size
# dataset and dataloader self.ds = Dataset(folder, image_size)
self.dl = cycle(data.DataLoader(self.ds, batch_size = train_batch_size, shuffle=True, pin_memory=True))
self.ds = Dataset(folder, self.image_size, augment_horizontal_flip = augment_horizontal_flip, convert_image_to = convert_image_to) self.opt = Adam(diffusion_model.parameters(), lr=train_lr)
dl = DataLoader(self.ds, batch_size = train_batch_size, shuffle = True, pin_memory = True, num_workers = cpu_count())
self.dl = cycle(dl)
# optimizer
self.opt = Adam(diffusion_model.parameters(), lr = train_lr)
# for logging results in a folder periodically
if self.accelerator.is_main_process:
self.ema = EMA(diffusion_model, beta = ema_decay, update_every = ema_update_every)
self.results_folder = Path(results_folder)
self.results_folder.mkdir(exist_ok = True)
# step counter state
self.step = 0 self.step = 0
# prepare model, dataloader, optimizer with accelerator assert not fp16 or fp16 and APEX_AVAILABLE, 'Apex must be installed in order for mixed precision training to be turned on'
self.model, self.dl, self.opt = self.accelerator.prepare(self.model, self.dl, self.opt) self.fp16 = fp16
if fp16:
(self.model, self.ema_model), self.opt = amp.initialize([self.model, self.ema_model], self.opt, opt_level='O1')
self.reset_parameters()
def reset_parameters(self):
self.ema_model.load_state_dict(self.model.state_dict())
def step_ema(self):
if self.step < 2000:
self.reset_parameters()
return
self.ema.update_model_average(self.ema_model, self.model)
def save(self, milestone): def save(self, milestone):
if not self.accelerator.is_local_main_process:
return
data = { data = {
'step': self.step, 'step': self.step,
'model': self.accelerator.get_state_dict(self.model), 'model': self.model.state_dict(),
'opt': self.opt.state_dict(), 'ema': self.ema_model.state_dict()
'ema': self.ema.state_dict(),
'scaler': self.accelerator.scaler.state_dict() if exists(self.accelerator.scaler) else None
} }
torch.save(data, f'./model-{milestone}.pt')
torch.save(data, str(self.results_folder / f'model-{milestone}.pt'))
def load(self, milestone): def load(self, milestone):
data = torch.load(str(self.results_folder / f'model-{milestone}.pt')) data = torch.load(f'./model-{milestone}.pt')
model = self.accelerator.unwrap_model(self.model)
model.load_state_dict(data['model'])
self.step = data['step'] self.step = data['step']
self.opt.load_state_dict(data['opt']) self.model.load_state_dict(data['model'])
self.ema.load_state_dict(data['ema']) self.ema_model.load_state_dict(data['ema'])
if exists(self.accelerator.scaler) and exists(data['scaler']):
self.accelerator.scaler.load_state_dict(data['scaler'])
def train(self): def train(self):
accelerator = self.accelerator backwards = partial(loss_backwards, self.fp16)
device = accelerator.device
with tqdm(initial = self.step, total = self.train_num_steps, disable = not accelerator.is_main_process) as pbar: while self.step < self.train_num_steps:
for i in range(self.gradient_accumulate_every):
data = next(self.dl).cuda()
loss = self.model(data)
print(f'{self.step}: {loss.item()}')
backwards(loss / self.gradient_accumulate_every, self.opt)
while self.step < self.train_num_steps: self.opt.step()
self.opt.zero_grad()
for _ in range(self.gradient_accumulate_every): if self.step % UPDATE_EMA_EVERY == 0:
data = next(self.dl).to(device) self.step_ema()
with self.accelerator.autocast(): if self.step % SAVE_AND_SAMPLE_EVERY == 0:
loss = self.model(data) milestone = self.step // SAVE_AND_SAMPLE_EVERY
self.accelerator.backward(loss / self.gradient_accumulate_every) all_images = self.ema_model.p_sample_loop((64, 3, self.image_size, self.image_size))
utils.save_image(all_images, f'./sample-{milestone}.png', nrow=8)
self.save(milestone)
pbar.set_description(f'loss: {loss.item():.4f}') self.step += 1
accelerator.wait_for_everyone() print('training completed')
self.opt.step()
self.opt.zero_grad()
accelerator.wait_for_everyone()
if accelerator.is_main_process:
self.ema.to(device)
self.ema.update()
if self.step != 0 and self.step % self.save_and_sample_every == 0:
self.ema.ema_model.eval()
with torch.no_grad():
milestone = self.step // self.save_and_sample_every
batches = num_to_groups(self.num_samples, self.batch_size)
all_images_list = list(map(lambda n: self.ema.ema_model.sample(batch_size=n), batches))
all_images = torch.cat(all_images_list, dim = 0)
utils.save_image(all_images, str(self.results_folder / f'sample-{milestone}.png'), nrow = int(math.sqrt(self.num_samples)))
self.save(milestone)
self.step += 1
pbar.update(1)
accelerator.print('training complete')
@@ -1,220 +0,0 @@
from math import sqrt
import torch
from torch import nn, einsum
import torch.nn.functional as F
from tqdm import tqdm
from einops import rearrange, repeat, reduce
# 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
# tensor helpers
def log(t, eps = 1e-20):
return torch.log(t.clamp(min = eps))
# 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
# main class
class ElucidatedDiffusion(nn.Module):
def __init__(
self,
net,
*,
image_size,
channels = 3,
num_sample_steps = 32, # number of sampling steps
sigma_min = 0.002, # min noise level
sigma_max = 80, # max noise level
sigma_data = 0.5, # standard deviation of data distribution
rho = 7, # controls the sampling schedule
P_mean = -1.2, # mean of log-normal distribution from which noise is drawn for training
P_std = 1.2, # standard deviation of log-normal distribution from which noise is drawn for training
S_churn = 80, # parameters for stochastic sampling - depends on dataset, Table 5 in apper
S_tmin = 0.05,
S_tmax = 50,
S_noise = 1.003,
):
super().__init__()
assert net.learned_sinusoidal_cond
self.net = net
# image dimensions
self.channels = channels
self.image_size = image_size
# parameters
self.sigma_min = sigma_min
self.sigma_max = sigma_max
self.sigma_data = sigma_data
self.rho = rho
self.P_mean = P_mean
self.P_std = P_std
self.num_sample_steps = num_sample_steps # otherwise known as N in the paper
self.S_churn = S_churn
self.S_tmin = S_tmin
self.S_tmax = S_tmax
self.S_noise = S_noise
@property
def device(self):
return next(self.net.parameters()).device
# derived preconditioning params - Table 1
def c_skip(self, sigma):
return (self.sigma_data ** 2) / (sigma ** 2 + self.sigma_data ** 2)
def c_out(self, sigma):
return sigma * self.sigma_data * (self.sigma_data ** 2 + sigma ** 2) ** -0.5
def c_in(self, sigma):
return 1 * (sigma ** 2 + self.sigma_data ** 2) ** -0.5
def c_noise(self, sigma):
return log(sigma) * 0.25
# preconditioned network output
# equation (7) in the paper
def preconditioned_network_forward(self, noised_images, sigma, clamp = False):
batch, device = noised_images.shape[0], noised_images.device
if isinstance(sigma, float):
sigma = torch.full((batch,), sigma, device = device)
padded_sigma = rearrange(sigma, 'b -> b 1 1 1')
net_out = self.net(
self.c_in(padded_sigma) * noised_images,
self.c_noise(sigma)
)
out = self.c_skip(padded_sigma) * noised_images + self.c_out(padded_sigma) * net_out
if clamp:
out = out.clamp(-1., 1.)
return out
# sampling
# sample schedule
# equation (5) in the paper
def sample_schedule(self, num_sample_steps = None):
num_sample_steps = default(num_sample_steps, self.num_sample_steps)
N = num_sample_steps
inv_rho = 1 / self.rho
steps = torch.arange(num_sample_steps, device = self.device, dtype = torch.float32)
sigmas = (self.sigma_max ** inv_rho + steps / (N - 1) * (self.sigma_min ** inv_rho - self.sigma_max ** inv_rho)) ** self.rho
sigmas = F.pad(sigmas, (0, 1), value = 0.) # last step is sigma value of 0.
return sigmas
@torch.no_grad()
def sample(self, batch_size = 16, num_sample_steps = None, clamp = True):
num_sample_steps = default(num_sample_steps, self.num_sample_steps)
shape = (batch_size, self.channels, self.image_size, self.image_size)
# get the schedule, which is returned as (sigma, gamma) tuple, and pair up with the next sigma and gamma
sigmas = self.sample_schedule(num_sample_steps)
gammas = torch.where(
(sigmas >= self.S_tmin) & (sigmas <= self.S_tmax),
min(self.S_churn / num_sample_steps, sqrt(2) - 1),
0.
)
sigmas_and_gammas = list(zip(sigmas[:-1], sigmas[1:], gammas[:-1]))
# images is noise at the beginning
init_sigma = sigmas[0]
images = init_sigma * torch.randn(shape, device = self.device)
# gradually denoise
for sigma, sigma_next, gamma in tqdm(sigmas_and_gammas, desc = 'sampling time step'):
sigma, sigma_next, gamma = map(lambda t: t.item(), (sigma, sigma_next, gamma))
eps = self.S_noise * torch.randn(shape, device = self.device) # stochastic sampling
sigma_hat = sigma + gamma * sigma
images_hat = images + sqrt(sigma_hat ** 2 - sigma ** 2) * eps
model_output = self.preconditioned_network_forward(images_hat, sigma_hat, clamp = clamp)
denoised_over_sigma = (images_hat - model_output) / sigma_hat
images_next = images_hat + (sigma_next - sigma_hat) * denoised_over_sigma
# second order correction, if not the last timestep
if sigma_next != 0:
model_output_next = self.preconditioned_network_forward(images_next, sigma_next, clamp = clamp)
denoised_prime_over_sigma = (images_next - model_output_next) / sigma_next
images_next = images_hat + 0.5 * (sigma_next - sigma_hat) * (denoised_over_sigma + denoised_prime_over_sigma)
images = images_next
images = images.clamp(-1., 1.)
return unnormalize_to_zero_to_one(images)
# training
def loss_weight(self, sigma):
return (sigma ** 2 + self.sigma_data ** 2) * (sigma * self.sigma_data) ** -2
def noise_distribution(self, batch_size):
return (self.P_mean + self.P_std * torch.randn((batch_size,), device = self.device)).exp()
def forward(self, images):
batch_size, c, h, w, device, image_size, channels = *images.shape, images.device, self.image_size, self.channels
assert h == image_size and w == image_size, f'height and width of image must be {image_size}'
assert c == channels, 'mismatch of image channels'
images = normalize_to_neg_one_to_one(images)
sigmas = self.noise_distribution(batch_size)
padded_sigmas = rearrange(sigmas, 'b -> b 1 1 1')
noise = torch.randn_like(images)
noised_images = images + padded_sigmas * noise # alphas are 1. in the paper
denoised = self.preconditioned_network_forward(noised_images, sigmas)
losses = F.mse_loss(denoised, images, reduction = 'none')
losses = reduce(losses, 'b ... -> b', 'mean')
losses = losses * self.loss_weight(sigmas)
return losses.mean()
@@ -1,149 +0,0 @@
import torch
from collections import namedtuple
from math import pi, sqrt, log as ln
from inspect import isfunction
from torch import nn, einsum
from einops import rearrange
from denoising_diffusion_pytorch.denoising_diffusion_pytorch import GaussianDiffusion, extract, unnormalize_to_zero_to_one
# constants
NAT = 1. / ln(2)
ModelPrediction = namedtuple('ModelPrediction', ['pred_noise', 'pred_x_start', 'pred_variance'])
# helper functions
def exists(x):
return x is not None
def default(val, d):
if exists(val):
return val
return d() if isfunction(d) else d
# tensor helpers
def log(t, eps = 1e-15):
return torch.log(t.clamp(min = eps))
def meanflat(x):
return x.mean(dim = tuple(range(1, len(x.shape))))
def normal_kl(mean1, logvar1, mean2, logvar2):
"""
KL divergence between normal distributions parameterized by mean and log-variance.
"""
return 0.5 * (-1.0 + logvar2 - logvar1 + torch.exp(logvar1 - logvar2) + ((mean1 - mean2) ** 2) * torch.exp(-logvar2))
def approx_standard_normal_cdf(x):
return 0.5 * (1.0 + torch.tanh(sqrt(2.0 / pi) * (x + 0.044715 * (x ** 3))))
def discretized_gaussian_log_likelihood(x, *, means, log_scales, thres = 0.999):
assert x.shape == means.shape == log_scales.shape
centered_x = x - means
inv_stdv = torch.exp(-log_scales)
plus_in = inv_stdv * (centered_x + 1. / 255.)
cdf_plus = approx_standard_normal_cdf(plus_in)
min_in = inv_stdv * (centered_x - 1. / 255.)
cdf_min = approx_standard_normal_cdf(min_in)
log_cdf_plus = log(cdf_plus)
log_one_minus_cdf_min = log(1. - cdf_min)
cdf_delta = cdf_plus - cdf_min
log_probs = torch.where(x < -thres,
log_cdf_plus,
torch.where(x > thres,
log_one_minus_cdf_min,
log(cdf_delta)))
return log_probs
# https://arxiv.org/abs/2102.09672
# i thought the results were questionable, if one were to focus only on FID
# but may as well get this in here for others to try, as GLIDE is using it (and DALL-E2 first stage of cascade)
# gaussian diffusion for learned variance + hybrid eps simple + vb loss
class LearnedGaussianDiffusion(GaussianDiffusion):
def __init__(
self,
model,
vb_loss_weight = 0.001, # lambda was 0.001 in the paper
*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`'
self.vb_loss_weight = vb_loss_weight
def model_predictions(self, x, t):
model_output = self.model(x, t)
model_output, pred_variance = model_output.chunk(2, dim = 1)
if self.objective == 'pred_noise':
pred_noise = model_output
x_start = self.predict_start_from_noise(x, t, model_output)
elif self.objective == 'pred_x0':
pred_noise = self.predict_noise_from_start(x, t, model_output)
x_start = model_output
return ModelPrediction(pred_noise, x_start, pred_variance)
def p_mean_variance(self, *, x, t, clip_denoised, model_output = None):
model_output = default(model_output, lambda: self.model(x, t))
pred_noise, var_interp_frac_unnormalized = model_output.chunk(2, dim = 1)
min_log = extract(self.posterior_log_variance_clipped, t, x.shape)
max_log = extract(torch.log(self.betas), t, x.shape)
var_interp_frac = unnormalize_to_zero_to_one(var_interp_frac_unnormalized)
model_log_variance = var_interp_frac * max_log + (1 - var_interp_frac) * min_log
model_variance = model_log_variance.exp()
x_start = self.predict_start_from_noise(x, t, pred_noise)
if clip_denoised:
x_start.clamp_(-1., 1.)
model_mean, _, _ = self.q_posterior(x_start, x, t)
return model_mean, model_variance, model_log_variance
def p_losses(self, x_start, t, noise = None, clip_denoised = False):
noise = default(noise, lambda: torch.randn_like(x_start))
x_t = self.q_sample(x_start = x_start, t = t, noise = noise)
# model output
model_output = self.model(x_t, t)
# calculating kl loss for learned variance (interpolation)
true_mean, _, true_log_variance_clipped = self.q_posterior(x_start = x_start, x_t = x_t, t = t)
model_mean, _, model_log_variance = self.p_mean_variance(x = x_t, t = t, clip_denoised = clip_denoised, model_output = model_output)
# kl loss with detached model predicted mean, for stability reasons as in paper
detached_model_mean = model_mean.detach()
kl = normal_kl(true_mean, true_log_variance_clipped, detached_model_mean, model_log_variance)
kl = meanflat(kl) * NAT
decoder_nll = -discretized_gaussian_log_likelihood(x_start, means = detached_model_mean, log_scales = 0.5 * model_log_variance)
decoder_nll = meanflat(decoder_nll) * NAT
# at the first timestep return the decoder NLL, otherwise return KL(q(x_{t-1}|x_t,x_0) || p(x_{t-1}|x_t))
vb_losses = torch.where(t == 0, decoder_nll, kl)
# simple loss - predicting noise, x0, or x_prev
pred_noise, _ = model_output.chunk(2, dim = 1)
simple_losses = self.loss_fn(pred_noise, noise)
return simple_losses + vb_losses.mean() * self.vb_loss_weight
@@ -1,81 +0,0 @@
import torch
from inspect import isfunction
from torch import nn, einsum
from einops import rearrange
from denoising_diffusion_pytorch.denoising_diffusion_pytorch import GaussianDiffusion
# helper functions
def exists(x):
return x is not None
def default(val, d):
if exists(val):
return val
return d() if isfunction(d) else d
# some improvisation on my end
# where i have the model learn to both predict noise and x0
# and learn the weighted sum for each depending on time step
class WeightedObjectiveGaussianDiffusion(GaussianDiffusion):
def __init__(
self,
model,
*args,
pred_noise_loss_weight = 0.1,
pred_x_start_loss_weight = 0.1,
**kwargs
):
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 self.is_ddim_sampling, 'ddim sampling cannot be used'
self.split_dims = (channels, channels, 2)
self.pred_noise_loss_weight = pred_noise_loss_weight
self.pred_x_start_loss_weight = pred_x_start_loss_weight
def p_mean_variance(self, *, x, t, clip_denoised, model_output = None):
model_output = self.model(x, t)
pred_noise, pred_x_start, weights = model_output.split(self.split_dims, dim = 1)
normalized_weights = weights.softmax(dim = 1)
x_start_from_noise = self.predict_start_from_noise(x, t = t, noise = pred_noise)
x_starts = torch.stack((x_start_from_noise, pred_x_start), dim = 1)
weighted_x_start = einsum('b j h w, b j c h w -> b c h w', normalized_weights, x_starts)
if clip_denoised:
weighted_x_start.clamp_(-1., 1.)
model_mean, model_variance, model_log_variance = self.q_posterior(weighted_x_start, x, t)
return model_mean, model_variance, model_log_variance
def p_losses(self, x_start, t, noise = None, clip_denoised = False):
noise = default(noise, lambda: torch.randn_like(x_start))
x_t = self.q_sample(x_start = x_start, t = t, noise = noise)
model_output = self.model(x_t, t)
pred_noise, pred_x_start, weights = model_output.split(self.split_dims, dim = 1)
# get loss for predicted noise and x_start
# with the loss weight given at initialization
noise_loss = self.loss_fn(noise, pred_noise) * self.pred_noise_loss_weight
x_start_loss = self.loss_fn(x_start, pred_x_start) * self.pred_x_start_loss_weight
# calculate x_start from predicted noise
# then do a weighted sum of the x_start prediction, weights also predicted by the model (softmax normalized)
x_start_from_pred_noise = self.predict_start_from_noise(x_t, t, pred_noise)
x_start_from_pred_noise = x_start_from_pred_noise.clamp(-2., 2.)
weighted_x_start = einsum('b j h w, b j c h w -> b c h w', weights.softmax(dim = 1), torch.stack((x_start_from_pred_noise, pred_x_start), dim = 1))
# main loss to x_start with the weighted one
weighted_x_start_loss = self.loss_fn(x_start, weighted_x_start)
return weighted_x_start_loss + x_start_loss + noise_loss
BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 842 KiB

After

Width:  |  Height:  |  Size: 1.3 MiB

+3 -5
View File
@@ -3,23 +3,21 @@ from setuptools import setup, find_packages
setup( setup(
name = 'denoising-diffusion-pytorch', name = 'denoising-diffusion-pytorch',
packages = find_packages(), packages = find_packages(),
version = '0.25.2', version = '0.2.3',
license='MIT', license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch', description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang', author = 'Phil Wang',
author_email = 'lucidrains@gmail.com', author_email = 'lucidrains@gmail.com',
url = 'https://github.com/lucidrains/denoising-diffusion-pytorch', url = 'https://github.com/lucidrains/denoising-diffusion-pytorch',
long_description_content_type = 'text/markdown',
keywords = [ keywords = [
'artificial intelligence', 'artificial intelligence',
'generative models' 'generative models'
], ],
install_requires=[ install_requires=[
'accelerate',
'einops', 'einops',
'ema-pytorch', 'numpy',
'pillow', 'pillow',
'torch', 'torch>=1.6',
'torchvision', 'torchvision',
'tqdm' 'tqdm'
], ],