Compare commits

...
8 Commits
8 changed files with 161 additions and 63 deletions
+16 -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,7 @@ 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> <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 +60,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
@@ -159,3 +160,13 @@ $ accelerate launch train.py
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}
}
```
@@ -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,9 @@ 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
self.denoise_fn = denoise_fn self.model = model
# image dimensions # image dimensions
@@ -170,7 +170,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 +195,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 +266,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')
@@ -4,6 +4,7 @@ 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 inspect import isfunction
from collections import namedtuple
from functools import partial from functools import partial
from torch.utils.data import Dataset, DataLoader from torch.utils.data import Dataset, DataLoader
@@ -22,6 +23,10 @@ from ema_pytorch import EMA
from accelerate import Accelerator from accelerate import Accelerator
# constants
ModelPrediction = namedtuple('ModelPrediction', ['pred_noise', 'pred_x_start'])
# helpers functions # helpers functions
def exists(x): def exists(x):
@@ -53,6 +58,9 @@ def convert_image_to(img_type, image):
return image.convert(img_type) return image.convert(img_type)
return image return image
def l2norm(t):
return F.normalize(t, dim = -1)
# normalization functions # normalization functions
def normalize_to_neg_one_to_one(img): def normalize_to_neg_one_to_one(img):
@@ -203,6 +211,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)
@@ -210,9 +220,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 = 16):
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)
@@ -222,10 +232,10 @@ 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)
@@ -383,25 +393,29 @@ 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, 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)
self.channels = channels self.channels = channels
self.image_size = image_size self.image_size = image_size
self.denoise_fn = denoise_fn self.model = model
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':
@@ -417,6 +431,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))
@@ -457,6 +479,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 +
@@ -466,15 +494,22 @@ 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):
model_output = self.denoise_fn(x, t) model_output = self.model(x, t)
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, model_output)
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
else:
raise ValueError(f'unknown objective {self.objective}') return ModelPrediction(pred_noise, x_start)
def p_mean_variance(self, x, t, clip_denoised: bool):
preds = self.model_predictions(x, t)
x_start = preds.pred_x_start
if clip_denoised: if clip_denoised:
x_start.clamp_(-1., 1.) x_start.clamp_(-1., 1.)
@@ -483,32 +518,63 @@ class GaussianDiffusion(nn.Module):
return model_mean, posterior_variance, posterior_log_variance return model_mean, posterior_variance, posterior_log_variance
@torch.no_grad() @torch.no_grad()
def p_sample(self, x, t, clip_denoised=True): def p_sample(self, x, t: int, clip_denoised = True):
b, *_, device = *x.shape, x.device 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 = self.p_mean_variance(x = x, t = batched_times, 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))) return model_mean + (0.5 * model_log_variance).exp() * noise
return model_mean + nonzero_mask * (0.5 * model_log_variance).exp() * noise
@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): for t in tqdm(reversed(range(0, self.num_timesteps)), desc = 'sampling loop time step'):
img = self.p_sample(img, torch.full((b,), i, device=device, dtype=torch.long)) img = self.p_sample(img, t)
img = unnormalize_to_zero_to_one(img)
return img
@torch.no_grad()
def ddim_sample(self, shape, clip_denoised = True):
batch, device, total_timesteps, sampling_timesteps, eta, objective = shape[0], self.betas.device, self.num_timesteps, self.sampling_timesteps, self.ddim_sampling_eta, self.objective
times = torch.linspace(0., total_timesteps, steps = sampling_timesteps + 2)[:-1]
times = list(reversed(times.int().tolist()))
time_pairs = list(zip(times[:-1], times[1:]))
img = torch.randn(shape, device = device)
for time, time_next in tqdm(time_pairs, desc = 'sampling loop time step'):
alpha = self.alphas_cumprod_prev[time]
alpha_next = self.alphas_cumprod_prev[time_next]
time_cond = torch.full((batch,), time, device = device, dtype = torch.long)
pred_noise, x_start, *_ = self.model_predictions(img, time_cond)
if clip_denoised:
x_start.clamp_(-1., 1.)
sigma = eta * ((1 - alpha / alpha_next) * (1 - alpha_next) / (1 - alpha)).sqrt()
c = ((1 - alpha_next) - sigma ** 2).sqrt()
noise = torch.randn_like(img) if time_next > 0 else 0.
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):
@@ -547,8 +613,8 @@ 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) x = self.q_sample(x_start = x_start, t = t, noise = noise)
model_out = self.denoise_fn(x, t) model_out = self.model(x, t)
if self.objective == 'pred_noise': if self.objective == 'pred_noise':
target = noise target = noise
@@ -620,6 +686,7 @@ class Trainer(object):
train_num_steps = 100000, train_num_steps = 100000,
ema_update_every = 10, ema_update_every = 10,
ema_decay = 0.995, ema_decay = 0.995,
adam_betas = (0.9, 0.99),
save_and_sample_every = 1000, save_and_sample_every = 1000,
num_samples = 25, num_samples = 25,
results_folder = './results', results_folder = './results',
@@ -654,11 +721,12 @@ 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
self.opt = Adam(diffusion_model.parameters(), lr = train_lr) self.opt = Adam(diffusion_model.parameters(), lr = train_lr, betas = adam_betas)
# for logging results in a folder periodically # for logging results in a folder periodically
@@ -674,18 +742,16 @@ 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_main_process: if not self.accelerator.is_local_main_process:
return return
opt = self.accelerator.unwrap_model(self.opt)
data = { data = {
'step': self.step, 'step': self.step,
'model': self.accelerator.get_state_dict(self.model), 'model': self.accelerator.get_state_dict(self.model),
'opt': opt.state_dict(), 'opt': self.opt.state_dict(),
'ema': self.ema.state_dict(), 'ema': self.ema.state_dict(),
'scaler': self.accelerator.scaler.state_dict() if exists(self.accelerator.scaler) else None 'scaler': self.accelerator.scaler.state_dict() if exists(self.accelerator.scaler) else None
} }
@@ -696,12 +762,10 @@ class Trainer(object):
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 = self.accelerator.unwrap_model(self.model)
opt = self.accelerator.unwrap_model(self.opt)
model.load_state_dict(data['model']) model.load_state_dict(data['model'])
opt.load_state_dict(data['opt'])
self.step = data['step'] self.step = data['step']
self.opt.load_state_dict(data['opt'])
self.ema.load_state_dict(data['ema']) self.ema.load_state_dict(data['ema'])
if exists(self.accelerator.scaler) and exists(data['scaler']): if exists(self.accelerator.scaler) and exists(data['scaler']):
@@ -715,14 +779,19 @@ class Trainer(object):
while self.step < self.train_num_steps: while self.step < self.train_num_steps:
total_loss = 0.
for _ in range(self.gradient_accumulate_every): for _ in range(self.gradient_accumulate_every):
data = next(self.dl).to(device) data = next(self.dl).to(device)
with self.accelerator.autocast(): with self.accelerator.autocast():
loss = self.model(data) loss = self.model(data)
self.accelerator.backward(loss / self.gradient_accumulate_every) loss = loss / self.gradient_accumulate_every
total_loss += loss.item()
pbar.set_description(f'loss: {loss.item():.4f}') self.accelerator.backward(loss)
pbar.set_description(f'loss: {total_loss:.4f}')
accelerator.wait_for_everyone() accelerator.wait_for_everyone()
@@ -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):
@@ -67,17 +70,31 @@ 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`'
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 +119,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,23 @@ 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 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 +59,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

+1 -1
View File
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
setup( setup(
name = 'denoising-diffusion-pytorch', name = 'denoising-diffusion-pytorch',
packages = find_packages(), packages = find_packages(),
version = '0.24.4', version = '0.26.4',
license='MIT', license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch', description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang', author = 'Phil Wang',