From 91cff459394b784f870f3998023439796c95b7ca Mon Sep 17 00:00:00 2001 From: Phil Wang Date: Tue, 25 Jan 2022 09:02:45 -0800 Subject: [PATCH] replace resnets with convnext blocks --- README.md | 15 ++- .../denoising_diffusion_pytorch.py | 112 ++++++++---------- setup.py | 2 +- 3 files changed, 66 insertions(+), 63 deletions(-) diff --git a/README.md b/README.md index bdef026..d96342a 100644 --- a/README.md +++ b/README.md @@ -2,7 +2,9 @@ ## Denoising Diffusion Probabilistic Model, in Pytorch -Implementation of Denoising Diffusion Probabilistic Model in Pytorch. It is a new approach to generative modeling that may have the potential to rival GANs. It uses denoising score matching to estimate the gradient of the data distribution, followed by Langevin sampling to sample from the true distribution. This implementation was transcribed from the official Tensorflow version here. +Implementation of Denoising Diffusion Probabilistic Model in Pytorch. It is a new approach to generative modeling that may have the potential to rival GANs. It uses denoising score matching to estimate the gradient of the data distribution, followed by Langevin sampling to sample from the true distribution. + +This implementation was transcribed from the official Tensorflow version here and then modified to use ConvNext blocks instead of Resnets. @@ -97,3 +99,14 @@ Samples and model checkpoints will be logged to `./results` periodically note = {under review} } ``` + +```bibtex +@misc{liu2022convnet, + title = {A ConvNet for the 2020s}, + author = {Zhuang Liu and Hanzi Mao and Chao-Yuan Wu and Christoph Feichtenhofer and Trevor Darrell and Saining Xie}, + year = {2022}, + eprint = {2201.03545}, + archivePrefix = {arXiv}, + primaryClass = {cs.CV} +} +``` diff --git a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py index 8ca6788..5f5f8f4 100644 --- a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py +++ b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py @@ -91,25 +91,11 @@ class SinusoidalPosEmb(nn.Module): emb = torch.cat((emb.sin(), emb.cos()), dim=-1) return emb -class Mish(nn.Module): - def forward(self, x): - return x * torch.tanh(F.softplus(x)) +def Upsample(dim): + return nn.ConvTranspose2d(dim, dim, 4, 2, 1) -class Upsample(nn.Module): - def __init__(self, dim): - super().__init__() - self.conv = nn.ConvTranspose2d(dim, dim, 4, 2, 1) - - def forward(self, x): - return self.conv(x) - -class Downsample(nn.Module): - def __init__(self, dim): - super().__init__() - self.conv = nn.Conv2d(dim, dim, 3, 2, 1) - - def forward(self, x): - return self.conv(x) +def Downsample(dim): + return nn.Conv2d(dim, dim, 3, 2, 1) class LayerNorm(nn.Module): def __init__(self, dim, eps = 1e-5): @@ -119,9 +105,9 @@ class LayerNorm(nn.Module): self.b = nn.Parameter(torch.zeros(1, dim, 1, 1)) def forward(self, x): - std = torch.var(x, dim = 1, unbiased = False, keepdim = True).sqrt() + var = torch.var(x, dim = 1, unbiased = False, keepdim = True) mean = torch.mean(x, dim = 1, keepdim = True) - return (x - mean) / (std + self.eps) * self.g + self.b + return (x - mean) / (var + self.eps).sqrt() * self.g + self.b class PreNorm(nn.Module): def __init__(self, dim, fn): @@ -135,41 +121,43 @@ class PreNorm(nn.Module): # building block modules -class Block(nn.Module): - def __init__(self, dim, dim_out, groups = 8): - super().__init__() - self.block = nn.Sequential( - nn.Conv2d(dim, dim_out, 3, padding=1), - nn.GroupNorm(groups, dim_out), - Mish() - ) - def forward(self, x): - return self.block(x) +class ConvNextBlock(nn.Module): + """ https://arxiv.org/abs/2201.03545 """ -class ResnetBlock(nn.Module): - def __init__(self, dim, dim_out, *, time_emb_dim = None, groups = 8): + def __init__(self, dim, dim_out, *, time_emb_dim = None, mult = 2, norm = True): super().__init__() self.mlp = nn.Sequential( - Mish(), - nn.Linear(time_emb_dim, dim_out) + nn.GELU(), + nn.Linear(time_emb_dim, dim) ) if exists(time_emb_dim) else None - self.block1 = Block(dim, dim_out) - self.block2 = Block(dim_out, dim_out) + self.ds_conv = nn.Conv2d(dim, dim, 7, padding = 3, groups = dim) + + self.net = nn.Sequential( + LayerNorm(dim) if norm else nn.Identity(), + nn.Conv2d(dim, dim_out * mult, 1), + nn.GELU(), + LayerNorm(dim_out * mult), + nn.Conv2d(dim_out * mult, dim_out, 1) + ) + self.res_conv = nn.Conv2d(dim, dim_out, 1) if dim != dim_out else nn.Identity() - def forward(self, x, time_emb): - h = self.block1(x) + def forward(self, x, time_emb = None): + h = self.ds_conv(x) if exists(self.mlp): - h += self.mlp(time_emb)[:, :, None, None] + assert exists(time_emb), 'time emb must be passed in' + condition = self.mlp(time_emb) + h = h + rearrange(condition, 'b c -> b c 1 1') - h = self.block2(h) + h = self.net(h) return h + self.res_conv(x) class LinearAttention(nn.Module): def __init__(self, dim, heads = 4, dim_head = 32): super().__init__() + self.scale = dim_head ** -0.5 self.heads = heads hidden_dim = dim_head * heads self.to_qkv = nn.Conv2d(dim, hidden_dim * 3, 1, bias = False) @@ -177,12 +165,15 @@ class LinearAttention(nn.Module): def forward(self, x): b, c, h, w = x.shape - qkv = self.to_qkv(x) - q, k, v = rearrange(qkv, 'b (qkv heads c) h w -> qkv b heads c (h w)', heads = self.heads, qkv=3) - k = k.softmax(dim=-1) - context = torch.einsum('bhdn,bhen->bhde', k, v) - out = torch.einsum('bhde,bhdn->bhen', context, q) - out = rearrange(out, 'b heads c (h w) -> b (heads c) h w', heads=self.heads, h=h, w=w) + qkv = self.to_qkv(x).chunk(3, dim = 1) + q, k, v = map(lambda t: rearrange(t, 'b (h c) x y -> b h c (x y)', h = self.heads), qkv) + q = q * self.scale + + k = k.softmax(dim = -1) + context = torch.einsum('b h d n, b h e n -> b h d e', k, v) + + out = torch.einsum('b h d e, b h d n -> b h e n', context, q) + out = rearrange(out, 'b h c (x y) -> b (h c) x y', h = self.heads, x = h, y = w) return self.to_out(out) # model @@ -193,7 +184,6 @@ class Unet(nn.Module): dim, out_dim = None, dim_mults=(1, 2, 4, 8), - groups = 8, channels = 3, with_time_emb = True ): @@ -208,7 +198,7 @@ class Unet(nn.Module): self.time_mlp = nn.Sequential( SinusoidalPosEmb(dim), nn.Linear(dim, dim * 4), - Mish(), + nn.GELU(), nn.Linear(dim * 4, dim) ) else: @@ -223,30 +213,30 @@ class Unet(nn.Module): is_last = ind >= (num_resolutions - 1) self.downs.append(nn.ModuleList([ - ResnetBlock(dim_in, dim_out, time_emb_dim = time_dim), - ResnetBlock(dim_out, dim_out, time_emb_dim = time_dim), + ConvNextBlock(dim_in, dim_out, time_emb_dim = time_dim, norm = ind != 0), + ConvNextBlock(dim_out, dim_out, time_emb_dim = time_dim), Residual(PreNorm(dim_out, LinearAttention(dim_out))), Downsample(dim_out) if not is_last else nn.Identity() ])) mid_dim = dims[-1] - self.mid_block1 = ResnetBlock(mid_dim, mid_dim, time_emb_dim = time_dim) + self.mid_block1 = ConvNextBlock(mid_dim, mid_dim, time_emb_dim = time_dim) self.mid_attn = Residual(PreNorm(mid_dim, LinearAttention(mid_dim))) - self.mid_block2 = ResnetBlock(mid_dim, mid_dim, time_emb_dim = time_dim) + self.mid_block2 = ConvNextBlock(mid_dim, mid_dim, time_emb_dim = time_dim) for ind, (dim_in, dim_out) in enumerate(reversed(in_out[1:])): is_last = ind >= (num_resolutions - 1) self.ups.append(nn.ModuleList([ - ResnetBlock(dim_out * 2, dim_in, time_emb_dim = time_dim), - ResnetBlock(dim_in, dim_in, time_emb_dim = time_dim), + ConvNextBlock(dim_out * 2, dim_in, time_emb_dim = time_dim), + ConvNextBlock(dim_in, dim_in, time_emb_dim = time_dim), Residual(PreNorm(dim_in, LinearAttention(dim_in))), Upsample(dim_in) if not is_last else nn.Identity() ])) out_dim = default(out_dim, channels) self.final_conv = nn.Sequential( - Block(dim, dim), + ConvNextBlock(dim, dim), nn.Conv2d(dim, out_dim, 1) ) @@ -255,9 +245,9 @@ class Unet(nn.Module): h = [] - for resnet, resnet2, attn, downsample in self.downs: - x = resnet(x, t) - x = resnet2(x, t) + for convnext, convnext2, attn, downsample in self.downs: + x = convnext(x, t) + x = convnext2(x, t) x = attn(x) h.append(x) x = downsample(x) @@ -266,10 +256,10 @@ class Unet(nn.Module): x = self.mid_attn(x) x = self.mid_block2(x, t) - for resnet, resnet2, attn, upsample in self.ups: + for convnext, convnext2, attn, upsample in self.ups: x = torch.cat((x, h.pop()), dim=1) - x = resnet(x, t) - x = resnet2(x, t) + x = convnext(x, t) + x = convnext2(x, t) x = attn(x) x = upsample(x) diff --git a/setup.py b/setup.py index 9749e76..1b9f790 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.9', + version = '0.7.0', license='MIT', description = 'Denoising Diffusion Probabilistic Models - Pytorch', author = 'Phil Wang',