Compare commits

..
21 Commits
Author SHA1 Message Date
Phil Wang ef2ca0b625 new paper suggests image linear attention is more effective without query normalization 2020-10-04 21:53:53 -07:00
Phil Wang 9f95a03c07 fix bug with rezero and linear attention 2020-09-21 20:11:34 -07:00
Phil Wang a4c68d3569 fix bug 2020-09-15 15:57:39 -07:00
Phil Wang b33a48e342 make sure when sampling, batch does not exceed training batch size 2020-09-15 15:15:10 -07:00
Phil Wang 8e5fb17063 add badge 2020-09-14 13:38:06 -07:00
Phil Wang 4bf28914bc allow for mixed precision training with fp16 flag 2020-09-08 17:26:23 -07:00
Phil Wang 88f83d0ff2 fix loading checkpoint 2020-09-08 15:15:59 -07:00
Phil Wang 26b5cab6c8 update ema more frequently 2020-09-08 15:03:38 -07:00
Phil Wang 1307b3115d add exponential moving average of model, also allow passing in custom noise schedule 2020-09-08 14:58:09 -07:00
Phil Wang 81fb2a0386 add sample 2020-09-08 10:17:48 -07:00
Phil Wang 698227ae13 fix small bug with gradient accumulation 2020-09-07 23:14:38 -07:00
Phil Wang c479adf960 add interpolation 2020-09-07 22:38:59 -07:00
Phil Wang d70fb08f8a small helper fn to make sampling more clear 2020-09-07 16:40:36 -07:00
Phil Wang e700a7c6de add image back 2020-09-07 10:37:58 -07:00
Phil Wang 9c758662a3 fix another stray bug 2020-09-06 23:25:30 -07:00
Phil Wang 1f1e42e9f9 update readme 2020-09-06 22:35:37 -07:00
Phil Wang e1800c1a8d remove wip, seems to be working 2020-09-06 14:26:23 -07:00
Phil Wang 11f27032ba offer training class to easily train model off an image directory 2020-09-06 14:22:44 -07:00
Phil Wang d8472a6220 update readme 2020-09-06 03:44:25 -07:00
Phil Wang d59d8b05f6 fix a small bug 2020-09-06 01:50:59 -07:00
Phil Wang 365ce0c335 update readme 2020-09-06 01:44:05 -07:00
6 changed files with 263 additions and 25 deletions
+44 -5
View File
@@ -1,6 +1,12 @@
## Denoising Diffusion Probabilistic Model, in Pytorch (wip)
<img src="./denoising-diffusion.png" width="500px"></img>
Implementation of <a href="https://arxiv.org/abs/2006.11239">Denoising Diffusion Probabilistic Model</a> in Pytorch
## 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>.
<img src="./sample.png" width="500px"><img>
[![PyPI version](https://badge.fury.io/py/denoising-diffusion-pytorch.svg)](https://badge.fury.io/py/denoising-diffusion-pytorch)
## Install
@@ -24,7 +30,7 @@ diffusion = GaussianDiffusion(
beta_start = 0.0001,
beta_end = 0.02,
num_diffusion_timesteps = 1000, # number of steps
loss_type = 'l1' # L1 or L2
loss_type = 'l1' # L1 or L2 (wavegrad paper claims l1 is better?)
)
training_images = torch.randn(8, 3, 128, 128)
@@ -32,8 +38,41 @@ loss = diffusion(training_images)
loss.backward()
# after a lot of training
sampled_images = diffusion.p_sample_loop((1, 3, 128, 128))
sampled_images.shape # (1, 3, 128, 128)
sampled_images = diffusion.sample(128, batch_size = 4)
sampled_images.shape # (4, 3, 128, 128)
```
Or, if you simply want to pass in a folder name and the desired image dimensions, you can use the `Trainer` class to easily train a model.
```python
from denoising_diffusion_pytorch import Unet, GaussianDiffusion, Trainer
model = Unet(
dim = 64,
dim_mults = (1, 2, 4, 8)
).cuda()
diffusion = GaussianDiffusion(
model,
beta_start = 0.0001,
beta_end = 0.02,
num_diffusion_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
gradient_accumulate_every = 2, # gradient accumulation steps
ema_decay = 0.995, # exponential moving average decay
fp16 = True # turn on mixed precision training with apex
)
trainer.train()
```
## Citations
Binary file not shown.

After

Width:  |  Height:  |  Size: 40 KiB

+1 -1
View File
@@ -1 +1 @@
from denoising_diffusion_pytorch.denoising_diffusion_pytorch import GaussianDiffusion, Unet
from denoising_diffusion_pytorch.denoising_diffusion_pytorch import GaussianDiffusion, Unet, Trainer
@@ -1,14 +1,33 @@
import math
import copy
import torch
from inspect import isfunction
from functools import partial
from torch import nn, einsum
import torch.nn.functional as F
from inspect import isfunction
from functools import partial
from torch.utils import data
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', 'png']
# helpers functions
def exists(x):
@@ -19,11 +38,43 @@ def default(val, d):
return val
return d() if isfunction(d) else d
def normal_kl(mean1, logvar1, mean2, logvar2):
return 0.5 * (-1. + logvar2 - logvar1 + torch.exp(logvar1 - logvar2) + torch.exp(-logvar2) * (mean1 - mean2) ** 2)
def cycle(dl):
while True:
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:
scaled_loss.backward(**kwargs)
else:
loss.backward(**kwargs)
# small helper modules
class EMA():
def __init__(self, beta):
super().__init__()
self.beta = beta
def update_model_average(self, ma_model, current_model):
for current_params, ma_params in zip(current_model.parameters(), ma_model.parameters()):
old_weight, up_weight = ma_params.data, current_params.data
ma_params.data = self.update_average(old_weight, up_weight)
def update_average(self, old, new):
if old is None:
return new
return old * self.beta + (1 - self.beta) * new
class Residual(nn.Module):
def __init__(self, fn):
super().__init__()
@@ -67,17 +118,18 @@ class Downsample(nn.Module):
return self.conv(x)
class Rezero(nn.Module):
def __init__(self, dim):
def __init__(self, fn):
super().__init__()
self.fn = fn
self.g = nn.Parameter(torch.zeros(1))
def forward(self, x):
return x * self.g
return self.fn(x) * self.g
# building block modules
class Block(nn.Module):
def __init__(self, dim, dim_out, groups = 32):
def __init__(self, dim, dim_out, groups = 8):
super().__init__()
self.block = nn.Sequential(
nn.Conv2d(dim, dim_out, 3, padding=1),
@@ -88,7 +140,7 @@ class Block(nn.Module):
return self.block(x)
class ResnetBlock(nn.Module):
def __init__(self, dim, dim_out, *, time_emb_dim, groups = 32):
def __init__(self, dim, dim_out, *, time_emb_dim, groups = 8):
super().__init__()
self.mlp = nn.Sequential(
Mish(),
@@ -106,18 +158,17 @@ class ResnetBlock(nn.Module):
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.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)
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)
@@ -127,7 +178,7 @@ class LinearAttention(nn.Module):
# model
class Unet(nn.Module):
def __init__(self, dim, out_dim = None, dim_mults=(1, 2, 4, 8), groups = 32):
def __init__(self, dim, out_dim = None, dim_mults=(1, 2, 4, 8), groups = 8):
super().__init__()
dims = [3, *map(lambda m: dim * m, dim_mults)]
in_out = list(zip(dims[:-1], dims[1:]))
@@ -148,6 +199,7 @@ class Unet(nn.Module):
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))),
Downsample(dim_out) if not is_last else nn.Identity()
]))
@@ -162,6 +214,7 @@ class Unet(nn.Module):
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))),
Upsample(dim_in) if not is_last else nn.Identity()
]))
@@ -178,8 +231,9 @@ class Unet(nn.Module):
h = []
for resnet, attn, downsample in self.downs:
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)
@@ -188,9 +242,10 @@ class Unet(nn.Module):
x = self.mid_attn(x)
x = self.mid_block2(x, t)
for resnet, attn, upsample in self.ups:
for resnet, resnet2, attn, upsample in self.ups:
x = torch.cat((x, h.pop()), dim=1)
x = resnet(x, t)
x = resnet2(x, t)
x = attn(x)
x = upsample(x)
@@ -209,11 +264,15 @@ def noise_like(shape, device, repeat=False):
return repeat_noise() if repeat else noise()
class GaussianDiffusion(nn.Module):
def __init__(self, denoise_fn, beta_start=0.0001, beta_end=0.02, num_diffusion_timesteps=1000, loss_type='l1'):
def __init__(self, denoise_fn, beta_start=0.0001, beta_end=0.02, num_diffusion_timesteps=1000, loss_type='l1', betas = None):
super().__init__()
self.denoise_fn = denoise_fn
self.np_betas = betas = np.linspace(beta_start, beta_end, num_diffusion_timesteps).astype(np.float64)
if exists(betas):
self.np_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
@@ -296,6 +355,26 @@ class GaussianDiffusion(nn.Module):
img = self.p_sample(img, torch.full((b,), i, device=device, dtype=torch.long))
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))
@torch.no_grad()
def interpolate(self, x1, x2, t = None, lam = 0.5):
b, *_, device = *x1.shape, x1.device
t = default(t, self.num_timesteps - 1)
assert x1.shape == x2.shape
t_batched = torch.stack([torch.tensor(t, device=device)] * b)
xt1, xt2 = map(lambda x: self.q_sample(x, t=t_batched), (x1, x2))
img = (1 - lam) * xt1 + lam * xt2
for i in tqdm(reversed(range(0, t)), desc='interpolation sample time step', total=t):
img = self.p_sample(img, torch.full((b,), i, device=device, dtype=torch.long))
return img
def q_sample(self, x_start, t, noise=None):
noise = default(noise, lambda: torch.randn_like(x_start))
@@ -322,5 +401,123 @@ class GaussianDiffusion(nn.Module):
def forward(self, x, *args, **kwargs):
b, *_, device = *x.shape, x.device
t = torch.randint(0, 1000, (b,), device=device).long()
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):
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.transform = transforms.Compose([
transforms.Resize(image_size),
transforms.RandomHorizontalFlip(),
transforms.CenterCrop(image_size),
transforms.ToTensor()
])
def __len__(self):
return len(self.paths)
def __getitem__(self, index):
path = self.paths[index]
img = Image.open(path)
return self.transform(img)
# trainer class
class Trainer(object):
def __init__(
self,
diffusion_model,
folder,
*,
ema_decay = 0.995,
image_size = 128,
train_batch_size = 32,
train_lr = 2e-5,
train_num_steps = 100000,
gradient_accumulate_every = 2,
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.step_start_ema = step_start_ema
self.batch_size = train_batch_size
self.image_size = image_size
self.gradient_accumulate_every = gradient_accumulate_every
self.train_num_steps = train_num_steps
self.ds = Dataset(folder, image_size)
self.dl = cycle(data.DataLoader(self.ds, batch_size = train_batch_size, shuffle=True, pin_memory=True))
self.opt = Adam(diffusion_model.parameters(), lr=train_lr)
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.reset_parameters()
def reset_parameters(self):
self.ema_model.load_state_dict(self.model.state_dict())
def step_ema(self):
if self.step < self.step_start_ema:
self.reset_parameters()
return
self.ema.update_model_average(self.ema_model, self.model)
def save(self, milestone):
data = {
'step': self.step,
'model': self.model.state_dict(),
'ema': self.ema_model.state_dict()
}
torch.save(data, f'./model-{milestone}.pt')
def load(self, milestone):
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'])
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()
self.opt.zero_grad()
if self.step % 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
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 = torch.cat(all_images_list, dim=0)
utils.save_image(all_images, f'./sample-{milestone}.png', nrow=6)
self.save(milestone)
self.step += 1
print('training completed')
BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 1.3 MiB

+3 -1
View File
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
setup(
name = 'denoising-diffusion-pytorch',
packages = find_packages(),
version = '0.0.1',
version = '0.4.0',
license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang',
@@ -16,7 +16,9 @@ setup(
install_requires=[
'einops',
'numpy',
'pillow',
'torch',
'torchvision',
'tqdm'
],
classifiers=[