mirror of
https://github.com/wassname/denoising-diffusion-pytorch.git
synced 2026-09-09 11:21:11 +08:00
first pass at ddpm with learned variance
This commit is contained in:
@@ -1 +1,2 @@
|
||||
from denoising_diffusion_pytorch.denoising_diffusion_pytorch import GaussianDiffusion, Unet, Trainer
|
||||
from denoising_diffusion_pytorch.learned_gaussian_diffusion import LearnedGaussianDiffusion
|
||||
|
||||
@@ -204,7 +204,8 @@ class Unet(nn.Module):
|
||||
dim_mults=(1, 2, 4, 8),
|
||||
channels = 3,
|
||||
with_time_emb = True,
|
||||
resnet_block_groups = 8
|
||||
resnet_block_groups = 8,
|
||||
learned_variance = False
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
@@ -265,10 +266,12 @@ class Unet(nn.Module):
|
||||
Upsample(dim_in) if not is_last else nn.Identity()
|
||||
]))
|
||||
|
||||
out_dim = default(out_dim, channels)
|
||||
default_out_dim = channels * (1 if not learned_variance else 2)
|
||||
self.out_dim = default(out_dim, default_out_dim)
|
||||
|
||||
self.final_conv = nn.Sequential(
|
||||
block_klass(dim, dim),
|
||||
nn.Conv2d(dim, out_dim, 1)
|
||||
nn.Conv2d(dim, self.out_dim, 1)
|
||||
)
|
||||
|
||||
def forward(self, x, time):
|
||||
@@ -320,7 +323,7 @@ def cosine_beta_schedule(timesteps, s = 0.008):
|
||||
alphas_cumprod = torch.cos(((x / timesteps) + s) / (1 + s) * torch.pi * 0.5) ** 2
|
||||
alphas_cumprod = alphas_cumprod / alphas_cumprod[0]
|
||||
betas = 1 - (alphas_cumprod[1:] / alphas_cumprod[:-1])
|
||||
return torch.clip(betas, 0, 0.9999)
|
||||
return torch.clip(betas, 0, 0.999)
|
||||
|
||||
class GaussianDiffusion(nn.Module):
|
||||
def __init__(
|
||||
@@ -333,6 +336,8 @@ class GaussianDiffusion(nn.Module):
|
||||
loss_type = 'l1'
|
||||
):
|
||||
super().__init__()
|
||||
assert not (type(self) == GaussianDiffusion and denoise_fn.channels != denoise_fn.out_dim)
|
||||
|
||||
self.channels = channels
|
||||
self.image_size = image_size
|
||||
self.denoise_fn = denoise_fn
|
||||
|
||||
@@ -0,0 +1,113 @@
|
||||
import torch
|
||||
from math import pi, sqrt, log as ln
|
||||
from inspect import isfunction
|
||||
from torch import nn, einsum
|
||||
from einops import rearrange
|
||||
|
||||
from denoising_diffusion_pytorch.denoising_diffusion_pytorch import GaussianDiffusion, extract
|
||||
|
||||
# constants
|
||||
|
||||
NAT = 1. / ln(2)
|
||||
|
||||
# helper functions
|
||||
|
||||
def exists(x):
|
||||
return x is not None
|
||||
|
||||
def default(val, d):
|
||||
if exists(val):
|
||||
return val
|
||||
return d() if isfunction(d) else d
|
||||
|
||||
# tensor helpers
|
||||
|
||||
def log(t, eps = 1e-12):
|
||||
return torch.log(t.clamp(min = eps))
|
||||
|
||||
def meanflat(x):
|
||||
return x.mean(dim = tuple(range(1, len(x.shape))))
|
||||
|
||||
def normal_kl(mean1, logvar1, mean2, logvar2):
|
||||
"""
|
||||
KL divergence between normal distributions parameterized by mean and log-variance.
|
||||
"""
|
||||
return 0.5 * (-1.0 + logvar2 - logvar1 + torch.exp(logvar1 - logvar2) + ((mean1 - mean2) ** 2) * torch.exp(-logvar2))
|
||||
|
||||
def approx_standard_normal_cdf(x):
|
||||
return 0.5 * (1.0 + torch.tanh(sqrt(2.0 / pi) * (x + 0.044715 * (x ** 3))))
|
||||
|
||||
def discretized_gaussian_log_likelihood(x, *, means, log_scales, thres = 0.999):
|
||||
assert x.shape == means.shape == log_scales.shape
|
||||
|
||||
centered_x = x - means
|
||||
inv_stdv = torch.exp(-log_scales)
|
||||
plus_in = inv_stdv * (centered_x + 1. / 255.)
|
||||
cdf_plus = approx_standard_normal_cdf(plus_in)
|
||||
min_in = inv_stdv * (centered_x - 1. / 255.)
|
||||
cdf_min = approx_standard_normal_cdf(min_in)
|
||||
log_cdf_plus = log(cdf_plus)
|
||||
log_one_minus_cdf_min = log(1. - cdf_min)
|
||||
cdf_delta = cdf_plus - cdf_min
|
||||
|
||||
log_probs = torch.where(x < -thres,
|
||||
log_cdf_plus,
|
||||
torch.where(x > thres,
|
||||
log_one_minus_cdf_min,
|
||||
log(cdf_delta)))
|
||||
|
||||
return log_probs
|
||||
|
||||
# gaussian diffusion for learned variance
|
||||
|
||||
class LearnedGaussianDiffusion(GaussianDiffusion):
|
||||
def __init__(
|
||||
self,
|
||||
denoise_fn,
|
||||
*args,
|
||||
**kwargs
|
||||
):
|
||||
super().__init__(denoise_fn, *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`'
|
||||
|
||||
def q_posterior_mean_variance(self, x_start, x_t, t):
|
||||
"""
|
||||
Compute the mean and variance of the diffusion posterior q(x_{t-1} | x_t, x_0)
|
||||
"""
|
||||
posterior_mean = (
|
||||
extract(self.posterior_mean_coef1, t, x_t.shape) * x_start +
|
||||
extract(self.posterior_mean_coef2, t, x_t.shape) * x_t
|
||||
)
|
||||
posterior_variance = extract(self.posterior_variance, 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
|
||||
|
||||
def predict_xstart_from_xprev(self, x_t, t, xprev):
|
||||
# (xprev - coef2*x_t) / coef1
|
||||
return (
|
||||
extract(1. / self.posterior_mean_coef1, t, x_t.shape) * xprev -
|
||||
extract(self.posterior_mean_coef2 / self.posterior_mean_coef1, t, x_t.shape) * x_t
|
||||
)
|
||||
|
||||
def p_mean_variance(self, *, x, t, clip_denoised):
|
||||
model_output = self.denoise_fn(x, t)
|
||||
model_output, model_log_variance = model_output.chunk(2, dim = 1)
|
||||
model_variance = model_log_variance.exp()
|
||||
return model_output, model_variance, model_log_variance
|
||||
|
||||
def p_losses(self, x_start, t, noise = None, clip_denoised = False):
|
||||
noise = default(noise, lambda: torch.randn_like(x_start))
|
||||
x_t = self.q_sample(x_start = x_start, t = t, noise = noise)
|
||||
|
||||
true_mean, _, true_log_variance_clipped = self.q_posterior_mean_variance(x_start = x_start, x_t = x_t, t = t)
|
||||
model_mean, _, model_log_variance = self.p_mean_variance(x = x_t, t = t, clip_denoised = clip_denoised)
|
||||
|
||||
kl = normal_kl(true_mean, true_log_variance_clipped, model_mean, model_log_variance)
|
||||
kl = meanflat(kl) * NAT
|
||||
|
||||
decoder_nll = -discretized_gaussian_log_likelihood(x_start, means = model_mean, log_scales = 0.5 * model_log_variance)
|
||||
decoder_nll = meanflat(decoder_nll) * NAT
|
||||
|
||||
# At the first timestep return the decoder NLL, otherwise return KL(q(x_{t-1}|x_t,x_0) || p(x_{t-1}|x_t))
|
||||
losses = torch.where(t == 0, decoder_nll, kl)
|
||||
return losses.mean()
|
||||
Reference in New Issue
Block a user