Compare commits

...
1 Commits
2 changed files with 21 additions and 11 deletions
@@ -118,20 +118,27 @@ class PreNorm(nn.Module):
class Block(nn.Module): class Block(nn.Module):
def __init__(self, dim, dim_out, groups = 8): def __init__(self, dim, dim_out, groups = 8):
super().__init__() super().__init__()
self.block = nn.Sequential( self.proj = nn.Conv2d(dim, dim_out, 3, padding = 1)
nn.Conv2d(dim, dim_out, 3, padding = 1), self.norm = nn.GroupNorm(groups, dim_out)
nn.GroupNorm(groups, dim_out), self.act = nn.SiLU()
nn.SiLU()
) def forward(self, x, scale_shift = None):
def forward(self, x): x = self.proj(x)
return self.block(x) x = self.norm(x)
if exists(scale_shift):
scale, shift = scale_shift
x = x * (scale + 1) + shift
x = self.act(x)
return x
class ResnetBlock(nn.Module): 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, groups = 8):
super().__init__() super().__init__()
self.mlp = nn.Sequential( self.mlp = nn.Sequential(
nn.SiLU(), nn.SiLU(),
nn.Linear(time_emb_dim, dim_out) nn.Linear(time_emb_dim, dim_out * 2)
) if exists(time_emb_dim) else None ) if exists(time_emb_dim) else None
self.block1 = Block(dim, dim_out, groups = groups) self.block1 = Block(dim, dim_out, groups = groups)
@@ -139,11 +146,14 @@ class ResnetBlock(nn.Module):
self.res_conv = nn.Conv2d(dim, dim_out, 1) if dim != dim_out else nn.Identity() self.res_conv = nn.Conv2d(dim, dim_out, 1) if dim != dim_out else nn.Identity()
def forward(self, x, time_emb = None): def forward(self, x, time_emb = None):
h = self.block1(x)
scale_shift = None
if exists(self.mlp) and exists(time_emb): if exists(self.mlp) and exists(time_emb):
time_emb = self.mlp(time_emb) time_emb = self.mlp(time_emb)
h = rearrange(time_emb, 'b c -> b c 1 1') + h time_emb = rearrange(time_emb, 'b c -> b c 1 1')
scale_shift = time_emb.chunk(2, dim = 1)
h = self.block1(x, scale_shift = scale_shift)
h = self.block2(h) h = self.block2(h)
return h + self.res_conv(x) return h + self.res_conv(x)
+1 -1
View File
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
setup( setup(
name = 'denoising-diffusion-pytorch', name = 'denoising-diffusion-pytorch',
packages = find_packages(), packages = find_packages(),
version = '0.15.7', version = '0.16.0',
license='MIT', license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch', description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang', author = 'Phil Wang',