mirror of
https://github.com/wassname/denoising-diffusion-pytorch.git
synced 2026-09-10 12:01:08 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
12079cadee | ||
|
|
0ffff59ca0 | ||
|
|
aadaa7d288 | ||
|
|
23fd887a5f |
@@ -195,3 +195,13 @@ $ accelerate launch train.py
|
|||||||
volume = {abs/1903.10520}
|
volume = {abs/1903.10520}
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
```bibtex
|
||||||
|
@article{Salimans2022ProgressiveDF,
|
||||||
|
title = {Progressive Distillation for Fast Sampling of Diffusion Models},
|
||||||
|
author = {Tim Salimans and Jonathan Ho},
|
||||||
|
journal = {ArXiv},
|
||||||
|
year = {2022},
|
||||||
|
volume = {abs/2202.00512}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|||||||
@@ -4,3 +4,4 @@ from denoising_diffusion_pytorch.learned_gaussian_diffusion import LearnedGaussi
|
|||||||
from denoising_diffusion_pytorch.continuous_time_gaussian_diffusion import ContinuousTimeGaussianDiffusion
|
from denoising_diffusion_pytorch.continuous_time_gaussian_diffusion import ContinuousTimeGaussianDiffusion
|
||||||
from denoising_diffusion_pytorch.weighted_objective_gaussian_diffusion import WeightedObjectiveGaussianDiffusion
|
from denoising_diffusion_pytorch.weighted_objective_gaussian_diffusion import WeightedObjectiveGaussianDiffusion
|
||||||
from denoising_diffusion_pytorch.elucidated_diffusion import ElucidatedDiffusion
|
from denoising_diffusion_pytorch.elucidated_diffusion import ElucidatedDiffusion
|
||||||
|
from denoising_diffusion_pytorch.v_param_continuous_time_gaussian_diffusion import VParamContinuousTimeGaussianDiffusion
|
||||||
|
|||||||
@@ -126,7 +126,7 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
|
|||||||
p2_loss_weight_k = 1
|
p2_loss_weight_k = 1
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
assert model.learned_sinusoidal_cond
|
assert model.random_or_learned_sinusoidal_cond
|
||||||
assert not model.self_condition, 'not supported yet'
|
assert not model.self_condition, 'not supported yet'
|
||||||
|
|
||||||
self.model = model
|
self.model = model
|
||||||
|
|||||||
@@ -140,15 +140,15 @@ class SinusoidalPosEmb(nn.Module):
|
|||||||
emb = torch.cat((emb.sin(), emb.cos()), dim=-1)
|
emb = torch.cat((emb.sin(), emb.cos()), dim=-1)
|
||||||
return emb
|
return emb
|
||||||
|
|
||||||
class LearnedSinusoidalPosEmb(nn.Module):
|
class RandomOrLearnedSinusoidalPosEmb(nn.Module):
|
||||||
""" following @crowsonkb 's lead with learned sinusoidal pos emb """
|
""" following @crowsonkb 's lead with random (learned optional) sinusoidal pos emb """
|
||||||
""" https://github.com/crowsonkb/v-diffusion-jax/blob/master/diffusion/models/danbooru_128.py#L8 """
|
""" https://github.com/crowsonkb/v-diffusion-jax/blob/master/diffusion/models/danbooru_128.py#L8 """
|
||||||
|
|
||||||
def __init__(self, dim):
|
def __init__(self, dim, is_random = False):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
assert (dim % 2) == 0
|
assert (dim % 2) == 0
|
||||||
half_dim = dim // 2
|
half_dim = dim // 2
|
||||||
self.weights = nn.Parameter(torch.randn(half_dim))
|
self.weights = nn.Parameter(torch.randn(half_dim), requires_grad = not is_random)
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
x = rearrange(x, 'b -> b 1')
|
x = rearrange(x, 'b -> b 1')
|
||||||
@@ -234,7 +234,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 = 10):
|
def __init__(self, dim, heads = 4, dim_head = 32):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.scale = dim_head ** -0.5
|
self.scale = dim_head ** -0.5
|
||||||
self.heads = heads
|
self.heads = heads
|
||||||
@@ -271,6 +271,7 @@ class Unet(nn.Module):
|
|||||||
resnet_block_groups = 8,
|
resnet_block_groups = 8,
|
||||||
learned_variance = False,
|
learned_variance = False,
|
||||||
learned_sinusoidal_cond = False,
|
learned_sinusoidal_cond = False,
|
||||||
|
random_fourier_features = False,
|
||||||
learned_sinusoidal_dim = 16
|
learned_sinusoidal_dim = 16
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
@@ -293,10 +294,10 @@ class Unet(nn.Module):
|
|||||||
|
|
||||||
time_dim = dim * 4
|
time_dim = dim * 4
|
||||||
|
|
||||||
self.learned_sinusoidal_cond = learned_sinusoidal_cond
|
self.random_or_learned_sinusoidal_cond = learned_sinusoidal_cond or random_fourier_features
|
||||||
|
|
||||||
if learned_sinusoidal_cond:
|
if self.random_or_learned_sinusoidal_cond:
|
||||||
sinu_pos_emb = LearnedSinusoidalPosEmb(learned_sinusoidal_dim)
|
sinu_pos_emb = RandomOrLearnedSinusoidalPosEmb(learned_sinusoidal_dim, random_fourier_features)
|
||||||
fourier_dim = learned_sinusoidal_dim + 1
|
fourier_dim = learned_sinusoidal_dim + 1
|
||||||
else:
|
else:
|
||||||
sinu_pos_emb = SinusoidalPosEmb(dim)
|
sinu_pos_emb = SinusoidalPosEmb(dim)
|
||||||
@@ -429,7 +430,7 @@ 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
|
assert not model.random_or_learned_sinusoidal_cond
|
||||||
|
|
||||||
self.model = model
|
self.model = model
|
||||||
self.channels = self.model.channels
|
self.channels = self.model.channels
|
||||||
@@ -439,7 +440,7 @@ class GaussianDiffusion(nn.Module):
|
|||||||
|
|
||||||
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', 'pred_v'}, 'objective must be either pred_noise (predict noise) or pred_x0 (predict image start) or pred_v (predict v [v-parameterization as defined in appendix D of progressive distillation paper, used in imagen-video successfully])'
|
||||||
|
|
||||||
if beta_schedule == 'linear':
|
if beta_schedule == 'linear':
|
||||||
betas = linear_beta_schedule(timesteps)
|
betas = linear_beta_schedule(timesteps)
|
||||||
@@ -510,6 +511,18 @@ class GaussianDiffusion(nn.Module):
|
|||||||
extract(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape)
|
extract(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def predict_v(self, x_start, t, noise):
|
||||||
|
return (
|
||||||
|
extract(self.sqrt_alphas_cumprod, t, x_start.shape) * noise -
|
||||||
|
extract(self.sqrt_one_minus_alphas_cumprod, t, x_start.shape) * x_start
|
||||||
|
)
|
||||||
|
|
||||||
|
def predict_start_from_v(self, x_t, t, v):
|
||||||
|
return (
|
||||||
|
extract(self.sqrt_alphas_cumprod, t, x_t.shape) * x_t -
|
||||||
|
extract(self.sqrt_one_minus_alphas_cumprod, t, x_t.shape) * v
|
||||||
|
)
|
||||||
|
|
||||||
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 +
|
||||||
@@ -533,6 +546,12 @@ class GaussianDiffusion(nn.Module):
|
|||||||
x_start = maybe_clip(x_start)
|
x_start = maybe_clip(x_start)
|
||||||
pred_noise = self.predict_noise_from_start(x, t, x_start)
|
pred_noise = self.predict_noise_from_start(x, t, x_start)
|
||||||
|
|
||||||
|
elif self.objective == 'pred_v':
|
||||||
|
v = model_output
|
||||||
|
x_start = self.predict_start_from_v(x, t, v)
|
||||||
|
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, x_self_cond = None, clip_denoised = True):
|
def p_mean_variance(self, x, t, x_self_cond = None, clip_denoised = True):
|
||||||
@@ -670,6 +689,9 @@ class GaussianDiffusion(nn.Module):
|
|||||||
target = noise
|
target = noise
|
||||||
elif self.objective == 'pred_x0':
|
elif self.objective == 'pred_x0':
|
||||||
target = x_start
|
target = x_start
|
||||||
|
elif self.objective == 'pred_v':
|
||||||
|
v = self.predict_v(x_start, t, noise)
|
||||||
|
target = v
|
||||||
else:
|
else:
|
||||||
raise ValueError(f'unknown objective {self.objective}')
|
raise ValueError(f'unknown objective {self.objective}')
|
||||||
|
|
||||||
|
|||||||
@@ -52,7 +52,7 @@ class ElucidatedDiffusion(nn.Module):
|
|||||||
S_noise = 1.003,
|
S_noise = 1.003,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
assert net.learned_sinusoidal_cond
|
assert net.random_or_learned_sinusoidal_cond
|
||||||
self.self_condition = net.self_condition
|
self.self_condition = net.self_condition
|
||||||
|
|
||||||
self.net = net
|
self.net = net
|
||||||
|
|||||||
@@ -0,0 +1,184 @@
|
|||||||
|
import math
|
||||||
|
import torch
|
||||||
|
from torch import sqrt
|
||||||
|
from torch import nn, einsum
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch.special import expm1
|
||||||
|
|
||||||
|
from tqdm import tqdm
|
||||||
|
from einops import rearrange, repeat, reduce
|
||||||
|
from einops.layers.torch import Rearrange
|
||||||
|
|
||||||
|
# helpers
|
||||||
|
|
||||||
|
def exists(val):
|
||||||
|
return val is not None
|
||||||
|
|
||||||
|
def default(val, d):
|
||||||
|
if exists(val):
|
||||||
|
return val
|
||||||
|
return d() if callable(d) else d
|
||||||
|
|
||||||
|
# normalization functions
|
||||||
|
|
||||||
|
def normalize_to_neg_one_to_one(img):
|
||||||
|
return img * 2 - 1
|
||||||
|
|
||||||
|
def unnormalize_to_zero_to_one(t):
|
||||||
|
return (t + 1) * 0.5
|
||||||
|
|
||||||
|
# diffusion helpers
|
||||||
|
|
||||||
|
def right_pad_dims_to(x, t):
|
||||||
|
padding_dims = x.ndim - t.ndim
|
||||||
|
if padding_dims <= 0:
|
||||||
|
return t
|
||||||
|
return t.view(*t.shape, *((1,) * padding_dims))
|
||||||
|
|
||||||
|
# continuous schedules
|
||||||
|
# log(snr) that approximates the original linear schedule
|
||||||
|
|
||||||
|
def log(t, eps = 1e-20):
|
||||||
|
return torch.log(t.clamp(min = eps))
|
||||||
|
|
||||||
|
def alpha_cosine_log_snr(t, s = 0.008):
|
||||||
|
return -log((torch.cos((t + s) / (1 + s) * math.pi * 0.5) ** -2) - 1, eps = 1e-5)
|
||||||
|
|
||||||
|
class VParamContinuousTimeGaussianDiffusion(nn.Module):
|
||||||
|
"""
|
||||||
|
a new type of parameterization in v-space proposed in https://arxiv.org/abs/2202.00512 that
|
||||||
|
(1) allows for improved distillation over noise prediction objective and
|
||||||
|
(2) noted in imagen-video to improve upsampling unets by removing the color shifting artifacts
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model,
|
||||||
|
*,
|
||||||
|
image_size,
|
||||||
|
channels = 3,
|
||||||
|
num_sample_steps = 500,
|
||||||
|
clip_sample_denoised = True,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
assert model.random_or_learned_sinusoidal_cond
|
||||||
|
assert not model.self_condition, 'not supported yet'
|
||||||
|
|
||||||
|
self.model = model
|
||||||
|
|
||||||
|
# image dimensions
|
||||||
|
|
||||||
|
self.channels = channels
|
||||||
|
self.image_size = image_size
|
||||||
|
|
||||||
|
# continuous noise schedule related stuff
|
||||||
|
|
||||||
|
self.log_snr = alpha_cosine_log_snr
|
||||||
|
|
||||||
|
# sampling
|
||||||
|
|
||||||
|
self.num_sample_steps = num_sample_steps
|
||||||
|
self.clip_sample_denoised = clip_sample_denoised
|
||||||
|
|
||||||
|
@property
|
||||||
|
def device(self):
|
||||||
|
return next(self.model.parameters()).device
|
||||||
|
|
||||||
|
def p_mean_variance(self, x, time, time_next):
|
||||||
|
# reviewer found an error in the equation in the paper (missing sigma)
|
||||||
|
# following - https://openreview.net/forum?id=2LdBqxc1Yv¬eId=rIQgH0zKsRt
|
||||||
|
|
||||||
|
log_snr = self.log_snr(time)
|
||||||
|
log_snr_next = self.log_snr(time_next)
|
||||||
|
c = -expm1(log_snr - log_snr_next)
|
||||||
|
|
||||||
|
squared_alpha, squared_alpha_next = log_snr.sigmoid(), log_snr_next.sigmoid()
|
||||||
|
squared_sigma, squared_sigma_next = (-log_snr).sigmoid(), (-log_snr_next).sigmoid()
|
||||||
|
|
||||||
|
alpha, sigma, alpha_next = map(sqrt, (squared_alpha, squared_sigma, squared_alpha_next))
|
||||||
|
|
||||||
|
batch_log_snr = repeat(log_snr, ' -> b', b = x.shape[0])
|
||||||
|
|
||||||
|
pred_v = self.model(x, batch_log_snr)
|
||||||
|
|
||||||
|
# shown in Appendix D in the paper
|
||||||
|
x_start = alpha * x - sigma * pred_v
|
||||||
|
|
||||||
|
if self.clip_sample_denoised:
|
||||||
|
x_start.clamp_(-1., 1.)
|
||||||
|
|
||||||
|
model_mean = alpha_next * (x * (1 - c) / alpha + c * x_start)
|
||||||
|
|
||||||
|
posterior_variance = squared_sigma_next * c
|
||||||
|
|
||||||
|
return model_mean, posterior_variance
|
||||||
|
|
||||||
|
# sampling related functions
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def p_sample(self, x, time, time_next):
|
||||||
|
batch, *_, device = *x.shape, x.device
|
||||||
|
|
||||||
|
model_mean, model_variance = self.p_mean_variance(x = x, time = time, time_next = time_next)
|
||||||
|
|
||||||
|
if time_next == 0:
|
||||||
|
return model_mean
|
||||||
|
|
||||||
|
noise = torch.randn_like(x)
|
||||||
|
return model_mean + sqrt(model_variance) * noise
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def p_sample_loop(self, shape):
|
||||||
|
batch = shape[0]
|
||||||
|
|
||||||
|
img = torch.randn(shape, device = self.device)
|
||||||
|
steps = torch.linspace(1., 0., self.num_sample_steps + 1, device = self.device)
|
||||||
|
|
||||||
|
for i in tqdm(range(self.num_sample_steps), desc = 'sampling loop time step', total = self.num_sample_steps):
|
||||||
|
times = steps[i]
|
||||||
|
times_next = steps[i + 1]
|
||||||
|
img = self.p_sample(img, times, times_next)
|
||||||
|
|
||||||
|
img.clamp_(-1., 1.)
|
||||||
|
img = unnormalize_to_zero_to_one(img)
|
||||||
|
return img
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def sample(self, batch_size = 16):
|
||||||
|
return self.p_sample_loop((batch_size, self.channels, self.image_size, self.image_size))
|
||||||
|
|
||||||
|
# training related functions - noise prediction
|
||||||
|
|
||||||
|
def q_sample(self, x_start, times, noise = None):
|
||||||
|
noise = default(noise, lambda: torch.randn_like(x_start))
|
||||||
|
|
||||||
|
log_snr = self.log_snr(times)
|
||||||
|
|
||||||
|
log_snr_padded = right_pad_dims_to(x_start, log_snr)
|
||||||
|
alpha, sigma = sqrt(log_snr_padded.sigmoid()), sqrt((-log_snr_padded).sigmoid())
|
||||||
|
x_noised = x_start * alpha + noise * sigma
|
||||||
|
|
||||||
|
return x_noised, log_snr, alpha, sigma
|
||||||
|
|
||||||
|
def random_times(self, batch_size):
|
||||||
|
return torch.zeros((batch_size,), device = self.device).float().uniform_(0, 1)
|
||||||
|
|
||||||
|
def p_losses(self, x_start, times, noise = None):
|
||||||
|
noise = default(noise, lambda: torch.randn_like(x_start))
|
||||||
|
|
||||||
|
x, log_snr, alpha, sigma = self.q_sample(x_start = x_start, times = times, noise = noise)
|
||||||
|
|
||||||
|
# described in section 4 as the prediction objective, with derivation in Appendix D
|
||||||
|
v = alpha * noise - sigma * x_start
|
||||||
|
|
||||||
|
model_out = self.model(x, log_snr)
|
||||||
|
|
||||||
|
return F.mse_loss(model_out, v)
|
||||||
|
|
||||||
|
def forward(self, img, *args, **kwargs):
|
||||||
|
b, c, h, w, device, img_size, = *img.shape, img.device, self.image_size
|
||||||
|
assert h == img_size and w == img_size, f'height and width of image must be {img_size}'
|
||||||
|
|
||||||
|
times = self.random_times(b)
|
||||||
|
img = normalize_to_neg_one_to_one(img)
|
||||||
|
return self.p_losses(img, times, *args, **kwargs)
|
||||||
@@ -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.28.0',
|
version = '0.30.0',
|
||||||
license='MIT',
|
license='MIT',
|
||||||
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
|
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
|
||||||
author = 'Phil Wang',
|
author = 'Phil Wang',
|
||||||
|
|||||||
Reference in New Issue
Block a user