import os import gc import json import time import math import wandb import torch import random import datetime import argparse import numpy as np import torch.nn as nn import torch.distributed as dist import torchaudio.functional as AF import torchaudio.transforms as AT from tqdm import tqdm from suno_utils.audio import Audio from colorama import Fore, Style from dataclasses import dataclass from datetime import timedelta from torch.nn import functional as F from suno_utils.utils.text import read_jsonl from typing import List, Tuple, Optional from torch.distributed import barrier, is_initialized, init_process_group RUN_START_TIME = time.strftime("%Y-%m-%d_%H-%M-%S") CHECKPOINT_DIR = f"/app2/suno/checkpoints/{RUN_START_TIME}_s{random.randint(0, 9999)}" os.makedirs(CHECKPOINT_DIR, exist_ok=True) def is_ddp(): return int(os.environ.get("RANK", -1)) != -1 def is_master(): if is_ddp(): return int(os.environ["RANK"]) == 0 return True def dist_barrier(): if is_initialized(): barrier() def print_with_time(content): """Print the content with the current time.""" print(f"[{datetime.datetime.now().strftime('%Y-%m-%d_%H:%M:%S')}]: {content}") def print_with_time_master(content): if is_master(): print_with_time(content) FULL_PRECISION_KEY_FRAGMENTS = ("pos_emb", "inv_freq") def convert_to_precision(module, weights_precision=torch.float16): for p_name, param in module.named_parameters(): if not any(s in p_name for s in FULL_PRECISION_KEY_FRAGMENTS): param.data = param.data.to(weights_precision) def setup_distributed(master_addr, master_port): # Set environment variables based on parsed arguments os.environ["MASTER_ADDR"] = str(master_addr) os.environ["MASTER_PORT"] = str(master_port) if "SLURM_PROCID" in os.environ: # Running on SLURM if int(os.environ["SLURM_NTASKS_PER_NODE"]) != torch.cuda.device_count(): raise ValueError( f"SLURM_NTASKS_PER_NODE ({os.environ['SLURM_NTASKS_PER_NODE']}) does not match" f" the number of CUDA devices ({torch.cuda.device_count()}) on node {os.environ['HOSTNAME']}" ) rank = int(os.environ["SLURM_PROCID"]) local_rank = int(os.environ["SLURM_LOCALID"]) world_size = int(os.environ["SLURM_JOB_NUM_NODES"]) * int( os.environ["SLURM_NTASKS_PER_NODE"] ) else: # Running locally rank = 0 local_rank = 0 world_size = 1 os.environ["RANK"] = str(rank) os.environ["LOCAL_RANK"] = str(local_rank) print(f"Initializing distributed process group on rank {local_rank}") torch.cuda.set_device(local_rank) try: dist.init_process_group( backend="nccl", timeout=timedelta(hours=6), rank=rank, world_size=world_size, device_id=torch.device(f"cuda:{local_rank}"), ) except Exception as e: print(f"Distributed error on rank {rank} with host {os.environ['HOSTNAME']}") raise e print(f"Done initializing on rank: {dist.get_rank()}") # barrier to check if nccl is working dist_barrier() print_with_time_master("distributed setup ready.") # ----------------- # dataset stuff # ----------------- def _fast_trim_mono( x: np.ndarray, # shape: (samples,), float32/64 in [-1, 1] sr: int, # sample rate (Hz) thresh_db_rel: float = -35, # keep where RMS > max_RMS + thresh (dB) win_ms: float = 20.0, # moving RMS window size (ms) pad_ms: float = 20.0, # pad around kept regions (ms) min_keep_ms: float = 40.0, # drop kept bits shorter than this (ms) ) -> Tuple[np.ndarray, List[Tuple[int, int]]]: """ Ultra-fast silence trimmer for mono audio. No convolutions, all O(n). Returns (trimmed_audio, kept_spans) with kept_spans in original sample indices. """ assert x.ndim == 1, "Expected mono waveform of shape (samples,)" n = x.size if n == 0: return x[:0], [] # --- Moving RMS via cumulative sums (box filter), O(n) --- # Compute moving average of power over a window, then sqrt. win = max(1, int(round(sr * win_ms / 1000.0))) if win > n: win = n # power and cumulative sum (use float64 for numeric safety) sq = x.astype(np.float64) ** 2 csum = np.empty(n + 1, dtype=np.float64) csum[0] = 0.0 np.cumsum(sq, out=csum[1:]) # csum[k] = sum_{i thresh_db_rel # True = keep # --- Turn mask into spans, expand by pad, merge, drop short --- pad = max(0, int(round(sr * pad_ms / 1000.0))) min_keep = max(1, int(round(sr * min_keep_ms / 1000.0))) # Find rising/falling edges m = mask.astype(np.int8) edges = np.flatnonzero(np.diff(m, prepend=0, append=0)) # edges come in pairs [start0, end0, start1, end1, ...] starts = edges[::2] ends = edges[1::2] if starts.size == 0: return x[:0], [] # Expand by pad and clamp starts = np.maximum(0, starts - pad) ends = np.minimum(n, ends + pad) # Merge overlaps and drop short spans spans: List[Tuple[int, int]] = [] s_prev = int(starts[0]) e_prev = int(ends[0]) for s, e in zip(starts[1:], ends[1:]): s = int(s) e = int(e) if s <= e_prev: # overlap/adjacent -> merge e_prev = max(e_prev, e) else: if (e_prev - s_prev) >= min_keep: spans.append((s_prev, e_prev)) s_prev, e_prev = s, e # last span if (e_prev - s_prev) >= min_keep: spans.append((s_prev, e_prev)) if not spans: return x[:0], [] # --- Concatenate kept spans (one pass) --- parts = [x[a:b] for (a, b) in spans] y = np.concatenate(parts, axis=0).astype(x.dtype) return y, spans def _get_segments( audio_np: np.ndarray, sample_rate: int, num_segments: int = 3, segment_duration_sec: float = 3.0, ): """ audio_np: np.ndarray of shape (samples,) sample_rate: int num_segments: int, number of random segments to crop and return segment_duration_sec: float, length (seconds) of each segment to return Returns: segments: list of np.ndarray of shape (segment_samples,) """ total_samples = len(audio_np) segment_samples = int(segment_duration_sec * sample_rate) if total_samples < segment_samples: # Pad if audio is too short padded = np.pad(audio_np, (0, segment_samples - total_samples), mode="constant") return [padded.copy() for _ in range(num_segments)] segments = [] for _ in range(num_segments): start_idx = np.random.randint(0, total_samples - segment_samples + 1) seg = audio_np[start_idx : start_idx + segment_samples] segments.append(seg.copy()) return segments # collate function def artist_segment_collate_fn(batch): """ Collate function for batching artist segment samples. Each element in `batch` is a tuple: (artist_id: str, artist_index: int, segments: List[np.ndarray]) Returns: flat_artist_ids: List[str] # len = batch_size * num_segments flat_artist_indices: torch.LongTensor # shape (batch_size * num_segments,) flat_segments: torch.FloatTensor # shape (batch_size * num_segments, segment_samples) """ flat_artist_ids = [] flat_artist_indices = [] flat_segments = [] for artist_id, artist_index, segments in batch: # segments: list of np.ndarray (num_segments, segment_samples) for seg in segments: flat_artist_ids.append(artist_id) flat_artist_indices.append(artist_index) flat_segments.append(torch.tensor(seg, dtype=torch.float32)) flat_artist_indices = torch.tensor(flat_artist_indices, dtype=torch.long) flat_segments = torch.stack( flat_segments, dim=0 ) # (batch_size * num_segments, segment_samples) return flat_artist_ids, flat_artist_indices, flat_segments class BasicIterableDataset(torch.utils.data.IterableDataset): def __init__(self, metas, num_segments: int = 3, segment_duration_s: float = 3.0): super(BasicIterableDataset, self).__init__() self.metas = metas self.num_segments = num_segments self.segment_duration_s = segment_duration_s # Create artist_id to index mapping for classifier self.artist_ids = list(metas.keys()) self.artist_id_to_index = { artist_id: idx for idx, artist_id in enumerate(self.artist_ids) } self.num_artists = len(self.artist_ids) print(f"Dataset initialized with {self.num_artists} artists") def __iter__(self): while True: # Randomly sample an artist artist_id = random.choice(self.artist_ids) artist_index = self.artist_id_to_index[artist_id] # Randomly sample a stem from this artist stems_dicts = self.metas[artist_id] stem_dict = random.choice(stems_dicts) # load this audio file audio = Audio.from_file(stem_dict["path"], n_channels=1) # trim silence audio_trim, _ = _fast_trim_mono(audio.array_float, audio.sample_rate) # get segments segments = _get_segments( audio_trim, audio.sample_rate, self.num_segments, self.segment_duration_s, ) # don't yield the segments separately # we will return list with the artist ids and indices # and then use a special collate function to merge them yield (artist_id, artist_index, segments) # ----------------- # model stuff # ----------------- # ---------------------------- # TDNN building blocks # ---------------------------- class TDNNBlock(nn.Module): """ 1D time-dilated conv (x-vector style) with ReLU+BN. Input: (B, C_in, T) Output: (B, C_out, T) """ def __init__(self, c_in: int, c_out: int, kernel: int = 5, dilation: int = 1): super().__init__() pad = dilation * (kernel // 2) self.conv = nn.Conv1d( c_in, c_out, kernel_size=kernel, dilation=dilation, padding=pad, bias=False ) self.bn = nn.BatchNorm1d(c_out) self.act = nn.ReLU(inplace=True) def forward(self, x): x = self.conv(x) x = self.bn(x) x = self.act(x) return x class StatsPooling(nn.Module): """ Mean+Std pooling over time. Input: (B, C, T) Output: (B, 2*C) """ def forward(self, x): # x: (B, C, T) mean = x.mean(dim=-1) std = x.std(dim=-1, unbiased=False) return torch.cat([mean, std], dim=1) # ---------------------------- # Speaker Encoder -> Embedding # ---------------------------- @dataclass class EncoderConfig: n_mels: int = 80 tdnn_channels: tuple = (256, 512, 1024, 2048, 2048) # 5 TDNN layers tdnn_kernels: tuple = (5, 5, 7, 1, 1) tdnn_dilations: tuple = (1, 2, 3, 1, 1) emb_hidden: int = 256 # penultimate projection before final embedding emb_dim: int = 128 # final embedding dimension dropout: float = 0.1 class SpeakerEncoder(nn.Module): """ Input: log-mel features (B, F, T); F = n_mels Output: L2-normalized embedding (B, emb_dim) """ def __init__(self, cfg: EncoderConfig): super().__init__() self.cfg = cfg C = [cfg.n_mels] + list(cfg.tdnn_channels) self.tdnn = nn.Sequential( *[ TDNNBlock( C[i], C[i + 1], kernel=cfg.tdnn_kernels[i], dilation=cfg.tdnn_dilations[i], ) for i in range(len(cfg.tdnn_channels)) ] ) self.pool = StatsPooling() # (B, 2*C_last) pooled_dim = 2 * cfg.tdnn_channels[-1] self.fc1 = nn.Linear(pooled_dim, cfg.emb_hidden, bias=False) self.bn1 = nn.BatchNorm1d(cfg.emb_hidden) self.drop = nn.Dropout(p=cfg.dropout) self.fc2 = nn.Linear(cfg.emb_hidden, cfg.emb_dim, bias=True) # Kaiming for convs; Xavier for linears is fine self.apply(self._init_weights) @staticmethod def _init_weights(m): if isinstance(m, nn.Conv1d): nn.init.kaiming_normal_(m.weight, nonlinearity="relu") elif isinstance(m, nn.Linear): nn.init.xavier_uniform_(m.weight) if m.bias is not None: nn.init.zeros_(m.bias) def forward(self, mels: torch.Tensor) -> torch.Tensor: """ mels: (B, F, T) returns: L2-normalized embeddings (B, emb_dim) """ x = mels # (B, F, T) x = self.tdnn(x) # (B, C, T) x = self.pool(x) # (B, 2C) x = self.fc1(x) # (B, H) x = self.bn1(x) x = F.relu(x, inplace=True) x = self.drop(x) x = self.fc2(x) # (B, D) # L2 normalize to put on the unit hypersphere x = F.normalize(x, p=2, dim=-1) return x @torch.no_grad() def embed(self, mels: torch.Tensor) -> torch.Tensor: self.eval() return self.forward(mels) # ---------------------------- # ArcFace / AAM-Softmax Head # ---------------------------- class ArcMarginProduct(nn.Module): """ Implements AAM-Softmax (ArcFace) logits on-the-fly. - Weight matrix W is L2-normalized per row. - Inputs are expected already L2-normalized. logits = s * cos(theta + m) for the target class, s * cos(theta) otherwise. Args: in_features: embedding dim D num_classes: number of speakers C s: scale (30 is common) m: angular margin (0.2~0.3 common) easy_margin: if True, use the easy-margin variant ls_eps: label smoothing epsilon (optional) """ def __init__( self, in_features: int, num_classes: int, s: float = 30.0, m: float = 0.2, easy_margin: bool = False, ls_eps: float = 0.0, ): super().__init__() self.in_features = in_features self.num_classes = num_classes self.s = s self.m = m self.easy_margin = easy_margin self.ls_eps = ls_eps self.weight = nn.Parameter(torch.empty(num_classes, in_features)) nn.init.xavier_uniform_(self.weight) # Precompute margin constants self.cos_m = math.cos(m) self.sin_m = math.sin(m) self.th = math.cos(math.pi - m) # cos(pi - m) self.mm = math.sin(math.pi - m) * m def forward(self, emb: torch.Tensor, labels: torch.Tensor) -> torch.Tensor: """ emb: (B, D) L2-normalized labels: (B,) int64 in [0, C-1] returns: scaled logits for CE loss, shape (B, C) """ # Normalize class weights W = F.normalize(self.weight, p=2, dim=1) # (C, D) # Cosine similarity between emb and each class weight # cos_theta: (B, C) cos_theta = torch.matmul(emb, W.t()).clamp(-1.0, 1.0) # Gather target cosine idx = torch.arange(emb.size(0), device=emb.device) cos_theta_y = cos_theta[idx, labels] # (B,) # Compute cos(theta + m) via trig identity sin_theta_y = torch.sqrt(torch.clamp(1.0 - cos_theta_y * cos_theta_y, min=0.0)) cos_theta_m = cos_theta_y * self.cos_m - sin_theta_y * self.sin_m # (B,) if self.easy_margin: # If easy margin: if cos(theta_y) > 0 use margin, else keep original cos cond = (cos_theta_y > 0).to(cos_theta.dtype) cos_theta_y_m = cond * cos_theta_m + (1 - cond) * cos_theta_y else: # Classic ArcFace margin decision cond = (cos_theta_y > self.th).to(cos_theta.dtype) cos_theta_y_m = cond * cos_theta_m + (1 - cond) * (cos_theta_y - self.mm) # Replace target logit logits = cos_theta.clone() logits[idx, labels] = cos_theta_y_m # Scale logits = logits * self.s # Optional label smoothing (applied in CE). We return logits; apply CE outside, # but provide a helper to build smoothed targets if needed. return logits # ---------------------------- # Full Model = Encoder + ArcFace head # ---------------------------- class SpeakerEmbeddingModel(nn.Module): """ Training: logits = model(mels, labels) -> (B, C) # pass to nn.CrossEntropyLoss Inference: emb = model.embed(mels) -> (B, D) L2-normalized """ def __init__( self, n_classes: int, cfg: Optional[EncoderConfig] = None, s: float = 30.0, m: float = 0.2, easy_margin: bool = False, ls_eps: float = 0.0, ): super().__init__() self.cfg = cfg or EncoderConfig() self.encoder = SpeakerEncoder(self.cfg) self.arcface = ArcMarginProduct( in_features=self.cfg.emb_dim, num_classes=n_classes, s=s, m=m, easy_margin=easy_margin, ls_eps=ls_eps, ) def forward(self, mels: torch.Tensor, labels: Optional[torch.Tensor] = None): """ mels: (B, F, T) float labels: (B,) long, required for training with ArcFace """ emb = self.encoder(mels) # (B, D), L2-normalized if labels is None: return emb logits = self.arcface(emb, labels) # (B, C) return logits @torch.no_grad() def embed(self, mels: torch.Tensor) -> torch.Tensor: return self.encoder.embed(mels) def wav_to_logmels( wav: torch.Tensor, sr: int, target_sr: int = 16_000, n_mels: int = 80, win_ms: float = 25, hop_ms: float = 10, fmin: float = 20.0, fmax: float | None = None, top_db: float = 80.0, cmvn: bool = True, ) -> torch.Tensor: # mix to mono if (B, C, T) if wav.dim() == 3: wav = wav.mean(1) # resample if needed if sr != target_sr: wav = AF.resample(wav, sr, target_sr) n_fft = int(target_sr * win_ms / 1000) win_length = n_fft hop_length = int(target_sr * hop_ms / 1000) # Ensure MelSpectrogram and its buffers are on the same device as input device = wav.device # Build MelSpectrogram on correct device & move module to input's device mel_spect = AT.MelSpectrogram( sample_rate=target_sr, n_fft=n_fft, win_length=win_length, hop_length=hop_length, f_min=fmin, f_max=fmax or target_sr / 2, n_mels=n_mels, window_fn=lambda window_length: torch.hann_window(window_length, device=device), power=2.0, mel_scale="slaney", norm="slaney", ) mel_spect = mel_spect.to(device) mel = mel_spect(wav) # (B, n_mels, T) # AmplitudeToDB -- ensure on same device as well db_xfm = AT.AmplitudeToDB(stype="power", top_db=top_db).to(device) logmel = db_xfm(mel) if cmvn: mu = logmel.mean(dim=(-1, -2), keepdim=True) sd = logmel.std(dim=(-1, -2), keepdim=True).clamp_min(1e-5) logmel = (logmel - mu) / sd return logmel # ----------------- # Training loop # ----------------- def train_step(model, batch, pretrain=False): model.train() pos_vae, neg_vae = batch pos_vae = pos_vae.cuda() neg_vae = neg_vae.cuda() if pretrain: acc = torch.tensor(0.0) batch_z = pos_vae if np.random.random() < 0.5 else neg_vae loss, _ = model(batch_z, mask_ratio=0.7) return loss, acc else: pos_logits = model(pos_vae) neg_logits = model(neg_vae) labels = torch.zeros_like(pos_logits) loss = bradley_terry_loss(pos_logits, neg_logits, labels) # compute accuracy preds = (pos_logits > neg_logits).float() acc = (preds == torch.ones_like(preds)).float().mean() return loss, acc def validate(model, dataloader, pretrain=False): model.eval() total_loss = 0 total_acc = 0 for idx, batch in enumerate(dataloader): with torch.no_grad(): loss, acc = train_step(model, batch, pretrain=pretrain) total_loss += loss.item() total_acc += acc.item() return total_loss / len(dataloader), total_acc / len(dataloader) def log_metrics(master_process, loss, grad_norm, acc, world_size, global_step, epoch): metrics = torch.tensor([loss.item(), grad_norm, acc]).to(loss.device) dist.all_reduce(metrics, op=dist.ReduceOp.SUM) metrics = metrics / world_size avg_loss, avg_grad_norm, avg_acc = metrics[:3] if master_process: log_data = { "train/loss": avg_loss.item(), "train/grad_norm": avg_grad_norm.item(), "train/acc": avg_acc.item(), "trainer/global_step": global_step, "trainer/epoch": epoch, } wandb.log(log_data) print_with_time_master( f"{Fore.GREEN}Training metrics:{Style.RESET_ALL} " + ", ".join( [ ( f"{Fore.YELLOW}{k.replace('train/', '')}:{Style.RESET_ALL} {v:.6f}" if isinstance(v, float) else f"{Fore.YELLOW}{k.replace('train/', '')}:{Style.RESET_ALL} {v}" ) for k, v in log_data.items() ] ) ) def log_val_metrics(master_process, loss, acc, global_step, epoch): if master_process: log_data = { "val/loss": loss, "val/acc": acc, "trainer/global_step": global_step, "trainer/epoch": epoch, } wandb.log(log_data) print_with_time_master( f"{Fore.RED}Validation metrics:{Style.RESET_ALL} " + ", ".join( [ f"{Fore.YELLOW}{k.replace('val/', '')}:{Style.RESET_ALL} {v:.6f}" for k, v in log_data.items() ] ) ) def save_checkpoint( model, optimizer, lr_scheduler, global_step, checkpoint_dir, run_config, epoch, ckpt_name="last_ckpt.pt", ): checkpoint = { "model": model.state_dict(), "optimizer": optimizer.state_dict(), "lr_scheduler": lr_scheduler.state_dict(), "global_step": global_step, "run_config": run_config, "epoch": epoch, } checkpoint_path = os.path.join(checkpoint_dir, ckpt_name) torch.save(checkpoint, checkpoint_path) print_with_time_master(f"Saved checkpoint to {checkpoint_path}") def parse_args(): parser = argparse.ArgumentParser(description="Train the diffusion model") parser.add_argument( "--master_addr", type=str, default="localhost", help="Master node address" ) parser.add_argument( "--master_port", type=str, default="12355", help="Master node port" ) parser.add_argument( "--config_path", type=str, default=None, help="Path to config file to override" ) return parser.parse_args() if __name__ == "__main__": args = parse_args() setup_distributed(args.master_addr, args.master_port) master_process = int(os.environ["RANK"]) == 0 # load run config with open(args.config_path, "r") as f: run_config = json.load(f) # Initialize datasets print_with_time_master("loading datasets...") train_dataset = BasicIterableDataset( run_config["dataset"]["train_metas"], num_segments=run_config["dataset"]["num_segments"], segment_duration_s=run_config["dataset"]["segment_duration_s"], ) val_dataset = BasicIterableDataset( run_config["dataset"]["val_metas"], num_segments=run_config["dataset"]["num_segments"], segment_duration_s=run_config["dataset"]["segment_duration_s"], ) # Setup ddp_rank = int(os.environ["RANK"]) ddp_local_rank = int(os.environ["LOCAL_RANK"]) world_size = dist.get_world_size() group_size = min(world_size, 8) device = f"cuda:{ddp_local_rank}" torch.cuda.set_device(device) master_process = ddp_rank == 0 # Initialize wandb for logging (it will be disabled if in debug mode) if master_process: wandb.init( project=run_config["training"]["wandb_project"], name=run_config["training"]["wandb_name"], config={ "run_config": run_config, "world_size": world_size, "slurm_id": os.environ.get("SLURM_JOB_ID"), "slurm_name": os.environ.get("SLURM_JOB_NAME"), "slurm_script_path": os.environ.get("SLURM_SCRIPT_PATH"), "checkpoint_dir": CHECKPOINT_DIR, }, ) wandb.run.log_code(".") # Prepare data loaders train_dataloader = torch.utils.data.DataLoader( train_dataset, batch_size=run_config["training"]["batch_size"], drop_last=True, num_workers=run_config["training"]["num_workers"], ) val_dataloader = torch.utils.data.DataLoader( val_dataset, batch_size=run_config["training"]["batch_size"], drop_last=True, num_workers=0, ) # create the model print_with_time_master("setting up model...") model = SpeakerEmbeddingModel( n_classes=run_config["model"]["n_classes"], cfg=EncoderConfig() ) model.to(device) num_params = sum(p.numel() for p in model.parameters()) num_trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad) print_with_time_master(f"Number of model parameters: {num_params:,}") print_with_time_master(f"Number of trainable parameters: {num_trainable_params:,}") if run_config["training"]["compile"]: print_with_time_master("compiling model...") model_copy_if_compiled = model model = torch.compile(model, dynamic=False) # create the optimizer optimizer = torch.optim.AdamW(model.parameters(), lr=run_config["training"]["lr"]) # create the lr scheduler lr_scheduler = torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, T_max=run_config["training"]["max_steps"] ) torch.cuda.empty_cache() gc.collect() dist_barrier() print_with_time_master("finished model initialization.") torch.manual_seed(ddp_rank) random.seed(ddp_rank) np.random.seed(ddp_rank) global_step = 0 epoch = 0 best_val_loss = np.inf # train print_with_time_master("starting training...") log_every = run_config["training"]["log_every"] val_every = run_config["training"]["val_every"] while global_step < run_config["training"]["max_steps"]: for batch in tqdm( train_dataloader, desc=f"{Fore.CYAN}Epoch {epoch + 1} - Training{Style.RESET_ALL}", disable=True, ): if val_every > 0: if global_step % val_every == 0: val_loss, val_acc = validate( model, val_dataloader, pretrain=run_config["training"]["pretrain"], ) log_val_metrics( master_process, val_loss, val_acc, global_step, epoch ) # Save checkpoint every ckpt_every steps if ( run_config["training"]["ckpt_every"] > 0 and global_step % run_config["training"]["ckpt_every"] == 0 ): save_checkpoint( model, optimizer, lr_scheduler, global_step, CHECKPOINT_DIR, run_config, epoch, ckpt_name="last_ckpt.pt", ) # Save checkpoint if val_acc improves if "val_loss" in locals() and val_loss < best_val_loss: best_val_loss = val_loss save_checkpoint( model, optimizer, lr_scheduler, global_step, CHECKPOINT_DIR, run_config, epoch, ckpt_name="best_ckpt.pt", ) loss, acc = train_step( model, batch, pretrain=run_config["training"]["pretrain"] ) optimizer.zero_grad() loss.backward() grad_norm = torch.nn.utils.clip_grad_norm_( model.parameters(), run_config["training"]["grad_clip"] ) optimizer.step() lr_scheduler.step() global_step += 1 if global_step % log_every == 0: log_metrics( master_process, loss, grad_norm, acc, world_size, global_step, epoch, ) epoch += 1 # Cleanup dist_barrier() dist.destroy_process_group() wandb.finish()