From d70fb08f8a9e60a4d020cc1d6f43eede5d7b9167 Mon Sep 17 00:00:00 2001 From: Phil Wang Date: Mon, 7 Sep 2020 16:40:36 -0700 Subject: [PATCH] small helper fn to make sampling more clear --- README.md | 4 ++-- denoising_diffusion_pytorch/denoising_diffusion_pytorch.py | 4 ++++ 2 files changed, 6 insertions(+), 2 deletions(-) diff --git a/README.md b/README.md index 247d685..393d7eb 100644 --- a/README.md +++ b/README.md @@ -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. diff --git a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py index 41f5359..fb256d8 100644 --- a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py +++ b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py @@ -309,6 +309,10 @@ 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)) + def q_sample(self, x_start, t, noise=None): noise = default(noise, lambda: torch.randn_like(x_start))