asr / src /train.py
shubhexists's picture
Add Zipformer-inspired ASR model: weights, tokenizer, config, and training code
ce3c8df verified
Raw History Blame Contribute Delete
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()
}
@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()