Compare commits

...
26 Commits
Author SHA1 Message Date
Phil Wang 9939a48139 make sure all versions of torch supported 2022-06-23 12:28:09 -07:00
Phil Wang 75ea49a7ef pass parameter for Trainer to EMA properly 2022-06-21 07:37:43 -07:00
Phil Wang 8c3609a6e3 move EMA logic out of the repository for clarity 2022-06-20 13:17:51 -07:00
Phil Wang 1586d1a8a0 just pluck the image size off the gaussian diffusion class 2022-06-17 13:54:41 -07:00
Phil Wang b4fb8804d2 conditioning on final resnet block 2022-06-17 10:38:17 -07:00
Phil Wang 9fd05f1b1f switch to learned sinsuoidal pos emb for the continuous case 2022-06-17 09:24:51 -07:00
Phil Wang ec2397f0ba add one more residual 2022-06-16 11:08:42 -07:00
Phil Wang 844e557dfb fix a missing residual needed at the top most resolution in the unet 2022-06-15 19:10:05 -07:00
Phil Wang 8b30be8042 add p2 loss reweighting for default ddpm as an option 2022-06-14 10:49:13 -07:00
Phil Wang f2f3994b92 link to Letitia 2022-06-12 14:57:44 -07:00
Phil Wang 8ec4ea56a5 link to yannic 2022-06-12 14:56:55 -07:00
Phil Wang 99cf9b5b96 link to ai educator 2022-06-12 14:54:43 -07:00
Phil Wang ecc6f30901 for https://github.com/lucidrains/denoising-diffusion-pytorch/issues/36 2022-06-11 10:51:13 -07:00
Phil Wang f900f40f14 allow for turning off horizontal flip augmentation 2022-06-09 20:59:59 -07:00
Phil Wang 479f60c178 add p2 loss weighting to SNR version of denoising diffusion, brought up by @Mut1nyJD, paper is https://arxiv.org/abs/2204.00227 2022-06-09 08:25:05 -07:00
Phil Wang 96bb2ff310 alpha cosine noise schedule is now working for continuous time gaussian diffusion 2022-06-08 23:07:01 -07:00
Phil Wang 582bfe275b successfully did some basic math and clipped the predicted x0 intermediate for the continuous time case 2022-06-08 17:59:41 -07:00
Phil Wang 4284c8840d clipping for continuous time diffusion not working 2022-06-08 16:26:18 -07:00
Phil Wang d4ffa3fced link to annotated ddpm 2022-06-08 12:55:15 -07:00
Phil Wang c44d3ea01d learned noise schedule seems to be working, allow for one to make the monotonic net learn a bit more slowly than the unet 2022-06-08 12:34:07 -07:00
Phil Wang c4991f576f allow for configuring the hidden dimension of the monotonic mlp parameterizing the noise schedule 2022-06-08 11:18:54 -07:00
Phil Wang a19331aa59 fix learned noise schedule 2022-06-08 10:19:02 -07:00
Phil Wang 94eabaca1a complete learned noise schedule for variational ddpm paper, still need to finish cosine alpha schedule in log(snr) form 2022-06-08 09:47:09 -07:00
Phil Wang eaf9d9fdc4 unet needs to be conditioned on log(snr) in p_mean_variance for continuous time gaussian diffusion 2022-06-08 00:41:41 -07:00
Phil Wang 3bf5e768c2 use a non-sinusoidal embedded condition for continuous time gaussian diffusion conditioned on log(snr) 2022-06-07 21:15:27 -07:00
Phil Wang 532178a6a3 assume when sampling all batch samples are at the same time, and do not noise for the last time step 2022-06-07 16:12:44 -07:00
4 changed files with 239 additions and 108 deletions
+14
View File
@@ -6,6 +6,10 @@ Implementation of <a href="https://arxiv.org/abs/2006.11239">Denoising Diffusion
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) [![PyPI version](https://badge.fury.io/py/denoising-diffusion-pytorch.svg)](https://badge.fury.io/py/denoising-diffusion-pytorch)
@@ -119,3 +123,13 @@ Samples and model checkpoints will be logged to `./results` periodically
url = {https://openreview.net/forum?id=2LdBqxc1Yv} 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}
}
```
@@ -1,3 +1,4 @@
import math
import torch import torch
from torch import sqrt from torch import sqrt
from torch import nn, einsum from torch import nn, einsum
@@ -5,7 +6,8 @@ import torch.nn.functional as F
from torch.special import expm1 from torch.special import expm1
from tqdm import tqdm from tqdm import tqdm
from einops import rearrange, repeat from einops import rearrange, repeat, reduce
from einops.layers.torch import Rearrange
# helpers # helpers
@@ -33,6 +35,24 @@ def right_pad_dims_to(x, t):
return t return t
return t.view(*t.shape, *((1,) * padding_dims)) 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 # continuous schedules
# equations are taken from https://openreview.net/attachment?id=2LdBqxc1Yv&name=supplementary_material # equations are taken from https://openreview.net/attachment?id=2LdBqxc1Yv&name=supplementary_material
@@ -40,17 +60,54 @@ def right_pad_dims_to(x, t):
# log(snr) that approximates the original linear schedule # log(snr) that approximates the original linear schedule
def beta_linear_log_snr(t): def log(t, eps = 1e-20):
return -torch.log(expm1(1e-4 + 10 * (t ** 2))) return torch.log(t.clamp(min = eps))
def alpha_cosine_log_snr(t): def beta_linear_log_snr(t):
raise NotImplementedError 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): class learned_noise_schedule(nn.Module):
def __init__(self): """ 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__() super().__init__()
raise NotImplementedError self.slope = log_snr_min - log_snr_max
# learned noise schedule, using learned monotonic MLP (weights kept positive) in the paper 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): class ContinuousTimeGaussianDiffusion(nn.Module):
def __init__( def __init__(
@@ -59,12 +116,17 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
*, *,
image_size, image_size,
channels = 3, channels = 3,
cond_scale = 500,
loss_type = 'l1', loss_type = 'l1',
noise_schedule = 'linear', noise_schedule = 'linear',
num_sample_steps = 500 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__() super().__init__()
assert denoise_fn.learned_sinusoidal_cond
self.denoise_fn = denoise_fn self.denoise_fn = denoise_fn
@@ -75,17 +137,36 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
# continuous noise schedule related stuff # continuous noise schedule related stuff
self.cond_scale = cond_scale # the log(snr) will be scaled by this value
self.loss_type = loss_type self.loss_type = loss_type
if noise_schedule == 'linear': if noise_schedule == 'linear':
self.log_snr = beta_linear_log_snr 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: else:
raise ValueError(f'unknown noise schedule {noise_schedule}') raise ValueError(f'unknown noise schedule {noise_schedule}')
# sampling # sampling
self.num_sample_steps = num_sample_steps 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 @property
def device(self): def device(self):
@@ -104,14 +185,6 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
# reviewer found an error in the equation in the paper (missing sigma) # reviewer found an error in the equation in the paper (missing sigma)
# following - https://openreview.net/forum?id=2LdBqxc1Yv&noteId=rIQgH0zKsRt # following - https://openreview.net/forum?id=2LdBqxc1Yv&noteId=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 = self.log_snr(time)
log_snr_next = self.log_snr(time_next) log_snr_next = self.log_snr(time_next)
c = -expm1(log_snr - log_snr_next) c = -expm1(log_snr - log_snr_next)
@@ -119,7 +192,21 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
squared_alpha, squared_alpha_next = log_snr.sigmoid(), log_snr_next.sigmoid() squared_alpha, squared_alpha_next = log_snr.sigmoid(), log_snr_next.sigmoid()
squared_sigma, squared_sigma_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) 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.denoise_fn(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 posterior_variance = squared_sigma_next * c
return model_mean, posterior_variance return model_mean, posterior_variance
@@ -131,6 +218,10 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
batch, *_, device = *x.shape, x.device batch, *_, device = *x.shape, x.device
model_mean, model_variance = self.p_mean_variance(x = x, time = time, time_next = time_next) 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) noise = torch.randn_like(x)
return model_mean + sqrt(model_variance) * noise return model_mean + sqrt(model_variance) * noise
@@ -146,6 +237,7 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
times_next = steps[i + 1] times_next = steps[i + 1]
img = self.p_sample(img, times, times_next) img = self.p_sample(img, times, times_next)
img.clamp_(-1., 1.)
img = unnormalize_to_zero_to_one(img) img = unnormalize_to_zero_to_one(img)
return img return img
@@ -174,9 +266,17 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
noise = default(noise, lambda: torch.randn_like(x_start)) noise = default(noise, lambda: torch.randn_like(x_start))
x, log_snr = self.q_sample(x_start = x_start, times = times, noise = noise) x, log_snr = self.q_sample(x_start = x_start, times = times, noise = noise)
model_out = self.denoise_fn(x, log_snr)
model_out = self.denoise_fn(x, log_snr * self.cond_scale) losses = self.loss_fn(model_out, noise, reduction = 'none')
return self.loss_fn(model_out, noise) 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): def forward(self, img, *args, **kwargs):
b, c, h, w, device, img_size, = *img.shape, img.device, self.image_size b, c, h, w, device, img_size, = *img.shape, img.device, self.image_size
@@ -7,6 +7,7 @@ from inspect import isfunction
from functools import partial from functools import partial
from torch.utils import data from torch.utils import data
from multiprocessing import cpu_count
from torch.cuda.amp import autocast, GradScaler from torch.cuda.amp import autocast, GradScaler
from pathlib import Path from pathlib import Path
@@ -15,7 +16,10 @@ from torchvision import transforms, utils
from PIL import Image from PIL import Image
from tqdm import tqdm from tqdm import tqdm
from einops import rearrange from einops import rearrange, reduce
from einops.layers.torch import Rearrange
from ema_pytorch import EMA
# helpers functions # helpers functions
@@ -48,21 +52,6 @@ def unnormalize_to_zero_to_one(t):
# 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__()
@@ -71,20 +60,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
class SinusoidalPosEmb(nn.Module):
def __init__(self, dim):
super().__init__()
self.dim = dim
def forward(self, x):
device = x.device
half_dim = self.dim // 2
emb = math.log(10000) / (half_dim - 1)
emb = torch.exp(torch.arange(half_dim, device=device) * -emb)
emb = x[:, None] * emb[None, :]
emb = torch.cat((emb.sin(), emb.cos()), dim=-1)
return emb
def Upsample(dim): def Upsample(dim):
return nn.ConvTranspose2d(dim, dim, 4, 2, 1) return nn.ConvTranspose2d(dim, dim, 4, 2, 1)
@@ -113,6 +88,39 @@ class PreNorm(nn.Module):
x = self.norm(x) x = self.norm(x)
return self.fn(x) return self.fn(x)
# sinusoidal positional embeds
class SinusoidalPosEmb(nn.Module):
def __init__(self, dim):
super().__init__()
self.dim = dim
def forward(self, x):
device = x.device
half_dim = self.dim // 2
emb = math.log(10000) / (half_dim - 1)
emb = torch.exp(torch.arange(half_dim, device=device) * -emb)
emb = x[:, None] * emb[None, :]
emb = torch.cat((emb.sin(), emb.cos()), dim=-1)
return emb
class LearnedSinusoidalPosEmb(nn.Module):
""" following @crowsonkb 's lead with learned sinusoidal pos emb """
""" https://github.com/crowsonkb/v-diffusion-jax/blob/master/diffusion/models/danbooru_128.py#L8 """
def __init__(self, dim):
super().__init__()
assert (dim % 2) == 0
half_dim = dim // 2
self.weights = nn.Parameter(torch.randn(half_dim))
def forward(self, x):
x = rearrange(x, 'b -> b 1')
freqs = x * rearrange(self.weights, 'd -> 1 d') * 2 * math.pi
fouriered = torch.cat((freqs.sin(), freqs.cos()), dim = -1)
fouriered = torch.cat((x, fouriered), dim = -1)
return fouriered
# building block modules # building block modules
class Block(nn.Module): class Block(nn.Module):
@@ -156,6 +164,7 @@ class ResnetBlock(nn.Module):
h = self.block1(x, scale_shift = scale_shift) 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):
@@ -219,9 +228,10 @@ 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,
with_time_emb = True,
resnet_block_groups = 8, resnet_block_groups = 8,
learned_variance = False learned_variance = False,
learned_sinusoidal_cond = False,
learned_sinusoidal_dim = 16
): ):
super().__init__() super().__init__()
@@ -229,7 +239,7 @@ class Unet(nn.Module):
self.channels = channels self.channels = channels
init_dim = default(init_dim, dim // 3 * 2) init_dim = default(init_dim, dim)
self.init_conv = nn.Conv2d(channels, init_dim, 7, padding = 3) self.init_conv = nn.Conv2d(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)]
@@ -239,17 +249,23 @@ class Unet(nn.Module):
# time embeddings # time embeddings
if with_time_emb: time_dim = dim * 4
time_dim = dim * 4
self.time_mlp = nn.Sequential( self.learned_sinusoidal_cond = learned_sinusoidal_cond
SinusoidalPosEmb(dim),
nn.Linear(dim, time_dim), if learned_sinusoidal_cond:
nn.GELU(), sinu_pos_emb = LearnedSinusoidalPosEmb(learned_sinusoidal_dim)
nn.Linear(time_dim, time_dim) fourier_dim = learned_sinusoidal_dim + 1
)
else: else:
time_dim = None sinu_pos_emb = SinusoidalPosEmb(dim)
self.time_mlp = None 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 # layers
@@ -272,8 +288,8 @@ class Unet(nn.Module):
self.mid_attn = Residual(PreNorm(mid_dim, Attention(mid_dim))) self.mid_attn = Residual(PreNorm(mid_dim, Attention(mid_dim)))
self.mid_block2 = block_klass(mid_dim, mid_dim, time_emb_dim = time_dim) self.mid_block2 = block_klass(mid_dim, mid_dim, time_emb_dim = time_dim)
for ind, (dim_in, dim_out) in enumerate(reversed(in_out[1:])): for ind, (dim_in, dim_out) in enumerate(reversed(in_out)):
is_last = ind >= (num_resolutions - 1) is_last = ind == (len(in_out) - 1)
self.ups.append(nn.ModuleList([ self.ups.append(nn.ModuleList([
block_klass(dim_out * 2, dim_in, time_emb_dim = time_dim), block_klass(dim_out * 2, dim_in, time_emb_dim = time_dim),
@@ -285,15 +301,14 @@ class Unet(nn.Module):
default_out_dim = channels * (1 if not learned_variance else 2) default_out_dim = channels * (1 if not learned_variance else 2)
self.out_dim = default(out_dim, default_out_dim) self.out_dim = default(out_dim, default_out_dim)
self.final_conv = nn.Sequential( self.final_res_block = block_klass(dim * 2, dim, time_emb_dim = time_dim)
block_klass(dim, dim), self.final_conv = nn.Conv2d(dim, self.out_dim, 1)
nn.Conv2d(dim, self.out_dim, 1)
)
def forward(self, x, time): def forward(self, x, time):
x = self.init_conv(x) x = self.init_conv(x)
r = x.clone()
t = self.time_mlp(time) if exists(self.time_mlp) else None t = self.time_mlp(time)
h = [] h = []
@@ -309,12 +324,15 @@ class Unet(nn.Module):
x = self.mid_block2(x, t) x = self.mid_block2(x, t)
for block1, block2, attn, upsample in self.ups: for block1, block2, 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 = block1(x, t)
x = block2(x, t) 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
@@ -337,7 +355,7 @@ def cosine_beta_schedule(timesteps, s = 0.008):
""" """
steps = timesteps + 1 steps = timesteps + 1
x = torch.linspace(0, timesteps, steps, dtype = torch.float64) x = torch.linspace(0, timesteps, steps, dtype = torch.float64)
alphas_cumprod = torch.cos(((x / timesteps) + s) / (1 + s) * torch.pi * 0.5) ** 2 alphas_cumprod = torch.cos(((x / timesteps) + s) / (1 + s) * math.pi * 0.5) ** 2
alphas_cumprod = alphas_cumprod / alphas_cumprod[0] alphas_cumprod = alphas_cumprod / alphas_cumprod[0]
betas = 1 - (alphas_cumprod[1:] / alphas_cumprod[:-1]) betas = 1 - (alphas_cumprod[1:] / alphas_cumprod[:-1])
return torch.clip(betas, 0, 0.999) return torch.clip(betas, 0, 0.999)
@@ -352,7 +370,9 @@ class GaussianDiffusion(nn.Module):
timesteps = 1000, timesteps = 1000,
loss_type = 'l1', loss_type = 'l1',
objective = 'pred_noise', objective = 'pred_noise',
beta_schedule = 'cosine' 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
): ):
super().__init__() super().__init__()
assert not (type(self) == GaussianDiffusion and denoise_fn.channels != denoise_fn.out_dim) assert not (type(self) == GaussianDiffusion and denoise_fn.channels != denoise_fn.out_dim)
@@ -407,6 +427,10 @@ class GaussianDiffusion(nn.Module):
register_buffer('posterior_mean_coef1', betas * torch.sqrt(alphas_cumprod_prev) / (1. - alphas_cumprod)) register_buffer('posterior_mean_coef1', betas * torch.sqrt(alphas_cumprod_prev) / (1. - alphas_cumprod))
register_buffer('posterior_mean_coef2', (1. - alphas_cumprod_prev) * torch.sqrt(alphas) / (1. - alphas_cumprod)) register_buffer('posterior_mean_coef2', (1. - alphas_cumprod_prev) * torch.sqrt(alphas) / (1. - alphas_cumprod))
# calculate p2 reweighting
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 (
extract(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t - extract(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t -
@@ -513,8 +537,11 @@ class GaussianDiffusion(nn.Module):
else: else:
raise ValueError(f'unknown objective {self.objective}') raise ValueError(f'unknown objective {self.objective}')
loss = self.loss_fn(model_out, target) 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)
return loss.mean()
def forward(self, img, *args, **kwargs): def forward(self, img, *args, **kwargs):
b, c, h, w, device, img_size, = *img.shape, img.device, self.image_size b, c, h, w, device, img_size, = *img.shape, img.device, self.image_size
@@ -527,7 +554,7 @@ class GaussianDiffusion(nn.Module):
# dataset classes # dataset classes
class Dataset(data.Dataset): class Dataset(data.Dataset):
def __init__(self, folder, image_size, exts = ['jpg', 'jpeg', 'png']): def __init__(self, folder, image_size, exts = ['jpg', 'jpeg', 'png'], augment_horizontal_flip = False):
super().__init__() super().__init__()
self.folder = folder self.folder = folder
self.image_size = image_size self.image_size = image_size
@@ -535,7 +562,7 @@ class Dataset(data.Dataset):
self.transform = transforms.Compose([ self.transform = transforms.Compose([
transforms.Resize(image_size), transforms.Resize(image_size),
transforms.RandomHorizontalFlip(), transforms.RandomHorizontalFlip() if augment_horizontal_flip else nn.Identity(),
transforms.CenterCrop(image_size), transforms.CenterCrop(image_size),
transforms.ToTensor() transforms.ToTensor()
]) ])
@@ -557,22 +584,22 @@ class Trainer(object):
folder, folder,
*, *,
ema_decay = 0.995, ema_decay = 0.995,
image_size = 128,
train_batch_size = 32, train_batch_size = 32,
train_lr = 1e-4, train_lr = 1e-4,
train_num_steps = 100000, train_num_steps = 100000,
gradient_accumulate_every = 2, gradient_accumulate_every = 2,
amp = False, amp = False,
step_start_ema = 2000, step_start_ema = 2000,
update_ema_every = 10, ema_update_every = 10,
save_and_sample_every = 1000, save_and_sample_every = 1000,
results_folder = './results' results_folder = './results',
augment_horizontal_flip = True
): ):
super().__init__() super().__init__()
self.image_size = diffusion_model.image_size
self.model = diffusion_model self.model = diffusion_model
self.ema = EMA(ema_decay) self.ema = EMA(diffusion_model, beta = ema_decay, update_every = ema_update_every)
self.ema_model = copy.deepcopy(self.model)
self.update_ema_every = update_ema_every
self.step_start_ema = step_start_ema self.step_start_ema = step_start_ema
self.save_and_sample_every = save_and_sample_every self.save_and_sample_every = save_and_sample_every
@@ -582,9 +609,9 @@ class Trainer(object):
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.ds = Dataset(folder, image_size) self.ds = Dataset(folder, self.image_size, augment_horizontal_flip = augment_horizontal_flip)
self.dl = cycle(data.DataLoader(self.ds, batch_size = train_batch_size, shuffle=True, pin_memory=True)) self.dl = cycle(data.DataLoader(self.ds, batch_size = train_batch_size, shuffle = True, pin_memory = True, num_workers = cpu_count()))
self.opt = Adam(diffusion_model.parameters(), lr=train_lr) self.opt = Adam(diffusion_model.parameters(), lr = train_lr)
self.step = 0 self.step = 0
@@ -594,22 +621,11 @@ class Trainer(object):
self.results_folder = Path(results_folder) self.results_folder = Path(results_folder)
self.results_folder.mkdir(exist_ok = True) self.results_folder.mkdir(exist_ok = True)
self.reset_parameters()
def reset_parameters(self):
self.ema_model.load_state_dict(self.model.state_dict())
def step_ema(self):
if self.step < self.step_start_ema:
self.reset_parameters()
return
self.ema.update_model_average(self.ema_model, self.model)
def save(self, milestone): def save(self, milestone):
data = { data = {
'step': self.step, 'step': self.step,
'model': self.model.state_dict(), 'model': self.model.state_dict(),
'ema': self.ema_model.state_dict(), 'ema': self.ema.state_dict(),
'scaler': self.scaler.state_dict() 'scaler': self.scaler.state_dict()
} }
torch.save(data, str(self.results_folder / f'model-{milestone}.pt')) torch.save(data, str(self.results_folder / f'model-{milestone}.pt'))
@@ -619,7 +635,7 @@ class Trainer(object):
self.step = data['step'] self.step = data['step']
self.model.load_state_dict(data['model']) self.model.load_state_dict(data['model'])
self.ema_model.load_state_dict(data['ema']) self.ema.load_state_dict(data['ema'])
self.scaler.load_state_dict(data['scaler']) self.scaler.load_state_dict(data['scaler'])
def train(self): def train(self):
@@ -639,15 +655,15 @@ class Trainer(object):
self.scaler.update() self.scaler.update()
self.opt.zero_grad() self.opt.zero_grad()
if self.step % self.update_ema_every == 0: self.ema.update()
self.step_ema()
if self.step != 0 and self.step % self.save_and_sample_every == 0: if self.step != 0 and self.step % self.save_and_sample_every == 0:
self.ema_model.eval() self.ema.ema_model.eval()
with torch.no_grad():
milestone = self.step // self.save_and_sample_every
batches = num_to_groups(36, self.batch_size)
all_images_list = list(map(lambda n: self.ema.ema_model.sample(batch_size=n), batches))
milestone = self.step // self.save_and_sample_every
batches = num_to_groups(36, self.batch_size)
all_images_list = list(map(lambda n: self.ema_model.sample(batch_size=n), batches))
all_images = torch.cat(all_images_list, dim=0) all_images = torch.cat(all_images_list, dim=0)
utils.save_image(all_images, str(self.results_folder / f'sample-{milestone}.png'), nrow = 6) utils.save_image(all_images, str(self.results_folder / f'sample-{milestone}.png'), nrow = 6)
self.save(milestone) self.save(milestone)
+2 -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.16.3', version = '0.21.2',
license='MIT', license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch', description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang', author = 'Phil Wang',
@@ -16,6 +16,7 @@ setup(
], ],
install_requires=[ install_requires=[
'einops', 'einops',
'ema-pytorch',
'pillow', 'pillow',
'torch', 'torch',
'torchvision', 'torchvision',