""" WER evaluation on a LibriSpeech split, via greedy CTC and greedy attention decode. Usage (from project root, venv active): python -m src.evaluate --config configs/zipformer_s.yaml \ --checkpoint checkpoints/latest.pt --split test-clean """ import os os.environ.setdefault("PYTORCH_ENABLE_MPS_FALLBACK", "1") import argparse import warnings import jiwer import torch import yaml from torch.utils.data import DataLoader from tqdm import tqdm warnings.filterwarnings("ignore", message=".*output with one or more elements was resized.*") from src.dataset import ASRCollate, LibriSpeechASR from src.model import ASRModel from src.tokenizer import ASRTokenizer, BOS_ID, EOS_ID def ctc_greedy_decode(log_probs: torch.Tensor, lengths: torch.Tensor, blank_id: int) -> list: """log_probs: (B, T, V). Returns a list of token-id lists, one per sample.""" ids = log_probs.argmax(dim=-1) # (B, T) results = [] for i in range(ids.size(0)): length = int(lengths[i].item()) seq = ids[i, :length].tolist() collapsed = [] prev = None for tok in seq: if tok != prev and tok != blank_id: collapsed.append(tok) prev = tok results.append(collapsed) return results @torch.no_grad() def evaluate(model: ASRModel, tokenizer: ASRTokenizer, loader: DataLoader, device: torch.device, max_batches=None): model.eval() refs = [] ctc_hyps = [] attn_hyps = [] for i, batch in enumerate(tqdm(loader, desc="evaluating")): if max_batches is not None and i >= max_batches: break waveforms = batch["waveforms"].to(device) wave_lengths = batch["wave_lengths"].to(device) enc_out, enc_lengths, ctc_log_probs = model.forward_eval(waveforms, wave_lengths) ctc_ids = ctc_greedy_decode(ctc_log_probs, enc_lengths, blank_id=model.blank_id) attn_ids = model.decoder.greedy_decode(enc_out, enc_lengths, bos_id=BOS_ID, eos_id=EOS_ID, max_len=200) for ref, ctc_id_seq, attn_id_seq in zip(batch["transcripts"], ctc_ids, attn_ids): refs.append(ref.lower()) ctc_hyps.append(tokenizer.decode(ctc_id_seq)) attn_hyps.append(tokenizer.decode(attn_id_seq)) ctc_wer = jiwer.wer(refs, ctc_hyps) attn_wer = jiwer.wer(refs, attn_hyps) return {"ctc_wer": ctc_wer, "attn_wer": attn_wer, "refs": refs, "ctc_hyps": ctc_hyps, "attn_hyps": attn_hyps} def main(): parser = argparse.ArgumentParser() parser.add_argument("--config", required=True) parser.add_argument("--checkpoint", required=True) parser.add_argument("--split", default="test-clean") parser.add_argument("--max-batches", type=int, default=None) parser.add_argument("--batch-size", type=int, default=None) parser.add_argument("--show-examples", type=int, default=5) 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") tokenizer = ASRTokenizer(cfg["tokenizer_model"]) model = ASRModel(vocab_size=tokenizer.vocab_size, **cfg["model"]).to(device) ckpt = torch.load(args.checkpoint, map_location=device) model.load_state_dict(ckpt["model"]) print(f"Loaded checkpoint {args.checkpoint} (epoch {ckpt.get('epoch')}, step {ckpt.get('step')})") ds = LibriSpeechASR(cfg["data_root"], [args.split], download=False) collate = ASRCollate(tokenizer) batch_size = args.batch_size or cfg["batch_size"] loader = DataLoader(ds, batch_size=batch_size, shuffle=False, collate_fn=collate) results = evaluate(model, tokenizer, loader, device, max_batches=args.max_batches) print(f"\nCTC (greedy) WER: {results['ctc_wer'] * 100:.2f}%") print(f"Attention (greedy) WER: {results['attn_wer'] * 100:.2f}%") n = min(args.show_examples, len(results["refs"])) for i in range(n): print(f"\nREF : {results['refs'][i]}") print(f"CTC : {results['ctc_hyps'][i]}") print(f"ATTN: {results['attn_hyps'][i]}") if __name__ == "__main__": main()