Compare commits

...
28 Commits
Author SHA1 Message Date
Phil Wang 6e8a0f2082 fix auto-conversion of images to mode in dataset 2022-09-20 19:29:35 -07:00
Phil Wang 8c36559295 0.27.10 2022-09-16 17:15:02 -07:00
Phil Wang f74f536339 Merge pull request #90 from kashif/patch-1
fix torch.cumprod
2022-09-16 17:14:48 -07:00
Kashif Rasul d85b8bbe2e fix torch.cumprod 2022-09-16 17:19:00 +02:00
Phil Wang e0a1bed31a 0.27.9 2022-09-05 02:43:11 -07:00
Phil Wang 82b67fc00a Merge pull request #85 from RyannDaGreat/main
Trainer.load can use GPU's other than cuda:0
2022-09-05 02:42:49 -07:00
Ryan Burgert 7c0cd05c27 Trainer.load can use GPU's other than cuda:0 2022-09-04 22:11:35 -04:00
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
9 changed files with 174 additions and 67 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
@@ -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>
<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)
@@ -170,3 +172,24 @@ $ accelerate launch train.py
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__()
assert model.learned_sinusoidal_cond
assert not model.self_condition, 'not supported yet'
self.model = model
@@ -1,23 +1,23 @@
import math
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
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.data import Dataset, DataLoader
from multiprocessing import cpu_count
from pathlib import Path
from torch.optim import Adam
from torchvision import transforms as T, utils
from PIL import Image
from einops import rearrange, reduce
from einops.layers.torch import Rearrange
from PIL import Image
from tqdm.auto import tqdm
from ema_pytorch import EMA
@@ -35,7 +35,10 @@ def exists(x):
def default(val, d):
if exists(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):
while True:
@@ -53,11 +56,14 @@ def num_to_groups(num, divisor):
arr.append(remainder)
return arr
def convert_image_to(img_type, image):
def convert_image_to_fn(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):
@@ -85,17 +91,31 @@ def Upsample(dim, dim_out = None):
def Downsample(dim, dim_out = None):
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):
def __init__(self, dim, eps = 1e-5):
def __init__(self, dim):
super().__init__()
self.eps = eps
self.g = nn.Parameter(torch.ones(1, dim, 1, 1))
self.b = nn.Parameter(torch.zeros(1, dim, 1, 1))
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)
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):
def __init__(self, dim, fn):
@@ -145,7 +165,7 @@ class LearnedSinusoidalPosEmb(nn.Module):
class Block(nn.Module):
def __init__(self, dim, dim_out, groups = 8):
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.act = nn.SiLU()
@@ -208,6 +228,8 @@ class LinearAttention(nn.Module):
k = k.softmax(dim = -1)
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)
out = torch.einsum('b h d e, b h d n -> b h e n', context, q)
@@ -215,9 +237,9 @@ class LinearAttention(nn.Module):
return self.to_out(out)
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__()
self.scale = dim_head ** -0.5
self.scale = scale
self.heads = heads
hidden_dim = dim_head * heads
self.to_qkv = nn.Conv2d(dim, hidden_dim * 3, 1, bias = False)
@@ -227,12 +249,11 @@ class Attention(nn.Module):
b, c, h, w = x.shape
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 = q * self.scale
sim = einsum('b h d i, b h d j -> b h i j', q, k)
sim = sim - sim.amax(dim = -1, keepdim = True).detach()
q, k = map(l2norm, (q, k))
sim = einsum('b h d i, b h d j -> b h i j', q, k) * self.scale
attn = sim.softmax(dim = -1)
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)
return self.to_out(out)
@@ -247,6 +268,7 @@ class Unet(nn.Module):
out_dim = None,
dim_mults=(1, 2, 4, 8),
channels = 3,
self_condition = False,
resnet_block_groups = 8,
learned_variance = False,
learned_sinusoidal_cond = False,
@@ -257,9 +279,11 @@ class Unet(nn.Module):
# determine dimensions
self.channels = channels
self.self_condition = self_condition
input_channels = channels * (2 if self_condition else 1)
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)]
in_out = list(zip(dims[:-1], dims[1:]))
@@ -323,7 +347,11 @@ class Unet(nn.Module):
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):
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)
r = x.clone()
@@ -391,7 +419,6 @@ class GaussianDiffusion(nn.Module):
model,
*,
image_size,
channels = 3,
timesteps = 1000,
sampling_timesteps = None,
loss_type = 'l1',
@@ -403,10 +430,14 @@ class GaussianDiffusion(nn.Module):
):
super().__init__()
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.channels = self.model.channels
self.self_condition = self.model.self_condition
self.image_size = image_size
self.objective = objective
assert objective in {'pred_noise', 'pred_x0'}, 'objective must be either pred_noise (predict noise) or pred_x0 (predict image start)'
@@ -419,7 +450,7 @@ class GaussianDiffusion(nn.Module):
raise ValueError(f'unknown beta schedule {beta_schedule}')
alphas = 1. - betas
alphas_cumprod = torch.cumprod(alphas, axis=0)
alphas_cumprod = torch.cumprod(alphas, dim=0)
alphas_cumprod_prev = F.pad(alphas_cumprod[:-1], (1, 0), value = 1.)
timesteps, = betas.shape
@@ -476,7 +507,7 @@ class GaussianDiffusion(nn.Module):
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_recip_alphas_cumprod, t, x_t.shape) * x_t - x0) / \
extract(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape)
)
@@ -489,36 +520,40 @@ 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 model_predictions(self, x, t):
model_output = self.model(x, t)
def model_predictions(self, x, t, x_self_cond = None, clip_x_start = False):
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':
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':
pred_noise = self.predict_noise_from_start(x, t, 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)
def p_mean_variance(self, x, t, clip_denoised: bool):
preds = self.model_predictions(x, t)
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:
x_start.clamp_(-1., 1.)
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()
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
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
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()
def p_sample_loop(self, shape):
@@ -526,8 +561,11 @@ class GaussianDiffusion(nn.Module):
img = torch.randn(shape, device=device)
for t in tqdm(reversed(range(0, self.num_timesteps)), desc = 'sampling loop time step'):
img = self.p_sample(img, t)
x_start = None
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
@@ -536,27 +574,30 @@ class GaussianDiffusion(nn.Module):
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 = 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:]))
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'):
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)
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)
if clip_denoised:
x_start.clamp_(-1., 1.)
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()
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() + \
c * pred_noise + \
@@ -578,11 +619,11 @@ class GaussianDiffusion(nn.Module):
assert x1.shape == x2.shape
t_batched = torch.stack([torch.tensor(t, device=device)] * b)
xt1, xt2 = map(lambda x: self.q_sample(x, t=t_batched), (x1, x2))
t_batched = torch.stack([torch.tensor(t, device = device)] * b)
xt1, xt2 = map(lambda x: self.q_sample(x, t = t_batched), (x1, x2))
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))
return img
@@ -608,8 +649,23 @@ class GaussianDiffusion(nn.Module):
b, c, h, w = x_start.shape
noise = default(noise, lambda: torch.randn_like(x_start))
# noise sample
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':
target = noise
@@ -648,7 +704,7 @@ class Dataset(Dataset):
self.image_size = image_size
self.paths = [p for ext in exts for p in Path(f'{folder}').glob(f'**/*.{ext}')]
maybe_convert_fn = partial(convert_image_to, convert_image_to) if exists(convert_image_to) else nn.Identity()
maybe_convert_fn = partial(convert_image_to_fn, convert_image_to) if exists(convert_image_to) else nn.Identity()
self.transform = T.Compose([
T.Lambda(maybe_convert_fn),
@@ -716,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)
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
@@ -736,7 +793,7 @@ class Trainer(object):
# 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):
if not self.accelerator.is_local_main_process:
@@ -753,7 +810,10 @@ class Trainer(object):
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'))
accelerator = self.accelerator
device = accelerator.device
data = torch.load(str(self.results_folder / f'model-{milestone}.pt'), map_location=device)
model = self.accelerator.unwrap_model(self.model)
model.load_state_dict(data['model'])
@@ -794,6 +854,7 @@ class Trainer(object):
accelerator.wait_for_everyone()
self.step += 1
if accelerator.is_main_process:
self.ema.to(device)
self.ema.update()
@@ -810,7 +871,6 @@ class Trainer(object):
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)
accelerator.print('training complete')
@@ -1,4 +1,5 @@
from math import sqrt
from random import random
import torch
from torch import nn, einsum
import torch.nn.functional as F
@@ -52,6 +53,7 @@ class ElucidatedDiffusion(nn.Module):
):
super().__init__()
assert net.learned_sinusoidal_cond
self.self_condition = net.self_condition
self.net = net
@@ -99,7 +101,7 @@ class ElucidatedDiffusion(nn.Module):
# preconditioned network output
# 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
if isinstance(sigma, float):
@@ -109,7 +111,8 @@ class ElucidatedDiffusion(nn.Module):
net_out = self.net(
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
@@ -160,6 +163,10 @@ class ElucidatedDiffusion(nn.Module):
images = init_sigma * torch.randn(shape, device = self.device)
# for self conditioning
x_start = None
# gradually denoise
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
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
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
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
images_next = images_hat + 0.5 * (sigma_next - sigma_hat) * (denoised_over_sigma + denoised_prime_over_sigma)
images = images_next
x_start = model_output
images = images.clamp(-1., 1.)
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
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 = reduce(losses, 'b ... -> b', 'mean')
@@ -77,6 +77,8 @@ class LearnedGaussianDiffusion(GaussianDiffusion):
):
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 not model.self_condition, 'not supported yet'
self.vb_loss_weight = vb_loss_weight
def model_predictions(self, x, t):
@@ -31,6 +31,7 @@ class WeightedObjectiveGaussianDiffusion(GaussianDiffusion):
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 model.self_condition, 'not supported yet'
assert not self.is_ddim_sampling, 'ddim sampling cannot be used'
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(
name = 'denoising-diffusion-pytorch',
packages = find_packages(),
version = '0.25.3',
version = '0.27.11',
license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang',
@@ -30,4 +30,4 @@ setup(
'License :: OSI Approved :: MIT License',
'Programming Language :: Python :: 3.6',
],
)
)