| 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. |
| """ |
| |
| 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) |
| |
| self.config = config |
| self.vocab_size = vocab_size |
| dim = config.hidden_size |
| cond_dim = config.cond_dim |
| 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 |
| 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) |
| emb = emb.reshape(emb.shape[0], 1, 2 * self.config.hidden_size) |
| emb = self.mlp_withcorr(emb) |
| emb = emb.reshape(emb.shape[0], 2, self.config.hidden_size) |
| emb = self.mlp_nocorr2(emb) |
| 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) |
| |
| self.config = config |
| self.vocab_size = vocab_size |
| dim = config.hidden_size |
| cond_dim = config.cond_dim |
| latent_dim = config.latent_dim |
| 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) |
|
|
| with torch.amp.autocast('cuda', dtype=torch.bfloat16): |
|
|
| emb = self.mlp_nocorr1(emb) |
| emb = emb.reshape(emb.shape[0], 1, 2 * self.config.hidden_size) |
| emb = self.mlp_withcorr(emb) |
| emb = emb.reshape(emb.shape[0], 2, self.config.hidden_size) |
| emb = self.mlp_nocorr2(emb) |
|
|
| emb_merged = torch.stack([ emb[: len(emb_x)] + emb[len(emb_x):] ], dim=1).mean(dim=(1,)) |
| emb_merged = emb_merged.reshape(emb_merged.shape[0], 1, 2 * self.config.hidden_size) |
| emb_merged = self.mlp_merged(emb_merged).squeeze(1) |
|
|
| 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 |
| latent_dim = config.model.latent_dim |
|
|
| 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 |
| |
| assert latent is not None |
| return self.decoder( |
| x, sigma, latent=latent) |
|
|
| else: |
| raise ValueError(f"Invalid mode: {mode}") |