Compare commits

...
4 Commits
3 changed files with 31 additions and 19 deletions
+1 -1
View File
@@ -64,7 +64,7 @@ trainer = Trainer(
diffusion, diffusion,
'path/to/your/images', 'path/to/your/images',
train_batch_size = 32, train_batch_size = 32,
train_lr = 2e-5, train_lr = 1e-4,
train_num_steps = 700000, # total training steps train_num_steps = 700000, # total training steps
gradient_accumulate_every = 2, # gradient accumulation steps gradient_accumulate_every = 2, # gradient accumulation steps
ema_decay = 0.995, # exponential moving average decay ema_decay = 0.995, # exponential moving average decay
@@ -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)
@@ -439,6 +449,8 @@ class GaussianDiffusion(nn.Module):
for i in tqdm(reversed(range(0, self.num_timesteps)), desc='sampling loop time step', total=self.num_timesteps): for i in tqdm(reversed(range(0, self.num_timesteps)), desc='sampling loop time step', total=self.num_timesteps):
img = self.p_sample(img, torch.full((b,), i, device=device, dtype=torch.long)) img = self.p_sample(img, torch.full((b,), i, device=device, dtype=torch.long))
img = unnormalize_to_zero_to_one(img)
return img return img
@torch.no_grad() @torch.no_grad()
@@ -497,11 +509,13 @@ class GaussianDiffusion(nn.Module):
loss = self.loss_fn(model_out, target) loss = self.loss_fn(model_out, target)
return loss return loss
def forward(self, x, *args, **kwargs): def forward(self, img, *args, **kwargs):
b, c, h, w, device, img_size, = *x.shape, x.device, self.image_size b, c, h, w, device, img_size, = *img.shape, img.device, self.image_size
assert h == img_size and w == img_size, f'height and width of image must be {img_size}' assert h == img_size and w == img_size, f'height and width of image must be {img_size}'
t = torch.randint(0, self.num_timesteps, (b,), device=device).long() t = torch.randint(0, self.num_timesteps, (b,), device=device).long()
return self.p_losses(x, t, *args, **kwargs)
img = normalize_to_neg_one_to_one(img)
return self.p_losses(img, t, *args, **kwargs)
# dataset classes # dataset classes
@@ -516,8 +530,7 @@ class Dataset(data.Dataset):
transforms.Resize(image_size), transforms.Resize(image_size),
transforms.RandomHorizontalFlip(), transforms.RandomHorizontalFlip(),
transforms.CenterCrop(image_size), transforms.CenterCrop(image_size),
transforms.ToTensor(), transforms.ToTensor()
transforms.Lambda(normalize_to_neg_one_to_one)
]) ])
def __len__(self): def __len__(self):
@@ -539,7 +552,7 @@ class Trainer(object):
ema_decay = 0.995, ema_decay = 0.995,
image_size = 128, image_size = 128,
train_batch_size = 32, train_batch_size = 32,
train_lr = 2e-5, train_lr = 1e-4,
train_num_steps = 100000, train_num_steps = 100000,
gradient_accumulate_every = 2, gradient_accumulate_every = 2,
amp = False, amp = False,
@@ -629,7 +642,6 @@ class Trainer(object):
batches = num_to_groups(36, self.batch_size) batches = num_to_groups(36, self.batch_size)
all_images_list = list(map(lambda n: self.ema_model.sample(batch_size=n), batches)) all_images_list = list(map(lambda n: self.ema_model.sample(batch_size=n), batches))
all_images = torch.cat(all_images_list, dim=0) all_images = torch.cat(all_images_list, dim=0)
all_images = unnormalize_to_zero_to_one(all_images)
utils.save_image(all_images, str(self.results_folder / f'sample-{milestone}.png'), nrow = 6) utils.save_image(all_images, str(self.results_folder / f'sample-{milestone}.png'), nrow = 6)
self.save(milestone) self.save(milestone)
+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.3', 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',