Compare commits

...
28 Commits
Author SHA1 Message Date
Phil Wang 6dda508ff6 in ddim, clip x0 before calculation of predicted noise, thanks to @lukovnikov again for pointing out this inconsistency with glides implementation 2022-09-01 09:34:50 -07:00
Phil Wang 9ec8d27217 0.27.7 2022-08-31 07:36:19 -07:00
Phil Wang 4b4ebab7c3 Merge pull request #83 from lukovnikov/fix_ddim
Fix ddim
2022-08-31 07:02:26 -07:00
lukovnikov e4a4e4acaa Revert "Revert "fix ddim sampling""
This reverts commit cd8329cdd7.
2022-08-31 15:28:35 +02:00
lukovnikov cd8329cdd7 Revert "fix ddim sampling"
This reverts commit aec2a26984.
2022-08-31 15:25:51 +02:00
lukovnikov aec2a26984 fix ddim sampling 2022-08-31 15:24:59 +02:00
Phil Wang c78709f887 0.27.6 2022-08-31 06:20:59 -07:00
Phil Wang 4436128a0b Merge pull request #82 from TheDudeFromCI/patch-1
Update step before training checkpoint
2022-08-31 06:20:31 -07:00
TheDudeFromCI e46a89e2bc Update step before training checkpoint 2022-08-31 02:19:12 -07:00
Phil Wang 42158d6248 fix ddim, for issue https://github.com/lucidrains/denoising-diffusion-pytorch/issues/81 2022-08-30 20:30:29 -07:00
Phil Wang 44f95e2e9d readme 2022-08-22 09:13:32 -07:00
Phil Wang d9275a744c add weight standardization prior to groupnorm, lessen cosine sim attention scale to 10 for fp16 2022-08-17 11:42:54 -07:00
Phil Wang beb2f2d8dd add self conditioning for elucidated ddpm 2022-08-10 13:15:45 -07:00
Phil Wang f0d59acdfd fix sampling ddpm tqdm 2022-08-10 12:02:56 -07:00
Phil Wang 689593a579 add the new self conditioning technique from hintons group from bit diffusion paper
0.27.0
2022-08-10 10:52:07 -07:00
Phil Wang eba44498d1 higher epsilon for fp16 in layernorm 2022-07-29 13:29:52 -07:00
Phil Wang 12f95b33d8 rescale values to prevent linear attention from overflowing in fp16 setting 2022-07-27 12:25:40 -07:00
Phil Wang 6b504c4ae9 fix accelerator prepare bug for dataloader 2022-07-25 08:13:46 -07:00
Phil Wang 37334ae824 fix a bug with ddim and predict x0 objective 2022-07-18 19:04:57 -07:00
Phil Wang 6eba6cdd50 refresh images 2022-07-18 11:30:34 -07:00
Phil Wang 555566c188 take a gamble on cosine sim attention 2022-07-18 11:29:20 -07:00
Phil Wang 2b742dd2cc move accelerator backward outside of autocast context, also calculate total loss correctly across gradient accumulated steps 2022-07-11 21:02:06 -07:00
Phil Wang 1345a8a41d do not noise at the last timestep for ddim 2022-07-09 18:36:45 -07:00
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
9 changed files with 416 additions and 133 deletions
+55 -5
View File
@@ -1,4 +1,4 @@
<img src="./denoising-diffusion.png" width="500px"></img> <img src="./images/denoising-diffusion.png" width="500px"></img>
## Denoising Diffusion Probabilistic Model, in Pytorch ## Denoising Diffusion Probabilistic Model, in Pytorch
@@ -10,7 +10,9 @@ Youtube AI Educators - <a href="https://www.youtube.com/watch?v=W-O7AZNzbzQ">Yan
<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> Update: Turns out none of the technicalities really matters at all | <a href="https://arxiv.org/abs/2208.09392">"Cold Diffusion" paper</a>
<img src="./images/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)
@@ -60,15 +62,16 @@ model = Unet(
diffusion = GaussianDiffusion( diffusion = GaussianDiffusion(
model, model,
image_size = 128, image_size = 128,
timesteps = 1000, # number of steps timesteps = 1000, # number of steps
loss_type = 'l1' # L1 or L2 sampling_timesteps = 250, # number of sampling timesteps (using ddim for faster inference [see citation for ddim paper])
loss_type = 'l1' # L1 or L2
).cuda() ).cuda()
trainer = Trainer( 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 = 8e-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
@@ -80,6 +83,22 @@ trainer.train()
Samples and model checkpoints will be logged to `./results` periodically 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 ## Citations
```bibtex ```bibtex
@@ -143,3 +162,34 @@ Samples and model checkpoints will be logged to `./results` periodically
volume = {abs/2206.00364} 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}
}
```
```bibtex
@misc{chen2022analog,
title = {Analog Bits: Generating Discrete Data using Diffusion Models with Self-Conditioning},
author = {Ting Chen and Ruixiang Zhang and Geoffrey Hinton},
year = {2022},
eprint = {2208.04202},
archivePrefix = {arXiv},
primaryClass = {cs.CV}
}
```
```bibtex
@article{Qiao2019WeightS,
title = {Weight Standardization},
author = {Siyuan Qiao and Huiyu Wang and Chenxi Liu and Wei Shen and Alan Loddon Yuille},
journal = {ArXiv},
year = {2019},
volume = {abs/1903.10520}
}
```
@@ -112,7 +112,7 @@ class learned_noise_schedule(nn.Module):
class ContinuousTimeGaussianDiffusion(nn.Module): class ContinuousTimeGaussianDiffusion(nn.Module):
def __init__( def __init__(
self, self,
denoise_fn, model,
*, *,
image_size, image_size,
channels = 3, channels = 3,
@@ -126,9 +126,10 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
p2_loss_weight_k = 1 p2_loss_weight_k = 1
): ):
super().__init__() super().__init__()
assert denoise_fn.learned_sinusoidal_cond assert model.learned_sinusoidal_cond
assert not model.self_condition, 'not supported yet'
self.denoise_fn = denoise_fn self.model = model
# image dimensions # image dimensions
@@ -170,7 +171,7 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
@property @property
def device(self): def device(self):
return next(self.denoise_fn.parameters()).device return next(self.model.parameters()).device
@property @property
def loss_fn(self): def loss_fn(self):
@@ -195,7 +196,7 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
alpha, sigma, alpha_next = map(sqrt, (squared_alpha, squared_sigma, squared_alpha_next)) alpha, sigma, alpha_next = map(sqrt, (squared_alpha, squared_sigma, squared_alpha_next))
batch_log_snr = repeat(log_snr, ' -> b', b = x.shape[0]) batch_log_snr = repeat(log_snr, ' -> b', b = x.shape[0])
pred_noise = self.denoise_fn(x, batch_log_snr) pred_noise = self.model(x, batch_log_snr)
if self.clip_sample_denoised: if self.clip_sample_denoised:
x_start = (x - sigma * pred_noise) / alpha x_start = (x - sigma * pred_noise) / alpha
@@ -266,7 +267,7 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
noise = default(noise, lambda: torch.randn_like(x_start)) noise = default(noise, lambda: torch.randn_like(x_start))
x, log_snr = self.q_sample(x_start = x_start, times = times, noise = noise) x, log_snr = self.q_sample(x_start = x_start, times = times, noise = noise)
model_out = self.denoise_fn(x, log_snr) model_out = self.model(x, log_snr)
losses = self.loss_fn(model_out, noise, reduction = 'none') losses = self.loss_fn(model_out, noise, reduction = 'none')
losses = reduce(losses, 'b ... -> b', 'mean') losses = reduce(losses, 'b ... -> b', 'mean')
@@ -1,26 +1,32 @@
import math import math
import copy import copy
from pathlib import Path
from random import random
from functools import partial
from collections import namedtuple
from multiprocessing import cpu_count
import torch import torch
from torch import nn, einsum from torch import nn, einsum
import torch.nn.functional as F import torch.nn.functional as F
from inspect import isfunction from torch.utils.data import Dataset, DataLoader
from functools import partial
from torch.utils import data
from multiprocessing import cpu_count
from torch.cuda.amp import autocast, GradScaler
from pathlib import Path
from torch.optim import Adam from torch.optim import Adam
from torchvision import transforms, utils from torchvision import transforms as T, utils
from PIL import Image
from einops import rearrange, reduce from einops import rearrange, reduce
from einops.layers.torch import Rearrange from einops.layers.torch import Rearrange
from PIL import Image
from tqdm.auto import tqdm from tqdm.auto import tqdm
from ema_pytorch import EMA from ema_pytorch import EMA
from accelerate import Accelerator
# constants
ModelPrediction = namedtuple('ModelPrediction', ['pred_noise', 'pred_x_start'])
# helpers functions # helpers functions
def exists(x): def exists(x):
@@ -29,13 +35,19 @@ def exists(x):
def default(val, d): def default(val, d):
if exists(val): if exists(val):
return val return val
return d() if isfunction(d) else d return d() if callable(d) else d
def identity(t, *args, **kwargs):
return t
def cycle(dl): def cycle(dl):
while True: while True:
for data in dl: for data in dl:
yield data yield data
def has_int_squareroot(num):
return (math.sqrt(num) ** 2) == num
def num_to_groups(num, divisor): def num_to_groups(num, divisor):
groups = num // divisor groups = num // divisor
remainder = num % divisor remainder = num % divisor
@@ -44,6 +56,16 @@ def num_to_groups(num, divisor):
arr.append(remainder) arr.append(remainder)
return arr return arr
def convert_image_to(img_type, image):
if image.mode != img_type:
return image.convert(img_type)
return image
def l2norm(t):
return F.normalize(t, dim = -1)
# normalization functions
def normalize_to_neg_one_to_one(img): def normalize_to_neg_one_to_one(img):
return img * 2 - 1 return img * 2 - 1
@@ -69,17 +91,31 @@ def Upsample(dim, dim_out = None):
def Downsample(dim, dim_out = None): def Downsample(dim, dim_out = None):
return nn.Conv2d(dim, default(dim_out, dim), 4, 2, 1) return nn.Conv2d(dim, default(dim_out, dim), 4, 2, 1)
class WeightStandardizedConv2d(nn.Conv2d):
"""
https://arxiv.org/abs/1903.10520
weight standardization purportedly works synergistically with group normalization
"""
def forward(self, x):
eps = 1e-5 if x.dtype == torch.float32 else 1e-3
weight = self.weight
mean = reduce(weight, 'o ... -> o 1 1 1', 'mean')
var = reduce(weight, 'o ... -> o 1 1 1', partial(torch.var, unbiased = False))
normalized_weight = (weight - mean) * (var + eps).rsqrt()
return F.conv2d(x, normalized_weight, self.bias, self.stride, self.padding, self.dilation, self.groups)
class LayerNorm(nn.Module): class LayerNorm(nn.Module):
def __init__(self, dim, eps = 1e-5): def __init__(self, dim):
super().__init__() super().__init__()
self.eps = eps
self.g = nn.Parameter(torch.ones(1, dim, 1, 1)) self.g = nn.Parameter(torch.ones(1, dim, 1, 1))
self.b = nn.Parameter(torch.zeros(1, dim, 1, 1))
def forward(self, x): def forward(self, x):
eps = 1e-5 if x.dtype == torch.float32 else 1e-3
var = torch.var(x, dim = 1, unbiased = False, keepdim = True) var = torch.var(x, dim = 1, unbiased = False, keepdim = True)
mean = torch.mean(x, dim = 1, keepdim = True) mean = torch.mean(x, dim = 1, keepdim = True)
return (x - mean) / (var + self.eps).sqrt() * self.g + self.b return (x - mean) * (var + eps).rsqrt() * self.g
class PreNorm(nn.Module): class PreNorm(nn.Module):
def __init__(self, dim, fn): def __init__(self, dim, fn):
@@ -129,7 +165,7 @@ class LearnedSinusoidalPosEmb(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.proj = WeightStandardizedConv2d(dim, dim_out, 3, padding = 1)
self.norm = nn.GroupNorm(groups, dim_out) self.norm = nn.GroupNorm(groups, dim_out)
self.act = nn.SiLU() self.act = nn.SiLU()
@@ -192,6 +228,8 @@ class LinearAttention(nn.Module):
k = k.softmax(dim = -1) k = k.softmax(dim = -1)
q = q * self.scale q = q * self.scale
v = v / (h * w)
context = torch.einsum('b h d n, b h e n -> b h d e', k, v) context = torch.einsum('b h d n, b h e n -> b h d e', k, v)
out = torch.einsum('b h d e, b h d n -> b h e n', context, q) out = torch.einsum('b h d e, b h d n -> b h e n', context, q)
@@ -199,9 +237,9 @@ class LinearAttention(nn.Module):
return self.to_out(out) return self.to_out(out)
class Attention(nn.Module): class Attention(nn.Module):
def __init__(self, dim, heads = 4, dim_head = 32): def __init__(self, dim, heads = 4, dim_head = 32, scale = 10):
super().__init__() super().__init__()
self.scale = dim_head ** -0.5 self.scale = scale
self.heads = heads self.heads = heads
hidden_dim = dim_head * heads hidden_dim = dim_head * heads
self.to_qkv = nn.Conv2d(dim, hidden_dim * 3, 1, bias = False) self.to_qkv = nn.Conv2d(dim, hidden_dim * 3, 1, bias = False)
@@ -211,12 +249,11 @@ class Attention(nn.Module):
b, c, h, w = x.shape b, c, h, w = x.shape
qkv = self.to_qkv(x).chunk(3, dim = 1) qkv = self.to_qkv(x).chunk(3, dim = 1)
q, k, v = map(lambda t: rearrange(t, 'b (h c) x y -> b h c (x y)', h = self.heads), qkv) q, k, v = map(lambda t: rearrange(t, 'b (h c) x y -> b h c (x y)', h = self.heads), qkv)
q = q * self.scale
sim = einsum('b h d i, b h d j -> b h i j', q, k) q, k = map(l2norm, (q, k))
sim = sim - sim.amax(dim = -1, keepdim = True).detach()
sim = einsum('b h d i, b h d j -> b h i j', q, k) * self.scale
attn = sim.softmax(dim = -1) attn = sim.softmax(dim = -1)
out = einsum('b h i j, b h d j -> b h i d', attn, v) out = einsum('b h i j, b h d j -> b h i d', attn, v)
out = rearrange(out, 'b h (x y) d -> b (h d) x y', x = h, y = w) out = rearrange(out, 'b h (x y) d -> b (h d) x y', x = h, y = w)
return self.to_out(out) return self.to_out(out)
@@ -231,6 +268,7 @@ 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,
self_condition = False,
resnet_block_groups = 8, resnet_block_groups = 8,
learned_variance = False, learned_variance = False,
learned_sinusoidal_cond = False, learned_sinusoidal_cond = False,
@@ -241,9 +279,11 @@ class Unet(nn.Module):
# determine dimensions # determine dimensions
self.channels = channels self.channels = channels
self.self_condition = self_condition
input_channels = channels * (2 if self_condition else 1)
init_dim = default(init_dim, dim) init_dim = default(init_dim, dim)
self.init_conv = nn.Conv2d(channels, init_dim, 7, padding = 3) self.init_conv = nn.Conv2d(input_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)]
in_out = list(zip(dims[:-1], dims[1:])) in_out = list(zip(dims[:-1], dims[1:]))
@@ -307,7 +347,11 @@ class Unet(nn.Module):
self.final_res_block = block_klass(dim * 2, dim, time_emb_dim = time_dim) self.final_res_block = block_klass(dim * 2, dim, time_emb_dim = time_dim)
self.final_conv = nn.Conv2d(dim, self.out_dim, 1) self.final_conv = nn.Conv2d(dim, self.out_dim, 1)
def forward(self, x, time): def forward(self, x, time, x_self_cond = None):
if self.self_condition:
x_self_cond = default(x_self_cond, lambda: torch.zeros_like(x))
x = torch.cat((x_self_cond, x), dim = 1)
x = self.init_conv(x) x = self.init_conv(x)
r = x.clone() r = x.clone()
@@ -372,25 +416,32 @@ def cosine_beta_schedule(timesteps, s = 0.008):
class GaussianDiffusion(nn.Module): class GaussianDiffusion(nn.Module):
def __init__( def __init__(
self, self,
denoise_fn, model,
*, *,
image_size, image_size,
channels = 3,
timesteps = 1000, timesteps = 1000,
sampling_timesteps = None,
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_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 p2_loss_weight_k = 1,
ddim_sampling_eta = 1.
): ):
super().__init__() 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)
assert not model.learned_sinusoidal_cond
self.model = model
self.channels = self.model.channels
self.self_condition = self.model.self_condition
self.channels = channels
self.image_size = image_size self.image_size = image_size
self.denoise_fn = denoise_fn
self.objective = objective 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': if beta_schedule == 'linear':
betas = linear_beta_schedule(timesteps) betas = linear_beta_schedule(timesteps)
elif beta_schedule == 'cosine': elif beta_schedule == 'cosine':
@@ -406,6 +457,14 @@ class GaussianDiffusion(nn.Module):
self.num_timesteps = int(timesteps) self.num_timesteps = int(timesteps)
self.loss_type = loss_type 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 # helper function to register buffer from float64 to float32
register_buffer = lambda name, val: self.register_buffer(name, val.to(torch.float32)) register_buffer = lambda name, val: self.register_buffer(name, val.to(torch.float32))
@@ -446,6 +505,12 @@ class GaussianDiffusion(nn.Module):
extract(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape) * noise extract(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape) * noise
) )
def predict_noise_from_start(self, x_t, t, x0):
return (
(extract(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t - x0) / \
extract(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape)
)
def q_posterior(self, x_start, x_t, t): def q_posterior(self, x_start, x_t, t):
posterior_mean = ( posterior_mean = (
extract(self.posterior_mean_coef1, t, x_t.shape) * x_start + extract(self.posterior_mean_coef1, t, x_t.shape) * x_start +
@@ -455,49 +520,97 @@ class GaussianDiffusion(nn.Module):
posterior_log_variance_clipped = extract(self.posterior_log_variance_clipped, t, x_t.shape) posterior_log_variance_clipped = extract(self.posterior_log_variance_clipped, t, x_t.shape)
return posterior_mean, posterior_variance, posterior_log_variance_clipped return posterior_mean, posterior_variance, posterior_log_variance_clipped
def p_mean_variance(self, x, t, clip_denoised: bool): def model_predictions(self, x, t, x_self_cond = None, clip_x_start = False):
model_output = self.denoise_fn(x, t) model_output = self.model(x, t, x_self_cond)
maybe_clip = partial(torch.clamp, min = -1., max = 1.) if clip_x_start else identity
if self.objective == 'pred_noise': 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, pred_noise)
x_start = maybe_clip(x_start)
elif self.objective == 'pred_x0': elif self.objective == 'pred_x0':
x_start = model_output x_start = model_output
else: x_start = maybe_clip(x_start)
raise ValueError(f'unknown objective {self.objective}') pred_noise = self.predict_noise_from_start(x, t, x_start)
return ModelPrediction(pred_noise, x_start)
def p_mean_variance(self, x, t, x_self_cond = None, clip_denoised = True):
preds = self.model_predictions(x, t, x_self_cond)
x_start = preds.pred_x_start
if clip_denoised: if clip_denoised:
x_start.clamp_(-1., 1.) x_start.clamp_(-1., 1.)
model_mean, posterior_variance, posterior_log_variance = self.q_posterior(x_start = x_start, x_t = x, t = t) model_mean, posterior_variance, posterior_log_variance = self.q_posterior(x_start = x_start, x_t = x, t = t)
return model_mean, posterior_variance, posterior_log_variance return model_mean, posterior_variance, posterior_log_variance, x_start
@torch.no_grad() @torch.no_grad()
def p_sample(self, x, t, clip_denoised=True): def p_sample(self, x, t: int, x_self_cond = None, clip_denoised = True):
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) batched_times = torch.full((x.shape[0],), t, device = x.device, dtype = torch.long)
noise = torch.randn_like(x) model_mean, _, model_log_variance, x_start = self.p_mean_variance(x = x, t = batched_times, x_self_cond = x_self_cond, clip_denoised = clip_denoised)
# no noise when t == 0 noise = torch.randn_like(x) if t > 0 else 0. # no noise if t == 0
nonzero_mask = (1 - (t == 0).float()).reshape(b, *((1,) * (len(x.shape) - 1))) pred_img = model_mean + (0.5 * model_log_variance).exp() * noise
return model_mean + nonzero_mask * (0.5 * model_log_variance).exp() * noise return pred_img, x_start
@torch.no_grad() @torch.no_grad()
def p_sample_loop(self, shape): 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) 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): x_start = None
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', total = self.num_timesteps):
self_cond = x_start if self.self_condition else None
img, x_start = self.p_sample(img, t, self_cond)
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(-1, total_timesteps - 1, steps=sampling_timesteps + 1) # [-1, 0, 1, 2, ..., T-1] when sampling_timesteps == total_timesteps
times = list(reversed(times.int().tolist()))
time_pairs = list(zip(times[:-1], times[1:])) # [(T-1, T-2), (T-2, T-3), ..., (1, 0), (0, -1)]
img = torch.randn(shape, device = device)
x_start = None
for time, time_next in tqdm(time_pairs, desc = 'sampling loop time step'):
time_cond = torch.full((batch,), time, device=device, dtype=torch.long)
self_cond = x_start if self.self_condition else None
pred_noise, x_start, *_ = self.model_predictions(img, time_cond, self_cond, clip_x_start = clip_denoised)
if time_next < 0:
img = x_start
continue
alpha = self.alphas_cumprod[time]
alpha_next = self.alphas_cumprod[time_next]
sigma = eta * ((1 - alpha / alpha_next) * (1 - alpha_next) / (1 - alpha)).sqrt()
c = (1 - alpha_next - sigma ** 2).sqrt()
noise = torch.randn_like(img)
img = x_start * alpha_next.sqrt() + \
c * pred_noise + \
sigma * noise
img = unnormalize_to_zero_to_one(img) img = unnormalize_to_zero_to_one(img)
return img return img
@torch.no_grad() @torch.no_grad()
def sample(self, batch_size = 16): def sample(self, batch_size = 16):
image_size = self.image_size image_size, channels = self.image_size, self.channels
channels = self.channels sample_fn = self.p_sample_loop if not self.is_ddim_sampling else self.ddim_sample
return self.p_sample_loop((batch_size, channels, image_size, image_size)) return sample_fn((batch_size, channels, image_size, image_size))
@torch.no_grad() @torch.no_grad()
def interpolate(self, x1, x2, t = None, lam = 0.5): def interpolate(self, x1, x2, t = None, lam = 0.5):
@@ -506,11 +619,11 @@ class GaussianDiffusion(nn.Module):
assert x1.shape == x2.shape assert x1.shape == x2.shape
t_batched = torch.stack([torch.tensor(t, device=device)] * b) t_batched = torch.stack([torch.tensor(t, device = device)] * b)
xt1, xt2 = map(lambda x: self.q_sample(x, t=t_batched), (x1, x2)) xt1, xt2 = map(lambda x: self.q_sample(x, t = t_batched), (x1, x2))
img = (1 - lam) * xt1 + lam * xt2 img = (1 - lam) * xt1 + lam * xt2
for i in tqdm(reversed(range(0, t)), desc='interpolation sample time step', total=t): for i in tqdm(reversed(range(0, t)), desc = 'interpolation sample time step', total = t):
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))
return img return img
@@ -536,8 +649,23 @@ class GaussianDiffusion(nn.Module):
b, c, h, w = x_start.shape b, c, h, w = x_start.shape
noise = default(noise, lambda: torch.randn_like(x_start)) noise = default(noise, lambda: torch.randn_like(x_start))
x = self.q_sample(x_start=x_start, t=t, noise=noise) # noise sample
model_out = self.denoise_fn(x, t)
x = self.q_sample(x_start = x_start, t = t, noise = noise)
# if doing self-conditioning, 50% of the time, predict x_start from current set of times
# and condition with unet with that
# this technique will slow down training by 25%, but seems to lower FID significantly
x_self_cond = None
if self.self_condition and random() < 0.5:
with torch.no_grad():
x_self_cond = self.model_predictions(x, t).pred_x_start
x_self_cond.detach_()
# predict and take gradient step
model_out = self.model(x, t, x_self_cond)
if self.objective == 'pred_noise': if self.objective == 'pred_noise':
target = noise target = noise
@@ -562,18 +690,28 @@ class GaussianDiffusion(nn.Module):
# dataset classes # dataset classes
class Dataset(data.Dataset): class Dataset(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', 'tiff'],
augment_horizontal_flip = False,
convert_image_to = None
):
super().__init__() super().__init__()
self.folder = folder self.folder = folder
self.image_size = image_size self.image_size = image_size
self.paths = [p for ext in exts for p in Path(f'{folder}').glob(f'**/*.{ext}')] self.paths = [p for ext in exts for p in Path(f'{folder}').glob(f'**/*.{ext}')]
self.transform = transforms.Compose([ maybe_convert_fn = partial(convert_image_to, convert_image_to) if exists(convert_image_to) else nn.Identity()
transforms.Resize(image_size),
transforms.RandomHorizontalFlip() if augment_horizontal_flip else nn.Identity(), self.transform = T.Compose([
transforms.CenterCrop(image_size), T.Lambda(maybe_convert_fn),
transforms.ToTensor() T.Resize(image_size),
T.RandomHorizontalFlip() if augment_horizontal_flip else nn.Identity(),
T.CenterCrop(image_size),
T.ToTensor()
]) ])
def __len__(self): def __len__(self):
@@ -592,92 +730,144 @@ class Trainer(object):
diffusion_model, diffusion_model,
folder, folder,
*, *,
ema_decay = 0.995, train_batch_size = 16,
train_batch_size = 32, gradient_accumulate_every = 1,
augment_horizontal_flip = True,
train_lr = 1e-4, train_lr = 1e-4,
train_num_steps = 100000, train_num_steps = 100000,
gradient_accumulate_every = 2,
amp = False,
step_start_ema = 2000,
ema_update_every = 10, ema_update_every = 10,
ema_decay = 0.995,
adam_betas = (0.9, 0.99),
save_and_sample_every = 1000, save_and_sample_every = 1000,
num_samples = 25,
results_folder = './results', results_folder = './results',
augment_horizontal_flip = True amp = False,
fp16 = False,
split_batches = True,
convert_image_to = None
): ):
super().__init__() super().__init__()
self.image_size = diffusion_model.image_size
self.accelerator = Accelerator(
split_batches = split_batches,
mixed_precision = 'fp16' if fp16 else 'no'
)
self.accelerator.native_amp = amp
self.model = diffusion_model self.model = diffusion_model
self.ema = EMA(diffusion_model, beta = ema_decay, update_every = ema_update_every)
self.step_start_ema = step_start_ema 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.save_and_sample_every = save_and_sample_every
self.batch_size = train_batch_size self.batch_size = train_batch_size
self.image_size = diffusion_model.image_size
self.gradient_accumulate_every = gradient_accumulate_every self.gradient_accumulate_every = gradient_accumulate_every
self.train_num_steps = train_num_steps
self.ds = Dataset(folder, self.image_size, augment_horizontal_flip = augment_horizontal_flip) self.train_num_steps = train_num_steps
self.dl = cycle(data.DataLoader(self.ds, batch_size = train_batch_size, shuffle = True, pin_memory = True, num_workers = cpu_count())) self.image_size = diffusion_model.image_size
self.opt = Adam(diffusion_model.parameters(), lr = train_lr)
# 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())
dl = self.accelerator.prepare(dl)
self.dl = cycle(dl)
# optimizer
self.opt = Adam(diffusion_model.parameters(), lr = train_lr, betas = adam_betas)
# 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.step = 0
self.amp = amp # prepare model, dataloader, optimizer with accelerator
self.scaler = GradScaler(enabled = amp)
self.results_folder = Path(results_folder) self.model, self.opt = self.accelerator.prepare(self.model, self.opt)
self.results_folder.mkdir(exist_ok = True)
def save(self, milestone): def save(self, milestone):
if not self.accelerator.is_local_main_process:
return
data = { data = {
'step': self.step, 'step': self.step,
'model': self.model.state_dict(), 'model': self.accelerator.get_state_dict(self.model),
'opt': self.opt.state_dict(),
'ema': self.ema.state_dict(), 'ema': self.ema.state_dict(),
'scaler': self.scaler.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')) torch.save(data, str(self.results_folder / f'model-{milestone}.pt'))
def load(self, milestone): def load(self, milestone):
data = torch.load(str(self.results_folder / f'model-{milestone}.pt')) 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.step = data['step']
self.model.load_state_dict(data['model']) self.opt.load_state_dict(data['opt'])
self.ema.load_state_dict(data['ema']) self.ema.load_state_dict(data['ema'])
self.scaler.load_state_dict(data['scaler'])
if exists(self.accelerator.scaler) and exists(data['scaler']):
self.accelerator.scaler.load_state_dict(data['scaler'])
def train(self): 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: 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): total_loss = 0.
for _ in range(self.gradient_accumulate_every):
data = next(self.dl).to(device)
with self.accelerator.autocast():
loss = self.model(data) loss = self.model(data)
self.scaler.scale(loss / self.gradient_accumulate_every).backward() loss = loss / self.gradient_accumulate_every
total_loss += loss.item()
pbar.set_description(f'loss: {loss.item():.4f}') self.accelerator.backward(loss)
self.scaler.step(self.opt) pbar.set_description(f'loss: {total_loss:.4f}')
self.scaler.update()
accelerator.wait_for_everyone()
self.opt.step()
self.opt.zero_grad() self.opt.zero_grad()
self.ema.update() accelerator.wait_for_everyone()
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(36, 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 = 6)
self.save(milestone)
self.step += 1 self.step += 1
if accelerator.is_main_process:
self.ema.to(device)
self.ema.update()
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)
pbar.update(1) pbar.update(1)
print('training complete') accelerator.print('training complete')
@@ -1,4 +1,5 @@
from math import sqrt from math import sqrt
from random import random
import torch import torch
from torch import nn, einsum from torch import nn, einsum
import torch.nn.functional as F import torch.nn.functional as F
@@ -52,6 +53,7 @@ class ElucidatedDiffusion(nn.Module):
): ):
super().__init__() super().__init__()
assert net.learned_sinusoidal_cond assert net.learned_sinusoidal_cond
self.self_condition = net.self_condition
self.net = net self.net = net
@@ -99,7 +101,7 @@ class ElucidatedDiffusion(nn.Module):
# preconditioned network output # preconditioned network output
# equation (7) in the paper # equation (7) in the paper
def preconditioned_network_forward(self, noised_images, sigma, clamp = False): def preconditioned_network_forward(self, noised_images, sigma, self_cond = None, clamp = False):
batch, device = noised_images.shape[0], noised_images.device batch, device = noised_images.shape[0], noised_images.device
if isinstance(sigma, float): if isinstance(sigma, float):
@@ -109,7 +111,8 @@ class ElucidatedDiffusion(nn.Module):
net_out = self.net( net_out = self.net(
self.c_in(padded_sigma) * noised_images, self.c_in(padded_sigma) * noised_images,
self.c_noise(sigma) self.c_noise(sigma),
self_cond
) )
out = self.c_skip(padded_sigma) * noised_images + self.c_out(padded_sigma) * net_out out = self.c_skip(padded_sigma) * noised_images + self.c_out(padded_sigma) * net_out
@@ -160,6 +163,10 @@ class ElucidatedDiffusion(nn.Module):
images = init_sigma * torch.randn(shape, device = self.device) images = init_sigma * torch.randn(shape, device = self.device)
# for self conditioning
x_start = None
# gradually denoise # gradually denoise
for sigma, sigma_next, gamma in tqdm(sigmas_and_gammas, desc = 'sampling time step'): for sigma, sigma_next, gamma in tqdm(sigmas_and_gammas, desc = 'sampling time step'):
@@ -170,7 +177,9 @@ class ElucidatedDiffusion(nn.Module):
sigma_hat = sigma + gamma * sigma sigma_hat = sigma + gamma * sigma
images_hat = images + sqrt(sigma_hat ** 2 - sigma ** 2) * eps images_hat = images + sqrt(sigma_hat ** 2 - sigma ** 2) * eps
model_output = self.preconditioned_network_forward(images_hat, sigma_hat, clamp = clamp) self_cond = x_start if self.self_condition else None
model_output = self.preconditioned_network_forward(images_hat, sigma_hat, self_cond, clamp = clamp)
denoised_over_sigma = (images_hat - model_output) / sigma_hat denoised_over_sigma = (images_hat - model_output) / sigma_hat
images_next = images_hat + (sigma_next - sigma_hat) * denoised_over_sigma images_next = images_hat + (sigma_next - sigma_hat) * denoised_over_sigma
@@ -178,11 +187,14 @@ class ElucidatedDiffusion(nn.Module):
# second order correction, if not the last timestep # second order correction, if not the last timestep
if sigma_next != 0: if sigma_next != 0:
model_output_next = self.preconditioned_network_forward(images_next, sigma_next, clamp = clamp) self_cond = model_output if self.self_condition else None
model_output_next = self.preconditioned_network_forward(images_next, sigma_next, self_cond, clamp = clamp)
denoised_prime_over_sigma = (images_next - model_output_next) / sigma_next 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_next = images_hat + 0.5 * (sigma_next - sigma_hat) * (denoised_over_sigma + denoised_prime_over_sigma)
images = images_next images = images_next
x_start = model_output
images = images.clamp(-1., 1.) images = images.clamp(-1., 1.)
return unnormalize_to_zero_to_one(images) return unnormalize_to_zero_to_one(images)
@@ -210,7 +222,15 @@ class ElucidatedDiffusion(nn.Module):
noised_images = images + padded_sigmas * noise # alphas are 1. in the paper noised_images = images + padded_sigmas * noise # alphas are 1. in the paper
denoised = self.preconditioned_network_forward(noised_images, sigmas) self_cond = None
if self.self_condition and random() < 0.5:
# from hinton's group's bit diffusion paper
with torch.no_grad():
self_cond = self.preconditioned_network_forward(noised_images, sigmas)
self_cond.detach_()
denoised = self.preconditioned_network_forward(noised_images, sigmas, self_cond)
losses = F.mse_loss(denoised, images, reduction = 'none') losses = F.mse_loss(denoised, images, reduction = 'none')
losses = reduce(losses, 'b ... -> b', 'mean') losses = reduce(losses, 'b ... -> b', 'mean')
@@ -1,4 +1,5 @@
import torch import torch
from collections import namedtuple
from math import pi, sqrt, log as ln from math import pi, sqrt, log as ln
from inspect import isfunction from inspect import isfunction
from torch import nn, einsum from torch import nn, einsum
@@ -10,6 +11,8 @@ from denoising_diffusion_pytorch.denoising_diffusion_pytorch import GaussianDiff
NAT = 1. / ln(2) NAT = 1. / ln(2)
ModelPrediction = namedtuple('ModelPrediction', ['pred_noise', 'pred_x_start', 'pred_variance'])
# helper functions # helper functions
def exists(x): def exists(x):
@@ -22,7 +25,7 @@ def default(val, d):
# tensor helpers # tensor helpers
def log(t, eps = 1e-12): def log(t, eps = 1e-15):
return torch.log(t.clamp(min = eps)) return torch.log(t.clamp(min = eps))
def meanflat(x): def meanflat(x):
@@ -67,17 +70,33 @@ def discretized_gaussian_log_likelihood(x, *, means, log_scales, thres = 0.999):
class LearnedGaussianDiffusion(GaussianDiffusion): class LearnedGaussianDiffusion(GaussianDiffusion):
def __init__( def __init__(
self, self,
denoise_fn, model,
vb_loss_weight = 0.001, # lambda was 0.001 in the paper vb_loss_weight = 0.001, # lambda was 0.001 in the paper
*args, *args,
**kwargs **kwargs
): ):
super().__init__(denoise_fn, *args, **kwargs) super().__init__(model, *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`' 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`'
assert not model.self_condition, 'not supported yet'
self.vb_loss_weight = vb_loss_weight 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): 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) pred_noise, var_interp_frac_unnormalized = model_output.chunk(2, dim = 1)
min_log = extract(self.posterior_log_variance_clipped, t, x.shape) min_log = extract(self.posterior_log_variance_clipped, t, x.shape)
@@ -102,7 +121,7 @@ class LearnedGaussianDiffusion(GaussianDiffusion):
# model output # model output
model_output = self.denoise_fn(x_t, t) model_output = self.model(x_t, t)
# calculating kl loss for learned variance (interpolation) # calculating kl loss for learned variance (interpolation)
@@ -22,22 +22,24 @@ def default(val, d):
class WeightedObjectiveGaussianDiffusion(GaussianDiffusion): class WeightedObjectiveGaussianDiffusion(GaussianDiffusion):
def __init__( def __init__(
self, self,
denoise_fn, model,
*args, *args,
pred_noise_loss_weight = 0.1, pred_noise_loss_weight = 0.1,
pred_x_start_loss_weight = 0.1, pred_x_start_loss_weight = 0.1,
**kwargs **kwargs
): ):
super().__init__(denoise_fn, *args, **kwargs) super().__init__(model, *args, **kwargs)
channels = denoise_fn.channels channels = model.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' 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 model.self_condition, 'not supported yet'
assert not self.is_ddim_sampling, 'ddim sampling cannot be used'
self.split_dims = (channels, channels, 2) self.split_dims = (channels, channels, 2)
self.pred_noise_loss_weight = pred_noise_loss_weight self.pred_noise_loss_weight = pred_noise_loss_weight
self.pred_x_start_loss_weight = pred_x_start_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): 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) pred_noise, pred_x_start, weights = model_output.split(self.split_dims, dim = 1)
normalized_weights = weights.softmax(dim = 1) normalized_weights = weights.softmax(dim = 1)
@@ -58,7 +60,7 @@ class WeightedObjectiveGaussianDiffusion(GaussianDiffusion):
noise = default(noise, lambda: torch.randn_like(x_start)) noise = default(noise, lambda: torch.randn_like(x_start))
x_t = self.q_sample(x_start = x_start, t = t, noise = noise) 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) pred_noise, pred_x_start, weights = model_output.split(self.split_dims, dim = 1)
# get loss for predicted noise and x_start # get loss for predicted noise and x_start

Before

Width:  |  Height:  |  Size: 40 KiB

After

Width:  |  Height:  |  Size: 40 KiB

View File

Before

Width:  |  Height:  |  Size: 842 KiB

After

Width:  |  Height:  |  Size: 842 KiB

+3 -2
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.23.4', version = '0.27.8',
license='MIT', license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch', description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang', author = 'Phil Wang',
@@ -15,6 +15,7 @@ setup(
'generative models' 'generative models'
], ],
install_requires=[ install_requires=[
'accelerate',
'einops', 'einops',
'ema-pytorch', 'ema-pytorch',
'pillow', 'pillow',
@@ -29,4 +30,4 @@ setup(
'License :: OSI Approved :: MIT License', 'License :: OSI Approved :: MIT License',
'Programming Language :: Python :: 3.6', 'Programming Language :: Python :: 3.6',
], ],
) )