import math import copy import torch from torch import nn, einsum import torch.nn.functional as F from inspect import isfunction from functools import partial from torch.utils import data from pathlib import Path from torch.optim import Adam from torchvision import transforms, utils from PIL import Image import numpy as np from tqdm import tqdm from einops import rearrange # constants SAVE_AND_SAMPLE_EVERY = 1000 UPDATE_EMA_EVERY = 10 EXTS = ['jpg', 'png'] # helpers functions def exists(x): return x is not None def default(val, d): if exists(val): return val return d() if isfunction(d) else d def cycle(dl): while True: for data in dl: yield data # small helper modules class EMA(): def __init__(self, beta): super().__init__() self.beta = beta def update_model_average(self, ma_model, current_model): for current_params, ma_params in zip(current_model.parameters(), ma_model.parameters()): old_weight, up_weight = ma_params.data, current_params.data ma_params.data = self.update_average(old_weight, up_weight) def update_average(self, old, new): if old is None: return new return old * self.beta + (1 - self.beta) * new class Residual(nn.Module): def __init__(self, fn): super().__init__() self.fn = fn def forward(self, x, *args, **kwargs): return self.fn(x, *args, **kwargs) + x class SinusoidalPosEmb(nn.Module): def __init__(self, dim): super().__init__() self.dim = dim def forward(self, x): device = x.device half_dim = self.dim // 2 emb = math.log(10000) / (half_dim - 1) emb = torch.exp(torch.arange(half_dim, device=device) * -emb) emb = x[:, None] * emb[None, :] 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)) 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) class Rezero(nn.Module): def __init__(self, dim): super().__init__() self.g = nn.Parameter(torch.zeros(1)) def forward(self, x): return x * self.g # building block modules class Block(nn.Module): def __init__(self, dim, dim_out, groups = 32): 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 ResnetBlock(nn.Module): def __init__(self, dim, dim_out, *, time_emb_dim, groups = 32): super().__init__() self.mlp = nn.Sequential( Mish(), nn.Linear(time_emb_dim, dim_out) ) self.block1 = Block(dim, dim_out) self.block2 = Block(dim_out, dim_out) 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) h += self.mlp(time_emb)[:, :, None, None] h = self.block2(h) return h + self.res_conv(x) class LinearAttention(nn.Module): def __init__(self, dim, heads = 8, dim_head = 32): super().__init__() self.heads = heads hidden_dim = dim_head * heads self.to_qkv = nn.Conv2d(dim, hidden_dim, 1, bias = False) self.to_out = nn.Conv2d(hidden_dim, dim, 1) 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) q = q.softmax(dim=-2) 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) return self.to_out(out) # model class Unet(nn.Module): def __init__(self, dim, out_dim = None, dim_mults=(1, 2, 4, 8), groups = 32): super().__init__() dims = [3, *map(lambda m: dim * m, dim_mults)] in_out = list(zip(dims[:-1], dims[1:])) self.time_pos_emb = SinusoidalPosEmb(dim) self.mlp = nn.Sequential( nn.Linear(dim, dim * 4), Mish(), nn.Linear(dim * 4, dim) ) self.downs = nn.ModuleList([]) self.ups = nn.ModuleList([]) num_resolutions = len(in_out) for ind, (dim_in, dim_out) in enumerate(in_out): is_last = ind >= (num_resolutions - 1) self.downs.append(nn.ModuleList([ ResnetBlock(dim_in, dim_out, time_emb_dim = dim), Residual(Rezero(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 = dim) self.mid_attn = Residual(Rezero(LinearAttention(mid_dim))) self.mid_block2 = ResnetBlock(mid_dim, mid_dim, time_emb_dim = 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 = dim), Residual(Rezero(LinearAttention(dim_in))), Upsample(dim_in) if not is_last else nn.Identity() ])) out_dim = default(out_dim, 3) self.final_conv = nn.Sequential( Block(dim, dim), nn.Conv2d(dim, out_dim, 1) ) def forward(self, x, time): t = self.time_pos_emb(time) t = self.mlp(t) h = [] for resnet, attn, downsample in self.downs: x = resnet(x, t) x = attn(x) h.append(x) x = downsample(x) x = self.mid_block1(x, t) x = self.mid_attn(x) x = self.mid_block2(x, t) for resnet, attn, upsample in self.ups: x = torch.cat((x, h.pop()), dim=1) x = resnet(x, t) x = attn(x) x = upsample(x) return self.final_conv(x) # gaussian diffusion trainer class def extract(a, t, x_shape): b, *_ = t.shape out = a.gather(-1, t) return out.reshape(b, *((1,) * (len(x_shape) - 1))) def noise_like(shape, device, repeat=False): repeat_noise = lambda: torch.randn((1, *shape[1:]), device=device).repeat(shape[0], *((1,) * (len(shape) - 1))) noise = lambda: torch.randn(shape, device=device) return repeat_noise() if repeat else noise() class GaussianDiffusion(nn.Module): def __init__(self, denoise_fn, beta_start=0.0001, beta_end=0.02, num_diffusion_timesteps=1000, loss_type='l1', betas = None): super().__init__() self.denoise_fn = denoise_fn if exists(betas): self.np_betas = betas.detach().cpu().numpy() if isinstance(betas, torch.Tensor) else betas else: self.np_betas = betas = np.linspace(beta_start, beta_end, num_diffusion_timesteps).astype(np.float64) timesteps, = betas.shape self.num_timesteps = int(timesteps) self.loss_type = loss_type alphas = 1. - betas alphas_cumprod = np.cumprod(alphas, axis=0) alphas_cumprod_prev = np.append(1., alphas_cumprod[:-1]) to_torch = partial(torch.tensor, dtype=torch.float32) self.register_buffer('betas', to_torch(betas)) self.register_buffer('alphas_cumprod', to_torch(alphas_cumprod)) self.register_buffer('alphas_cumprod_prev', to_torch(alphas_cumprod_prev)) # calculations for diffusion q(x_t | x_{t-1}) and others self.register_buffer('sqrt_alphas_cumprod', to_torch(np.sqrt(alphas_cumprod))) self.register_buffer('sqrt_one_minus_alphas_cumprod', to_torch(np.sqrt(1. - alphas_cumprod))) self.register_buffer('log_one_minus_alphas_cumprod', to_torch(np.log(1. - alphas_cumprod))) self.register_buffer('sqrt_recip_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod))) self.register_buffer('sqrt_recipm1_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod - 1))) # calculations for posterior q(x_{t-1} | x_t, x_0) posterior_variance = betas * (1. - alphas_cumprod_prev) / (1. - alphas_cumprod) # above: equal to 1. / (1. / (1. - alpha_cumprod_tm1) + alpha_t / beta_t) self.register_buffer('posterior_variance', to_torch(posterior_variance)) # below: log calculation clipped because the posterior variance is 0 at the beginning of the diffusion chain self.register_buffer('posterior_log_variance_clipped', to_torch(np.log(np.maximum(posterior_variance, 1e-20)))) self.register_buffer('posterior_mean_coef1', to_torch( betas * np.sqrt(alphas_cumprod_prev) / (1. - alphas_cumprod))) self.register_buffer('posterior_mean_coef2', to_torch( (1. - alphas_cumprod_prev) * np.sqrt(alphas) / (1. - alphas_cumprod))) def q_mean_variance(self, x_start, t): mean = extract(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start variance = extract(1. - self.alphas_cumprod, t, x_start.shape) log_variance = extract(self.log_one_minus_alphas_cumprod, t, x_start.shape) return mean, variance, log_variance def predict_start_from_noise(self, x_t, t, noise): return ( extract(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t - extract(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape) * noise ) def q_posterior(self, x_start, x_t, t): posterior_mean = ( extract(self.posterior_mean_coef1, t, x_t.shape) * x_start + extract(self.posterior_mean_coef2, t, x_t.shape) * x_t ) posterior_variance = extract(self.posterior_variance, t, x_t.shape) posterior_log_variance_clipped = extract(self.posterior_log_variance_clipped, t, x_t.shape) return posterior_mean, posterior_variance, posterior_log_variance_clipped def p_mean_variance(self, x, t, clip_denoised: bool): x_recon = self.predict_start_from_noise(x, t=t, noise=self.denoise_fn(x, t)) if clip_denoised: x_recon.clamp_(-1., 1.) model_mean, posterior_variance, posterior_log_variance = self.q_posterior(x_start=x_recon, x_t=x, t=t) return model_mean, posterior_variance, posterior_log_variance @torch.no_grad() def p_sample(self, x, t, clip_denoised=True, repeat_noise=False): b, *_, device = *x.shape, x.device model_mean, _, model_log_variance = self.p_mean_variance(x=x, t=t, clip_denoised=clip_denoised) noise = noise_like(x.shape, device, repeat_noise) # no noise when t == 0 nonzero_mask = (1 - (t == 0).float()).reshape(b, *((1,) * (len(x.shape) - 1))) return model_mean + nonzero_mask * (0.5 * model_log_variance).exp() * noise @torch.no_grad() def p_sample_loop(self, shape): device = self.betas.device b = shape[0] img = torch.randn(shape, device=device) 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)) return img @torch.no_grad() def sample(self, image_size, batch_size = 16): return self.p_sample_loop((16, 3, image_size, image_size)) @torch.no_grad() def interpolate(self, x1, x2, t = None, lam = 0.5): b, *_, device = *x1.shape, x1.device t = default(t, self.num_timesteps - 1) assert x1.shape == x2.shape t_batched = torch.stack([torch.tensor(t, device=device)] * b) xt1, xt2 = map(lambda x: self.q_sample(x, t=t_batched), (x1, x2)) img = (1 - lam) * xt1 + lam * xt2 for i in tqdm(reversed(range(0, t)), desc='interpolation sample time step', total=t): img = self.p_sample(img, torch.full((b,), i, device=device, dtype=torch.long)) return img def q_sample(self, x_start, t, noise=None): noise = default(noise, lambda: torch.randn_like(x_start)) return ( extract(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start + extract(self.sqrt_one_minus_alphas_cumprod, t, x_start.shape) * noise ) def p_losses(self, x_start, t, noise = None): b, c, h, w = x_start.shape noise = default(noise, lambda: torch.randn_like(x_start)) x_noisy = self.q_sample(x_start=x_start, t=t, noise=noise) x_recon = self.denoise_fn(x_noisy, t) if self.loss_type == 'l1': loss = (noise - x_recon).abs().mean() elif self.loss_type == 'l2': loss = F.mse_loss(noise, x_recon) else: raise NotImplementedError() return loss def forward(self, x, *args, **kwargs): b, *_, device = *x.shape, x.device t = torch.randint(0, self.num_timesteps, (b,), device=device).long() return self.p_losses(x, t, *args, **kwargs) # dataset classes class Dataset(data.Dataset): def __init__(self, folder, image_size): super().__init__() self.folder = folder self.image_size = image_size self.paths = [p for ext in EXTS for p in Path(f'{folder}').glob(f'**/*.{ext}')] self.transform = transforms.Compose([ transforms.Resize(image_size), transforms.RandomHorizontalFlip(), transforms.CenterCrop(image_size), transforms.ToTensor() ]) def __len__(self): return len(self.paths) def __getitem__(self, index): path = self.paths[index] img = Image.open(path) return self.transform(img) # trainer class class Trainer(object): def __init__( self, diffusion_model, folder, *, ema_decay = 0.995, image_size = 128, train_batch_size = 32, train_lr = 2e-5, train_num_steps = 100000, gradient_accumulate_every = 2, ): super().__init__() self.model = diffusion_model self.image_size = image_size self.gradient_accumulate_every = gradient_accumulate_every self.train_num_steps = train_num_steps self.ema = EMA(ema_decay) self.ema_model = copy.deepcopy(self.model) self.ds = Dataset(folder, image_size) self.dl = cycle(data.DataLoader(self.ds, batch_size = train_batch_size, shuffle=True, pin_memory=True)) self.opt = Adam(diffusion_model.parameters(), lr=train_lr) self.step = 0 def save(self, milestone): data = { 'step': self.step, 'model': self.model.state_dict(), 'ema': self.ema_model.state_dict() } torch.save(data, f'./model-{milestone}.pt') def load(self, milestone): data = torch.load(f'./model-{milestone}.pt') self.step = data['step'] self.model.load_state_dict(data['model']) self.ema_model.load_state_dict(data['ema']) def train(self): while self.step < self.train_num_steps: for i in range(self.gradient_accumulate_every): data = next(self.dl).cuda() loss = self.model(data) print(f'{self.step}: {loss.item()}') (loss / self.gradient_accumulate_every).backward() self.opt.step() self.opt.zero_grad() if self.step % UPDATE_EMA_EVERY == 0: self.ema.update_model_average(self.ema_model, self.model) if self.step % SAVE_AND_SAMPLE_EVERY == 0: milestone = self.step // SAVE_AND_SAMPLE_EVERY all_images = self.ema_model.p_sample_loop((64, 3, self.image_size, self.image_size)) utils.save_image(all_images, f'./sample-{milestone}.png', nrow=8) self.save(milestone) self.step += 1 print('training completed')