Text Generation
PyTorch
English
diffusion-language-modeling
SDLLM-AR-1.7B-Base / trainer_base.py
jlemercier's picture
Release SDLLM inference package
8901f3e verified
Raw
History Blame Contribute Delete
25 kB
import itertools
from dataclasses import dataclass
import random
import hydra.utils
import lightning as L
import numpy as np
import torch
import torch.nn.functional as F
import transformers
# import mup
import dataloader
import metrics
import models
import utils
from samplers import _get_sampler
@dataclass
class Loss:
loss: torch.FloatTensor
nlls: torch.FloatTensor
reconstruction_loss: torch.FloatTensor
num_tokens: torch.FloatTensor
class LogLinear(torch.nn.Module):
def __init__(self, alpha_0=1):
super().__init__()
self.eps = 1e-3 # To be consistent with SEDD: https://github.com/louaaron/Score-Entropy-Discrete-Diffusion/blob/0605786da5ccb5747545e26d66fdf477187598b6/noise_lib.py#L56
self.alpha_0 = alpha_0
def forward(self, t):
t = (1 - self.eps) * t
alpha_t = self.alpha_0 * (1 - t)
dalpha_t = - self.alpha_0 * (1 - self.eps)
return dalpha_t, alpha_t
def sample_categorical(categorical_probs):
gumbel_norm = (
1e-10
- (torch.rand_like(categorical_probs) + 1e-10).log())
return (categorical_probs / gumbel_norm).argmax(dim=-1)
def _unsqueeze(x, reference):
return x.view(
* x.shape,
* ((1,) * (len(reference.shape) - len(x.shape))))
class TrainerBase(L.LightningModule):
def __init__(self, config, tokenizer):
super().__init__()
self.save_hyperparameters()
self.config = config
if hasattr(self.config.algo, 'loss_type'):
self.loss_type = config.algo.loss_type
self.tokenizer = tokenizer
self.vocab_size = len(tokenizer)
if (not hasattr(tokenizer, 'mask_token')
or tokenizer.mask_token is None):
self.mask_index = self.vocab_size
self.vocab_size += 1
else:
self.mask_index = tokenizer.mask_token_id
self.sampler = self.config.sampling.predictor
self.antithetic_sampling = self.config.training.antithetic_sampling
self.parameterization = self.config.algo.parameterization
if self.config.algo.backbone in ['dit', 'dit_legacy']:
self.backbone = models.dit.DiT(
self.config, vocab_size=self.vocab_size)
elif self.config.algo.backbone in ['esolm_dit', 'esolm_dit_legacy']:
self.backbone = models.dit.EsoLMDiT(
self.config, vocab_size=self.vocab_size,
mask_index=self.mask_index) # mask_index is defined in the child class
elif self.config.algo.backbone == 'hf_dit':
self.backbone = transformers.AutoModelForMaskedLM.from_pretrained(
config.eval.checkpoint_path, trust_remote_code=True)
elif self.config.algo.backbone == 'vdlm_dit':
self.backbone = models.dit.VDLMDiT(
self.config, vocab_size=self.vocab_size,
mask_index=self.mask_index)
elif self.config.algo.backbone == 'vdlm_dit_vq':
self.backbone = models.dit.VDLMDiTVQ(
self.config, vocab_size=self.vocab_size,
mask_index=self.mask_index)
elif self.config.algo.backbone == 'llada_dit':
self.backbone = transformers.AutoModel.from_pretrained(
self.config.algo.model_path, trust_remote_code=True, torch_dtype=torch.bfloat16)
elif self.config.algo.backbone == 'mlp':
# self.backbone = models.mlp.MLPDenoiser(
self.backbone = models.mlp.MLP(
self.config, vocab_size=self.vocab_size)
elif self.config.algo.backbone == 'vdlm_mlp':
self.backbone = models.mlp.VDLMMLP(
self.config, vocab_size=self.vocab_size)
self.T = self.config.algo.T
self.num_tokens = self.config.model.length
self.softplus = torch.nn.Softplus()
self.noise = LogLinear()
self.p_nucleus = self.config.sampling.p_nucleus
self.metrics = metrics.Metrics(
gen_ppl_eval_model_name_or_path=self.config.eval.gen_ppl_eval_model_name_or_path,
eval_ppl_batch_size=self.config.eval.perplexity_batch_size)
if self.config.training.ema > 0:
self.ema = models.ema.ExponentialMovingAverage(
self._get_parameters(),
decay=self.config.training.ema)
else:
self.ema = None
self.lr = self.config.optim.lr
self.sampling_eps = float(self.config.training.sampling_eps)
self.time_conditioning = self.config.algo.time_conditioning
self.neg_infinity = -1000000.0
self.fast_forward_epochs = None
self.fast_forward_batches = None
self.train_start_file_idx = 0
def setup(self, stage=None):
# different randomness for different ranks
# does not affect dataloading seed
del stage
new_seed = self.config.seed + self.trainer.global_rank
torch.manual_seed(new_seed)
np.random.seed(new_seed)
random.seed(new_seed)
def _validate_configuration(self):
assert self.config.algo.backbone in {'dit', 'dit_legacy', 'hf_dit', 'llada_dit',
'esolm_dit', 'esolm_dit_legacy',
'vdlm_dit', 'vdlm_dit_vq',
'mlp', 'vdlm_mlp'}
if self.config.algo.parameterization == 'ar':
assert not self.config.algo.time_conditioning
assert self.config.prior.type == 'none'
if self.parameterization in {'score', 'mean'}:
assert self.time_conditioning
if self.T > 0:
assert self.parameterization != 'score'
def to(self, *args, **kwargs):
self = super().to(*args, **kwargs)
self.metrics.to(*args, **kwargs)
return self
def q_xt(self, x, alpha_t):
raise NotImplementedError
def _get_parameters(self):
return itertools.chain(self.backbone.parameters(),
self.noise.parameters())
def _eval_mode(self):
if self.ema:
self.ema.store(self._get_parameters())
self.ema.copy_to(self._get_parameters())
self.backbone.eval()
self.noise.eval()
def _train_mode(self):
if self.ema:
self.ema.restore(self._get_parameters())
self.backbone.train()
self.noise.train()
def on_load_checkpoint(self, checkpoint):
if self.ema and 'ema' in checkpoint.keys():
self.ema.load_state_dict(checkpoint['ema'])
# Copied from:
# https://github.com/Dao-AILab/flash-attention/blob/main/training/src/datamodules/language_modeling_hf.py#L41
self.fast_forward_epochs = checkpoint['loops'][
'fit_loop']['epoch_progress']['current']['completed']
self.fast_forward_batches = checkpoint['loops'][
'fit_loop']['epoch_loop.batch_progress'][
'current']['completed']
if 'train_start_file_idx' in checkpoint.keys(): # Backward compatibility
self.train_start_file_idx = checkpoint['train_start_file_idx']
def on_save_checkpoint(self, checkpoint):
if self.ema:
checkpoint['ema'] = self.ema.state_dict()
checkpoint['train_start_file_idx'] = self.train_start_file_idx
# Copied from:
# https://github.com/Dao-AILab/flash-attention/blob/main/training/src/tasks/seq.py
# ['epoch_loop.batch_progress']['total']['completed']
# is 1 iteration behind, so we're using the optimizer's progress.
checkpoint['loops']['fit_loop'][
'epoch_loop.batch_progress']['total'][
'completed'] = checkpoint['loops']['fit_loop'][
'epoch_loop.automatic_optimization.optim_progress'][
'optimizer']['step']['total'][
'completed'] * self.trainer.accumulate_grad_batches
checkpoint['loops']['fit_loop'][
'epoch_loop.batch_progress']['current'][
'completed'] = checkpoint['loops']['fit_loop'][
'epoch_loop.automatic_optimization.optim_progress'][
'optimizer']['step']['current'][
'completed'] * self.trainer.accumulate_grad_batches
# _batches_that_stepped tracks the number of global steps,
# not the number of local steps, so we don't multiply with
# self.trainer.accumulate_grad_batches here.
checkpoint['loops']['fit_loop'][
'epoch_loop.state_dict'][
'_batches_that_stepped'] = checkpoint['loops']['fit_loop'][
'epoch_loop.automatic_optimization.optim_progress'][
'optimizer']['step']['total']['completed']
if 'sampler' not in checkpoint.keys():
checkpoint['sampler'] = {}
if hasattr(self.trainer.train_dataloader.sampler,
'state_dict'):
sampler_state_dict = self.trainer.\
train_dataloader.sampler.state_dict()
checkpoint['sampler'][
'random_state'] = sampler_state_dict.get(
'random_state', None)
else:
checkpoint['sampler']['random_state'] = None
def on_train_start(self):
if self.ema:
self.ema.move_shadow_params_to_device(self.device)
# Adapted from:
# https://github.com/Dao-AILab/flash-attention/blob/main/training/src/datamodules/language_modeling_hf.py
updated_dls = []
for dl in self.trainer.fit_loop._combined_loader.flattened:
updated_dls.append(
torch.utils.data.DataLoader(
dl.dataset,
batch_size=self.config.loader.batch_size,
num_workers=self.config.loader.num_workers,
pin_memory=self.config.loader.pin_memory,
# sampler=dl_sampler,
shuffle=False,
persistent_workers=True))
self.trainer.fit_loop._combined_loader.flattened = updated_dls
def optimizer_step(self, *args, **kwargs):
super().optimizer_step(*args, **kwargs)
if self.ema:
self.ema.update(self._get_parameters())
def _process_sigma(self, sigma):
raise NotImplementedError
def _process_model_output(self, model_output, xt, sigma):
raise NotImplementedError
@torch.no_grad()
def augment_batch_cfg(self, batch, prompt_index, sigma=None, sort_idx=None):
if self.config.sampling.cfg > 0.:
if isinstance(prompt_index, list):
prompt_index = torch.tensor(prompt_index)
if prompt_index.ndim < 2:
prompt_index = prompt_index.unsqueeze(
0).repeat(batch.shape[0], 1)
assert prompt_index.shape[1] == batch.shape[1], f"Expected prompt_index length {prompt_index.shape[1]} to be equal to sequence length {batch.shape[1]}"
assert prompt_index.shape[0] == batch.shape[0], f"Expected prompt_index batch size {prompt_index.shape[0]} to be equal to sequence batch size {batch.shape[0]}"
un_batch = batch.clone()
prior_batch = self.prior_sample(*batch.shape)
un_batch[prompt_index] = prior_batch[prompt_index]
batch = torch.cat([batch, un_batch])
if sigma is not None:
sigma = torch.cat([sigma, sigma])
if sort_idx is not None:
sort_idx = torch.cat([sort_idx, sort_idx])
return batch, sigma, sort_idx
@torch.no_grad()
def reduce_batch_cfg(self, batch, logits):
if self.config.sampling.cfg > 0.:
logits, un_logits = torch.chunk(logits, 2, dim=0)
logits = un_logits + \
(self.config.sampling.cfg + 1) * \
(logits - un_logits)
return logits
def forward(self, xt, sigma, sort_idx=None, x0=None, latent=None, attn_mask=None, prompt_index=None, dynamic=False):
sigma = self._process_sigma(sigma)
xt, sigma, sort_idx = self.augment_batch_cfg(xt, prompt_index, sigma=sigma, sort_idx=sort_idx)
with torch.amp.autocast('cuda', dtype=torch.float32):
logits = self.backbone(xt, sigma, sort_idx, x0, latent=latent, attn_mask=attn_mask, dynamic=dynamic)
logits = self.reduce_batch_cfg(xt, logits)
return self._process_model_output(
model_output=logits, xt=xt, sigma=sigma)
def on_train_epoch_start(self):
self.metrics.reset()
self.metrics.to(self.device)
assert self.metrics.train_nlls.nll.mean_value == 0
assert self.metrics.train_nlls.nll.weight == 0
def training_step(self, batch, batch_idx):
torch.compiler.cudagraph_mark_step_begin()
current_accumulation_step = (
batch_idx % self.trainer.accumulate_grad_batches)
input_tokens = batch['input_ids']
if 'file_idx' in batch.keys():
self.train_start_file_idx = batch['file_idx'].max(
).cpu().item()
if (self.config.algo.name != 'ar'
and torch.rand(1) < self.config.training.ssl_ratio):
length = torch.randint(1, input_tokens.shape[-1] + 1, (1,))
input_tokens = input_tokens[:, :length]
attention_mask = torch.ones_like(input_tokens)
losses = self._loss(input_tokens, attention_mask,
current_accumulation_step,
train_mode=True)
if torch.isnan(losses.loss).any():
raise ValueError('Loss is nan')
self.metrics.update_train(losses.nlls,
losses.reconstruction_loss,
losses.num_tokens)
for key in losses.__dict__.keys():
if isinstance(getattr(losses, key), torch.Tensor):
self.log(name=f'trainer/{key}',
value=getattr(losses, key).item(),
on_step=True,
on_epoch=False,
sync_dist=True)
elif isinstance(getattr(losses, key), dict):
for k, v in getattr(losses, key).items():
self.log(name=f'trainer/{key}/{k}',
value=v.item(),
on_step=True,
on_epoch=False,
sync_dist=True)
self.log(name='trainer/train_start_file_idx',
value=self.train_start_file_idx,
on_step=True,
on_epoch=False,
sync_dist=True)
return losses.loss
def on_train_epoch_end(self):
for k, v in self.metrics.train_nlls.items():
self.log(name=k, value=v.compute(), on_step=False,
on_epoch=True, sync_dist=True)
def on_validation_epoch_start(self):
self.metrics.reset()
self._eval_mode()
assert self.metrics.valid_nlls.nll.mean_value == 0
assert self.metrics.valid_nlls.nll.weight == 0
def validation_step(self, batch, batch_idx):
del batch_idx
input_tokens = batch['input_ids']
attention_mask = torch.ones_like(input_tokens)
losses = self._loss(input_tokens, attention_mask)
self.metrics.update_valid(losses.nlls,
losses.reconstruction_loss,
losses.num_tokens)
return losses.loss
def on_validation_epoch_end(self):
for k, v in self.metrics.valid_nlls.items():
self.log(name=k, value=v.compute(), on_step=False,
on_epoch=True, sync_dist=True)
if ((self.config.eval.compute_perplexity_on_sanity
or not self.trainer.sanity_checking)
and self.config.eval.generate_samples):
samples, text_samples = None, None
num_sample_batches = self.config.sampling.num_samples // (
self.trainer.num_nodes * self.trainer.num_devices
* self.config.loader.eval_batch_size)
for _ in range(max(num_sample_batches, 1)):
samples = self.generate_samples(
num_samples=self.config.loader.eval_batch_size)
self.metrics.record_entropy(samples)
# Decode the samples to be re-tokenized by eval model
text_samples = self.tokenizer.batch_decode(samples)
if self.config.eval.compute_generative_perplexity:
self.metrics.record_generative_perplexity(
text_samples, self.num_tokens, self.device)
if text_samples is not None:
if self.trainer.global_rank == 0 and hasattr(
self.trainer.logger, 'log_table'):
# Log the last generated samples
text_samples = text_samples[
: self.config.sampling.num_log_samples]
self.trainer.logger.log_table(
key=f'samples@global_step{self.global_step}',
columns=['Generated Samples'],
data=[[s] for s in text_samples])
if self.config.eval.compute_generative_perplexity:
self.log('val/gen_ppl',
self.metrics.gen_ppl.compute(),
on_epoch=True,
on_step=False,
sync_dist=True)
self.log('val/sample_entropy',
self.metrics.sample_entropy.compute(),
on_epoch=True,
on_step=False,
sync_dist=True)
self._train_mode()
def configure_optimizers(self):
optimizer = torch.optim.AdamW(
self._get_parameters(),
lr=self.config.optim.lr,
betas=(self.config.optim.beta1,
self.config.optim.beta2),
eps=self.config.optim.eps,
weight_decay=self.config.optim.weight_decay)
scheduler = hydra.utils.instantiate(
self.config.lr_scheduler, optimizer=optimizer)
scheduler_dict = {'scheduler': scheduler,
'interval': 'step',
'monitor': 'val/loss',
'name': 'trainer/lr'}
return [optimizer], [scheduler_dict]
@torch.no_grad()
def generate_samples(self, num_samples, eps=1e-5, condition=None):
"""Generate samples from the model."""
# Lightning auto-casting is not working in this method for some reason
sampler = _get_sampler(
self.config, self, self.tokenizer)
return sampler(
num_samples=num_samples,
eps=eps,
condition=condition)
def restore_model_and_sample(self, eps=1e-5, condition=None):
"""Generate samples from the model."""
# Lightning auto-casting is not working in this method for some reason
self._eval_mode()
samples = self.generate_samples(
num_samples=self.config.loader.eval_batch_size,
eps=eps,
condition=condition)
self._train_mode()
return samples
def _process_model_input(self, x0, valid_tokens):
raise NotImplementedError
def nll(self, input_tokens, output_tokens,
current_accumulation_step=None, train_mode=False):
raise NotImplementedError
def _loss(self, x0, valid_tokens,
current_accumulation_step=None,
train_mode=False):
(input_tokens, output_tokens,
valid_tokens) = self._process_model_input(
x0, valid_tokens)
loss = self.nll(input_tokens, output_tokens,
current_accumulation_step, train_mode)
assert loss.ndim == 2
nlls = (loss * valid_tokens).sum()
num_tokens = valid_tokens.sum()
token_nll = nlls / num_tokens
return Loss(loss=token_nll,
nlls=nlls,
reconstruction_loss=torch.tensor(0),
num_tokens=num_tokens)
### TODO: test that and implement with resume etc.
# def predict_dataloader(self):
# world_size = (self.trainer.num_nodes
# * self.trainer.num_devices)
# total_batches = math.ceil(
# self.config.sampling.num_samples / (
# world_size * self.config.loader.eval_batch_size))
# # Ensure the dataset length is a multiple of world size
# # so each rank receives at least one sample with the
# # unrepeated distributed sampler.
# return torch.utils.data.DataLoader(
# torch.utils.data.TensorDataset(
# torch.arange(total_batches * world_size)),
# batch_size=1,
# num_workers=0,
# pin_memory=False,
# shuffle=False)
# def on_predict_start(self):
# self.metrics.gen_ppl.reset()
# self.metrics.sample_entropy.reset()
# self._eval_mode()
# def predict_step(self, batch, batch_idx, dataloader_idx=0):
# del batch, batch_idx, dataloader_idx
# samples = self.generate_samples(
# num_samples=self.config.loader.eval_batch_size)
# self.metrics.record_entropy(samples)
# # Decode the samples to be re-tokenized by eval model
# text_samples = self.tokenizer.batch_decode(samples)
# if self.config.eval.compute_generative_perplexity:
# self.metrics.record_generative_perplexity(
# text_samples, self.num_tokens, self.device)
# return text_samples
class Diffusion(TrainerBase):
def _validate_configuration(self):
super()._validate_configuration()
assert self.config.sampling.noise_removal in {
'none', 'ancestral', 'greedy'}
assert self.loss_type in {'elbo', 'low_var', 'delbo'}
if self.config.sampling.noise_removal == 'greedy':
assert self.sampler != 'analytic'
assert self.parameterization in {'mean', 'subs'}
def _process_model_input(self, x0, valid_tokens):
return x0, None, valid_tokens
def _process_sigma(self, sigma):
assert sigma.ndim == 2
sigma = sigma.mean(-1).squeeze()
if sigma.ndim == 0:
sigma = sigma.unsqueeze(0)
if not self.time_conditioning:
sigma = torch.zeros_like(sigma)
assert sigma.ndim == 1, sigma.shape
return sigma
def _sample_t(self, n, accum_step):
if accum_step is not None:
# During training
batch_dim = n
n = self.config.loader.global_batch_size
_eps_t = torch.rand(n, device=self.device)
if self.antithetic_sampling:
offset = torch.arange(n, device=self.device) / n
_eps_t = (_eps_t / n + offset) % 1
t = (1 - self.sampling_eps) * _eps_t + self.sampling_eps
if accum_step is not None:
t = t.chunk(self.trainer.num_nodes)[
self.trainer.node_rank]
t = t.chunk(self.trainer.num_devices)[
self.trainer.local_rank]
t = t.chunk(self.trainer.accumulate_grad_batches)[
accum_step]
# corner case for the last datapoint
t = t[:batch_dim]
return t
def _sigma_from_alphat(self, alpha_t):
return -torch.log(alpha_t)
def _reconstruction_loss(self, x0):
t0 = torch.zeros(1, x0.shape[0], dtype=self.dtype,
device=self.device)
sigma_t0 = self._sigma_from_alphat(self.noise(t0)[1])
model_output_t0 = self.forward(x0, sigma_t0)
return - torch.gather(input=model_output_t0,
dim=-1,
index=x0[:, :, None]).squeeze(-1)
def nll_per_token(self, model_output, xt, x0, alpha_t,
dalpha_t, low_var):
raise NotImplementedError
def nll(self, x0, output_tokens,
current_accumulation_step=None, train_mode=False):
del output_tokens
t = self._sample_t(x0.shape[0],
current_accumulation_step)
assert t.shape[0] == x0.shape[0]
if self.T > 0:
t = (t * self.T).to(torch.int)
t = t / self.T
# t \in {1/T, 2/T, ..., 1}
t += (1 / self.T)
dalpha_t, alpha_t = self.noise(t)
alpha_t = alpha_t.unsqueeze(-1)
assert alpha_t.ndim == 2
sigma = self._sigma_from_alphat(alpha_t)
xt = self.q_xt(x0, alpha_t)
# print("xt.shape", xt.shape)
# print("Masked xt", (xt == self.mask_index).sum() / xt.numel())
log_x_theta = self.forward(xt, sigma=sigma)
# print("log_x_theta.shape", log_x_theta.shape)
utils.print_nans(log_x_theta, 'model_output')
return self.nll_per_token(
log_x_theta=log_x_theta,
xt=xt,
x0=x0,
alpha_t=alpha_t,
dalpha_t=dalpha_t,
low_var=train_mode and self.loss_type == 'low_var')
def _get_score(self, **kwargs):
del kwargs
raise NotImplementedError
def _denoiser_update(self, x, t):
raise NotImplementedError
def _analytic_update(self, x, t, dt):
raise NotImplementedError
def _ancestral_update(self, x, t, dt, p_x0, noise_removal_step):
raise NotImplementedError
class AbsorbingState(Diffusion):
def __init__(self, config, tokenizer):
self.subs_masking = config.algo.subs_masking
super().__init__(config, tokenizer)
self.save_hyperparameters()
def _validate_configuration(self):
super()._validate_configuration()
if self.parameterization in {'score', 'mean'}:
assert self.time_conditioning
assert not (self.parameterization == 'mean'
and self.T == 0)
if self.T > 0:
assert self.parameterization in {'mean', 'subs'}
if self.subs_masking:
assert self.parameterization == 'mean'
def q_xt(self, x, alpha_t):
"""Computes the noisy sample xt.
Args:
x: int torch.Tensor with shape (batch_size,
diffusion_model_input_length), input.
alpha_t: float torch.Tensor with shape (batch_size, 1).
"""
move_indices = torch.rand(
* x.shape, device=x.device) < 1 - alpha_t
xt = torch.where(move_indices, self.mask_index, x)
return xt
def prior_sample(self, *batch_dims):
return self.mask_index * torch.ones(
* batch_dims, dtype=torch.int64, device=self.device)
class UniformState(Diffusion):
def _validate_configuration(self):
super()._validate_configuration()
assert self.time_conditioning
assert self.parameterization == 'mean'
if self.config.algo.name != 'distillation':
assert self.T == 0
def q_xt(self, x, alpha_t):
"""Computes the noisy sample xt.
Args:
x: int torch.Tensor with shape (batch_size,
diffusion_model_input_length), input.
move_chance: float torch.Tensor with shape
(batch_size, 1).
"""
move_indices = torch.rand(
*x.shape, device=x.device) < 1 - alpha_t
uniform_tensor = torch.randint(
0, self.vocab_size, x.shape, device=x.device)
xt = torch.where(move_indices, uniform_tensor, x)
return xt
def prior_sample(self, *batch_dims):
return torch.randint(
0, self.vocab_size, batch_dims, dtype=torch.int64,
device=self.device)