Compare commits

...
46 Commits
Author SHA1 Message Date
Phil Wang 931a5af2c3 bring in ddim sampling 2022-07-09 16:10:23 -07:00
Phil Wang a0c3443eaa optimizer should be saved and loaded 2022-07-08 17:43:41 -07:00
Phil Wang 662172851b add convert_image_to keyword argument, for forcing images being loaded to be converted to some format, greyscale, rgb, rgba, whatever 2022-07-08 09:22:59 -07:00
Phil Wang 0248b5e4d3 also make sure grad scaler actually exists in the saved pt file 2022-07-06 11:48:04 -07:00
Phil Wang 6b56af08a2 support multi-gpu training using huggingface accelerate, addressing https://github.com/lucidrains/denoising-diffusion-pytorch/pull/54 2022-07-06 11:45:56 -07:00
Phil Wang d4420248f1 tqdm auto instead 2022-06-29 19:48:49 -07:00
Phil Wang a536e5bee9 reorganize elucidating code 2022-06-29 08:49:38 -07:00
Phil Wang 1b85379d3a add clamping option to elucidated diffusion 2022-06-29 08:32:10 -07:00
Phil Wang 8859864f63 patch 2022-06-29 07:55:11 -07:00
Phil Wang 32657f035f Merge pull request #52 from AryaAftab/patch-1
Solve issue #26
2022-06-29 07:54:48 -07:00
Arya Aftab d97bc0278c Solve issue #26 2022-06-29 11:42:54 +04:30
Phil Wang 8408775cfc fix bug in elucidating sampling 2022-06-28 17:52:02 -07:00
Phil Wang 86fcb6785b release elucidating diffusion 2022-06-28 17:39:48 -07:00
Phil Wang c535d31fc5 Merge pull request #51 from lucidrains/pw/elucidating-ddpm
elucidating diffusion, first pass
2022-06-28 17:29:28 -07:00
Phil Wang 5db64fec4b refactor sigmas and gamma generation 2022-06-28 17:26:58 -07:00
Phil Wang b87ea27781 fix off by one 2022-06-28 16:05:21 -07:00
Phil Wang a8403b83fe no clamping when training from sigmas drawn from log normal distribution, clamp final images being sampled 2022-06-28 16:00:16 -07:00
Phil Wang f4b1d7a67c complete a first pass of elucidated ddpm 2022-06-28 15:15:21 -07:00
Phil Wang 618493714f clamp the sigma coming out of the log normal distribution 2022-06-28 14:34:39 -07:00
Phil Wang 76b79aa847 take care of equation 7 in the paper 2022-06-28 14:16:05 -07:00
Phil Wang be2bd8d320 cleanup again 2022-06-28 13:42:41 -07:00
Phil Wang c3d1607019 cleanup 2022-06-28 13:38:37 -07:00
Phil Wang 06b2e52645 get training working 2022-06-28 13:35:06 -07:00
Phil Wang 09b8a1c805 some basic scaffold for elucidating diffusion and derived values 2022-06-28 13:11:56 -07:00
Phil Wang d26acbcae6 more skip connections, as in guided diffusion 2022-06-27 13:23:32 -07:00
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
8 changed files with 674 additions and 211 deletions
+53 -2
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>
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>
[![PyPI version](https://badge.fury.io/py/denoising-diffusion-pytorch.svg)](https://badge.fury.io/py/denoising-diffusion-pytorch)
@@ -56,8 +60,9 @@ model = Unet(
diffusion = GaussianDiffusion(
model,
image_size = 128,
timesteps = 1000, # number of steps
loss_type = 'l1' # L1 or L2
timesteps = 1000, # number of steps
sampling_timesteps = 250, # number of sampling timesteps (using ddim for faster inference [see citation for ddim paper])
loss_type = 'l1' # L1 or L2
).cuda()
trainer = Trainer(
@@ -76,6 +81,22 @@ trainer.train()
Samples and model checkpoints will be logged to `./results` periodically
## 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
```bibtex
@@ -119,3 +140,33 @@ Samples and model checkpoints will be logged to `./results` periodically
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}
}
```
+1
View File
@@ -3,3 +3,4 @@ from denoising_diffusion_pytorch.denoising_diffusion_pytorch import GaussianDiff
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,3 +1,4 @@
import math
import torch
from torch import sqrt
from torch import nn, einsum
@@ -5,7 +6,7 @@ import torch.nn.functional as F
from torch.special import expm1
from tqdm import tqdm
from einops import rearrange, repeat
from einops import rearrange, repeat, reduce
from einops.layers.torch import Rearrange
# helpers
@@ -59,11 +60,14 @@ class MonotonicLinear(nn.Module):
# log(snr) that approximates the original linear schedule
def beta_linear_log_snr(t):
return -torch.log(expm1(1e-4 + 10 * (t ** 2)))
def log(t, eps = 1e-20):
return torch.log(t.clamp(min = eps))
def alpha_cosine_log_snr(t):
raise NotImplementedError
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 """
@@ -73,7 +77,8 @@ class learned_noise_schedule(nn.Module):
*,
log_snr_max,
log_snr_min,
hidden_dim = 1024
hidden_dim = 1024,
frac_gradient = 1.
):
super().__init__()
self.slope = log_snr_min - log_snr_max
@@ -90,7 +95,10 @@ class learned_noise_schedule(nn.Module):
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))
@@ -98,25 +106,29 @@ class learned_noise_schedule(nn.Module):
x = self.net(x)
normalized = self.slope * ((x - out_zero) / (out_one - out_zero)) + self.intercept
return normalized
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,
model,
*,
image_size,
channels = 3,
loss_type = 'l1',
noise_schedule = 'linear',
num_sample_steps = 500,
clip_sample_after_noise = False
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
assert model.learned_sinusoidal_cond
self.denoise_fn = denoise_fn
self.model = model
# image dimensions
@@ -129,12 +141,16 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
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
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}')
@@ -142,14 +158,19 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
# sampling
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_sample_after_noise = clip_sample_after_noise
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
return next(self.model.parameters()).device
@property
def loss_fn(self):
@@ -164,9 +185,6 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
# reviewer found an error in the equation in the paper (missing sigma)
# 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_next = self.log_snr(time_next)
c = -expm1(log_snr - log_snr_next)
@@ -174,10 +192,21 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
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)
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)
model_mean = sqrt(squared_alpha_next / squared_alpha) * (x - c * sqrt(squared_sigma) * pred_noise)
posterior_variance = squared_sigma_next * c
return model_mean, posterior_variance
@@ -208,12 +237,9 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
times_next = steps[i + 1]
img = self.p_sample(img, times, times_next)
if self.clip_sample_after_noise:
# clip after noise is added. perhaps this is sufficient for Imagen dynamic thresholding?
img.clamp_(-1., 1.)
img.clamp_(-1., 1.)
img = unnormalize_to_zero_to_one(img)
return img.clamp(0., 1.)
return img
@torch.no_grad()
def sample(self, batch_size = 16):
@@ -240,9 +266,17 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
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)
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):
b, c, h, w, device, img_size, = *img.shape, img.device, self.image_size
@@ -4,20 +4,29 @@ import torch
from torch import nn, einsum
import torch.nn.functional as F
from inspect import isfunction
from collections import namedtuple
from functools import partial
from torch.utils import data
from torch.cuda.amp import autocast, GradScaler
from torch.utils.data import Dataset, DataLoader
from multiprocessing import cpu_count
from pathlib import Path
from torch.optim import Adam
from torchvision import transforms, utils
from torchvision import transforms as T, utils
from PIL import Image
from tqdm import tqdm
from einops import rearrange
from einops import rearrange, reduce
from einops.layers.torch import Rearrange
from tqdm.auto import tqdm
from ema_pytorch import EMA
from accelerate import Accelerator
# constants
ModelPrediction = namedtuple('ModelPrediction', ['pred_noise', 'pred_x_start'])
# helpers functions
def exists(x):
@@ -33,6 +42,9 @@ def cycle(dl):
for data in dl:
yield data
def has_int_squareroot(num):
return (math.sqrt(num) ** 2) == num
def num_to_groups(num, divisor):
groups = num // divisor
remainder = num % divisor
@@ -41,6 +53,13 @@ def num_to_groups(num, divisor):
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
@@ -49,21 +68,6 @@ def unnormalize_to_zero_to_one(t):
# 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):
def __init__(self, fn):
super().__init__()
@@ -72,25 +76,14 @@ class Residual(nn.Module):
def forward(self, x, *args, **kwargs):
return self.fn(x, *args, **kwargs) + x
class SinusoidalPosEmb(nn.Module):
def __init__(self, dim):
super().__init__()
self.dim = dim
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 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):
return nn.ConvTranspose2d(dim, dim, 4, 2, 1)
def Downsample(dim):
return nn.Conv2d(dim, dim, 4, 2, 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):
@@ -114,6 +107,39 @@ class PreNorm(nn.Module):
x = self.norm(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
class Block(nn.Module):
@@ -157,6 +183,7 @@ class ResnetBlock(nn.Module):
h = self.block1(x, scale_shift = scale_shift)
h = self.block2(h)
return h + self.res_conv(x)
class LinearAttention(nn.Module):
@@ -212,18 +239,6 @@ class Attention(nn.Module):
# 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):
def __init__(
self,
@@ -234,7 +249,8 @@ class Unet(nn.Module):
channels = 3,
resnet_block_groups = 8,
learned_variance = False,
sinusoidal_cond_mlp = True
learned_sinusoidal_cond = False,
learned_sinusoidal_dim = 16
):
super().__init__()
@@ -242,7 +258,7 @@ class Unet(nn.Module):
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)
dims = [init_dim, *map(lambda m: dim * m, dim_mults)]
@@ -254,17 +270,21 @@ class Unet(nn.Module):
time_dim = dim * 4
self.sinusoidal_cond_mlp = sinusoidal_cond_mlp
self.learned_sinusoidal_cond = learned_sinusoidal_cond
if sinusoidal_cond_mlp:
self.time_mlp = nn.Sequential(
SinusoidalPosEmb(dim),
nn.Linear(dim, time_dim),
nn.GELU(),
nn.Linear(time_dim, time_dim)
)
if learned_sinusoidal_cond:
sinu_pos_emb = LearnedSinusoidalPosEmb(learned_sinusoidal_dim)
fourier_dim = learned_sinusoidal_dim + 1
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
@@ -276,10 +296,10 @@ class Unet(nn.Module):
is_last = ind >= (num_resolutions - 1)
self.downs.append(nn.ModuleList([
block_klass(dim_in, dim_out, time_emb_dim = time_dim),
block_klass(dim_out, dim_out, time_emb_dim = time_dim),
Residual(PreNorm(dim_out, LinearAttention(dim_out))),
Downsample(dim_out) if not is_last else nn.Identity()
block_klass(dim_in, dim_in, time_emb_dim = time_dim),
block_klass(dim_in, dim_in, time_emb_dim = time_dim),
Residual(PreNorm(dim_in, LinearAttention(dim_in))),
Downsample(dim_in, dim_out) if not is_last else nn.Conv2d(dim_in, dim_out, 3, padding = 1)
]))
mid_dim = dims[-1]
@@ -287,35 +307,38 @@ class Unet(nn.Module):
self.mid_attn = Residual(PreNorm(mid_dim, Attention(mid_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:])):
is_last = ind >= (num_resolutions - 1)
for ind, (dim_in, dim_out) in enumerate(reversed(in_out)):
is_last = ind == (len(in_out) - 1)
self.ups.append(nn.ModuleList([
block_klass(dim_out * 2, dim_in, time_emb_dim = time_dim),
block_klass(dim_in, dim_in, time_emb_dim = time_dim),
Residual(PreNorm(dim_in, LinearAttention(dim_in))),
Upsample(dim_in) if not is_last else nn.Identity()
block_klass(dim_out + dim_in, dim_out, time_emb_dim = time_dim),
block_klass(dim_out + dim_in, dim_out, time_emb_dim = time_dim),
Residual(PreNorm(dim_out, LinearAttention(dim_out))),
Upsample(dim_out, dim_in) if not is_last else nn.Conv2d(dim_out, dim_in, 3, padding = 1)
]))
default_out_dim = channels * (1 if not learned_variance else 2)
self.out_dim = default(out_dim, default_out_dim)
self.final_conv = nn.Sequential(
block_klass(dim, dim),
nn.Conv2d(dim, self.out_dim, 1)
)
self.final_res_block = block_klass(dim * 2, dim, time_emb_dim = time_dim)
self.final_conv = nn.Conv2d(dim, self.out_dim, 1)
def forward(self, x, time):
x = self.init_conv(x)
r = x.clone()
t = self.time_mlp(time)
h = []
for block1, block2, attn, downsample in self.downs:
x = block1(x, t)
h.append(x)
x = block2(x, t)
x = attn(x)
h.append(x)
x = downsample(x)
x = self.mid_block1(x, t)
@@ -323,12 +346,18 @@ class Unet(nn.Module):
x = self.mid_block2(x, t)
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 = torch.cat((x, h.pop()), dim = 1)
x = block2(x, t)
x = attn(x)
x = upsample(x)
x = torch.cat((x, r), dim = 1)
x = self.final_res_block(x, t)
return self.final_conv(x)
# gaussian diffusion trainer class
@@ -351,7 +380,7 @@ def cosine_beta_schedule(timesteps, s = 0.008):
"""
steps = timesteps + 1
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]
betas = 1 - (alphas_cumprod[1:] / alphas_cumprod[:-1])
return torch.clip(betas, 0, 0.999)
@@ -359,23 +388,29 @@ def cosine_beta_schedule(timesteps, s = 0.008):
class GaussianDiffusion(nn.Module):
def __init__(
self,
denoise_fn,
model,
*,
image_size,
channels = 3,
timesteps = 1000,
sampling_timesteps = None,
loss_type = 'l1',
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,
ddim_sampling_eta = 1.
):
super().__init__()
assert not (type(self) == GaussianDiffusion and denoise_fn.channels != denoise_fn.out_dim)
assert not (type(self) == GaussianDiffusion and model.channels != model.out_dim)
self.channels = channels
self.image_size = image_size
self.denoise_fn = denoise_fn
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':
@@ -391,6 +426,14 @@ class GaussianDiffusion(nn.Module):
self.num_timesteps = int(timesteps)
self.loss_type = loss_type
# sampling related parameters
self.sampling_timesteps = default(sampling_timesteps, timesteps) # default num sampling timesteps to number of timesteps at training
assert self.sampling_timesteps <= timesteps
self.is_ddim_sampling = self.sampling_timesteps < timesteps
self.ddim_sampling_eta = ddim_sampling_eta
# helper function to register buffer from float64 to float32
register_buffer = lambda name, val: self.register_buffer(name, val.to(torch.float32))
@@ -421,12 +464,22 @@ class GaussianDiffusion(nn.Module):
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))
# 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):
return (
extract(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t -
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):
posterior_mean = (
extract(self.posterior_mean_coef1, t, x_t.shape) * x_start +
@@ -436,15 +489,22 @@ class GaussianDiffusion(nn.Module):
posterior_log_variance_clipped = extract(self.posterior_log_variance_clipped, t, x_t.shape)
return posterior_mean, posterior_variance, posterior_log_variance_clipped
def p_mean_variance(self, x, t, clip_denoised: bool):
model_output = self.denoise_fn(x, t)
def model_predictions(self, x, t):
model_output = self.model(x, t)
if self.objective == 'pred_noise':
x_start = self.predict_start_from_noise(x, t = t, noise = model_output)
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
else:
raise ValueError(f'unknown objective {self.objective}')
return ModelPrediction(pred_noise, x_start)
def p_mean_variance(self, x, t, clip_denoised: bool):
preds = self.model_predictions(x, t)
x_start = preds.pred_x_start
if clip_denoised:
x_start.clamp_(-1., 1.)
@@ -453,32 +513,62 @@ class GaussianDiffusion(nn.Module):
return model_mean, posterior_variance, posterior_log_variance
@torch.no_grad()
def p_sample(self, x, t, clip_denoised=True):
def p_sample(self, x, t: int, clip_denoised = True):
b, *_, device = *x.shape, x.device
model_mean, _, model_log_variance = self.p_mean_variance(x=x, t=t, clip_denoised=clip_denoised)
noise = torch.randn_like(x)
# no noise when t == 0
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
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 = batched_times, clip_denoised = clip_denoised)
noise = torch.randn_like(x) if t > 0 else 0. # no noise if t == 0
return model_mean + (0.5 * model_log_variance).exp() * noise
@torch.no_grad()
def p_sample_loop(self, shape):
device = self.betas.device
batch, device = shape[0], self.betas.device
b = shape[0]
img = torch.randn(shape, device=device)
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))
for t in tqdm(reversed(range(0, self.num_timesteps)), desc = 'sampling loop time step'):
img = self.p_sample(img, t)
img = unnormalize_to_zero_to_one(img)
return img
@torch.no_grad()
def ddim_sample(self, shape, clip_denoised = True):
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
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.)
c1 = eta * ((1 - alpha / alpha_next) * (1 - alpha_next) / (1 - alpha)).sqrt()
c2 = ((1 - alpha_next) - torch.square(c1)).sqrt()
img = x_start * alpha_next.sqrt() + \
c1 * torch.randn_like(img) + \
c2 * pred_noise
img = unnormalize_to_zero_to_one(img)
return img
@torch.no_grad()
def sample(self, batch_size = 16):
image_size = self.image_size
channels = self.channels
return self.p_sample_loop((batch_size, channels, image_size, image_size))
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()
def interpolate(self, x1, x2, t = None, lam = 0.5):
@@ -517,8 +607,8 @@ class GaussianDiffusion(nn.Module):
b, c, h, w = x_start.shape
noise = default(noise, lambda: torch.randn_like(x_start))
x = self.q_sample(x_start=x_start, t=t, noise=noise)
model_out = self.denoise_fn(x, t)
x = self.q_sample(x_start = x_start, t = t, noise = noise)
model_out = self.model(x, t)
if self.objective == 'pred_noise':
target = noise
@@ -527,8 +617,11 @@ class GaussianDiffusion(nn.Module):
else:
raise ValueError(f'unknown objective {self.objective}')
loss = self.loss_fn(model_out, target)
return loss
loss = self.loss_fn(model_out, target, reduction = 'none')
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):
b, c, h, w, device, img_size, = *img.shape, img.device, self.image_size
@@ -540,18 +633,28 @@ class GaussianDiffusion(nn.Module):
# dataset classes
class Dataset(data.Dataset):
def __init__(self, folder, image_size, exts = ['jpg', 'jpeg', 'png']):
class Dataset(Dataset):
def __init__(
self,
folder,
image_size,
exts = ['jpg', 'jpeg', 'png', 'tiff'],
augment_horizontal_flip = False,
convert_image_to = None
):
super().__init__()
self.folder = folder
self.image_size = image_size
self.paths = [p for ext in exts for p in Path(f'{folder}').glob(f'**/*.{ext}')]
self.transform = transforms.Compose([
transforms.Resize(image_size),
transforms.RandomHorizontalFlip(),
transforms.CenterCrop(image_size),
transforms.ToTensor()
maybe_convert_fn = partial(convert_image_to, convert_image_to) if exists(convert_image_to) else nn.Identity()
self.transform = T.Compose([
T.Lambda(maybe_convert_fn),
T.Resize(image_size),
T.RandomHorizontalFlip() if augment_horizontal_flip else nn.Identity(),
T.CenterCrop(image_size),
T.ToTensor()
])
def __len__(self):
@@ -570,103 +673,137 @@ class Trainer(object):
diffusion_model,
folder,
*,
ema_decay = 0.995,
image_size = 128,
train_batch_size = 32,
train_batch_size = 16,
gradient_accumulate_every = 1,
augment_horizontal_flip = True,
train_lr = 1e-4,
train_num_steps = 100000,
gradient_accumulate_every = 2,
amp = False,
step_start_ema = 2000,
update_ema_every = 10,
ema_update_every = 10,
ema_decay = 0.995,
save_and_sample_every = 1000,
results_folder = './results'
num_samples = 25,
results_folder = './results',
amp = False,
fp16 = False,
split_batches = True,
convert_image_to = None
):
super().__init__()
self.model = diffusion_model
self.ema = EMA(ema_decay)
self.ema_model = copy.deepcopy(self.model)
self.update_ema_every = update_ema_every
self.step_start_ema = step_start_ema
self.accelerator = Accelerator(
split_batches = split_batches,
mixed_precision = 'fp16' if fp16 else 'no'
)
self.accelerator.native_amp = amp
self.model = diffusion_model
assert has_int_squareroot(num_samples), 'number of samples must have an integer square root'
self.num_samples = num_samples
self.save_and_sample_every = save_and_sample_every
self.batch_size = train_batch_size
self.image_size = diffusion_model.image_size
self.gradient_accumulate_every = gradient_accumulate_every
self.train_num_steps = train_num_steps
self.ds = Dataset(folder, image_size)
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.train_num_steps = train_num_steps
self.image_size = diffusion_model.image_size
# dataset and dataloader
self.ds = Dataset(folder, self.image_size, augment_horizontal_flip = augment_horizontal_flip, convert_image_to = convert_image_to)
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.amp = amp
self.scaler = GradScaler(enabled = amp)
# prepare model, dataloader, optimizer with accelerator
self.results_folder = Path(results_folder)
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)
self.model, self.dl, self.opt = self.accelerator.prepare(self.model, self.dl, self.opt)
def save(self, milestone):
if not self.accelerator.is_local_main_process:
return
data = {
'step': self.step,
'model': self.model.state_dict(),
'ema': self.ema_model.state_dict(),
'scaler': self.scaler.state_dict()
'model': self.accelerator.get_state_dict(self.model),
'opt': self.opt.state_dict(),
'ema': self.ema.state_dict(),
'scaler': self.accelerator.scaler.state_dict() if exists(self.accelerator.scaler) else None
}
torch.save(data, str(self.results_folder / f'model-{milestone}.pt'))
def load(self, milestone):
data = torch.load(str(self.results_folder / f'model-{milestone}.pt'))
model = self.accelerator.unwrap_model(self.model)
model.load_state_dict(data['model'])
self.step = data['step']
self.model.load_state_dict(data['model'])
self.ema_model.load_state_dict(data['ema'])
self.scaler.load_state_dict(data['scaler'])
self.opt.load_state_dict(data['opt'])
self.ema.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):
with tqdm(initial = self.step, total = self.train_num_steps) as pbar:
accelerator = self.accelerator
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()
with autocast(enabled = self.amp):
for _ in range(self.gradient_accumulate_every):
data = next(self.dl).to(device)
with self.accelerator.autocast():
loss = self.model(data)
self.scaler.scale(loss / self.gradient_accumulate_every).backward()
self.accelerator.backward(loss / self.gradient_accumulate_every)
pbar.set_description(f'loss: {loss.item():.4f}')
pbar.set_description(f'loss: {loss.item():.4f}')
self.scaler.step(self.opt)
self.scaler.update()
accelerator.wait_for_everyone()
self.opt.step()
self.opt.zero_grad()
if self.step % self.update_ema_every == 0:
self.step_ema()
accelerator.wait_for_everyone()
if self.step != 0 and self.step % self.save_and_sample_every == 0:
self.ema_model.eval()
if accelerator.is_main_process:
self.ema.to(device)
self.ema.update()
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)
utils.save_image(all_images, str(self.results_folder / f'sample-{milestone}.png'), nrow = 6)
self.save(milestone)
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)
print('training complete')
accelerator.print('training complete')
@@ -0,0 +1,220 @@
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,4 +1,5 @@
import torch
from collections import namedtuple
from math import pi, sqrt, log as ln
from inspect import isfunction
from torch import nn, einsum
@@ -10,6 +11,8 @@ from denoising_diffusion_pytorch.denoising_diffusion_pytorch import GaussianDiff
NAT = 1. / ln(2)
ModelPrediction = namedtuple('ModelPrediction', ['pred_noise', 'pred_x_start', 'pred_variance'])
# helper functions
def exists(x):
@@ -22,7 +25,7 @@ def default(val, d):
# tensor helpers
def log(t, eps = 1e-12):
def log(t, eps = 1e-15):
return torch.log(t.clamp(min = eps))
def meanflat(x):
@@ -67,17 +70,31 @@ def discretized_gaussian_log_likelihood(x, *, means, log_scales, thres = 0.999):
class LearnedGaussianDiffusion(GaussianDiffusion):
def __init__(
self,
denoise_fn,
model,
vb_loss_weight = 0.001, # lambda was 0.001 in the paper
*args,
**kwargs
):
super().__init__(denoise_fn, *args, **kwargs)
assert denoise_fn.out_dim == (denoise_fn.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`'
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.denoise_fn(x, t))
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)
@@ -102,7 +119,7 @@ class LearnedGaussianDiffusion(GaussianDiffusion):
# model output
model_output = self.denoise_fn(x_t, t)
model_output = self.model(x_t, t)
# calculating kl loss for learned variance (interpolation)
@@ -22,22 +22,23 @@ def default(val, d):
class WeightedObjectiveGaussianDiffusion(GaussianDiffusion):
def __init__(
self,
denoise_fn,
model,
*args,
pred_noise_loss_weight = 0.1,
pred_x_start_loss_weight = 0.1,
**kwargs
):
super().__init__(denoise_fn, *args, **kwargs)
channels = denoise_fn.channels
assert denoise_fn.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'
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.denoise_fn(x, t)
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)
@@ -58,7 +59,7 @@ class WeightedObjectiveGaussianDiffusion(GaussianDiffusion):
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.denoise_fn(x_t, t)
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
+3 -1
View File
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
setup(
name = 'denoising-diffusion-pytorch',
packages = find_packages(),
version = '0.17.1',
version = '0.25.1',
license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang',
@@ -15,7 +15,9 @@ setup(
'generative models'
],
install_requires=[
'accelerate',
'einops',
'ema-pytorch',
'pillow',
'torch',
'torchvision',