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