import math import torch import torch.nn.functional as F import time from tqdm import tqdm import torch.distributed as dist from pathlib import Path from queue import Queue, Empty from threading import Thread, Event try: import wandb except ImportError: wandb = None class DummyOptimizer(torch.optim.Optimizer): def __init__(self): self.param_groups = [{"lr": 0}] None def step(self): None def zero_grad(self, set_to_none=False): None class DummyScheduler: def step(self): None class TrainState: """Track number of steps, examples, and tokens processed""" step: int = 0 # Steps in the current epoch accum_step: int = 0 # Number of gradient accumulation steps samples: int = 0 # total # of examples used tokens: int = 0 # total # of tokens processed best_eval_acc: float = 0 # best eval accuracy so far skipped_updates: int = 0 # number of updates skipped due to NaN gradients def make_perplexity_loss(pad_idx): def perplexity(out, target): x = out.detach().contiguous().view(-1, out.size(-1)) y = target.detach().contiguous().view(-1) return F.cross_entropy(x, y, ignore_index=pad_idx, reduction="sum") return perplexity def run_epoch( data_iter, make_eval_data_iter, model, loss_compute, optimizer, scheduler, batch_size, mode="train", num_batches=None, num_eval_batches=None, accum_iter=1, eval_iter=1, enable_checkpoint=False, epoch=0, train_state=TrainState(), quiet=False, scaler=None, max_examples_per_rank=None, checkpoint_path="./checkpoints", world_size=1, ): """Train a single epoch""" last_eval_acc = None num_kept_checkpoints = 5 def checkpoint(step, acc): # only keep the last num_kept_checkpoints checkpoints files_by_mtime = list( sorted((p.stat().st_mtime, p) for p in Path(checkpoint_path).glob("*.pt")) ) if len(files_by_mtime) > num_kept_checkpoints: for _, p in files_by_mtime[:-num_kept_checkpoints]: p.unlink() torch.save( { "model_state_dict": model.state_dict(), "optimizer_state_dict": optimizer.state_dict(), "scheduler_state_dict": scheduler.state_dict(), "train_state": train_state, }, f"{checkpoint_path}/epoch_{epoch}_step_{step}_accuracy_{acc}.pt", ) def verify_ranks_are_equal(): dist.barrier() for p in model.parameters(): if dist.get_rank() == 0: lst = [torch.zeros_like(p.data) for _ in range(world_size)] dist.gather(p.data, gather_list=lst) for i in range(1, world_size): if not torch.allclose(p.data, lst[i]): raise RuntimeError("ranks are not equal") else: dist.gather(p.data) def eval(): nonlocal last_eval_acc model.eval() with torch.no_grad(): verify_ranks_are_equal() eval_loss, eval_tokens, eval_accurate_count, _ = run_epoch( make_eval_data_iter(), None, model, make_perplexity_loss(model.module.params.padding_idx), DummyOptimizer(), DummyScheduler(), batch_size, num_batches=num_eval_batches, mode="eval", quiet=quiet, scaler=scaler, world_size=world_size, checkpoint_path=checkpoint_path, ) dist.all_reduce(eval_loss) dist.all_reduce(eval_tokens) dist.all_reduce(eval_accurate_count) dist.barrier() perplexity = eval_loss.item() / eval_tokens.item() accuracy = eval_accurate_count.item() / eval_tokens.item() if not quiet: tqdm.write(f"eval_perplexity: {perplexity}, eval_accuracy: {accuracy}") if wandb.run is not None: wandb.log( {"eval_perplexity": perplexity, "eval_accuracy": accuracy}, commit=False, ) last_eval_acc = accuracy model.train() start = time.time() total_tokens = torch.tensor([0], device="cuda", dtype=torch.long) total_loss = torch.tensor([0], device="cuda", dtype=torch.float) total_accurate_count = torch.tensor([0], device="cuda", dtype=torch.long) accum_tokens = torch.tensor([0], device="cuda", dtype=torch.long) display_tokens = torch.tensor([0], device="cuda", dtype=torch.long) display_tokens_postdesc = torch.tensor([0], device="cuda", dtype=torch.long) display_loss = torch.tensor([0], device="cuda", dtype=torch.float) display_accurate_count = torch.tensor([0], device="cuda", dtype=torch.long) display_accurate_count_postdesc = torch.tensor([0], device="cuda", dtype=torch.long) display_overflow_examples = torch.tensor([0], device="cuda", dtype=torch.long) display_batch_underruns = torch.tensor([0], device="cuda", dtype=torch.long) display_batch_fill = torch.tensor([0], device="cuda", dtype=torch.float) n_accum = 0 tqdm_total = num_batches if max_examples_per_rank is not None: tqdm_total = min( tqdm_total if tqdm_total is not None else +math.inf, max_examples_per_rank // batch_size, ) iterator = (i for i in tqdm(data_iter, total=tqdm_total, disable=quiet)) batch_q = Queue(20) batch_evt = Event() def batch_thd_entry(): while True: try: nextbatch = next(iterator) except StopIteration: nextbatch = None except Exception as e: nextbatch = e batch_q.put(nextbatch) if batch_evt.is_set(): iterator.close() break if nextbatch is None: break batch_thd = Thread(target=batch_thd_entry) batch_thd.start() try: i = 0 verify_ranks_are_equal() while True: try: batch = batch_q.get(block=False) except Empty: display_batch_underruns[0] += 1 batch = batch_q.get() if isinstance(batch, Exception): raise batch if batch is None: break out = model( batch.tgt, encoder_input_ids=batch.encoder_input_ids, encoder_attention_mask=batch.encoder_attention_mask, ) loss_node = loss_compute(out, batch.tgt_y) accum_tokens += batch.ntokens if mode == "train" or mode == "train+log": if scaler is not None: scaler.scale(loss_node).backward() else: loss_node.backward() train_state.step += 1 train_state.samples += batch.tgt.shape[0] train_state.tokens += batch.ntokens if i % accum_iter == 0: if scaler is not None: scaler.unscale_(optimizer) # scale gradients by total number of non-pad tokens accumulated dist.all_reduce(accum_tokens) # compensate for DDP dividing by number of workers (no way to disable) if accum_tokens > 0: factor = world_size / accum_tokens for p in model.parameters(): if p.grad is not None: p.grad *= factor total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) if torch.logical_or(total_norm.isnan(), total_norm.isinf()): train_state.skipped_updates += 1 if scaler is not None: scaler.step(optimizer) scaler.update() else: optimizer.step() if eval_iter is not None and i // accum_iter % eval_iter == 0: eval() if ( enable_checkpoint and last_eval_acc > train_state.best_eval_acc ): checkpoint(i, last_eval_acc) train_state.best_eval_acc = last_eval_acc if eval_iter is None and enable_checkpoint and i % 1000 == 0: checkpoint(i, 0) optimizer.zero_grad(set_to_none=True) accum_tokens[0] = 0 n_accum += 1 train_state.accum_step += 1 scheduler.step() accurate_count = torch.count_nonzero( torch.logical_and( torch.argmax(out, dim=-1) == batch.tgt_y, batch.tgt[:, :, 0] != model.module.params.padding_idx, ) ) def map_zero_to_high(x): x[x == 0] = 1e6 return x description_anchor_argmax = map_zero_to_high( torch.argmax( ( batch.tgt[:, :, 0] == model.module.vocab.description_anchor.index ).to(torch.long), dim=-1, keepdim=True, ) ) postdesc_mask = torch.logical_and( torch.arange(batch.tgt.shape[1], dtype=torch.long, device="cuda") .unsqueeze(0) .expand(batch.tgt.shape[0], -1) > description_anchor_argmax, batch.tgt[:, :, 0] != model.module.params.padding_idx, ) tokens_postdesc = torch.count_nonzero(postdesc_mask) accurate_count_postdesc = torch.count_nonzero( torch.logical_and( torch.argmax(out, dim=-1) == batch.tgt_y, postdesc_mask ) ) del postdesc_mask batch_size = batch.tgt.shape[0] total_loss += loss_node.detach() total_tokens += batch.ntokens total_accurate_count += accurate_count display_tokens += batch.ntokens display_tokens_postdesc += tokens_postdesc display_loss += loss_node.detach() display_accurate_count += accurate_count display_accurate_count_postdesc += accurate_count_postdesc display_overflow_examples += torch.count_nonzero( description_anchor_argmax == 1e6 ) del description_anchor_argmax display_batch_fill += batch.fill if i % 40 == 1 and (mode == "train" or mode == "train+log"): lr = optimizer.param_groups[0]["lr"] dist.all_reduce(display_tokens) dist.all_reduce(display_tokens_postdesc) dist.all_reduce(display_loss) dist.all_reduce(display_accurate_count) dist.all_reduce(display_accurate_count_postdesc) dist.all_reduce(display_overflow_examples) dist.all_reduce(display_batch_underruns) dist.all_reduce(display_batch_fill) elapsed = time.time() - start tokens_per_second = display_tokens.item() / elapsed loss_per_token = display_loss.item() / display_tokens.item() accuracy = display_accurate_count.item() / display_tokens.item() accuracy_postdesc = display_accurate_count_postdesc.item() / ( display_tokens_postdesc.item() + 1e-6 ) batch_underruns = display_batch_underruns.item() batch_fill = display_batch_fill.item() / (40 * world_size) overflow_pd_ratio = display_overflow_examples.item() / ( 40 * world_size * batch_size ) if not quiet: tqdm.write( ( "Epo Step: %6d / %d | Acc Step: %3d | Loss: %6.2f " + "| Acc: %6.2f | AccPD: %6.2f | Tok / Sec: %7.1f | LR: %6.1e " + "| Fill: %6.2f | OverflowPD: %6.2f | Underruns: %d" ) % ( i, num_batches, n_accum, loss_per_token, accuracy, accuracy_postdesc, tokens_per_second, lr, batch_fill, overflow_pd_ratio, batch_underruns, ) ) start = time.time() display_tokens[0] = 0 display_tokens_postdesc[0] = 0 display_loss[0] = 0 display_accurate_count[0] = 0 display_accurate_count_postdesc[0] = 0 display_overflow_examples[0] = 0 display_batch_underruns[0] = 0 display_batch_fill[0] = 0 if wandb.run is not None: log_obj = { "loss": loss_per_token, "accuracy": accuracy, "accuracy_postdesc": accuracy_postdesc, "lr": lr, "tokens_per_second": tokens_per_second, "epoch": 1.0 + epoch + i / num_batches, "batch_fill": batch_fill, "overflow_examples": overflow_pd_ratio, "seq_len": batch.length, "batch_underruns": batch_underruns, } if scaler is not None: log_obj["scale"] = scaler.get_scale() log_obj["skipped_updates"] = train_state.skipped_updates wandb.log(log_obj) train_state.skipped_updates = 0 del loss_node if ( max_examples_per_rank is not None and train_state.step * batch_size > max_examples_per_rank ): break i += 1 return total_loss, total_tokens, total_accurate_count, train_state finally: batch_evt.set() try: batch_q.get(timeout=2) except Empty: pass def rate(step, model_size, factor, min_factor, steps_in_epoch): """ we have to default the step to 1 for LambdaLR function to avoid zero raising to negative power. """ max_lr = factor * model_size ** (-0.5) if step > steps_in_epoch: return min_factor * max_lr return max_lr * ( min_factor + (1 + math.cos(step * math.pi / steps_in_epoch)) / 2 * (1 - min_factor) )