diff --git a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py index 0063542..5236788 100644 --- a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py +++ b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py @@ -192,6 +192,7 @@ class Unet(nn.Module): def __init__( self, dim, + init_dim = None, out_dim = None, dim_mults=(1, 2, 4, 8), channels = 3, @@ -200,7 +201,10 @@ class Unet(nn.Module): super().__init__() self.channels = channels - dims = [channels, *map(lambda m: dim * m, dim_mults)] + init_dim = default(init_dim, dim // 3 * 2) + self.init_conv = nn.Conv2d(channels, init_dim, 7, padding = 3) + + dims = [init_dim, *map(lambda m: dim * m, dim_mults)] in_out = list(zip(dims[:-1], dims[1:])) if with_time_emb: @@ -251,6 +255,8 @@ class Unet(nn.Module): ) def forward(self, x, time): + x = self.init_conv(x) + t = self.time_mlp(time) if exists(self.time_mlp) else None h = [] diff --git a/setup.py b/setup.py index 0db2fc7..ea1fb61 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.10.0', + version = '0.10.1', license='MIT', description = 'Denoising Diffusion Probabilistic Models - Pytorch', author = 'Phil Wang',