Compare commits

...
12 Commits
4 changed files with 101 additions and 95 deletions
+2
View File
@@ -6,6 +6,8 @@ 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> <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>
@@ -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
@@ -125,7 +125,7 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
p2_loss_weight_k = 1 p2_loss_weight_k = 1
): ):
super().__init__() super().__init__()
assert not denoise_fn.sinusoidal_cond_mlp assert denoise_fn.learned_sinusoidal_cond
self.denoise_fn = denoise_fn self.denoise_fn = denoise_fn
@@ -268,7 +268,7 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
model_out = self.denoise_fn(x, log_snr) model_out = self.denoise_fn(x, log_snr)
losses = self.loss_fn(model_out, noise, reduction = 'none') losses = self.loss_fn(model_out, noise, reduction = 'none')
losses = losses.mean(dim = tuple(range(1, losses.ndim))) losses = reduce(losses, 'b ... -> b', 'mean')
if self.p2_loss_weight_gamma >= 0: if self.p2_loss_weight_gamma >= 0:
# following eq 8. in https://arxiv.org/abs/2204.00227 # following eq 8. in https://arxiv.org/abs/2204.00227
@@ -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,9 +16,11 @@ 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
from ema_pytorch import EMA
# helpers functions # helpers functions
def exists(x): def exists(x):
@@ -49,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__()
@@ -72,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)
@@ -114,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):
@@ -157,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):
@@ -212,18 +220,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,
@@ -234,7 +230,8 @@ class Unet(nn.Module):
channels = 3, channels = 3,
resnet_block_groups = 8, resnet_block_groups = 8,
learned_variance = False, learned_variance = False,
sinusoidal_cond_mlp = True learned_sinusoidal_cond = False,
learned_sinusoidal_dim = 16
): ):
super().__init__() super().__init__()
@@ -242,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)]
@@ -254,17 +251,21 @@ class Unet(nn.Module):
time_dim = dim * 4 time_dim = dim * 4
self.sinusoidal_cond_mlp = sinusoidal_cond_mlp self.learned_sinusoidal_cond = learned_sinusoidal_cond
if sinusoidal_cond_mlp: if learned_sinusoidal_cond:
self.time_mlp = nn.Sequential( sinu_pos_emb = LearnedSinusoidalPosEmb(learned_sinusoidal_dim)
SinusoidalPosEmb(dim), fourier_dim = learned_sinusoidal_dim + 1
nn.Linear(dim, time_dim),
nn.GELU(),
nn.Linear(time_dim, time_dim)
)
else: else:
self.time_mlp = MLP(1, time_dim) 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 # layers
@@ -287,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),
@@ -300,13 +301,13 @@ 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) t = self.time_mlp(time)
h = [] h = []
@@ -323,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
@@ -366,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)
@@ -421,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 -
@@ -527,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
@@ -541,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
@@ -549,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()
]) ])
@@ -571,7 +584,6 @@ 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,
@@ -580,12 +592,14 @@ 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.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)
self.ema_model = copy.deepcopy(self.model)
self.update_ema_every = update_ema_every self.update_ema_every = update_ema_every
self.step_start_ema = step_start_ema self.step_start_ema = step_start_ema
@@ -596,9 +610,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
@@ -608,22 +622,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'))
@@ -633,7 +636,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):
@@ -653,15 +656,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.18.0', version = '0.21.0',
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',