Compare commits

...
30 Commits
Author SHA1 Message Date
Phil Wang 402b7c26df calculate noise schedule with float64 for numerical accuracy 2022-05-10 15:23:34 -07:00
Phil Wang 09613a40f3 cleanup 2022-05-07 05:47:21 -07:00
Phil Wang c6966ae95a Merge pull request #24 from kashif/patch-1
updated citation in README
2022-05-07 05:32:51 -07:00
Kashif Rasul 73591cf1ad updated citation in README 2022-05-07 11:23:45 +02:00
Phil Wang 989f0fcb8e remove convnext blocks, they do not work well, validated in video diffusion repository 2022-05-05 07:03:55 -07:00
Phil Wang 84731bb03d groupnorm groups should be actually configurable 2022-05-04 10:38:29 -07:00
Phil Wang c6ecca555b allow for configuring expansion factor in convnext 2022-05-04 10:33:23 -07:00
Phil Wang 1f5c233072 bring back resnet blocks, make convnext blocks an experimental option 2022-05-04 10:30:09 -07:00
Phil Wang de378158e5 readme 2022-05-01 13:16:06 -07:00
Phil Wang e274fb305a give an initial conv 2022-05-01 08:49:38 -07:00
Phil Wang f39b3b1d3f make sure time embedding dimension is kept at 4 x dimension (thanks @borisdayma) 2022-04-29 14:55:12 -07:00
Phil Wang 782c904d3b fix cosine beta schedule, thanks to @Zhengxinyang 2022-04-19 20:51:50 -07:00
Phil Wang 71953ebd22 fix bug, thanks to @jihoonerd 2022-04-15 06:37:31 -07:00
Phil Wang 0b8cdb4c8b remove outdated apex in favor of native pytorch AMP 2022-04-13 08:59:18 -07:00
Phil Wang e504e0e554 cleanup 2022-04-12 13:02:18 -07:00
Phil Wang bd1e3b676e get rid of numpy 2022-04-12 11:58:46 -07:00
Phil Wang f4615599bc use full attention at the center of the unet 2022-04-04 09:03:41 -07:00
Phil Wang eb6e1b508e greater kernel size in convnext blocks 2022-01-31 17:13:27 -08:00
Phil Wang 91cff45939 replace resnets with convnext blocks 2022-01-25 09:02:45 -08:00
Phil Wang 7b51e30da7 fix layernorm 2021-08-24 14:28:15 -07:00
Phil Wang dadbf20154 remove stray print 2021-07-16 15:18:20 -07:00
Phil Wang 7706bdfc6f use pre-layernorm with linear attention, and also allow for turning off time embedding 2021-06-25 10:48:11 -07:00
Phil Wang 183e5f3cc5 move all constants into configurable class init parameters 2021-06-25 10:37:36 -07:00
Phil Wang 16c9ae7bb3 fix data not being normalized to range of -1 to 1 2021-06-21 18:51:36 -07:00
Phil Wang f5916111f8 0.6.3 2021-06-21 17:41:16 -07:00
Phil Wang ad9e303ff3 fix channels 2021-06-11 15:28:36 -07:00
Phil Wang ae42f48f6a prepare so that unet can work with a channel of one, and also make it so image size is hard coded in diffusion class. preparing for training on protein distograms 2021-06-11 14:06:39 -07:00
Phil Wang 5989f4c77e recommit sample 2020-10-13 09:32:56 -07:00
Phil Wang 2082046888 set higher num train steps, so non-practitioners do not think it is completed 2020-10-11 13:49:49 -07:00
Phil Wang 3c5b7e2d56 update readme 2020-10-10 10:37:44 -07:00
5 changed files with 246 additions and 159 deletions
+3
View File
@@ -1,3 +1,6 @@
# Generation results
results/
# Byte-compiled / optimized / DLL files
__pycache__/
*.py[cod]
+34 -21
View File
@@ -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>
@@ -27,16 +29,17 @@ model = Unet(
diffusion = GaussianDiffusion(
model,
image_size = 128,
timesteps = 1000, # number of steps
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
sampled_images = diffusion.sample(128, batch_size = 4)
sampled_images = diffusion.sample(batch_size = 4)
sampled_images.shape # (4, 3, 128, 128)
```
@@ -52,6 +55,7 @@ model = Unet(
diffusion = GaussianDiffusion(
model,
image_size = 128,
timesteps = 1000, # number of steps
loss_type = 'l1' # L1 or L2
).cuda()
@@ -59,39 +63,48 @@ diffusion = GaussianDiffusion(
trainer = Trainer(
diffusion,
'path/to/your/images',
image_size = 128,
train_batch_size = 32,
train_lr = 2e-5,
train_num_steps = 100000, # total training steps
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()
```
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},
}
```
@@ -7,30 +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
# constants
SAVE_AND_SAMPLE_EVERY = 1000
UPDATE_EMA_EVERY = 10
EXTS = ['jpg', 'jpeg', 'png']
RESULTS_FOLDER = Path('./results')
RESULTS_FOLDER.mkdir(exist_ok = True)
# helpers functions
def exists(x):
@@ -54,13 +40,6 @@ 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():
@@ -100,34 +79,33 @@ 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):
def Downsample(dim):
return nn.Conv2d(dim, dim, 4, 2, 1)
class LayerNorm(nn.Module):
def __init__(self, dim, eps = 1e-5):
super().__init__()
self.conv = nn.ConvTranspose2d(dim, dim, 4, 2, 1)
self.eps = eps
self.g = nn.Parameter(torch.ones(1, dim, 1, 1))
self.b = nn.Parameter(torch.zeros(1, dim, 1, 1))
def forward(self, x):
return self.conv(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
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):
class PreNorm(nn.Module):
def __init__(self, dim, fn):
super().__init__()
self.fn = fn
self.g = nn.Parameter(torch.zeros(1))
self.norm = LayerNorm(dim)
def forward(self, x):
return self.fn(x) * self.g
x = self.norm(x)
return self.fn(x)
# building block modules
@@ -135,34 +113,67 @@ 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)
class ResnetBlock(nn.Module):
def __init__(self, dim, dim_out, *, time_emb_dim, groups = 8):
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)
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)
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)
@@ -170,28 +181,60 @@ 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
class Unet(nn.Module):
def __init__(self, dim, out_dim = None, dim_mults=(1, 2, 4, 8), groups = 8):
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
):
super().__init__()
dims = [3, *map(lambda m: dim * m, dim_mults)]
# 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)]
in_out = list(zip(dims[:-1], dims[1:]))
self.time_pos_emb = SinusoidalPosEmb(dim)
self.mlp = nn.Sequential(
nn.Linear(dim, dim * 4),
Mish(),
nn.Linear(dim * 4, dim)
)
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.downs = nn.ModuleList([])
self.ups = nn.ModuleList([])
@@ -201,42 +244,43 @@ class Unet(nn.Module):
is_last = ind >= (num_resolutions - 1)
self.downs.append(nn.ModuleList([
ResnetBlock(dim_in, dim_out, time_emb_dim = dim),
ResnetBlock(dim_out, dim_out, time_emb_dim = dim),
Residual(Rezero(LinearAttention(dim_out))),
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 = dim)
self.mid_attn = Residual(Rezero(LinearAttention(mid_dim)))
self.mid_block2 = ResnetBlock(mid_dim, mid_dim, time_emb_dim = 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 = dim),
ResnetBlock(dim_in, dim_in, time_emb_dim = dim),
Residual(Rezero(LinearAttention(dim_in))),
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, 3)
out_dim = default(out_dim, channels)
self.final_conv = nn.Sequential(
Block(dim, dim),
block_klass(dim, dim),
nn.Conv2d(dim, out_dim, 1)
)
def forward(self, x, time):
t = self.time_pos_emb(time)
t = self.mlp(t)
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)
@@ -245,10 +289,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)
@@ -272,53 +316,66 @@ 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.9999)
class GaussianDiffusion(nn.Module):
def __init__(self, denoise_fn, timesteps=1000, loss_type='l1', betas = None):
def __init__(
self,
denoise_fn,
*,
image_size,
channels = 3,
timesteps = 1000,
loss_type = 'l1'
):
super().__init__()
self.channels = channels
self.image_size = image_size
self.denoise_fn = denoise_fn
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))
register_buffer('posterior_variance', 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)))
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 q_mean_variance(self, x_start, t):
mean = extract(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start
@@ -371,8 +428,10 @@ class GaussianDiffusion(nn.Module):
return img
@torch.no_grad()
def sample(self, image_size, batch_size = 16):
return self.p_sample_loop((batch_size, 3, image_size, image_size))
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))
@torch.no_grad()
def interpolate(self, x1, x2, t = None, lam = 0.5):
@@ -415,24 +474,26 @@ class GaussianDiffusion(nn.Module):
return loss
def forward(self, x, *args, **kwargs):
b, *_, device = *x.shape, x.device
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}'
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):
def __init__(self, folder, image_size, exts = ['jpg', 'jpeg', 'png']):
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.ToTensor(),
transforms.Lambda(lambda t: (t * 2) - 1)
])
def __len__(self):
@@ -457,17 +518,23 @@ class Trainer(object):
train_lr = 2e-5,
train_num_steps = 100000,
gradient_accumulate_every = 2,
fp16 = False,
step_start_ema = 2000
amp = False,
step_start_ema = 2000,
update_ema_every = 10,
save_and_sample_every = 1000,
results_folder = './results'
):
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 = image_size
self.image_size = diffusion_model.image_size
self.gradient_accumulate_every = gradient_accumulate_every
self.train_num_steps = train_num_steps
@@ -477,11 +544,11 @@ 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.amp = amp
self.scaler = GradScaler(enabled = amp)
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.results_folder = Path(results_folder)
self.results_folder.mkdir(exist_ok = True)
self.reset_parameters()
@@ -498,39 +565,44 @@ 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(RESULTS_FOLDER / f'model-{milestone}.pt'))
torch.save(data, str(self.results_folder / f'model-{milestone}.pt'))
def load(self, milestone):
data = torch.load(str(RESULTS_FOLDER / f'model-{milestone}.pt'))
data = torch.load(str(self.results_folder / 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()
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 % UPDATE_EMA_EVERY == 0:
if self.step % self.update_ema_every == 0:
self.step_ema()
if self.step != 0 and self.step % SAVE_AND_SAMPLE_EVERY == 0:
milestone = self.step // SAVE_AND_SAMPLE_EVERY
if self.step != 0 and self.step % self.save_and_sample_every == 0:
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(self.image_size, batch_size=n), batches))
all_images_list = list(map(lambda n: self.ema_model.sample(batch_size=n), batches))
all_images = torch.cat(all_images_list, dim=0)
utils.save_image(all_images, str(RESULTS_FOLDER / f'sample-{milestone}.png'), nrow=6)
all_images = (all_images + 1) * 0.5
utils.save_image(all_images, str(self.results_folder / f'sample-{milestone}.png'), nrow = 6)
self.save(milestone)
self.step += 1
BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 842 KiB

+1 -2
View File
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
setup(
name = 'denoising-diffusion-pytorch',
packages = find_packages(),
version = '0.5.2',
version = '0.12.1',
license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang',
@@ -15,7 +15,6 @@ setup(
],
install_requires=[
'einops',
'numpy',
'pillow',
'torch',
'torchvision',