mirror of
https://github.com/wassname/denoising-diffusion-pytorch.git
synced 2026-09-10 12:01:08 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f2765c4614 |
@@ -1,6 +1,3 @@
|
||||
# Generation results
|
||||
results/
|
||||
|
||||
# Byte-compiled / optimized / DLL files
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
|
||||
@@ -2,9 +2,7 @@
|
||||
|
||||
## 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> and then modified to use <a href="https://arxiv.org/abs/2201.03545">ConvNext</a> blocks instead of Resnets.
|
||||
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>
|
||||
|
||||
@@ -29,17 +27,16 @@ model = Unet(
|
||||
|
||||
diffusion = GaussianDiffusion(
|
||||
model,
|
||||
image_size = 128,
|
||||
timesteps = 1000, # number of steps
|
||||
loss_type = 'l1' # L1 or L2
|
||||
loss_type = 'l1' # L1 or L2
|
||||
)
|
||||
|
||||
training_images = torch.randn(8, 3, 128, 128) # your images need to be normalized from a range of -1 to +1
|
||||
training_images = torch.randn(8, 3, 128, 128)
|
||||
loss = diffusion(training_images)
|
||||
loss.backward()
|
||||
# after a lot of training
|
||||
|
||||
sampled_images = diffusion.sample(batch_size = 4)
|
||||
sampled_images = diffusion.sample(128, batch_size = 4)
|
||||
sampled_images.shape # (4, 3, 128, 128)
|
||||
```
|
||||
|
||||
@@ -55,7 +52,6 @@ model = Unet(
|
||||
|
||||
diffusion = GaussianDiffusion(
|
||||
model,
|
||||
image_size = 128,
|
||||
timesteps = 1000, # number of steps
|
||||
loss_type = 'l1' # L1 or L2
|
||||
).cuda()
|
||||
@@ -63,39 +59,39 @@ diffusion = GaussianDiffusion(
|
||||
trainer = Trainer(
|
||||
diffusion,
|
||||
'path/to/your/images',
|
||||
image_size = 128,
|
||||
train_batch_size = 32,
|
||||
train_lr = 2e-5,
|
||||
train_num_steps = 700000, # total training steps
|
||||
train_num_steps = 100000, # total training steps
|
||||
gradient_accumulate_every = 2, # gradient accumulation steps
|
||||
ema_decay = 0.995, # exponential moving average decay
|
||||
amp = True # turn on mixed precision
|
||||
fp16 = True # turn on mixed precision training with apex
|
||||
)
|
||||
|
||||
trainer.train()
|
||||
```
|
||||
|
||||
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}
|
||||
title={Denoising Diffusion Probabilistic Models},
|
||||
author={Jonathan Ho and Ajay Jain and Pieter Abbeel},
|
||||
year={2020},
|
||||
eprint={2006.11239},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.LG}
|
||||
}
|
||||
```
|
||||
|
||||
```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{
|
||||
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}
|
||||
}
|
||||
```
|
||||
|
||||
@@ -7,16 +7,27 @@ 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
|
||||
|
||||
# constants
|
||||
|
||||
SAVE_AND_SAMPLE_EVERY = 1000
|
||||
UPDATE_EMA_EVERY = 10
|
||||
EXTS = ['jpg', 'jpeg', 'png']
|
||||
|
||||
# helpers functions
|
||||
|
||||
def exists(x):
|
||||
@@ -40,6 +51,13 @@ 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)
|
||||
|
||||
# small helper modules
|
||||
|
||||
class EMA():
|
||||
@@ -79,33 +97,34 @@ class SinusoidalPosEmb(nn.Module):
|
||||
emb = torch.cat((emb.sin(), emb.cos()), dim=-1)
|
||||
return emb
|
||||
|
||||
def Upsample(dim):
|
||||
return nn.ConvTranspose2d(dim, dim, 4, 2, 1)
|
||||
class Mish(nn.Module):
|
||||
def forward(self, x):
|
||||
return x * torch.tanh(F.softplus(x))
|
||||
|
||||
def Downsample(dim):
|
||||
return nn.Conv2d(dim, dim, 4, 2, 1)
|
||||
|
||||
class LayerNorm(nn.Module):
|
||||
def __init__(self, dim, eps = 1e-5):
|
||||
class Upsample(nn.Module):
|
||||
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))
|
||||
self.conv = nn.ConvTranspose2d(dim, dim, 4, 2, 1)
|
||||
|
||||
def forward(self, x):
|
||||
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 self.conv(x)
|
||||
|
||||
class PreNorm(nn.Module):
|
||||
def __init__(self, dim, fn):
|
||||
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)
|
||||
|
||||
class Rezero(nn.Module):
|
||||
def __init__(self, fn):
|
||||
super().__init__()
|
||||
self.fn = fn
|
||||
self.norm = LayerNorm(dim)
|
||||
self.g = nn.Parameter(torch.zeros(1))
|
||||
|
||||
def forward(self, x):
|
||||
x = self.norm(x)
|
||||
return self.fn(x)
|
||||
return self.fn(x) * self.g
|
||||
|
||||
# building block modules
|
||||
|
||||
@@ -113,67 +132,34 @@ 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),
|
||||
nn.SiLU()
|
||||
Mish()
|
||||
)
|
||||
def forward(self, x):
|
||||
return self.block(x)
|
||||
|
||||
class ResnetBlock(nn.Module):
|
||||
def __init__(self, dim, dim_out, *, time_emb_dim = None, groups = 8):
|
||||
def __init__(self, dim, dim_out, *, time_emb_dim, groups = 8):
|
||||
super().__init__()
|
||||
self.mlp = nn.Sequential(
|
||||
nn.SiLU(),
|
||||
Mish(),
|
||||
nn.Linear(time_emb_dim, dim_out)
|
||||
) if exists(time_emb_dim) else None
|
||||
)
|
||||
|
||||
self.block1 = Block(dim, dim_out, groups = groups)
|
||||
self.block2 = Block(dim_out, dim_out, groups = groups)
|
||||
self.block1 = Block(dim, dim_out)
|
||||
self.block2 = Block(dim_out, dim_out)
|
||||
self.res_conv = nn.Conv2d(dim, dim_out, 1) if dim != dim_out else nn.Identity()
|
||||
|
||||
def forward(self, x, time_emb = None):
|
||||
def forward(self, x, time_emb):
|
||||
h = self.block1(x)
|
||||
|
||||
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.mlp(time_emb)[:, :, None, None]
|
||||
h = self.block2(h)
|
||||
return h + self.res_conv(x)
|
||||
|
||||
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)
|
||||
@@ -181,60 +167,28 @@ class Attention(nn.Module):
|
||||
|
||||
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 * 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)
|
||||
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)
|
||||
return self.to_out(out)
|
||||
|
||||
# model
|
||||
|
||||
class Unet(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
init_dim = None,
|
||||
out_dim = None,
|
||||
dim_mults=(1, 2, 4, 8),
|
||||
channels = 3,
|
||||
with_time_emb = True,
|
||||
resnet_block_groups = 8
|
||||
):
|
||||
def __init__(self, dim, out_dim = None, dim_mults=(1, 2, 4, 8), groups = 8):
|
||||
super().__init__()
|
||||
|
||||
# determine dimensions
|
||||
|
||||
self.channels = channels
|
||||
|
||||
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)]
|
||||
dims = [3, *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 * 4
|
||||
self.time_mlp = nn.Sequential(
|
||||
SinusoidalPosEmb(dim),
|
||||
nn.Linear(dim, time_dim),
|
||||
nn.GELU(),
|
||||
nn.Linear(time_dim, time_dim)
|
||||
)
|
||||
else:
|
||||
time_dim = None
|
||||
self.time_mlp = None
|
||||
|
||||
# layers
|
||||
self.time_pos_emb = SinusoidalPosEmb(dim)
|
||||
self.mlp = nn.Sequential(
|
||||
nn.Linear(dim, dim * 4),
|
||||
Mish(),
|
||||
nn.Linear(dim * 4, dim)
|
||||
)
|
||||
|
||||
self.downs = nn.ModuleList([])
|
||||
self.ups = nn.ModuleList([])
|
||||
@@ -244,43 +198,42 @@ class Unet(nn.Module):
|
||||
is_last = ind >= (num_resolutions - 1)
|
||||
|
||||
self.downs.append(nn.ModuleList([
|
||||
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))),
|
||||
ResnetBlock(dim_in, dim_out, time_emb_dim = dim),
|
||||
ResnetBlock(dim_out, dim_out, time_emb_dim = dim),
|
||||
Residual(Rezero(LinearAttention(dim_out))),
|
||||
Downsample(dim_out) if not is_last else nn.Identity()
|
||||
]))
|
||||
|
||||
mid_dim = dims[-1]
|
||||
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)
|
||||
self.mid_block1 = ResnetBlock(mid_dim, mid_dim, time_emb_dim = dim)
|
||||
self.mid_attn = Residual(Rezero(LinearAttention(mid_dim)))
|
||||
self.mid_block2 = ResnetBlock(mid_dim, mid_dim, time_emb_dim = dim)
|
||||
|
||||
for ind, (dim_in, dim_out) in enumerate(reversed(in_out[1:])):
|
||||
is_last = ind >= (num_resolutions - 1)
|
||||
|
||||
self.ups.append(nn.ModuleList([
|
||||
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))),
|
||||
ResnetBlock(dim_out * 2, dim_in, time_emb_dim = dim),
|
||||
ResnetBlock(dim_in, dim_in, time_emb_dim = dim),
|
||||
Residual(Rezero(LinearAttention(dim_in))),
|
||||
Upsample(dim_in) if not is_last else nn.Identity()
|
||||
]))
|
||||
|
||||
out_dim = default(out_dim, channels)
|
||||
out_dim = default(out_dim, 3)
|
||||
self.final_conv = nn.Sequential(
|
||||
block_klass(dim, dim),
|
||||
Block(dim, dim),
|
||||
nn.Conv2d(dim, 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
|
||||
t = self.time_pos_emb(time)
|
||||
t = self.mlp(t)
|
||||
|
||||
h = []
|
||||
|
||||
for block1, block2, attn, downsample in self.downs:
|
||||
x = block1(x, t)
|
||||
x = block2(x, t)
|
||||
for resnet, resnet2, attn, downsample in self.downs:
|
||||
x = resnet(x, t)
|
||||
x = resnet2(x, t)
|
||||
x = attn(x)
|
||||
h.append(x)
|
||||
x = downsample(x)
|
||||
@@ -289,10 +242,10 @@ class Unet(nn.Module):
|
||||
x = self.mid_attn(x)
|
||||
x = self.mid_block2(x, t)
|
||||
|
||||
for block1, block2, attn, upsample in self.ups:
|
||||
for resnet, resnet2, attn, upsample in self.ups:
|
||||
x = torch.cat((x, h.pop()), dim=1)
|
||||
x = block1(x, t)
|
||||
x = block2(x, t)
|
||||
x = resnet(x, t)
|
||||
x = resnet2(x, t)
|
||||
x = attn(x)
|
||||
x = upsample(x)
|
||||
|
||||
@@ -316,62 +269,53 @@ def cosine_beta_schedule(timesteps, s = 0.008):
|
||||
as proposed in https://openreview.net/forum?id=-NEXDKk8gZ
|
||||
"""
|
||||
steps = timesteps + 1
|
||||
x = torch.linspace(0, timesteps, steps)
|
||||
alphas_cumprod = torch.cos(((x / timesteps) + s) / (1 + s) * torch.pi * 0.5) ** 2
|
||||
x = np.linspace(0, steps, steps)
|
||||
alphas_cumprod = np.cos(((x / steps) + s) / (1 + s) * np.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.999)
|
||||
return np.clip(betas, a_min = 0, a_max = 0.999)
|
||||
|
||||
class GaussianDiffusion(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
denoise_fn,
|
||||
*,
|
||||
image_size,
|
||||
channels = 3,
|
||||
timesteps = 1000,
|
||||
loss_type = 'l1'
|
||||
):
|
||||
def __init__(self, denoise_fn, timesteps=1000, loss_type='l1', betas = None):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
self.image_size = image_size
|
||||
self.denoise_fn = denoise_fn
|
||||
|
||||
betas = cosine_beta_schedule(timesteps)
|
||||
if exists(betas):
|
||||
betas = betas.detach().cpu().numpy() if isinstance(betas, torch.Tensor) else betas
|
||||
else:
|
||||
betas = cosine_beta_schedule(timesteps)
|
||||
|
||||
alphas = 1. - betas
|
||||
alphas_cumprod = torch.cumprod(alphas, axis=0)
|
||||
alphas_cumprod_prev = F.pad(alphas_cumprod[:-1], (1, 0), value = 1.)
|
||||
alphas_cumprod = np.cumprod(alphas, axis=0)
|
||||
alphas_cumprod_prev = np.append(1., alphas_cumprod[:-1])
|
||||
|
||||
timesteps, = betas.shape
|
||||
self.num_timesteps = int(timesteps)
|
||||
self.loss_type = loss_type
|
||||
|
||||
self.register_buffer('betas', betas)
|
||||
self.register_buffer('alphas_cumprod', alphas_cumprod)
|
||||
self.register_buffer('alphas_cumprod_prev', alphas_cumprod_prev)
|
||||
to_torch = partial(torch.tensor, dtype=torch.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))
|
||||
|
||||
# calculations for diffusion q(x_t | x_{t-1}) and others
|
||||
|
||||
self.register_buffer('sqrt_alphas_cumprod', torch.sqrt(alphas_cumprod))
|
||||
self.register_buffer('sqrt_one_minus_alphas_cumprod', torch.sqrt(1. - alphas_cumprod))
|
||||
self.register_buffer('log_one_minus_alphas_cumprod', torch.log(1. - alphas_cumprod))
|
||||
self.register_buffer('sqrt_recip_alphas_cumprod', torch.sqrt(1. / alphas_cumprod))
|
||||
self.register_buffer('sqrt_recipm1_alphas_cumprod', torch.sqrt(1. / alphas_cumprod - 1))
|
||||
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)))
|
||||
|
||||
# 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', posterior_variance)
|
||||
|
||||
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', torch.log(posterior_variance.clamp(min =1e-20)))
|
||||
self.register_buffer('posterior_mean_coef1', betas * torch.sqrt(alphas_cumprod_prev) / (1. - alphas_cumprod))
|
||||
self.register_buffer('posterior_mean_coef2', (1. - alphas_cumprod_prev) * torch.sqrt(alphas) / (1. - alphas_cumprod))
|
||||
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
|
||||
@@ -424,10 +368,8 @@ class GaussianDiffusion(nn.Module):
|
||||
return img
|
||||
|
||||
@torch.no_grad()
|
||||
def sample(self, batch_size = 16):
|
||||
image_size = self.image_size
|
||||
channels = self.channels
|
||||
return self.p_sample_loop((batch_size, channels, image_size, image_size))
|
||||
def sample(self, image_size, batch_size = 16):
|
||||
return self.p_sample_loop((batch_size, 3, image_size, image_size))
|
||||
|
||||
@torch.no_grad()
|
||||
def interpolate(self, x1, x2, t = None, lam = 0.5):
|
||||
@@ -470,26 +412,24 @@ class GaussianDiffusion(nn.Module):
|
||||
return loss
|
||||
|
||||
def forward(self, x, *args, **kwargs):
|
||||
b, c, h, w, device, img_size, = *x.shape, x.device, self.image_size
|
||||
assert h == img_size and w == img_size, f'height and width of image must be {img_size}'
|
||||
b, *_, device = *x.shape, x.device
|
||||
t = torch.randint(0, self.num_timesteps, (b,), device=device).long()
|
||||
return self.p_losses(x, t, *args, **kwargs)
|
||||
|
||||
# dataset classes
|
||||
|
||||
class Dataset(data.Dataset):
|
||||
def __init__(self, folder, image_size, exts = ['jpg', 'jpeg', 'png']):
|
||||
def __init__(self, folder, image_size):
|
||||
super().__init__()
|
||||
self.folder = folder
|
||||
self.image_size = image_size
|
||||
self.paths = [p for ext in exts for p in Path(f'{folder}').glob(f'**/*.{ext}')]
|
||||
self.paths = [p for ext in EXTS for p in Path(f'{folder}').glob(f'**/*.{ext}')]
|
||||
|
||||
self.transform = transforms.Compose([
|
||||
transforms.Resize(image_size),
|
||||
transforms.RandomHorizontalFlip(),
|
||||
transforms.CenterCrop(image_size),
|
||||
transforms.ToTensor(),
|
||||
transforms.Lambda(lambda t: (t * 2) - 1)
|
||||
transforms.ToTensor()
|
||||
])
|
||||
|
||||
def __len__(self):
|
||||
@@ -514,23 +454,17 @@ class Trainer(object):
|
||||
train_lr = 2e-5,
|
||||
train_num_steps = 100000,
|
||||
gradient_accumulate_every = 2,
|
||||
amp = False,
|
||||
step_start_ema = 2000,
|
||||
update_ema_every = 10,
|
||||
save_and_sample_every = 1000,
|
||||
results_folder = './results'
|
||||
fp16 = False,
|
||||
step_start_ema = 2000
|
||||
):
|
||||
super().__init__()
|
||||
self.model = diffusion_model
|
||||
self.ema = EMA(ema_decay)
|
||||
self.ema_model = copy.deepcopy(self.model)
|
||||
self.update_ema_every = update_ema_every
|
||||
|
||||
self.step_start_ema = step_start_ema
|
||||
self.save_and_sample_every = save_and_sample_every
|
||||
|
||||
self.batch_size = train_batch_size
|
||||
self.image_size = diffusion_model.image_size
|
||||
self.image_size = image_size
|
||||
self.gradient_accumulate_every = gradient_accumulate_every
|
||||
self.train_num_steps = train_num_steps
|
||||
|
||||
@@ -540,11 +474,11 @@ class Trainer(object):
|
||||
|
||||
self.step = 0
|
||||
|
||||
self.amp = amp
|
||||
self.scaler = GradScaler(enabled = amp)
|
||||
assert not fp16 or fp16 and APEX_AVAILABLE, 'Apex must be installed in order for mixed precision training to be turned on'
|
||||
|
||||
self.results_folder = Path(results_folder)
|
||||
self.results_folder.mkdir(exist_ok = True)
|
||||
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.reset_parameters()
|
||||
|
||||
@@ -561,44 +495,39 @@ class Trainer(object):
|
||||
data = {
|
||||
'step': self.step,
|
||||
'model': self.model.state_dict(),
|
||||
'ema': self.ema_model.state_dict(),
|
||||
'scaler': self.scaler.state_dict()
|
||||
'ema': self.ema_model.state_dict()
|
||||
}
|
||||
torch.save(data, str(self.results_folder / f'model-{milestone}.pt'))
|
||||
torch.save(data, f'./model-{milestone}.pt')
|
||||
|
||||
def load(self, milestone):
|
||||
data = torch.load(str(self.results_folder / f'model-{milestone}.pt'))
|
||||
data = torch.load(f'./model-{milestone}.pt')
|
||||
|
||||
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()
|
||||
|
||||
with autocast(enabled = self.amp):
|
||||
loss = self.model(data)
|
||||
self.scaler.scale(loss / self.gradient_accumulate_every).backward()
|
||||
|
||||
loss = self.model(data)
|
||||
print(f'{self.step}: {loss.item()}')
|
||||
backwards(loss / self.gradient_accumulate_every, self.opt)
|
||||
|
||||
self.scaler.step(self.opt)
|
||||
self.scaler.update()
|
||||
self.opt.step()
|
||||
self.opt.zero_grad()
|
||||
|
||||
if self.step % self.update_ema_every == 0:
|
||||
if self.step % UPDATE_EMA_EVERY == 0:
|
||||
self.step_ema()
|
||||
|
||||
if self.step != 0 and self.step % self.save_and_sample_every == 0:
|
||||
milestone = self.step // self.save_and_sample_every
|
||||
if self.step != 0 and self.step % SAVE_AND_SAMPLE_EVERY == 0:
|
||||
milestone = self.step // 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_list = list(map(lambda n: self.ema_model.sample(self.image_size, batch_size=n), batches))
|
||||
all_images = torch.cat(all_images_list, dim=0)
|
||||
all_images = (all_images + 1) * 0.5
|
||||
utils.save_image(all_images, str(self.results_folder / f'sample-{milestone}.png'), nrow = 6)
|
||||
utils.save_image(all_images, f'./sample-{milestone}.png', nrow=6)
|
||||
self.save(milestone)
|
||||
|
||||
self.step += 1
|
||||
|
||||
BIN
Binary file not shown.
|
Before Width: | Height: | Size: 842 KiB After Width: | Height: | Size: 1.3 MiB |
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
|
||||
setup(
|
||||
name = 'denoising-diffusion-pytorch',
|
||||
packages = find_packages(),
|
||||
version = '0.12.0',
|
||||
version = '0.5.0',
|
||||
license='MIT',
|
||||
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
|
||||
author = 'Phil Wang',
|
||||
@@ -15,6 +15,7 @@ setup(
|
||||
],
|
||||
install_requires=[
|
||||
'einops',
|
||||
'numpy',
|
||||
'pillow',
|
||||
'torch',
|
||||
'torchvision',
|
||||
|
||||
Reference in New Issue
Block a user