Compare commits

...
7 Commits
Author SHA1 Message Date
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
3 changed files with 30 additions and 10 deletions
+5 -5
View File
@@ -1,6 +1,6 @@
<img src="./denoising-diffusion.png" width="500px"></img>
## Denoising Diffusion Probabilistic Model, in Pytorch (wip)
## 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>.
@@ -34,8 +34,8 @@ 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.
@@ -61,9 +61,9 @@ trainer = Trainer(
'path/to/your/images',
image_size = 128,
train_batch_size = 32,
train_lr = 3e-4,
train_lr = 2e-5,
train_num_steps = 100000,
gradient_accumulate_every = 1
gradient_accumulate_every = 2
)
trainer.train()
@@ -309,6 +309,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((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):
noise = default(noise, lambda: torch.randn_like(x_start))
@@ -372,9 +392,9 @@ class Trainer(object):
*,
image_size = 128,
train_batch_size = 32,
train_lr = 3e-4,
train_lr = 2e-5,
train_num_steps = 100000,
gradient_accumulate_every = 1
gradient_accumulate_every = 2
):
super().__init__()
self.model = diffusion_model
@@ -394,7 +414,7 @@ class Trainer(object):
data = next(self.dl).cuda()
loss = self.model(data)
print(f'{ind}: {loss.item()}')
loss.backward()
(loss / self.gradient_accumulate_every).backward()
self.opt.step()
self.opt.zero_grad()
@@ -403,7 +423,7 @@ class Trainer(object):
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(model.state_dict(), f'./model-{milestone}.pt')
torch.save(self.model.state_dict(), f'./model-{milestone}.pt')
ind += 1
+1 -1
View File
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
setup(
name = 'denoising-diffusion-pytorch',
packages = find_packages(),
version = '0.1.0',
version = '0.1.4',
license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang',