|
|
|
@@ -118,20 +118,27 @@ class PreNorm(nn.Module):
|
|
|
|
|
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),
|
|
|
|
|
nn.SiLU()
|
|
|
|
|
)
|
|
|
|
|
def forward(self, x):
|
|
|
|
|
return self.block(x)
|
|
|
|
|
self.proj = nn.Conv2d(dim, dim_out, 3, padding = 1)
|
|
|
|
|
self.norm = nn.GroupNorm(groups, dim_out)
|
|
|
|
|
self.act = nn.SiLU()
|
|
|
|
|
|
|
|
|
|
def forward(self, x, scale_shift = None):
|
|
|
|
|
x = self.proj(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):
|
|
|
|
|
def __init__(self, dim, dim_out, *, time_emb_dim = None, groups = 8):
|
|
|
|
|
super().__init__()
|
|
|
|
|
self.mlp = nn.Sequential(
|
|
|
|
|
nn.SiLU(),
|
|
|
|
|
nn.Linear(time_emb_dim, dim_out)
|
|
|
|
|
nn.Linear(time_emb_dim, dim_out * 2)
|
|
|
|
|
) if exists(time_emb_dim) else None
|
|
|
|
|
|
|
|
|
|
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()
|
|
|
|
|
|
|
|
|
|
def forward(self, x, time_emb = None):
|
|
|
|
|
h = self.block1(x)
|
|
|
|
|
|
|
|
|
|
scale_shift = None
|
|
|
|
|
if exists(self.mlp) and exists(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)
|
|
|
|
|
return h + self.res_conv(x)
|
|
|
|
@@ -319,6 +329,12 @@ def noise_like(shape, device, repeat=False):
|
|
|
|
|
noise = lambda: torch.randn(shape, device=device)
|
|
|
|
|
return repeat_noise() if repeat else noise()
|
|
|
|
|
|
|
|
|
|
def linear_beta_schedule(timesteps):
|
|
|
|
|
scale = 1000 / timesteps
|
|
|
|
|
beta_start = scale * 0.0001
|
|
|
|
|
beta_end = scale * 0.02
|
|
|
|
|
return torch.linspace(beta_start, beta_end, timesteps, dtype = torch.float64)
|
|
|
|
|
|
|
|
|
|
def cosine_beta_schedule(timesteps, s = 0.008):
|
|
|
|
|
"""
|
|
|
|
|
cosine schedule
|
|
|
|
@@ -340,7 +356,8 @@ class GaussianDiffusion(nn.Module):
|
|
|
|
|
channels = 3,
|
|
|
|
|
timesteps = 1000,
|
|
|
|
|
loss_type = 'l1',
|
|
|
|
|
objective = 'pred_noise'
|
|
|
|
|
objective = 'pred_noise',
|
|
|
|
|
beta_schedule = 'cosine'
|
|
|
|
|
):
|
|
|
|
|
super().__init__()
|
|
|
|
|
assert not (type(self) == GaussianDiffusion and denoise_fn.channels != denoise_fn.out_dim)
|
|
|
|
@@ -350,7 +367,12 @@ class GaussianDiffusion(nn.Module):
|
|
|
|
|
self.denoise_fn = denoise_fn
|
|
|
|
|
self.objective = objective
|
|
|
|
|
|
|
|
|
|
betas = cosine_beta_schedule(timesteps)
|
|
|
|
|
if beta_schedule == 'linear':
|
|
|
|
|
betas = linear_beta_schedule(timesteps)
|
|
|
|
|
elif beta_schedule == 'cosine':
|
|
|
|
|
betas = cosine_beta_schedule(timesteps)
|
|
|
|
|
else:
|
|
|
|
|
raise ValueError(f'unknown beta schedule {beta_schedule}')
|
|
|
|
|
|
|
|
|
|
alphas = 1. - betas
|
|
|
|
|
alphas_cumprod = torch.cumprod(alphas, axis=0)
|
|
|
|
@@ -439,6 +461,8 @@ class GaussianDiffusion(nn.Module):
|
|
|
|
|
|
|
|
|
|
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 = unnormalize_to_zero_to_one(img)
|
|
|
|
|
return img
|
|
|
|
|
|
|
|
|
|
@torch.no_grad()
|
|
|
|
@@ -497,11 +521,13 @@ class GaussianDiffusion(nn.Module):
|
|
|
|
|
loss = self.loss_fn(model_out, target)
|
|
|
|
|
return loss
|
|
|
|
|
|
|
|
|
|
def forward(self, x, *args, **kwargs):
|
|
|
|
|
b, c, h, w, device, img_size, = *x.shape, x.device, self.image_size
|
|
|
|
|
def forward(self, img, *args, **kwargs):
|
|
|
|
|
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}'
|
|
|
|
|
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
|
|
|
|
|
|
|
|
|
@@ -516,8 +542,7 @@ class Dataset(data.Dataset):
|
|
|
|
|
transforms.Resize(image_size),
|
|
|
|
|
transforms.RandomHorizontalFlip(),
|
|
|
|
|
transforms.CenterCrop(image_size),
|
|
|
|
|
transforms.ToTensor(),
|
|
|
|
|
transforms.Lambda(normalize_to_neg_one_to_one)
|
|
|
|
|
transforms.ToTensor()
|
|
|
|
|
])
|
|
|
|
|
|
|
|
|
|
def __len__(self):
|
|
|
|
@@ -539,7 +564,7 @@ class Trainer(object):
|
|
|
|
|
ema_decay = 0.995,
|
|
|
|
|
image_size = 128,
|
|
|
|
|
train_batch_size = 32,
|
|
|
|
|
train_lr = 2e-5,
|
|
|
|
|
train_lr = 1e-4,
|
|
|
|
|
train_num_steps = 100000,
|
|
|
|
|
gradient_accumulate_every = 2,
|
|
|
|
|
amp = False,
|
|
|
|
@@ -629,7 +654,6 @@ class Trainer(object):
|
|
|
|
|
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 = 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)
|
|
|
|
|
self.save(milestone)
|
|
|
|
|
|
|
|
|
|