mirror of
https://github.com/wassname/denoising-diffusion-pytorch.git
synced 2026-09-10 12:01:08 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
91cff45939 | ||
|
|
7b51e30da7 | ||
|
|
dadbf20154 | ||
|
|
7706bdfc6f | ||
|
|
183e5f3cc5 | ||
|
|
16c9ae7bb3 | ||
|
|
f5916111f8 | ||
|
|
ad9e303ff3 | ||
|
|
ae42f48f6a | ||
|
|
5989f4c77e | ||
|
|
2082046888 | ||
|
|
3c5b7e2d56 | ||
|
|
d4ce9f6c38 | ||
|
|
ff451f697e | ||
|
|
3d96532c60 | ||
|
|
ef2ca0b625 | ||
|
|
9f95a03c07 | ||
|
|
a4c68d3569 | ||
|
|
b33a48e342 | ||
|
|
8e5fb17063 |
@@ -1,3 +1,6 @@
|
||||
# Generation results
|
||||
results/
|
||||
|
||||
# Byte-compiled / optimized / DLL files
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
|
||||
@@ -2,10 +2,14 @@
|
||||
|
||||
## 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> and then modified to use <a href="https://arxiv.org/abs/2201.03545">ConvNext</a> blocks instead of Resnets.
|
||||
|
||||
<img src="./sample.png" width="500px"><img>
|
||||
|
||||
[](https://badge.fury.io/py/denoising-diffusion-pytorch)
|
||||
|
||||
## Install
|
||||
|
||||
```bash
|
||||
@@ -25,10 +29,9 @@ model = Unet(
|
||||
|
||||
diffusion = GaussianDiffusion(
|
||||
model,
|
||||
beta_start = 0.0001,
|
||||
beta_end = 0.02,
|
||||
num_diffusion_timesteps = 1000, # number of steps
|
||||
loss_type = 'l1' # L1 or L2 (wavegrad paper claims l1 is better?)
|
||||
image_size = 128,
|
||||
timesteps = 1000, # number of steps
|
||||
loss_type = 'l1' # L1 or L2
|
||||
)
|
||||
|
||||
training_images = torch.randn(8, 3, 128, 128)
|
||||
@@ -36,7 +39,7 @@ 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,19 +55,17 @@ model = Unet(
|
||||
|
||||
diffusion = GaussianDiffusion(
|
||||
model,
|
||||
beta_start = 0.0001,
|
||||
beta_end = 0.02,
|
||||
num_diffusion_timesteps = 1000, # number of steps
|
||||
loss_type = 'l1' # L1 or L2
|
||||
image_size = 128,
|
||||
timesteps = 1000, # number of steps
|
||||
loss_type = 'l1' # L1 or L2
|
||||
).cuda()
|
||||
|
||||
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
|
||||
@@ -73,17 +74,39 @@ trainer = Trainer(
|
||||
trainer.train()
|
||||
```
|
||||
|
||||
Todo: Command line tool for one-line training
|
||||
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}
|
||||
}
|
||||
```
|
||||
|
||||
```bibtex
|
||||
@misc{liu2022convnet,
|
||||
title = {A ConvNet for the 2020s},
|
||||
author = {Zhuang Liu and Hanzi Mao and Chao-Yuan Wu and Christoph Feichtenhofer and Trevor Darrell and Saining Xie},
|
||||
year = {2022},
|
||||
eprint = {2201.03545},
|
||||
archivePrefix = {arXiv},
|
||||
primaryClass = {cs.CV}
|
||||
}
|
||||
```
|
||||
|
||||
@@ -22,12 +22,6 @@ try:
|
||||
except:
|
||||
APEX_AVAILABLE = False
|
||||
|
||||
# constants
|
||||
|
||||
SAVE_AND_SAMPLE_EVERY = 1000
|
||||
UPDATE_EMA_EVERY = 10
|
||||
EXTS = ['jpg', 'png']
|
||||
|
||||
# helpers functions
|
||||
|
||||
def exists(x):
|
||||
@@ -43,6 +37,14 @@ def cycle(dl):
|
||||
for data in dl:
|
||||
yield data
|
||||
|
||||
def num_to_groups(num, divisor):
|
||||
groups = num // divisor
|
||||
remainder = num % divisor
|
||||
arr = [divisor] * groups
|
||||
if remainder > 0:
|
||||
arr.append(remainder)
|
||||
return arr
|
||||
|
||||
def loss_backwards(fp16, loss, optimizer, **kwargs):
|
||||
if fp16:
|
||||
with amp.scale_loss(loss, optimizer) as scaled_loss:
|
||||
@@ -89,98 +91,119 @@ 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, 3, 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):
|
||||
class PreNorm(nn.Module):
|
||||
def __init__(self, dim, fn):
|
||||
super().__init__()
|
||||
self.conv = nn.Conv2d(dim, dim, 3, 2, 1)
|
||||
self.fn = fn
|
||||
self.norm = LayerNorm(dim)
|
||||
|
||||
def forward(self, x):
|
||||
return self.conv(x)
|
||||
|
||||
class Rezero(nn.Module):
|
||||
def __init__(self, dim):
|
||||
super().__init__()
|
||||
self.g = nn.Parameter(torch.zeros(1))
|
||||
|
||||
def forward(self, x):
|
||||
return x * self.g
|
||||
x = self.norm(x)
|
||||
return self.fn(x)
|
||||
|
||||
# building block modules
|
||||
|
||||
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.GroupNorm(groups, dim_out),
|
||||
Mish()
|
||||
)
|
||||
def forward(self, x):
|
||||
return self.block(x)
|
||||
class ConvNextBlock(nn.Module):
|
||||
""" https://arxiv.org/abs/2201.03545 """
|
||||
|
||||
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, mult = 2, norm = True):
|
||||
super().__init__()
|
||||
self.mlp = nn.Sequential(
|
||||
Mish(),
|
||||
nn.Linear(time_emb_dim, dim_out)
|
||||
nn.GELU(),
|
||||
nn.Linear(time_emb_dim, dim)
|
||||
) if exists(time_emb_dim) else None
|
||||
|
||||
self.ds_conv = nn.Conv2d(dim, dim, 7, padding = 3, groups = dim)
|
||||
|
||||
self.net = nn.Sequential(
|
||||
LayerNorm(dim) if norm else nn.Identity(),
|
||||
nn.Conv2d(dim, dim_out * mult, 1),
|
||||
nn.GELU(),
|
||||
LayerNorm(dim_out * mult),
|
||||
nn.Conv2d(dim_out * mult, dim_out, 1)
|
||||
)
|
||||
|
||||
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):
|
||||
h = self.block1(x)
|
||||
h += self.mlp(time_emb)[:, :, None, None]
|
||||
h = self.block2(h)
|
||||
def forward(self, x, time_emb = None):
|
||||
h = self.ds_conv(x)
|
||||
|
||||
if exists(self.mlp):
|
||||
assert exists(time_emb), 'time emb must be passed in'
|
||||
condition = self.mlp(time_emb)
|
||||
h = h + rearrange(condition, 'b c -> b c 1 1')
|
||||
|
||||
h = self.net(h)
|
||||
return h + self.res_conv(x)
|
||||
|
||||
class LinearAttention(nn.Module):
|
||||
def __init__(self, dim, heads = 8, dim_head = 32):
|
||||
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, 1, bias = False)
|
||||
self.to_qkv = nn.Conv2d(dim, hidden_dim * 3, 1, bias = False)
|
||||
self.to_out = nn.Conv2d(hidden_dim, dim, 1)
|
||||
|
||||
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)
|
||||
q = q.softmax(dim=-2)
|
||||
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
|
||||
|
||||
k = k.softmax(dim = -1)
|
||||
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)
|
||||
|
||||
# model
|
||||
|
||||
class Unet(nn.Module):
|
||||
def __init__(self, dim, out_dim = None, dim_mults=(1, 2, 4, 8), groups = 8):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
out_dim = None,
|
||||
dim_mults=(1, 2, 4, 8),
|
||||
channels = 3,
|
||||
with_time_emb = True
|
||||
):
|
||||
super().__init__()
|
||||
dims = [3, *map(lambda m: dim * m, dim_mults)]
|
||||
self.channels = channels
|
||||
|
||||
dims = [channels, *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)
|
||||
)
|
||||
if with_time_emb:
|
||||
time_dim = dim
|
||||
self.time_mlp = nn.Sequential(
|
||||
SinusoidalPosEmb(dim),
|
||||
nn.Linear(dim, dim * 4),
|
||||
nn.GELU(),
|
||||
nn.Linear(dim * 4, dim)
|
||||
)
|
||||
else:
|
||||
time_dim = None
|
||||
self.time_mlp = None
|
||||
|
||||
self.downs = nn.ModuleList([])
|
||||
self.ups = nn.ModuleList([])
|
||||
@@ -190,42 +213,41 @@ 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))),
|
||||
ConvNextBlock(dim_in, dim_out, time_emb_dim = time_dim, norm = ind != 0),
|
||||
ConvNextBlock(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 = ConvNextBlock(mid_dim, mid_dim, time_emb_dim = time_dim)
|
||||
self.mid_attn = Residual(PreNorm(mid_dim, LinearAttention(mid_dim)))
|
||||
self.mid_block2 = ConvNextBlock(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))),
|
||||
ConvNextBlock(dim_out * 2, dim_in, time_emb_dim = time_dim),
|
||||
ConvNextBlock(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),
|
||||
ConvNextBlock(dim, dim),
|
||||
nn.Conv2d(dim, out_dim, 1)
|
||||
)
|
||||
|
||||
def forward(self, x, time):
|
||||
t = self.time_pos_emb(time)
|
||||
t = self.mlp(t)
|
||||
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 convnext, convnext2, attn, downsample in self.downs:
|
||||
x = convnext(x, t)
|
||||
x = convnext2(x, t)
|
||||
x = attn(x)
|
||||
h.append(x)
|
||||
x = downsample(x)
|
||||
@@ -234,10 +256,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 convnext, convnext2, attn, upsample in self.ups:
|
||||
x = torch.cat((x, h.pop()), dim=1)
|
||||
x = resnet(x, t)
|
||||
x = resnet2(x, t)
|
||||
x = convnext(x, t)
|
||||
x = convnext2(x, t)
|
||||
x = attn(x)
|
||||
x = upsample(x)
|
||||
|
||||
@@ -255,24 +277,47 @@ def noise_like(shape, device, repeat=False):
|
||||
noise = lambda: torch.randn(shape, device=device)
|
||||
return repeat_noise() if repeat else noise()
|
||||
|
||||
def cosine_beta_schedule(timesteps, s = 0.008):
|
||||
"""
|
||||
cosine schedule
|
||||
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
|
||||
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)
|
||||
|
||||
class GaussianDiffusion(nn.Module):
|
||||
def __init__(self, denoise_fn, beta_start=0.0001, beta_end=0.02, num_diffusion_timesteps=1000, loss_type='l1', betas = None):
|
||||
def __init__(
|
||||
self,
|
||||
denoise_fn,
|
||||
*,
|
||||
image_size,
|
||||
channels = 3,
|
||||
timesteps = 1000,
|
||||
loss_type = 'l1',
|
||||
betas = None
|
||||
):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
self.image_size = image_size
|
||||
self.denoise_fn = denoise_fn
|
||||
|
||||
if exists(betas):
|
||||
self.np_betas = betas.detach().cpu().numpy() if isinstance(betas, torch.Tensor) else betas
|
||||
betas = betas.detach().cpu().numpy() if isinstance(betas, torch.Tensor) else betas
|
||||
else:
|
||||
self.np_betas = betas = np.linspace(beta_start, beta_end, num_diffusion_timesteps).astype(np.float64)
|
||||
|
||||
timesteps, = betas.shape
|
||||
self.num_timesteps = int(timesteps)
|
||||
self.loss_type = loss_type
|
||||
betas = cosine_beta_schedule(timesteps)
|
||||
|
||||
alphas = 1. - betas
|
||||
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
|
||||
|
||||
to_torch = partial(torch.tensor, dtype=torch.float32)
|
||||
|
||||
self.register_buffer('betas', to_torch(betas))
|
||||
@@ -348,8 +393,10 @@ class GaussianDiffusion(nn.Module):
|
||||
return img
|
||||
|
||||
@torch.no_grad()
|
||||
def sample(self, image_size, batch_size = 16):
|
||||
return self.p_sample_loop((16, 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):
|
||||
@@ -392,24 +439,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):
|
||||
@@ -434,14 +483,23 @@ class Trainer(object):
|
||||
train_lr = 2e-5,
|
||||
train_num_steps = 100000,
|
||||
gradient_accumulate_every = 2,
|
||||
fp16 = False
|
||||
fp16 = 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.image_size = image_size
|
||||
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.gradient_accumulate_every = gradient_accumulate_every
|
||||
self.train_num_steps = train_num_steps
|
||||
|
||||
@@ -457,13 +515,16 @@ class Trainer(object):
|
||||
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()
|
||||
|
||||
def reset_parameters(self):
|
||||
self.ema_model.load_state_dict(self.model.state_dict())
|
||||
|
||||
def step_ema(self):
|
||||
if self.step < 2000:
|
||||
if self.step < self.step_start_ema:
|
||||
self.reset_parameters()
|
||||
return
|
||||
self.ema.update_model_average(self.ema_model, self.model)
|
||||
@@ -474,10 +535,10 @@ class Trainer(object):
|
||||
'model': self.model.state_dict(),
|
||||
'ema': self.ema_model.state_dict()
|
||||
}
|
||||
torch.save(data, f'./model-{milestone}.pt')
|
||||
torch.save(data, str(self.results_folder / f'model-{milestone}.pt'))
|
||||
|
||||
def load(self, milestone):
|
||||
data = torch.load(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'])
|
||||
@@ -496,13 +557,16 @@ class Trainer(object):
|
||||
self.opt.step()
|
||||
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 % SAVE_AND_SAMPLE_EVERY == 0:
|
||||
milestone = self.step // SAVE_AND_SAMPLE_EVERY
|
||||
all_images = self.ema_model.p_sample_loop((64, 3, self.image_size, self.image_size))
|
||||
utils.save_image(all_images, f'./sample-{milestone}.png', nrow=8)
|
||||
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(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)
|
||||
self.save(milestone)
|
||||
|
||||
self.step += 1
|
||||
|
||||
BIN
Binary file not shown.
|
Before Width: | Height: | Size: 1.3 MiB After Width: | Height: | Size: 842 KiB |
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
|
||||
setup(
|
||||
name = 'denoising-diffusion-pytorch',
|
||||
packages = find_packages(),
|
||||
version = '0.2.4',
|
||||
version = '0.7.0',
|
||||
license='MIT',
|
||||
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
|
||||
author = 'Phil Wang',
|
||||
|
||||
Reference in New Issue
Block a user