mirror of
https://github.com/wassname/denoising-diffusion-pytorch.git
synced 2026-09-10 12:01:08 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
27853452a2 | ||
|
|
84ebb9ad13 | ||
|
|
caa5af170d | ||
|
|
55c658b967 | ||
|
|
e0f26677d6 | ||
|
|
e147839d74 | ||
|
|
62e8490385 | ||
|
|
d412d8816b | ||
|
|
402b7c26df | ||
|
|
09613a40f3 | ||
|
|
c6966ae95a | ||
|
|
73591cf1ad | ||
|
|
989f0fcb8e | ||
|
|
84731bb03d | ||
|
|
c6ecca555b | ||
|
|
1f5c233072 | ||
|
|
de378158e5 | ||
|
|
e274fb305a | ||
|
|
f39b3b1d3f | ||
|
|
782c904d3b | ||
|
|
71953ebd22 | ||
|
|
0b8cdb4c8b | ||
|
|
e504e0e554 | ||
|
|
bd1e3b676e | ||
|
|
f4615599bc | ||
|
|
eb6e1b508e | ||
|
|
91cff45939 |
@@ -2,7 +2,9 @@
|
||||
|
||||
## Denoising Diffusion Probabilistic Model, in Pytorch
|
||||
|
||||
Implementation of <a href="https://arxiv.org/abs/2006.11239">Denoising Diffusion Probabilistic Model</a> in Pytorch. It is a new approach to generative modeling that may <a href="https://ajolicoeur.wordpress.com/the-new-contender-to-gans-score-matching-with-langevin-sampling/">have the potential</a> to rival GANs. It uses denoising score matching to estimate the gradient of the data distribution, followed by Langevin sampling to sample from the true distribution. This implementation was transcribed from the official Tensorflow version <a href="https://github.com/hojonathanho/diffusion">here</a>.
|
||||
Implementation of <a href="https://arxiv.org/abs/2006.11239">Denoising Diffusion Probabilistic Model</a> in Pytorch. It is a new approach to generative modeling that may <a href="https://ajolicoeur.wordpress.com/the-new-contender-to-gans-score-matching-with-langevin-sampling/">have the potential</a> to rival GANs. It uses denoising score matching to estimate the gradient of the data distribution, followed by Langevin sampling to sample from the true distribution.
|
||||
|
||||
This implementation was transcribed from the official Tensorflow version <a href="https://github.com/hojonathanho/diffusion">here</a>
|
||||
|
||||
<img src="./sample.png" width="500px"><img>
|
||||
|
||||
@@ -32,7 +34,7 @@ diffusion = GaussianDiffusion(
|
||||
loss_type = 'l1' # L1 or L2
|
||||
)
|
||||
|
||||
training_images = torch.randn(8, 3, 128, 128)
|
||||
training_images = torch.randn(8, 3, 128, 128) # your images need to be normalized from a range of -1 to +1
|
||||
loss = diffusion(training_images)
|
||||
loss.backward()
|
||||
# after a lot of training
|
||||
@@ -66,7 +68,7 @@ trainer = Trainer(
|
||||
train_num_steps = 700000, # total training steps
|
||||
gradient_accumulate_every = 2, # gradient accumulation steps
|
||||
ema_decay = 0.995, # exponential moving average decay
|
||||
fp16 = True # turn on mixed precision training with apex
|
||||
amp = True # turn on mixed precision
|
||||
)
|
||||
|
||||
trainer.train()
|
||||
@@ -77,23 +79,32 @@ Samples and model checkpoints will be logged to `./results` periodically
|
||||
## Citations
|
||||
|
||||
```bibtex
|
||||
@misc{ho2020denoising,
|
||||
title = {Denoising Diffusion Probabilistic Models},
|
||||
author = {Jonathan Ho and Ajay Jain and Pieter Abbeel},
|
||||
year = {2020},
|
||||
eprint = {2006.11239},
|
||||
archivePrefix = {arXiv},
|
||||
primaryClass = {cs.LG}
|
||||
@inproceedings{NEURIPS2020_4c5bcfec,
|
||||
author = {Ho, Jonathan and Jain, Ajay and Abbeel, Pieter},
|
||||
booktitle = {Advances in Neural Information Processing Systems},
|
||||
editor = {H. Larochelle and M. Ranzato and R. Hadsell and M.F. Balcan and H. Lin},
|
||||
pages = {6840--6851},
|
||||
publisher = {Curran Associates, Inc.},
|
||||
title = {Denoising Diffusion Probabilistic Models},
|
||||
url = {https://proceedings.neurips.cc/paper/2020/file/4c5bcfec8584af0d967f1ab10179ca4b-Paper.pdf},
|
||||
volume = {33},
|
||||
year = {2020}
|
||||
}
|
||||
```
|
||||
|
||||
```bibtex
|
||||
@inproceedings{anonymous2021improved,
|
||||
title = {Improved Denoising Diffusion Probabilistic Models},
|
||||
author = {Anonymous},
|
||||
booktitle = {Submitted to International Conference on Learning Representations},
|
||||
year = {2021},
|
||||
url = {https://openreview.net/forum?id=-NEXDKk8gZ},
|
||||
note = {under review}
|
||||
@InProceedings{pmlr-v139-nichol21a,
|
||||
title = {Improved Denoising Diffusion Probabilistic Models},
|
||||
author = {Nichol, Alexander Quinn and Dhariwal, Prafulla},
|
||||
booktitle = {Proceedings of the 38th International Conference on Machine Learning},
|
||||
pages = {8162--8171},
|
||||
year = {2021},
|
||||
editor = {Meila, Marina and Zhang, Tong},
|
||||
volume = {139},
|
||||
series = {Proceedings of Machine Learning Research},
|
||||
month = {18--24 Jul},
|
||||
publisher = {PMLR},
|
||||
pdf = {http://proceedings.mlr.press/v139/nichol21a/nichol21a.pdf},
|
||||
url = {https://proceedings.mlr.press/v139/nichol21a.html},
|
||||
}
|
||||
```
|
||||
|
||||
@@ -1 +1,4 @@
|
||||
from denoising_diffusion_pytorch.denoising_diffusion_pytorch import GaussianDiffusion, Unet, Trainer
|
||||
|
||||
from denoising_diffusion_pytorch.learned_gaussian_diffusion import LearnedGaussianDiffusion
|
||||
from denoising_diffusion_pytorch.weighted_objective_gaussian_diffusion import WeightedObjectiveGaussianDiffusion
|
||||
|
||||
@@ -7,21 +7,16 @@ from inspect import isfunction
|
||||
from functools import partial
|
||||
|
||||
from torch.utils import data
|
||||
from torch.cuda.amp import autocast, GradScaler
|
||||
|
||||
from pathlib import Path
|
||||
from torch.optim import Adam
|
||||
from torchvision import transforms, utils
|
||||
from PIL import Image
|
||||
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
from einops import rearrange
|
||||
|
||||
try:
|
||||
from apex import amp
|
||||
APEX_AVAILABLE = True
|
||||
except:
|
||||
APEX_AVAILABLE = False
|
||||
|
||||
# helpers functions
|
||||
|
||||
def exists(x):
|
||||
@@ -45,12 +40,11 @@ def num_to_groups(num, divisor):
|
||||
arr.append(remainder)
|
||||
return arr
|
||||
|
||||
def loss_backwards(fp16, loss, optimizer, **kwargs):
|
||||
if fp16:
|
||||
with amp.scale_loss(loss, optimizer) as scaled_loss:
|
||||
scaled_loss.backward(**kwargs)
|
||||
else:
|
||||
loss.backward(**kwargs)
|
||||
def normalize_to_neg_one_to_one(img):
|
||||
return img * 2 - 1
|
||||
|
||||
def unnormalize_to_zero_to_one(t):
|
||||
return (t + 1) * 0.5
|
||||
|
||||
# small helper modules
|
||||
|
||||
@@ -91,25 +85,11 @@ class SinusoidalPosEmb(nn.Module):
|
||||
emb = torch.cat((emb.sin(), emb.cos()), dim=-1)
|
||||
return emb
|
||||
|
||||
class Mish(nn.Module):
|
||||
def forward(self, x):
|
||||
return x * torch.tanh(F.softplus(x))
|
||||
def Upsample(dim):
|
||||
return nn.ConvTranspose2d(dim, dim, 4, 2, 1)
|
||||
|
||||
class Upsample(nn.Module):
|
||||
def __init__(self, dim):
|
||||
super().__init__()
|
||||
self.conv = nn.ConvTranspose2d(dim, dim, 4, 2, 1)
|
||||
|
||||
def forward(self, x):
|
||||
return self.conv(x)
|
||||
|
||||
class Downsample(nn.Module):
|
||||
def __init__(self, dim):
|
||||
super().__init__()
|
||||
self.conv = nn.Conv2d(dim, dim, 3, 2, 1)
|
||||
|
||||
def forward(self, x):
|
||||
return self.conv(x)
|
||||
def Downsample(dim):
|
||||
return nn.Conv2d(dim, dim, 4, 2, 1)
|
||||
|
||||
class LayerNorm(nn.Module):
|
||||
def __init__(self, dim, eps = 1e-5):
|
||||
@@ -119,9 +99,9 @@ class LayerNorm(nn.Module):
|
||||
self.b = nn.Parameter(torch.zeros(1, dim, 1, 1))
|
||||
|
||||
def forward(self, x):
|
||||
std = torch.var(x, dim = 1, unbiased = False, keepdim = True).sqrt()
|
||||
var = torch.var(x, dim = 1, unbiased = False, keepdim = True)
|
||||
mean = torch.mean(x, dim = 1, keepdim = True)
|
||||
return (x - mean) / (std + self.eps) * self.g + self.b
|
||||
return (x - mean) / (var + self.eps).sqrt() * self.g + self.b
|
||||
|
||||
class PreNorm(nn.Module):
|
||||
def __init__(self, dim, fn):
|
||||
@@ -139,9 +119,9 @@ class Block(nn.Module):
|
||||
def __init__(self, dim, dim_out, groups = 8):
|
||||
super().__init__()
|
||||
self.block = nn.Sequential(
|
||||
nn.Conv2d(dim, dim_out, 3, padding=1),
|
||||
nn.Conv2d(dim, dim_out, 3, padding = 1),
|
||||
nn.GroupNorm(groups, dim_out),
|
||||
Mish()
|
||||
nn.SiLU()
|
||||
)
|
||||
def forward(self, x):
|
||||
return self.block(x)
|
||||
@@ -150,19 +130,20 @@ class ResnetBlock(nn.Module):
|
||||
def __init__(self, dim, dim_out, *, time_emb_dim = None, groups = 8):
|
||||
super().__init__()
|
||||
self.mlp = nn.Sequential(
|
||||
Mish(),
|
||||
nn.SiLU(),
|
||||
nn.Linear(time_emb_dim, dim_out)
|
||||
) if exists(time_emb_dim) else None
|
||||
|
||||
self.block1 = Block(dim, dim_out)
|
||||
self.block2 = Block(dim_out, dim_out)
|
||||
self.block1 = Block(dim, dim_out, groups = groups)
|
||||
self.block2 = Block(dim_out, dim_out, groups = groups)
|
||||
self.res_conv = nn.Conv2d(dim, dim_out, 1) if dim != dim_out else nn.Identity()
|
||||
|
||||
def forward(self, x, time_emb):
|
||||
def forward(self, x, time_emb = None):
|
||||
h = self.block1(x)
|
||||
|
||||
if exists(self.mlp):
|
||||
h += self.mlp(time_emb)[:, :, None, None]
|
||||
if exists(self.mlp) and exists(time_emb):
|
||||
time_emb = self.mlp(time_emb)
|
||||
h = rearrange(time_emb, 'b c -> b c 1 1') + h
|
||||
|
||||
h = self.block2(h)
|
||||
return h + self.res_conv(x)
|
||||
@@ -170,6 +151,35 @@ class ResnetBlock(nn.Module):
|
||||
class LinearAttention(nn.Module):
|
||||
def __init__(self, dim, heads = 4, dim_head = 32):
|
||||
super().__init__()
|
||||
self.scale = dim_head ** -0.5
|
||||
self.heads = heads
|
||||
hidden_dim = dim_head * heads
|
||||
self.to_qkv = nn.Conv2d(dim, hidden_dim * 3, 1, bias = False)
|
||||
|
||||
self.to_out = nn.Sequential(
|
||||
nn.Conv2d(hidden_dim, dim, 1),
|
||||
LayerNorm(dim)
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
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.softmax(dim = -2)
|
||||
k = k.softmax(dim = -1)
|
||||
|
||||
q = q * self.scale
|
||||
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 = rearrange(out, 'b h c (x y) -> b (h c) x y', h = self.heads, x = h, y = w)
|
||||
return self.to_out(out)
|
||||
|
||||
class Attention(nn.Module):
|
||||
def __init__(self, dim, heads = 4, dim_head = 32):
|
||||
super().__init__()
|
||||
self.scale = dim_head ** -0.5
|
||||
self.heads = heads
|
||||
hidden_dim = dim_head * heads
|
||||
self.to_qkv = nn.Conv2d(dim, hidden_dim * 3, 1, bias = False)
|
||||
@@ -177,12 +187,16 @@ class LinearAttention(nn.Module):
|
||||
|
||||
def forward(self, x):
|
||||
b, c, h, w = x.shape
|
||||
qkv = self.to_qkv(x)
|
||||
q, k, v = rearrange(qkv, 'b (qkv heads c) h w -> qkv b heads c (h w)', heads = self.heads, qkv=3)
|
||||
k = k.softmax(dim=-1)
|
||||
context = torch.einsum('bhdn,bhen->bhde', k, v)
|
||||
out = torch.einsum('bhde,bhdn->bhen', context, q)
|
||||
out = rearrange(out, 'b heads c (h w) -> b (heads c) h w', heads=self.heads, h=h, w=w)
|
||||
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()
|
||||
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)
|
||||
|
||||
# model
|
||||
@@ -191,30 +205,44 @@ class Unet(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
init_dim = None,
|
||||
out_dim = None,
|
||||
dim_mults=(1, 2, 4, 8),
|
||||
groups = 8,
|
||||
channels = 3,
|
||||
with_time_emb = True
|
||||
with_time_emb = True,
|
||||
resnet_block_groups = 8,
|
||||
learned_variance = False
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
# determine dimensions
|
||||
|
||||
self.channels = channels
|
||||
|
||||
dims = [channels, *map(lambda m: dim * m, dim_mults)]
|
||||
init_dim = default(init_dim, dim // 3 * 2)
|
||||
self.init_conv = nn.Conv2d(channels, init_dim, 7, padding = 3)
|
||||
|
||||
dims = [init_dim, *map(lambda m: dim * m, dim_mults)]
|
||||
in_out = list(zip(dims[:-1], dims[1:]))
|
||||
|
||||
block_klass = partial(ResnetBlock, groups = resnet_block_groups)
|
||||
|
||||
# time embeddings
|
||||
|
||||
if with_time_emb:
|
||||
time_dim = dim
|
||||
time_dim = dim * 4
|
||||
self.time_mlp = nn.Sequential(
|
||||
SinusoidalPosEmb(dim),
|
||||
nn.Linear(dim, dim * 4),
|
||||
Mish(),
|
||||
nn.Linear(dim * 4, dim)
|
||||
nn.Linear(dim, time_dim),
|
||||
nn.GELU(),
|
||||
nn.Linear(time_dim, time_dim)
|
||||
)
|
||||
else:
|
||||
time_dim = None
|
||||
self.time_mlp = None
|
||||
|
||||
# layers
|
||||
|
||||
self.downs = nn.ModuleList([])
|
||||
self.ups = nn.ModuleList([])
|
||||
num_resolutions = len(in_out)
|
||||
@@ -223,41 +251,45 @@ class Unet(nn.Module):
|
||||
is_last = ind >= (num_resolutions - 1)
|
||||
|
||||
self.downs.append(nn.ModuleList([
|
||||
ResnetBlock(dim_in, dim_out, time_emb_dim = time_dim),
|
||||
ResnetBlock(dim_out, dim_out, time_emb_dim = time_dim),
|
||||
block_klass(dim_in, dim_out, time_emb_dim = time_dim),
|
||||
block_klass(dim_out, dim_out, time_emb_dim = time_dim),
|
||||
Residual(PreNorm(dim_out, LinearAttention(dim_out))),
|
||||
Downsample(dim_out) if not is_last else nn.Identity()
|
||||
]))
|
||||
|
||||
mid_dim = dims[-1]
|
||||
self.mid_block1 = ResnetBlock(mid_dim, mid_dim, time_emb_dim = time_dim)
|
||||
self.mid_attn = Residual(PreNorm(mid_dim, LinearAttention(mid_dim)))
|
||||
self.mid_block2 = ResnetBlock(mid_dim, mid_dim, time_emb_dim = time_dim)
|
||||
self.mid_block1 = block_klass(mid_dim, mid_dim, time_emb_dim = time_dim)
|
||||
self.mid_attn = Residual(PreNorm(mid_dim, Attention(mid_dim)))
|
||||
self.mid_block2 = block_klass(mid_dim, mid_dim, time_emb_dim = time_dim)
|
||||
|
||||
for ind, (dim_in, dim_out) in enumerate(reversed(in_out[1:])):
|
||||
is_last = ind >= (num_resolutions - 1)
|
||||
|
||||
self.ups.append(nn.ModuleList([
|
||||
ResnetBlock(dim_out * 2, dim_in, time_emb_dim = time_dim),
|
||||
ResnetBlock(dim_in, dim_in, time_emb_dim = time_dim),
|
||||
block_klass(dim_out * 2, dim_in, time_emb_dim = time_dim),
|
||||
block_klass(dim_in, dim_in, time_emb_dim = time_dim),
|
||||
Residual(PreNorm(dim_in, LinearAttention(dim_in))),
|
||||
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(dim, dim),
|
||||
nn.Conv2d(dim, out_dim, 1)
|
||||
block_klass(dim, dim),
|
||||
nn.Conv2d(dim, self.out_dim, 1)
|
||||
)
|
||||
|
||||
def forward(self, x, time):
|
||||
x = self.init_conv(x)
|
||||
|
||||
t = self.time_mlp(time) if exists(self.time_mlp) else None
|
||||
|
||||
h = []
|
||||
|
||||
for resnet, resnet2, attn, downsample in self.downs:
|
||||
x = resnet(x, t)
|
||||
x = resnet2(x, t)
|
||||
for block1, block2, attn, downsample in self.downs:
|
||||
x = block1(x, t)
|
||||
x = block2(x, t)
|
||||
x = attn(x)
|
||||
h.append(x)
|
||||
x = downsample(x)
|
||||
@@ -266,10 +298,10 @@ class Unet(nn.Module):
|
||||
x = self.mid_attn(x)
|
||||
x = self.mid_block2(x, t)
|
||||
|
||||
for resnet, resnet2, attn, upsample in self.ups:
|
||||
for block1, block2, attn, upsample in self.ups:
|
||||
x = torch.cat((x, h.pop()), dim=1)
|
||||
x = resnet(x, t)
|
||||
x = resnet2(x, t)
|
||||
x = block1(x, t)
|
||||
x = block2(x, t)
|
||||
x = attn(x)
|
||||
x = upsample(x)
|
||||
|
||||
@@ -293,11 +325,11 @@ def cosine_beta_schedule(timesteps, s = 0.008):
|
||||
as proposed in https://openreview.net/forum?id=-NEXDKk8gZ
|
||||
"""
|
||||
steps = timesteps + 1
|
||||
x = np.linspace(0, steps, steps)
|
||||
alphas_cumprod = np.cos(((x / steps) + s) / (1 + s) * np.pi * 0.5) ** 2
|
||||
x = torch.linspace(0, timesteps, steps, dtype = torch.float64)
|
||||
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 np.clip(betas, a_min = 0, a_max = 0.999)
|
||||
return torch.clip(betas, 0, 0.999)
|
||||
|
||||
class GaussianDiffusion(nn.Module):
|
||||
def __init__(
|
||||
@@ -308,55 +340,55 @@ class GaussianDiffusion(nn.Module):
|
||||
channels = 3,
|
||||
timesteps = 1000,
|
||||
loss_type = 'l1',
|
||||
betas = None
|
||||
objective = 'pred_noise'
|
||||
):
|
||||
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
|
||||
self.objective = objective
|
||||
|
||||
if exists(betas):
|
||||
betas = betas.detach().cpu().numpy() if isinstance(betas, torch.Tensor) else betas
|
||||
else:
|
||||
betas = cosine_beta_schedule(timesteps)
|
||||
betas = cosine_beta_schedule(timesteps)
|
||||
|
||||
alphas = 1. - betas
|
||||
alphas_cumprod = np.cumprod(alphas, axis=0)
|
||||
alphas_cumprod_prev = np.append(1., alphas_cumprod[:-1])
|
||||
alphas_cumprod = torch.cumprod(alphas, axis=0)
|
||||
alphas_cumprod_prev = F.pad(alphas_cumprod[:-1], (1, 0), value = 1.)
|
||||
|
||||
timesteps, = betas.shape
|
||||
self.num_timesteps = int(timesteps)
|
||||
self.loss_type = loss_type
|
||||
|
||||
to_torch = partial(torch.tensor, dtype=torch.float32)
|
||||
# helper function to register buffer from float64 to float32
|
||||
|
||||
self.register_buffer('betas', to_torch(betas))
|
||||
self.register_buffer('alphas_cumprod', to_torch(alphas_cumprod))
|
||||
self.register_buffer('alphas_cumprod_prev', to_torch(alphas_cumprod_prev))
|
||||
register_buffer = lambda name, val: self.register_buffer(name, val.to(torch.float32))
|
||||
|
||||
register_buffer('betas', betas)
|
||||
register_buffer('alphas_cumprod', alphas_cumprod)
|
||||
register_buffer('alphas_cumprod_prev', alphas_cumprod_prev)
|
||||
|
||||
# calculations for diffusion q(x_t | x_{t-1}) and others
|
||||
self.register_buffer('sqrt_alphas_cumprod', to_torch(np.sqrt(alphas_cumprod)))
|
||||
self.register_buffer('sqrt_one_minus_alphas_cumprod', to_torch(np.sqrt(1. - alphas_cumprod)))
|
||||
self.register_buffer('log_one_minus_alphas_cumprod', to_torch(np.log(1. - alphas_cumprod)))
|
||||
self.register_buffer('sqrt_recip_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod)))
|
||||
self.register_buffer('sqrt_recipm1_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod - 1)))
|
||||
|
||||
register_buffer('sqrt_alphas_cumprod', torch.sqrt(alphas_cumprod))
|
||||
register_buffer('sqrt_one_minus_alphas_cumprod', torch.sqrt(1. - alphas_cumprod))
|
||||
register_buffer('log_one_minus_alphas_cumprod', torch.log(1. - alphas_cumprod))
|
||||
register_buffer('sqrt_recip_alphas_cumprod', torch.sqrt(1. / alphas_cumprod))
|
||||
register_buffer('sqrt_recipm1_alphas_cumprod', torch.sqrt(1. / alphas_cumprod - 1))
|
||||
|
||||
# calculations for posterior q(x_{t-1} | x_t, x_0)
|
||||
posterior_variance = betas * (1. - alphas_cumprod_prev) / (1. - alphas_cumprod)
|
||||
# above: equal to 1. / (1. / (1. - alpha_cumprod_tm1) + alpha_t / beta_t)
|
||||
self.register_buffer('posterior_variance', to_torch(posterior_variance))
|
||||
# below: log calculation clipped because the posterior variance is 0 at the beginning of the diffusion chain
|
||||
self.register_buffer('posterior_log_variance_clipped', to_torch(np.log(np.maximum(posterior_variance, 1e-20))))
|
||||
self.register_buffer('posterior_mean_coef1', to_torch(
|
||||
betas * np.sqrt(alphas_cumprod_prev) / (1. - alphas_cumprod)))
|
||||
self.register_buffer('posterior_mean_coef2', to_torch(
|
||||
(1. - alphas_cumprod_prev) * np.sqrt(alphas) / (1. - alphas_cumprod)))
|
||||
|
||||
def q_mean_variance(self, x_start, t):
|
||||
mean = extract(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start
|
||||
variance = extract(1. - self.alphas_cumprod, t, x_start.shape)
|
||||
log_variance = extract(self.log_one_minus_alphas_cumprod, t, x_start.shape)
|
||||
return mean, variance, log_variance
|
||||
posterior_variance = betas * (1. - alphas_cumprod_prev) / (1. - alphas_cumprod)
|
||||
|
||||
# above: equal to 1. / (1. / (1. - alpha_cumprod_tm1) + alpha_t / beta_t)
|
||||
|
||||
register_buffer('posterior_variance', posterior_variance)
|
||||
|
||||
# below: log calculation clipped because the posterior variance is 0 at the beginning of the diffusion chain
|
||||
|
||||
register_buffer('posterior_log_variance_clipped', torch.log(posterior_variance.clamp(min =1e-20)))
|
||||
register_buffer('posterior_mean_coef1', betas * torch.sqrt(alphas_cumprod_prev) / (1. - alphas_cumprod))
|
||||
register_buffer('posterior_mean_coef2', (1. - alphas_cumprod_prev) * torch.sqrt(alphas) / (1. - alphas_cumprod))
|
||||
|
||||
def predict_start_from_noise(self, x_t, t, noise):
|
||||
return (
|
||||
@@ -374,12 +406,19 @@ class GaussianDiffusion(nn.Module):
|
||||
return posterior_mean, posterior_variance, posterior_log_variance_clipped
|
||||
|
||||
def p_mean_variance(self, x, t, clip_denoised: bool):
|
||||
x_recon = self.predict_start_from_noise(x, t=t, noise=self.denoise_fn(x, t))
|
||||
model_output = self.denoise_fn(x, t)
|
||||
|
||||
if self.objective == 'pred_noise':
|
||||
x_start = self.predict_start_from_noise(x, t = t, noise = model_output)
|
||||
elif self.objective == 'pred_x0':
|
||||
x_start = model_output
|
||||
else:
|
||||
raise ValueError(f'unknown objective {self.objective}')
|
||||
|
||||
if clip_denoised:
|
||||
x_recon.clamp_(-1., 1.)
|
||||
x_start.clamp_(-1., 1.)
|
||||
|
||||
model_mean, posterior_variance, posterior_log_variance = self.q_posterior(x_start=x_recon, 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
|
||||
|
||||
@torch.no_grad()
|
||||
@@ -432,20 +471,30 @@ class GaussianDiffusion(nn.Module):
|
||||
extract(self.sqrt_one_minus_alphas_cumprod, t, x_start.shape) * noise
|
||||
)
|
||||
|
||||
@property
|
||||
def loss_fn(self):
|
||||
if self.loss_type == 'l1':
|
||||
return F.l1_loss
|
||||
elif self.loss_type == 'l2':
|
||||
return F.mse_loss
|
||||
else:
|
||||
raise ValueError(f'invalid loss type {self.loss_type}')
|
||||
|
||||
def p_losses(self, x_start, t, noise = None):
|
||||
b, c, h, w = x_start.shape
|
||||
noise = default(noise, lambda: torch.randn_like(x_start))
|
||||
|
||||
x_noisy = self.q_sample(x_start=x_start, t=t, noise=noise)
|
||||
x_recon = self.denoise_fn(x_noisy, t)
|
||||
x = self.q_sample(x_start=x_start, t=t, noise=noise)
|
||||
model_out = self.denoise_fn(x, t)
|
||||
|
||||
if self.loss_type == 'l1':
|
||||
loss = (noise - x_recon).abs().mean()
|
||||
elif self.loss_type == 'l2':
|
||||
loss = F.mse_loss(noise, x_recon)
|
||||
if self.objective == 'pred_noise':
|
||||
target = noise
|
||||
elif self.objective == 'pred_x0':
|
||||
target = x_start
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
raise ValueError(f'unknown objective {self.objective}')
|
||||
|
||||
loss = self.loss_fn(model_out, target)
|
||||
return loss
|
||||
|
||||
def forward(self, x, *args, **kwargs):
|
||||
@@ -468,7 +517,7 @@ class Dataset(data.Dataset):
|
||||
transforms.RandomHorizontalFlip(),
|
||||
transforms.CenterCrop(image_size),
|
||||
transforms.ToTensor(),
|
||||
transforms.Lambda(lambda t: (t * 2) - 1)
|
||||
transforms.Lambda(normalize_to_neg_one_to_one)
|
||||
])
|
||||
|
||||
def __len__(self):
|
||||
@@ -493,7 +542,7 @@ class Trainer(object):
|
||||
train_lr = 2e-5,
|
||||
train_num_steps = 100000,
|
||||
gradient_accumulate_every = 2,
|
||||
fp16 = False,
|
||||
amp = False,
|
||||
step_start_ema = 2000,
|
||||
update_ema_every = 10,
|
||||
save_and_sample_every = 1000,
|
||||
@@ -519,11 +568,8 @@ class Trainer(object):
|
||||
|
||||
self.step = 0
|
||||
|
||||
assert not fp16 or fp16 and APEX_AVAILABLE, 'Apex must be installed in order for mixed precision training to be turned on'
|
||||
|
||||
self.fp16 = fp16
|
||||
if fp16:
|
||||
(self.model, self.ema_model), self.opt = amp.initialize([self.model, self.ema_model], self.opt, opt_level='O1')
|
||||
self.amp = amp
|
||||
self.scaler = GradScaler(enabled = amp)
|
||||
|
||||
self.results_folder = Path(results_folder)
|
||||
self.results_folder.mkdir(exist_ok = True)
|
||||
@@ -543,7 +589,8 @@ class Trainer(object):
|
||||
data = {
|
||||
'step': self.step,
|
||||
'model': self.model.state_dict(),
|
||||
'ema': self.ema_model.state_dict()
|
||||
'ema': self.ema_model.state_dict(),
|
||||
'scaler': self.scaler.state_dict()
|
||||
}
|
||||
torch.save(data, str(self.results_folder / f'model-{milestone}.pt'))
|
||||
|
||||
@@ -553,29 +600,34 @@ class Trainer(object):
|
||||
self.step = data['step']
|
||||
self.model.load_state_dict(data['model'])
|
||||
self.ema_model.load_state_dict(data['ema'])
|
||||
self.scaler.load_state_dict(data['scaler'])
|
||||
|
||||
def train(self):
|
||||
backwards = partial(loss_backwards, self.fp16)
|
||||
|
||||
while self.step < self.train_num_steps:
|
||||
for i in range(self.gradient_accumulate_every):
|
||||
data = next(self.dl).cuda()
|
||||
loss = self.model(data)
|
||||
print(f'{self.step}: {loss.item()}')
|
||||
backwards(loss / self.gradient_accumulate_every, self.opt)
|
||||
|
||||
self.opt.step()
|
||||
with autocast(enabled = self.amp):
|
||||
loss = self.model(data)
|
||||
self.scaler.scale(loss / self.gradient_accumulate_every).backward()
|
||||
|
||||
print(f'{self.step}: {loss.item()}')
|
||||
|
||||
self.scaler.step(self.opt)
|
||||
self.scaler.update()
|
||||
self.opt.zero_grad()
|
||||
|
||||
if self.step % self.update_ema_every == 0:
|
||||
self.step_ema()
|
||||
|
||||
if self.step != 0 and self.step % self.save_and_sample_every == 0:
|
||||
self.ema_model.eval()
|
||||
|
||||
milestone = self.step // self.save_and_sample_every
|
||||
batches = num_to_groups(36, self.batch_size)
|
||||
all_images_list = list(map(lambda n: self.ema_model.sample(batch_size=n), batches))
|
||||
all_images = torch.cat(all_images_list, dim=0)
|
||||
all_images = (all_images + 1) * 0.5
|
||||
all_images = unnormalize_to_zero_to_one(all_images)
|
||||
utils.save_image(all_images, str(self.results_folder / f'sample-{milestone}.png'), nrow = 6)
|
||||
self.save(milestone)
|
||||
|
||||
|
||||
@@ -0,0 +1,132 @@
|
||||
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, unnormalize_to_zero_to_one
|
||||
|
||||
# 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
|
||||
|
||||
# https://arxiv.org/abs/2102.09672
|
||||
|
||||
# i thought the results were questionable, if one were to focus only on FID
|
||||
# but may as well get this in here for others to try, as GLIDE is using it (and DALL-E2 first stage of cascade)
|
||||
# gaussian diffusion for learned variance + hybrid eps simple + vb loss
|
||||
|
||||
class LearnedGaussianDiffusion(GaussianDiffusion):
|
||||
def __init__(
|
||||
self,
|
||||
denoise_fn,
|
||||
vb_loss_weight = 0.001, # lambda was 0.001 in the paper
|
||||
*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`'
|
||||
self.vb_loss_weight = vb_loss_weight
|
||||
|
||||
def p_mean_variance(self, *, x, t, clip_denoised, model_output = None):
|
||||
model_output = default(model_output, lambda: self.denoise_fn(x, t))
|
||||
pred_noise, var_interp_frac_unnormalized = model_output.chunk(2, dim = 1)
|
||||
|
||||
min_log = extract(self.posterior_log_variance_clipped, t, x.shape)
|
||||
max_log = extract(torch.log(self.betas), t, x.shape)
|
||||
var_interp_frac = unnormalize_to_zero_to_one(var_interp_frac_unnormalized)
|
||||
|
||||
model_log_variance = var_interp_frac * max_log + (1 - var_interp_frac) * min_log
|
||||
model_variance = model_log_variance.exp()
|
||||
|
||||
x_start = self.predict_start_from_noise(x, t, pred_noise)
|
||||
|
||||
if clip_denoised:
|
||||
x_start.clamp_(-1., 1.)
|
||||
|
||||
model_mean, _, _ = self.q_posterior(x_start, x, t)
|
||||
|
||||
return model_mean, 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)
|
||||
|
||||
# model output
|
||||
|
||||
model_output = self.denoise_fn(x_t, t)
|
||||
|
||||
# calculating kl loss for learned variance (interpolation)
|
||||
|
||||
true_mean, _, true_log_variance_clipped = self.q_posterior(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, model_output = model_output)
|
||||
|
||||
# kl loss with detached model predicted mean, for stability reasons as in paper
|
||||
|
||||
detached_model_mean = model_mean.detach()
|
||||
|
||||
kl = normal_kl(true_mean, true_log_variance_clipped, detached_model_mean, model_log_variance)
|
||||
kl = meanflat(kl) * NAT
|
||||
|
||||
decoder_nll = -discretized_gaussian_log_likelihood(x_start, means = detached_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))
|
||||
|
||||
vb_losses = torch.where(t == 0, decoder_nll, kl)
|
||||
|
||||
# simple loss - predicting noise, x0, or x_prev
|
||||
|
||||
pred_noise, _ = model_output.chunk(2, dim = 1)
|
||||
|
||||
simple_losses = self.loss_fn(pred_noise, noise)
|
||||
|
||||
return simple_losses + vb_losses.mean() * self.vb_loss_weight
|
||||
@@ -0,0 +1,80 @@
|
||||
import torch
|
||||
from inspect import isfunction
|
||||
from torch import nn, einsum
|
||||
from einops import rearrange
|
||||
|
||||
from denoising_diffusion_pytorch.denoising_diffusion_pytorch import GaussianDiffusion, extract, unnormalize_to_zero_to_one
|
||||
|
||||
# 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
|
||||
|
||||
# some improvisation on my end
|
||||
# where i have the model learn to both predict noise and x0
|
||||
# and learn the weighted sum for each depending on time step
|
||||
|
||||
class WeightedObjectiveGaussianDiffusion(GaussianDiffusion):
|
||||
def __init__(
|
||||
self,
|
||||
denoise_fn,
|
||||
*args,
|
||||
pred_noise_loss_weight = 0.1,
|
||||
pred_x_start_loss_weight = 0.1,
|
||||
**kwargs
|
||||
):
|
||||
super().__init__(denoise_fn, *args, **kwargs)
|
||||
channels = denoise_fn.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'
|
||||
|
||||
self.split_dims = (channels, channels, 2)
|
||||
self.pred_noise_loss_weight = pred_noise_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):
|
||||
model_output = self.denoise_fn(x, t)
|
||||
|
||||
pred_noise, pred_x_start, weights = model_output.split(self.split_dims, dim = 1)
|
||||
normalized_weights = weights.softmax(dim = 1)
|
||||
|
||||
x_start_from_noise = self.predict_start_from_noise(x, t = t, noise = pred_noise)
|
||||
|
||||
x_starts = torch.stack((x_start_from_noise, pred_x_start), dim = 1)
|
||||
weighted_x_start = einsum('b j h w, b j c h w -> b c h w', normalized_weights, x_starts)
|
||||
|
||||
if clip_denoised:
|
||||
weighted_x_start.clamp_(-1., 1.)
|
||||
|
||||
model_mean, model_variance, model_log_variance = self.q_posterior(weighted_x_start, x, t)
|
||||
|
||||
return model_mean, 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)
|
||||
|
||||
model_output = self.denoise_fn(x_t, t)
|
||||
pred_noise, pred_x_start, weights = model_output.split(self.split_dims, dim = 1)
|
||||
|
||||
# get loss for predicted noise and x_start
|
||||
# with the loss weight given at initialization
|
||||
|
||||
noise_loss = self.loss_fn(noise, pred_noise) * self.pred_noise_loss_weight
|
||||
x_start_loss = self.loss_fn(x_start, pred_x_start) * self.pred_x_start_loss_weight
|
||||
|
||||
# calculate x_start from predicted noise
|
||||
# then do a weighted sum of the x_start prediction, weights also predicted by the model (softmax normalized)
|
||||
|
||||
x_start_from_pred_noise = self.predict_start_from_noise(x_t, t, pred_noise)
|
||||
x_start_from_pred_noise = x_start_from_pred_noise.clamp(-2., 2.)
|
||||
weighted_x_start = einsum('b j h w, b j c h w -> b c h w', weights.softmax(dim = 1), torch.stack((x_start_from_pred_noise, pred_x_start), dim = 1))
|
||||
|
||||
# main loss to x_start with the weighted one
|
||||
|
||||
weighted_x_start_loss = self.loss_fn(x_start, weighted_x_start)
|
||||
return weighted_x_start_loss + x_start_loss + noise_loss
|
||||
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
|
||||
setup(
|
||||
name = 'denoising-diffusion-pytorch',
|
||||
packages = find_packages(),
|
||||
version = '0.6.9',
|
||||
version = '0.15.1',
|
||||
license='MIT',
|
||||
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
|
||||
author = 'Phil Wang',
|
||||
@@ -15,7 +15,6 @@ setup(
|
||||
],
|
||||
install_requires=[
|
||||
'einops',
|
||||
'numpy',
|
||||
'pillow',
|
||||
'torch',
|
||||
'torchvision',
|
||||
|
||||
Reference in New Issue
Block a user