Compare commits

..
2 Commits
2 changed files with 6 additions and 5 deletions
@@ -128,8 +128,8 @@ class ResnetBlock(nn.Module):
nn.Linear(time_emb_dim, dim_out)
) if exists(time_emb_dim) else None
self.block1 = Block(dim, dim_out)
self.block2 = Block(dim_out, dim_out)
self.block1 = Block(dim, dim_out, groups = groups)
self.block2 = Block(dim_out, dim_out, groups = groups)
self.res_conv = nn.Conv2d(dim, dim_out, 1) if dim != dim_out else nn.Identity()
def forward(self, x, time_emb = None):
@@ -238,7 +238,8 @@ class Unet(nn.Module):
channels = 3,
with_time_emb = True,
use_convnext = False,
resnet_block_groups = 8
resnet_block_groups = 8,
convnext_mult = 2
):
super().__init__()
@@ -255,7 +256,7 @@ class Unet(nn.Module):
# resnet or convnext
if use_convnext:
block_klass = ConvNextBlock
block_klass = partial(ConvNextBlock, mult = convnext_mult)
else:
block_klass = partial(ResnetBlock, groups = resnet_block_groups)
+1 -1
View File
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
setup(
name = 'denoising-diffusion-pytorch',
packages = find_packages(),
version = '0.11.0',
version = '0.11.2',
license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang',