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