""" Training entrypoint, driven by a YAML config (see configs/zipformer_s.yaml). Usage (from project root, venv active): python -m src.train --config configs/zipformer_s.yaml python -m src.train --config configs/zipformer_s.yaml --resume checkpoints/latest.pt """ import os # Must be set before any CTC loss call: torch.ctc_loss has no MPS kernel, so # this falls back to CPU for that one op instead of crashing. os.environ.setdefault("PYTORCH_ENABLE_MPS_FALLBACK", "1") import argparse import math import pathlib import warnings import torch import yaml from torch.utils.data import DataLoader from torch.utils.tensorboard import SummaryWriter from tqdm import tqdm from src.dataset import ASRCollate, LibriSpeechASR from src.model import ASRModel, count_parameters from src.tokenizer import ASRTokenizer # Cosmetic: torch.stft warns on every MelSpectrogram call. Harmless log noise. warnings.filterwarnings("ignore", message=".*output with one or more elements was resized.*") def build_lr_lambda(warmup_steps: int, total_steps: int): def lr_lambda(step): step = step + 1 if step < warmup_steps: return step / warmup_steps progress = (step - warmup_steps) / max(1, total_steps - warmup_steps) progress = min(1.0, progress) return 0.5 * (1.0 + math.cos(math.pi * progress)) return lr_lambda def batch_to_device(batch: dict, device: torch.device) -> dict: return { k: (v.to(device) if torch.is_tensor(v) else v) for k, v in batch.items() } @torch.no_grad() def validate(model: ASRModel, dev_loader: DataLoader, device: torch.device, max_batches: int = 50) -> float: model.eval() total_loss = 0.0 count = 0 for i, batch in enumerate(dev_loader): if i >= max_batches: break b = batch_to_device(batch, device) out = model.forward_val_loss( b["waveforms"], b["wave_lengths"], b["ctc_targets"], b["ctc_target_lengths"], b["decoder_input"], b["decoder_input_lengths"], b["decoder_target"], ) total_loss += out["loss"].item() count += 1 model.train() return total_loss / max(1, count) def save_checkpoint(path, model, optimizer, scheduler, epoch, step): torch.save( { "model": model.state_dict(), "optimizer": optimizer.state_dict(), "scheduler": scheduler.state_dict(), "epoch": epoch, "step": step, }, path, ) def main(): parser = argparse.ArgumentParser() parser.add_argument("--config", required=True) parser.add_argument("--resume", default=None, help="Checkpoint path to resume from") args = parser.parse_args() with open(args.config) as f: cfg = yaml.safe_load(f) device = torch.device("mps" if torch.backends.mps.is_available() else "cpu") print(f"Using device: {device}") tokenizer = ASRTokenizer(cfg["tokenizer_model"]) print(f"Tokenizer vocab size: {tokenizer.vocab_size}") train_ds = LibriSpeechASR(cfg["data_root"], cfg["train_splits"], download=False) dev_ds = LibriSpeechASR(cfg["data_root"], cfg["dev_splits"], download=False) print(f"Train utterances: {len(train_ds)} | dev utterances: {len(dev_ds)}") collate = ASRCollate(tokenizer) train_loader = DataLoader( train_ds, batch_size=cfg["batch_size"], shuffle=True, collate_fn=collate, num_workers=cfg.get("num_workers", 4), drop_last=True, ) dev_loader = DataLoader( dev_ds, batch_size=cfg["batch_size"], shuffle=False, collate_fn=collate, num_workers=cfg.get("num_workers", 2), ) model = ASRModel(vocab_size=tokenizer.vocab_size, **cfg["model"]).to(device) print(f"Model parameters: {count_parameters(model) / 1e6:.2f}M") opt = torch.optim.AdamW( model.parameters(), lr=cfg["lr"], weight_decay=cfg.get("weight_decay", 0.01) ) steps_per_epoch = len(train_loader) total_steps = steps_per_epoch * cfg["epochs"] warmup_steps = cfg.get("warmup_steps", 2000) scheduler = torch.optim.lr_scheduler.LambdaLR(opt, build_lr_lambda(warmup_steps, total_steps)) ckpt_dir = pathlib.Path(cfg.get("checkpoint_dir", "checkpoints")) ckpt_dir.mkdir(parents=True, exist_ok=True) log_dir = pathlib.Path(cfg.get("log_dir", "logs")) log_dir.mkdir(parents=True, exist_ok=True) writer = SummaryWriter(log_dir=str(log_dir)) start_epoch = 0 global_step = 0 if args.resume: ckpt = torch.load(args.resume, map_location=device) model.load_state_dict(ckpt["model"]) opt.load_state_dict(ckpt["optimizer"]) scheduler.load_state_dict(ckpt["scheduler"]) start_epoch = ckpt["epoch"] global_step = ckpt["step"] print(f"Resumed from {args.resume} at epoch {start_epoch}, step {global_step}") grad_clip = cfg.get("grad_clip", 5.0) log_every = cfg.get("log_every", 50) val_every = cfg.get("val_every", 1000) save_every = cfg.get("save_every", 1000) use_amp = cfg.get("use_amp", True) and device.type == "mps" model.train() for epoch in range(start_epoch, cfg["epochs"]): pbar = tqdm(train_loader, desc=f"epoch {epoch}") for batch in pbar: b = batch_to_device(batch, device) opt.zero_grad() with torch.autocast(device_type="mps", dtype=torch.bfloat16, enabled=use_amp): out = model.forward_train( b["waveforms"], b["wave_lengths"], b["ctc_targets"], b["ctc_target_lengths"], b["decoder_input"], b["decoder_input_lengths"], b["decoder_target"], ) loss = out["loss"] loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip) opt.step() scheduler.step() global_step += 1 if global_step % log_every == 0: lr = scheduler.get_last_lr()[0] pbar.set_postfix( loss=f"{out['loss'].item():.3f}", ctc=f"{out['ctc_loss'].item():.3f}", ce=f"{out['ce_loss'].item():.3f}", lr=f"{lr:.2e}", ) writer.add_scalar("train/loss", out["loss"].item(), global_step) writer.add_scalar("train/ctc_loss", out["ctc_loss"].item(), global_step) writer.add_scalar("train/ce_loss", out["ce_loss"].item(), global_step) writer.add_scalar("train/cr_loss", out["cr_loss"].item(), global_step) writer.add_scalar("train/lr", lr, global_step) if global_step % val_every == 0: val_loss = validate(model, dev_loader, device) print(f"\n[step {global_step}] val_loss={val_loss:.4f}") writer.add_scalar("val/loss", val_loss, global_step) if global_step % save_every == 0: save_checkpoint(ckpt_dir / f"step{global_step}.pt", model, opt, scheduler, epoch, global_step) save_checkpoint(ckpt_dir / "latest.pt", model, opt, scheduler, epoch, global_step) save_checkpoint(ckpt_dir / f"epoch{epoch}.pt", model, opt, scheduler, epoch + 1, global_step) save_checkpoint(ckpt_dir / "latest.pt", model, opt, scheduler, epoch + 1, global_step) writer.close() if __name__ == "__main__": main()