Compare commits

...
20 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
9 changed files with 160 additions and 59 deletions
+25 -2
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)
@@ -170,3 +172,24 @@ $ accelerate launch train.py
volume = {abs/2010.02502} 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}
}
```
@@ -127,6 +127,7 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
): ):
super().__init__() super().__init__()
assert model.learned_sinusoidal_cond assert model.learned_sinusoidal_cond
assert not model.self_condition, 'not supported yet'
self.model = model self.model = model
@@ -1,23 +1,23 @@
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 collections import namedtuple
from functools import partial
from torch.utils.data import Dataset, DataLoader from torch.utils.data import Dataset, DataLoader
from multiprocessing import cpu_count
from pathlib import Path
from torch.optim import Adam from torch.optim import Adam
from torchvision import transforms as T, 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
@@ -35,7 +35,10 @@ 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:
@@ -88,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):
@@ -148,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()
@@ -211,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)
@@ -218,7 +237,7 @@ 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, scale = 16): def __init__(self, dim, heads = 4, dim_head = 32, scale = 10):
super().__init__() super().__init__()
self.scale = scale self.scale = scale
self.heads = heads self.heads = heads
@@ -235,7 +254,6 @@ class Attention(nn.Module):
sim = einsum('b h d i, b h d j -> b h i j', q, k) * self.scale 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)
@@ -250,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,
@@ -260,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:]))
@@ -326,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()
@@ -394,7 +419,6 @@ class GaussianDiffusion(nn.Module):
model, model,
*, *,
image_size, image_size,
channels = 3,
timesteps = 1000, timesteps = 1000,
sampling_timesteps = None, sampling_timesteps = None,
loss_type = 'l1', loss_type = 'l1',
@@ -406,10 +430,14 @@ class GaussianDiffusion(nn.Module):
): ):
super().__init__() super().__init__()
assert not (type(self) == GaussianDiffusion and model.channels != model.out_dim) assert not (type(self) == GaussianDiffusion and model.channels != model.out_dim)
assert not model.learned_sinusoidal_cond
self.channels = channels
self.image_size = image_size
self.model = model self.model = model
self.channels = self.model.channels
self.self_condition = self.model.self_condition
self.image_size = image_size
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)' assert objective in {'pred_noise', 'pred_x0'}, 'objective must be either pred_noise (predict noise) or pred_x0 (predict image start)'
@@ -479,7 +507,7 @@ class GaussianDiffusion(nn.Module):
def predict_noise_from_start(self, x_t, t, x0): def predict_noise_from_start(self, x_t, t, x0):
return ( return (
(x0 - extract(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t) / \ (extract(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t - x0) / \
extract(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape) extract(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape)
) )
@@ -492,36 +520,40 @@ 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 model_predictions(self, x, t): def model_predictions(self, x, t, x_self_cond = None, clip_x_start = False):
model_output = self.model(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':
pred_noise = model_output pred_noise = model_output
x_start = self.predict_start_from_noise(x, t, 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':
pred_noise = self.predict_noise_from_start(x, t, model_output)
x_start = model_output x_start = model_output
x_start = maybe_clip(x_start)
pred_noise = self.predict_noise_from_start(x, t, x_start)
return ModelPrediction(pred_noise, x_start) return ModelPrediction(pred_noise, x_start)
def p_mean_variance(self, x, t, clip_denoised: bool): def p_mean_variance(self, x, t, x_self_cond = None, clip_denoised = True):
preds = self.model_predictions(x, t) preds = self.model_predictions(x, t, x_self_cond)
x_start = preds.pred_x_start 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: int, 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
batched_times = torch.full((x.shape[0],), t, device = x.device, dtype = torch.long) 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) 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)
noise = torch.randn_like(x) if t > 0 else 0. # no noise if t == 0 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 pred_img = model_mean + (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):
@@ -529,8 +561,11 @@ class GaussianDiffusion(nn.Module):
img = torch.randn(shape, device=device) img = torch.randn(shape, device=device)
for t in tqdm(reversed(range(0, self.num_timesteps)), desc = 'sampling loop time step'): x_start = None
img = self.p_sample(img, t)
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) img = unnormalize_to_zero_to_one(img)
return img return img
@@ -539,27 +574,30 @@ class GaussianDiffusion(nn.Module):
def ddim_sample(self, shape, clip_denoised = True): 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 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 = 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())) times = list(reversed(times.int().tolist()))
time_pairs = list(zip(times[:-1], times[1:])) 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) img = torch.randn(shape, device = device)
x_start = None
for time, time_next in tqdm(time_pairs, desc = 'sampling loop time step'): for time, time_next in tqdm(time_pairs, desc = 'sampling loop time step'):
alpha = self.alphas_cumprod_prev[time] time_cond = torch.full((batch,), time, device=device, dtype=torch.long)
alpha_next = self.alphas_cumprod_prev[time_next] 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)
time_cond = torch.full((batch,), time, device = device, dtype = torch.long) if time_next < 0:
img = x_start
continue
pred_noise, x_start, *_ = self.model_predictions(img, time_cond) alpha = self.alphas_cumprod[time]
alpha_next = self.alphas_cumprod[time_next]
if clip_denoised:
x_start.clamp_(-1., 1.)
sigma = eta * ((1 - alpha / alpha_next) * (1 - alpha_next) / (1 - alpha)).sqrt() sigma = eta * ((1 - alpha / alpha_next) * (1 - alpha_next) / (1 - alpha)).sqrt()
c = ((1 - alpha_next) - sigma ** 2).sqrt() c = (1 - alpha_next - sigma ** 2).sqrt()
noise = torch.randn_like(img) if time_next > 0 else 0. noise = torch.randn_like(img)
img = x_start * alpha_next.sqrt() + \ img = x_start * alpha_next.sqrt() + \
c * pred_noise + \ c * pred_noise + \
@@ -581,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
@@ -611,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))
# noise sample
x = self.q_sample(x_start = x_start, t = t, noise = noise) x = self.q_sample(x_start = x_start, t = t, noise = noise)
model_out = self.model(x, t)
# 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
@@ -719,6 +772,7 @@ class Trainer(object):
self.ds = Dataset(folder, self.image_size, augment_horizontal_flip = augment_horizontal_flip, convert_image_to = convert_image_to) 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 = 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) self.dl = cycle(dl)
# optimizer # optimizer
@@ -739,7 +793,7 @@ class Trainer(object):
# prepare model, dataloader, optimizer with accelerator # prepare model, dataloader, optimizer with accelerator
self.model, self.dl, self.opt = self.accelerator.prepare(self.model, self.dl, self.opt) self.model, self.opt = self.accelerator.prepare(self.model, self.opt)
def save(self, milestone): def save(self, milestone):
if not self.accelerator.is_local_main_process: if not self.accelerator.is_local_main_process:
@@ -797,6 +851,7 @@ class Trainer(object):
accelerator.wait_for_everyone() accelerator.wait_for_everyone()
self.step += 1
if accelerator.is_main_process: if accelerator.is_main_process:
self.ema.to(device) self.ema.to(device)
self.ema.update() self.ema.update()
@@ -813,7 +868,6 @@ class Trainer(object):
utils.save_image(all_images, str(self.results_folder / f'sample-{milestone}.png'), nrow = int(math.sqrt(self.num_samples))) utils.save_image(all_images, str(self.results_folder / f'sample-{milestone}.png'), nrow = int(math.sqrt(self.num_samples)))
self.save(milestone) self.save(milestone)
self.step += 1
pbar.update(1) pbar.update(1)
accelerator.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')
@@ -77,6 +77,8 @@ class LearnedGaussianDiffusion(GaussianDiffusion):
): ):
super().__init__(model, *args, **kwargs) 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`' 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): def model_predictions(self, x, t):
@@ -31,6 +31,7 @@ class WeightedObjectiveGaussianDiffusion(GaussianDiffusion):
super().__init__(model, *args, **kwargs) super().__init__(model, *args, **kwargs)
channels = model.channels 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 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' assert not self.is_ddim_sampling, 'ddim sampling cannot be used'
self.split_dims = (channels, channels, 2) self.split_dims = (channels, channels, 2)

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

+2 -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.26.0', 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',
@@ -30,4 +30,4 @@ setup(
'License :: OSI Approved :: MIT License', 'License :: OSI Approved :: MIT License',
'Programming Language :: Python :: 3.6', 'Programming Language :: Python :: 3.6',
], ],
) )