small helper fn to make sampling more clear

This commit is contained in:
Phil Wang
2020-09-07 16:40:36 -07:00
parent e700a7c6de
commit d70fb08f8a
2 changed files with 6 additions and 2 deletions
+2 -2
View File
@@ -34,8 +34,8 @@ 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. 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,10 @@ 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))
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))