mirror of
https://github.com/wassname/denoising-diffusion-pytorch.git
synced 2026-09-10 12:01:08 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
698227ae13 | ||
|
|
c479adf960 | ||
|
|
d70fb08f8a | ||
|
|
e700a7c6de | ||
|
|
9c758662a3 | ||
|
|
1f1e42e9f9 | ||
|
|
e1800c1a8d |
@@ -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
|
||||
|
||||
|
||||
@@ -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',
|
||||
|
||||
Reference in New Issue
Block a user