From 3bf5e768c253670c60959413a9610924b7e7e4a6 Mon Sep 17 00:00:00 2001 From: Phil Wang Date: Tue, 7 Jun 2022 21:15:27 -0700 Subject: [PATCH] use a non-sinusoidal embedded condition for continuous time gaussian diffusion conditioned on log(snr) --- .../continuous_time_gaussian_diffusion.py | 1 + .../denoising_diffusion_pytorch.py | 30 ++++++++++++++----- setup.py | 2 +- 3 files changed, 24 insertions(+), 9 deletions(-) diff --git a/denoising_diffusion_pytorch/continuous_time_gaussian_diffusion.py b/denoising_diffusion_pytorch/continuous_time_gaussian_diffusion.py index 063be8b..74a41b5 100644 --- a/denoising_diffusion_pytorch/continuous_time_gaussian_diffusion.py +++ b/denoising_diffusion_pytorch/continuous_time_gaussian_diffusion.py @@ -65,6 +65,7 @@ class ContinuousTimeGaussianDiffusion(nn.Module): num_sample_steps = 500 ): super().__init__() + assert not denoise_fn.sinusoidal_cond_mlp self.denoise_fn = denoise_fn diff --git a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py index 5958d15..e9115c8 100644 --- a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py +++ b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py @@ -16,6 +16,7 @@ from PIL import Image from tqdm import tqdm from einops import rearrange +from einops.layers.torch import Rearrange # helpers functions @@ -211,6 +212,18 @@ class Attention(nn.Module): # model +def MLP(dim_in, dim_hidden): + return nn.Sequential( + Rearrange('... -> ... 1'), + nn.Linear(1, dim_hidden), + nn.GELU(), + nn.LayerNorm(dim_hidden), + nn.Linear(dim_hidden, dim_hidden), + nn.GELU(), + nn.LayerNorm(dim_hidden), + nn.Linear(dim_hidden, dim_hidden) + ) + class Unet(nn.Module): def __init__( self, @@ -219,9 +232,9 @@ class Unet(nn.Module): out_dim = None, dim_mults=(1, 2, 4, 8), channels = 3, - with_time_emb = True, resnet_block_groups = 8, - learned_variance = False + learned_variance = False, + sinusoidal_cond_mlp = True ): super().__init__() @@ -239,8 +252,11 @@ class Unet(nn.Module): # time embeddings - if with_time_emb: - time_dim = dim * 4 + time_dim = dim * 4 + + self.sinusoidal_cond_mlp = sinusoidal_cond_mlp + + if sinusoidal_cond_mlp: self.time_mlp = nn.Sequential( SinusoidalPosEmb(dim), nn.Linear(dim, time_dim), @@ -248,8 +264,7 @@ class Unet(nn.Module): nn.Linear(time_dim, time_dim) ) else: - time_dim = None - self.time_mlp = None + self.time_mlp = MLP(1, time_dim) # layers @@ -292,8 +307,7 @@ 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 + t = self.time_mlp(time) h = [] diff --git a/setup.py b/setup.py index 22fe6a8..d78b332 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.16.4', + version = '0.16.5', license='MIT', description = 'Denoising Diffusion Probabilistic Models - Pytorch', author = 'Phil Wang',