Compare commits

..
1 Commits
6 changed files with 71 additions and 431 deletions
+2 -27
View File
@@ -6,10 +6,6 @@ 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)
@@ -38,7 +34,7 @@ diffusion = GaussianDiffusion(
loss_type = 'l1' # L1 or L2 loss_type = 'l1' # L1 or L2
) )
training_images = torch.randn(8, 3, 128, 128) # images are normalized from 0 to 1 training_images = torch.randn(8, 3, 128, 128) # your images need to be normalized from a range of -1 to +1
loss = diffusion(training_images) loss = diffusion(training_images)
loss.backward() loss.backward()
# after a lot of training # after a lot of training
@@ -68,7 +64,7 @@ trainer = Trainer(
diffusion, diffusion,
'path/to/your/images', 'path/to/your/images',
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 = 700000, # 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
@@ -112,24 +108,3 @@ Samples and model checkpoints will be logged to `./results` periodically
url = {https://proceedings.mlr.press/v139/nichol21a.html}, 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}
}
```
-1
View File
@@ -1,5 +1,4 @@
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.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.weighted_objective_gaussian_diffusion import WeightedObjectiveGaussianDiffusion
@@ -1,286 +0,0 @@
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) * torch.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,
denoise_fn,
*,
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 not denoise_fn.sinusoidal_cond_mlp
self.denoise_fn = denoise_fn
# 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.denoise_fn.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.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
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.denoise_fn(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)
@@ -7,7 +7,6 @@ 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
@@ -16,8 +15,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, reduce from einops import rearrange
from einops.layers.torch import Rearrange
# helpers functions # helpers functions
@@ -120,27 +118,20 @@ class PreNorm(nn.Module):
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),
nn.SiLU()
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 = None, groups = 8):
super().__init__() super().__init__()
self.mlp = nn.Sequential( self.mlp = nn.Sequential(
nn.SiLU(), nn.SiLU(),
nn.Linear(time_emb_dim, dim_out * 2) nn.Linear(time_emb_dim, dim_out)
) if exists(time_emb_dim) else None ) if exists(time_emb_dim) else None
self.block1 = Block(dim, dim_out, groups = groups) self.block1 = Block(dim, dim_out, groups = groups)
@@ -148,14 +139,11 @@ class ResnetBlock(nn.Module):
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 = None):
h = self.block1(x)
scale_shift = None
if exists(self.mlp) and exists(time_emb): if exists(self.mlp) and exists(time_emb):
time_emb = self.mlp(time_emb) time_emb = self.mlp(time_emb)
time_emb = rearrange(time_emb, 'b c -> b c 1 1') h = rearrange(time_emb, 'b c -> b c 1 1') + h
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)
@@ -213,18 +201,6 @@ class Attention(nn.Module):
# model # model
def MLP(dim_in, dim_hidden):
return nn.Sequential(
Rearrange('... -> ... 1'),
nn.Linear(1, dim_hidden),
nn.GELU(),
nn.LayerNorm(dim_hidden),
nn.Linear(dim_hidden, dim_hidden),
nn.GELU(),
nn.LayerNorm(dim_hidden),
nn.Linear(dim_hidden, dim_hidden)
)
class Unet(nn.Module): class Unet(nn.Module):
def __init__( def __init__(
self, self,
@@ -233,9 +209,9 @@ 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
sinusoidal_cond_mlp = True
): ):
super().__init__() super().__init__()
@@ -243,7 +219,7 @@ class Unet(nn.Module):
self.channels = channels self.channels = channels
init_dim = default(init_dim, dim) init_dim = default(init_dim, dim // 3 * 2)
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)]
@@ -253,11 +229,8 @@ class Unet(nn.Module):
# time embeddings # time embeddings
time_dim = dim * 4 if with_time_emb:
time_dim = dim * 4
self.sinusoidal_cond_mlp = sinusoidal_cond_mlp
if sinusoidal_cond_mlp:
self.time_mlp = nn.Sequential( self.time_mlp = nn.Sequential(
SinusoidalPosEmb(dim), SinusoidalPosEmb(dim),
nn.Linear(dim, time_dim), nn.Linear(dim, time_dim),
@@ -265,7 +238,8 @@ class Unet(nn.Module):
nn.Linear(time_dim, time_dim) nn.Linear(time_dim, time_dim)
) )
else: else:
self.time_mlp = MLP(1, time_dim) time_dim = None
self.time_mlp = None
# layers # layers
@@ -288,8 +262,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)): 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 * 2, dim_in, time_emb_dim = time_dim), block_klass(dim_out * 2, dim_in, time_emb_dim = time_dim),
@@ -308,7 +282,8 @@ class Unet(nn.Module):
def forward(self, x, time): def forward(self, x, time):
x = self.init_conv(x) x = self.init_conv(x)
t = self.time_mlp(time)
t = self.time_mlp(time) if exists(self.time_mlp) else None
h = [] h = []
@@ -324,7 +299,7 @@ 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)
@@ -339,11 +314,10 @@ 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): def cosine_beta_schedule(timesteps, s = 0.008):
""" """
@@ -366,10 +340,7 @@ class GaussianDiffusion(nn.Module):
channels = 3, channels = 3,
timesteps = 1000, timesteps = 1000,
loss_type = 'l1', loss_type = 'l1',
objective = 'pred_noise', 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
): ):
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)
@@ -379,12 +350,7 @@ class GaussianDiffusion(nn.Module):
self.denoise_fn = denoise_fn self.denoise_fn = denoise_fn
self.objective = objective self.objective = objective
if beta_schedule == 'linear': betas = cosine_beta_schedule(timesteps)
betas = linear_beta_schedule(timesteps)
elif beta_schedule == 'cosine':
betas = cosine_beta_schedule(timesteps)
else:
raise ValueError(f'unknown beta schedule {beta_schedule}')
alphas = 1. - betas alphas = 1. - betas
alphas_cumprod = torch.cumprod(alphas, axis=0) alphas_cumprod = torch.cumprod(alphas, axis=0)
@@ -424,10 +390,6 @@ 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 -
@@ -460,10 +422,10 @@ class GaussianDiffusion(nn.Module):
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, 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
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=t, clip_denoised=clip_denoised)
noise = torch.randn_like(x) noise = noise_like(x.shape, device, repeat_noise)
# no noise when t == 0 # no noise when t == 0
nonzero_mask = (1 - (t == 0).float()).reshape(b, *((1,) * (len(x.shape) - 1))) 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 return model_mean + nonzero_mask * (0.5 * model_log_variance).exp() * noise
@@ -477,8 +439,6 @@ class GaussianDiffusion(nn.Module):
for i in tqdm(reversed(range(0, self.num_timesteps)), desc='sampling loop time step', total=self.num_timesteps): for i in tqdm(reversed(range(0, self.num_timesteps)), desc='sampling loop time step', total=self.num_timesteps):
img = self.p_sample(img, torch.full((b,), i, device=device, dtype=torch.long)) img = self.p_sample(img, torch.full((b,), i, device=device, dtype=torch.long))
img = unnormalize_to_zero_to_one(img)
return img return img
@torch.no_grad() @torch.no_grad()
@@ -534,24 +494,19 @@ 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, reduction = 'none') loss = self.loss_fn(model_out, target)
loss = reduce(loss, 'b ... -> b (...)', 'mean') return loss
loss = loss * extract(self.p2_loss_weight, t, loss.shape) def forward(self, x, *args, **kwargs):
return loss.mean() b, c, h, w, device, img_size, = *x.shape, x.device, self.image_size
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}' 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(data.Dataset): class Dataset(data.Dataset):
def __init__(self, folder, image_size, exts = ['jpg', 'jpeg', 'png'], augment_horizontal_flip = False): def __init__(self, folder, image_size, exts = ['jpg', 'jpeg', 'png']):
super().__init__() super().__init__()
self.folder = folder self.folder = folder
self.image_size = image_size self.image_size = image_size
@@ -559,9 +514,10 @@ class Dataset(data.Dataset):
self.transform = transforms.Compose([ self.transform = transforms.Compose([
transforms.Resize(image_size), transforms.Resize(image_size),
transforms.RandomHorizontalFlip() if augment_horizontal_flip else nn.Identity(), transforms.RandomHorizontalFlip(),
transforms.CenterCrop(image_size), transforms.CenterCrop(image_size),
transforms.ToTensor() transforms.ToTensor(),
transforms.Lambda(normalize_to_neg_one_to_one)
]) ])
def __len__(self): def __len__(self):
@@ -583,15 +539,14 @@ class Trainer(object):
ema_decay = 0.995, ema_decay = 0.995,
image_size = 128, image_size = 128,
train_batch_size = 32, train_batch_size = 32,
train_lr = 1e-4, train_lr = 2e-5,
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, 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
@@ -607,8 +562,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, augment_horizontal_flip = augment_horizontal_flip) self.ds = Dataset(folder, image_size)
self.dl = cycle(data.DataLoader(self.ds, batch_size = train_batch_size, shuffle = True, pin_memory = True, num_workers = cpu_count())) self.dl = cycle(data.DataLoader(self.ds, batch_size = train_batch_size, shuffle=True, pin_memory=True))
self.opt = Adam(diffusion_model.parameters(), lr=train_lr) self.opt = Adam(diffusion_model.parameters(), lr=train_lr)
self.step = 0 self.step = 0
@@ -648,36 +603,34 @@ class Trainer(object):
self.scaler.load_state_dict(data['scaler']) self.scaler.load_state_dict(data['scaler'])
def train(self): def train(self):
with tqdm(initial = self.step, total = self.train_num_steps) as pbar: while self.step < self.train_num_steps:
for i in range(self.gradient_accumulate_every):
data = next(self.dl).cuda()
while self.step < self.train_num_steps: with autocast(enabled = self.amp):
for i in range(self.gradient_accumulate_every): loss = self.model(data)
data = next(self.dl).cuda() self.scaler.scale(loss / self.gradient_accumulate_every).backward()
with autocast(enabled = self.amp): print(f'{self.step}: {loss.item()}')
loss = self.model(data)
self.scaler.scale(loss / self.gradient_accumulate_every).backward()
pbar.set_description(f'loss: {loss.item():.4f}') self.scaler.step(self.opt)
self.scaler.update()
self.opt.zero_grad()
self.scaler.step(self.opt) if self.step % self.update_ema_every == 0:
self.scaler.update() self.step_ema()
self.opt.zero_grad()
if self.step % self.update_ema_every == 0: if self.step != 0 and self.step % self.save_and_sample_every == 0:
self.step_ema() self.ema_model.eval()
if self.step != 0 and self.step % self.save_and_sample_every == 0: milestone = self.step // self.save_and_sample_every
self.ema_model.eval() 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 = unnormalize_to_zero_to_one(all_images)
utils.save_image(all_images, str(self.results_folder / f'sample-{milestone}.png'), nrow = 6)
self.save(milestone)
milestone = self.step // self.save_and_sample_every self.step += 1
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)
utils.save_image(all_images, str(self.results_folder / f'sample-{milestone}.png'), nrow = 6)
self.save(milestone)
self.step += 1 print('training completed')
pbar.update(1)
print('training complete')
@@ -3,7 +3,7 @@ from inspect import isfunction
from torch import nn, einsum from torch import nn, einsum
from einops import rearrange from einops import rearrange
from denoising_diffusion_pytorch.denoising_diffusion_pytorch import GaussianDiffusion from denoising_diffusion_pytorch.denoising_diffusion_pytorch import GaussianDiffusion, extract, unnormalize_to_zero_to_one
# helper functions # helper functions
+1 -2
View File
@@ -3,13 +3,12 @@ from setuptools import setup, find_packages
setup( setup(
name = 'denoising-diffusion-pytorch', name = 'denoising-diffusion-pytorch',
packages = find_packages(), packages = find_packages(),
version = '0.19.1', version = '0.15.1',
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'