Download src/train.py from shubhexists/asr: direct link, hf CLI and curl.
- Browser
- Download file 7.49 kB
-
https://huggingface.co/shubhexists/asr/resolve/main/src/train.py
- Command line
-
hf download hf://shubhexists/asr/src/train.py
-
curl -L -o train.py https://huggingface.co/shubhexists/asr/resolve/main/src/train.py
7.49 kB
| """ | |
| 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() | |
| } | |
| 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() | |