From ad9e303ff33767880bc4c6534b32b87cf8641492 Mon Sep 17 00:00:00 2001 From: Phil Wang Date: Fri, 11 Jun 2021 15:28:36 -0700 Subject: [PATCH] fix channels --- .../denoising_diffusion_pytorch.py | 9 +++++++-- setup.py | 2 +- 2 files changed, 8 insertions(+), 3 deletions(-) diff --git a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py index ad6ed14..a396c4e 100644 --- a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py +++ b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py @@ -190,6 +190,8 @@ class Unet(nn.Module): channels = 3 ): super().__init__() + self.channels = channels + dims = [channels, *map(lambda m: dim * m, dim_mults)] in_out = list(zip(dims[:-1], dims[1:])) @@ -229,7 +231,7 @@ class Unet(nn.Module): Upsample(dim_in) if not is_last else nn.Identity() ])) - out_dim = default(out_dim, 3) + out_dim = default(out_dim, channels) self.final_conv = nn.Sequential( Block(dim, dim), nn.Conv2d(dim, out_dim, 1) @@ -291,11 +293,13 @@ class GaussianDiffusion(nn.Module): denoise_fn, *, image_size, + channels = 3, timesteps = 1000, loss_type = 'l1', betas = None ): super().__init__() + self.channels = channels self.image_size = image_size self.denoise_fn = denoise_fn @@ -389,7 +393,8 @@ class GaussianDiffusion(nn.Module): @torch.no_grad() def sample(self, batch_size = 16): image_size = self.image_size - return self.p_sample_loop((batch_size, 3, image_size, image_size)) + channels = self.channels + return self.p_sample_loop((batch_size, channels, image_size, image_size)) @torch.no_grad() def interpolate(self, x1, x2, t = None, lam = 0.5): diff --git a/setup.py b/setup.py index 7b3107f..6de8a28 100644 --- a/setup.py +++ b/setup.py @@ -3,7 +3,7 @@ from setuptools import setup, find_packages setup( name = 'denoising-diffusion-pytorch', packages = find_packages(), - version = '0.6.0', + version = '0.6.1', license='MIT', description = 'Denoising Diffusion Probabilistic Models - Pytorch', author = 'Phil Wang',