Compare commits

..
4 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
3 changed files with 24 additions and 4 deletions
+2 -2
View File
@@ -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.
@@ -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))
@@ -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.1',
version = '0.1.3',
license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang',