Text Generation
PyTorch
English
diffusion-language-modeling
jlemercier's picture
Release SDLLM inference package
057ef5c verified
Raw
History Blame Contribute Delete
9.22 kB
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}")