Compare commits

...
13 Commits
4 changed files with 93 additions and 36 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}
}
```
@@ -5,7 +5,7 @@ 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 from einops.layers.torch import Rearrange
# helpers # helpers
@@ -59,11 +59,14 @@ class MonotonicLinear(nn.Module):
# 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) * torch.pi * 0.5) ** -2) - 1, eps = 1e-5)
class learned_noise_schedule(nn.Module): class learned_noise_schedule(nn.Module):
""" described in section H and then I.2 of the supplementary material for variational ddpm paper """ """ described in section H and then I.2 of the supplementary material for variational ddpm paper """
@@ -73,7 +76,8 @@ class learned_noise_schedule(nn.Module):
*, *,
log_snr_max, log_snr_max,
log_snr_min, log_snr_min,
hidden_dim = 1024 hidden_dim = 1024,
frac_gradient = 1.
): ):
super().__init__() super().__init__()
self.slope = log_snr_min - log_snr_max self.slope = log_snr_min - log_snr_max
@@ -90,7 +94,10 @@ class learned_noise_schedule(nn.Module):
Rearrange('... 1 -> ...'), Rearrange('... 1 -> ...'),
) )
self.frac_gradient = frac_gradient
def forward(self, x): def forward(self, x):
frac_gradient = self.frac_gradient
device = x.device device = x.device
out_zero = self.net(torch.zeros_like(x)) out_zero = self.net(torch.zeros_like(x))
@@ -98,8 +105,8 @@ class learned_noise_schedule(nn.Module):
x = self.net(x) x = self.net(x)
normalized = self.slope * ((x - out_zero) / (out_one - out_zero)) + self.intercept normed = self.slope * ((x - out_zero) / (out_one - out_zero)) + self.intercept
return normalized return normed * frac_gradient + normed.detach() * (1 - frac_gradient)
class ContinuousTimeGaussianDiffusion(nn.Module): class ContinuousTimeGaussianDiffusion(nn.Module):
def __init__( def __init__(
@@ -111,8 +118,11 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
loss_type = 'l1', loss_type = 'l1',
noise_schedule = 'linear', noise_schedule = 'linear',
num_sample_steps = 500, num_sample_steps = 500,
clip_after_noising_during_sampling = False, clip_sample_denoised = True,
learned_schedule_net_hidden_dim = 1024 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 not denoise_fn.sinusoidal_cond_mlp assert not denoise_fn.sinusoidal_cond_mlp
@@ -130,13 +140,16 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
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': elif noise_schedule == 'learned':
log_snr_max, log_snr_min = [beta_linear_log_snr(torch.tensor([time])).item() for time in (0., 1.)] 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( self.log_snr = learned_noise_schedule(
log_snr_max = log_snr_max, log_snr_max = log_snr_max,
log_snr_min = log_snr_min, log_snr_min = log_snr_min,
hidden_dim = learned_schedule_net_hidden_dim 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}')
@@ -144,10 +157,15 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
# sampling # sampling
self.num_sample_steps = num_sample_steps self.num_sample_steps = num_sample_steps
self.clip_sample_denoised = clip_sample_denoised
# clipping related hyperparameters # p2 loss weight
# proposed https://arxiv.org/abs/2204.00227
self.clip_after_noising_during_sampling = clip_after_noising_during_sampling 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):
@@ -166,9 +184,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
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)
@@ -176,10 +191,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()
alpha, sigma, alpha_next = map(sqrt, (squared_alpha, squared_sigma, squared_alpha_next))
batch_log_snr = repeat(log_snr, ' -> b', b = x.shape[0]) batch_log_snr = repeat(log_snr, ' -> b', b = x.shape[0])
pred_noise = self.denoise_fn(x, batch_log_snr) pred_noise = self.denoise_fn(x, batch_log_snr)
model_mean = sqrt(squared_alpha_next / squared_alpha) * (x - c * sqrt(squared_sigma) * pred_noise) 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
@@ -210,12 +236,9 @@ 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)
if self.clip_after_noising_during_sampling: img.clamp_(-1., 1.)
# clip after noise is added. perhaps this is sufficient for Imagen dynamic thresholding?
img.clamp_(-1., 1.)
img = unnormalize_to_zero_to_one(img) img = unnormalize_to_zero_to_one(img)
return img.clamp(0., 1.) return img
@torch.no_grad() @torch.no_grad()
def sample(self, batch_size = 16): def sample(self, batch_size = 16):
@@ -242,9 +265,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)
return self.loss_fn(model_out, noise)
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): 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,7 @@ 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 einops.layers.torch import Rearrange
# helpers functions # helpers functions
@@ -301,7 +302,7 @@ class Unet(nn.Module):
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_conv = nn.Sequential(
block_klass(dim, dim), block_klass(dim * 2, dim),
nn.Conv2d(dim, self.out_dim, 1) nn.Conv2d(dim, self.out_dim, 1)
) )
@@ -323,12 +324,13 @@ 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, h.pop()), dim = 1)
return self.final_conv(x) return self.final_conv(x)
# gaussian diffusion trainer class # gaussian diffusion trainer class
@@ -366,7 +368,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)
@@ -421,6 +425,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 -
@@ -527,8 +535,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
@@ -541,7 +552,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
@@ -549,7 +560,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()
]) ])
@@ -580,7 +591,8 @@ class Trainer(object):
step_start_ema = 2000, step_start_ema = 2000,
update_ema_every = 10, update_ema_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.model = diffusion_model self.model = diffusion_model
@@ -596,8 +608,8 @@ 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, 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
+1 -1
View File
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
setup( setup(
name = 'denoising-diffusion-pytorch', name = 'denoising-diffusion-pytorch',
packages = find_packages(), packages = find_packages(),
version = '0.17.2', version = '0.19.0',
license='MIT', license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch', description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang', author = 'Phil Wang',