Compare commits

..
8 Commits
Author SHA1 Message Date
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
5 changed files with 151 additions and 21 deletions
+39 -15
View File
@@ -1,6 +1,8 @@
## 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>.
## Install ## Install
@@ -32,10 +34,43 @@ loss = diffusion(training_images)
loss.backward() loss.backward()
# after a lot of training # after a lot of training
sampled_images = diffusion.p_sample_loop((1, 3, 128, 128)) sampled_images = diffusion.sample(128, batch_size = 4)
sampled_images.shape # (1, 3, 128, 128) 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,
gradient_accumulate_every = 2
)
trainer.train()
```
Todo: Command line tool for one-line training
## Citations ## Citations
```bibtex ```bibtex
@@ -48,14 +83,3 @@ sampled_images.shape # (1, 3, 128, 128)
primaryClass={cs.LG} primaryClass={cs.LG}
} }
``` ```
```bibtex
@misc{chen2020wavegrad,
title={WaveGrad: Estimating Gradients for Waveform Generation},
author={Nanxin Chen and Yu Zhang and Heiga Zen and Ron J. Weiss and Mohammad Norouzi and William Chan},
year={2020},
eprint={2009.00713},
archivePrefix={arXiv},
primaryClass={eess.AS}
}
```
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,25 @@
import math import math
import torch import torch
from inspect import isfunction
from functools import partial
from torch import nn, einsum from torch import nn, einsum
import torch.nn.functional as F 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 import numpy as np
from tqdm import tqdm from tqdm import tqdm
from einops import rearrange from einops import rearrange
# constants
SAVE_AND_SAMPLE_EVERY = 1000
EXTS = ['jpg', 'png']
# helpers functions # helpers functions
def exists(x): def exists(x):
@@ -19,8 +30,10 @@ def default(val, d):
return val return val
return d() if isfunction(d) else d return d() if isfunction(d) else d
def normal_kl(mean1, logvar1, mean2, logvar2): def cycle(dl):
return 0.5 * (-1. + logvar2 - logvar1 + torch.exp(logvar1 - logvar2) + torch.exp(-logvar2) * (mean1 - mean2) ** 2) while True:
for data in dl:
yield data
# small helper modules # small helper modules
@@ -296,6 +309,26 @@ class GaussianDiffusion(nn.Module):
img = self.p_sample(img, torch.full((b,), i, device=device, dtype=torch.long)) img = self.p_sample(img, torch.full((b,), i, device=device, dtype=torch.long))
return img return img
@torch.no_grad()
def sample(self, image_size, batch_size = 16):
return self.p_sample_loop((16, 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): def q_sample(self, x_start, t, noise=None):
noise = default(noise, lambda: torch.randn_like(x_start)) noise = default(noise, lambda: torch.randn_like(x_start))
@@ -324,3 +357,74 @@ class GaussianDiffusion(nn.Module):
b, *_, device = *x.shape, x.device b, *_, device = *x.shape, x.device
t = torch.randint(0, self.num_timesteps, (b,), device=device).long() t = torch.randint(0, self.num_timesteps, (b,), device=device).long()
return self.p_losses(x, t, *args, **kwargs) 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,
*,
image_size = 128,
train_batch_size = 32,
train_lr = 2e-5,
train_num_steps = 100000,
gradient_accumulate_every = 2
):
super().__init__()
self.model = diffusion_model
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)
def train(self):
ind = 0
while ind < self.train_num_steps:
for i in range(self.gradient_accumulate_every):
data = next(self.dl).cuda()
loss = self.model(data)
print(f'{ind}: {loss.item()}')
loss.backward()
self.opt.step()
self.opt.zero_grad()
if ind % SAVE_AND_SAMPLE_EVERY == 0:
milestone = ind // SAVE_AND_SAMPLE_EVERY
all_images = self.model.p_sample_loop((64, 3, self.image_size, self.image_size))
utils.save_image(all_images, f'./sample-{milestone}.png', nrow=8)
torch.save(self.model.state_dict(), f'./model-{milestone}.pt')
ind += 1
print('training completed')
+3 -1
View File
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
setup( setup(
name = 'denoising-diffusion-pytorch', name = 'denoising-diffusion-pytorch',
packages = find_packages(), packages = find_packages(),
version = '0.0.2', version = '0.1.3',
license='MIT', license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch', description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang', author = 'Phil Wang',
@@ -16,7 +16,9 @@ setup(
install_requires=[ install_requires=[
'einops', 'einops',
'numpy', 'numpy',
'pillow',
'torch', 'torch',
'torchvision',
'tqdm' 'tqdm'
], ],
classifiers=[ classifiers=[