import math import typing import einops import omegaconf import torch import torch.nn as nn import torch.nn.functional as F import os from utils import _save_tensor torch.backends.cuda.matmul.fp32_precision = 'tf32' torch.set_float32_matmul_precision('high') torch.backends.cudnn.benchmark = True class EmbeddingLayer(nn.Module): def __init__(self, dim, vocab_dim): super().__init__() self.embedding = nn.Parameter(torch.empty((vocab_dim, dim))) torch.nn.init.kaiming_uniform_(self.embedding, a=math.sqrt(5)) def forward(self, x): if x.ndim == 2: return self.embedding[x] assert x.ndim == 3 return torch.einsum( "blv,ve->ble", torch.nn.functional.softmax(x, dim=-1).float(), self.embedding.float()).to(x.dtype) class TimestepEmbedder(nn.Module): """ Embeds scalar timesteps into vector representations. """ def __init__(self, hidden_size, frequency_embedding_size=1024): super().__init__() self.mlp = nn.Sequential( nn.Linear(frequency_embedding_size, hidden_size, bias=True), nn.ELU(), nn.Linear(hidden_size, hidden_size, bias=False), nn.ELU(), nn.Linear(hidden_size, hidden_size, bias=False), nn.ELU(), ) self.frequency_embedding_size = frequency_embedding_size @staticmethod def timestep_embedding(t, dim, max_period=10000): """ Create sinusoidal timestep embeddings. :param t: a 1-D Tensor of N indices, one per batch element. These may be fractional. :param dim: the dimension of the output. :param max_period: controls the minimum frequency of the embeddings. :return: an (N, D) Tensor of positional embeddings. """ # https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py half = dim // 2 freqs = torch.exp( - math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32, device=t.device) / half) args = t[:, None].float() * freqs[None] embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1) if dim % 2: embedding = torch.cat( [embedding, torch.zeros_like(embedding[:, :1])], dim=-1) return embedding def forward(self, t): t_freq = self.timestep_embedding(t, self.frequency_embedding_size) t_emb = self.mlp(t_freq) return t_emb class MLPDenoiser(nn.Module): def __init__(self, config, vocab_size: int): super().__init__() if type(config) == dict: config = omegaconf.OmegaConf.create(config) # Override config for encoder to let flexibility of complexity between encoder and decoder self.config = config self.vocab_size = vocab_size dim = config.hidden_size # 512 cond_dim = config.cond_dim # 512 self.sigma_map = TimestepEmbedder(cond_dim) self.vocab_embed = EmbeddingLayer(dim, vocab_size) if hasattr(config, 'latent_dim') and config.latent_dim is not None: latent_dim = config.latent_dim # 2 self.latent_embed = nn.Sequential(nn.Linear(latent_dim, dim), nn.ELU()) self.mlp_nocorr1 = nn.Sequential( *([nn.Sequential(nn.Linear(dim, dim), nn.ELU()) for _ in range(2)]) ) self.mlp_withcorr = nn.Sequential( nn.Linear(2 * dim, dim), nn.ELU(), nn.Linear(dim, dim), nn.ELU(), nn.Linear(dim, 2 * dim), nn.ELU(), ) self.mlp_nocorr2 = nn.Sequential( *([nn.Sequential(nn.Linear(dim, dim), nn.ELU()) for _ in range(2)]) ) self.output_layer = nn.Linear(dim, vocab_size) def forward(self, x, sigma, sort_idx=None, x0=None, latent=None, attn_mask=None, kv_cache=False, dynamic=False): x = self.vocab_embed(x) t_cond = self.sigma_map(sigma) if hasattr(self, 'latent_embed'): if latent.ndim == 2: latent = latent.unsqueeze(1) z = self.latent_embed(latent) emb = x + t_cond.unsqueeze(1) + z else: emb = x + t_cond.unsqueeze(1) with torch.amp.autocast('cuda', dtype=torch.bfloat16): emb = self.mlp_nocorr1(emb) # B, 2, D emb = emb.reshape(emb.shape[0], 1, 2 * self.config.hidden_size) # B, 1, 2 * D emb = self.mlp_withcorr(emb) # B, 1, 2 * D emb = emb.reshape(emb.shape[0], 2, self.config.hidden_size) # B, 1, 2 * D emb = self.mlp_nocorr2(emb) # B, 2, D logits = self.output_layer(emb) return logits class MLPEncoder(nn.Module): def __init__(self, config, vocab_size: int): super().__init__() if type(config) == dict: config = omegaconf.OmegaConf.create(config) # Override config for encoder to let flexibility of complexity between encoder and decoder self.config = config self.vocab_size = vocab_size dim = config.hidden_size # 512 cond_dim = config.cond_dim # 512 latent_dim = config.latent_dim # 2 self.sigma_map = TimestepEmbedder(cond_dim) self.vocab_embed = EmbeddingLayer(dim, vocab_size) self.mlp_nocorr1 = nn.Sequential( *([nn.Sequential(nn.Linear(dim, dim), nn.ELU()) for _ in range(2)]) ) self.mlp_withcorr = nn.Sequential( nn.Linear(2 * dim, dim), nn.ELU(), nn.Linear(dim, 2 * dim), nn.ELU(), ) self.mlp_nocorr2 = nn.Sequential( *([nn.Sequential(nn.Linear(dim, dim), nn.ELU()) for _ in range(2)]) ) self.mlp_merged = nn.Sequential( nn.Linear(2 * dim, 2 * dim), nn.ELU(), nn.Linear(2 * dim, dim), nn.ELU(), nn.Linear(dim, dim), nn.ELU(), ) def forward(self, x, x0, sigma): x = self.vocab_embed(x) x0 = self.vocab_embed(x0) t_cond = self.sigma_map(sigma) emb_x = x + t_cond.unsqueeze(1) emb_x0 = x0 + t_cond.unsqueeze(1) emb = torch.cat([emb_x, emb_x0], dim=0) #batch processing (independent emb_x and emb_x0) with torch.amp.autocast('cuda', dtype=torch.bfloat16): emb = self.mlp_nocorr1(emb) # 2B, 2, D emb = emb.reshape(emb.shape[0], 1, 2 * self.config.hidden_size) # 2B, 1, 2 * D emb = self.mlp_withcorr(emb) # 2B, 1, 2 * D emb = emb.reshape(emb.shape[0], 2, self.config.hidden_size) # 2B, 2, D emb = self.mlp_nocorr2(emb) # 2B, 2, D emb_merged = torch.stack([ emb[: len(emb_x)] + emb[len(emb_x):] ], dim=1).mean(dim=(1,)) # B, 2, D emb_merged = emb_merged.reshape(emb_merged.shape[0], 1, 2 * self.config.hidden_size) #B, 1, 2 * D emb_merged = self.mlp_merged(emb_merged).squeeze(1) #B, D return emb_merged class MLP(nn.Module): def __init__(self, config, vocab_size: int): super().__init__() if type(config) == dict: config = omegaconf.OmegaConf.create(config) self.vocab_size = vocab_size self.decoder = MLPDenoiser(config.model, vocab_size) self.config = config def forward(self, x, sigma, sort_idx=None, x0=None, latent=None, attn_mask=None, kv_cache=False, dynamic=False): assert latent is None return self.decoder(x, sigma) class VDLMMLP(nn.Module): def __init__(self, config, vocab_size: int): super().__init__() if type(config) == dict: config = omegaconf.OmegaConf.create(config) self.config = config self.vocab_size = vocab_size self.encoder = MLPEncoder(config.model.encoder, vocab_size) self.decoder = MLPDenoiser(config.model.denoiser, vocab_size) self.surrogate_posterior = None dim = config.model.encoder.hidden_size # 512 latent_dim = config.model.latent_dim # 2 self.encoder_mean = nn.Linear(dim, latent_dim) self.encoder_logvar = nn.Linear(dim, latent_dim) def reparameterize(self, mu: torch.Tensor, logvar: torch.Tensor) -> torch.Tensor: """ Will a single z be enough to compute the expectation for the loss?? :param mu: (Tensor) Mean of the latent Gaussian :param logvar: (Tensor) Standard deviation of the latent Gaussian :return: """ std = torch.exp(0.5 * logvar) eps = torch.randn_like(std) return eps * std + mu def forward(self, x, sigma, x0=None, latent=None, mode="diffusion_training"): if mode == "diffusion_training": assert x0 is not None assert latent is None with torch.amp.autocast('cuda', dtype=torch.float32): latent_raw = self.encoder( x, x0, sigma) latent_mean = self.encoder_mean(latent_raw) latent_logvar = self.encoder_logvar(latent_raw) latent = self.reparameterize( latent_mean, latent_logvar) log_x_theta = self.decoder(x, sigma, latent=latent) return log_x_theta, latent_mean, latent_logvar, latent elif mode == "posterior_training": assert x0 is not None assert latent is None assert self.surrogate_posterior is not None, "Surrogate posterior must be used for posterior training" with torch.amp.autocast('cuda', dtype=torch.float32): latent_raw = self.encoder( x, x0, sigma) latent_mean = self.encoder_mean(latent_raw) latent_logvar = self.encoder_logvar(latent_raw) latent = self.reparameterize( latent_mean, latent_logvar) return latent elif mode == "sampling": assert x0 is None # Here latent will be either the prior or the posterior! assert latent is not None return self.decoder( x, sigma, latent=latent) else: raise ValueError(f"Invalid mode: {mode}")