Phase-1: from-scratch Zipformer-M CTC streaming (Hindi/Hinglish) + full training scripts
e146811 verified | #!/usr/bin/env python3 | |
| """Word-level WER eval for the char-CTC zipformer (full-context greedy). | |
| Loads an averaged checkpoint, runs CTC greedy over each eval CutSet, reconstructs words | |
| by joining char tokens and splitting on ▁, and computes true word WER vs the reference. | |
| Reports per-set WER and writes a JSON summary. | |
| NOTE: encoder runs in full-context (non-streaming) mode -> this is the model-quality WER | |
| (optimistic vs live streaming). The streaming server demonstrates real chunked decoding. | |
| """ | |
| import argparse | |
| import json | |
| import math | |
| import os | |
| import sys | |
| import torch | |
| from torch.utils.data import DataLoader | |
| sys.path.insert(0, "/root/icefall/egs/hindi/ASR/zipformer") | |
| sys.path.insert(0, "/root/icefall") | |
| from train import add_model_arguments, get_model, get_params # noqa | |
| from icefall.lexicon import Lexicon # noqa | |
| from icefall.checkpoint import ( # noqa | |
| average_checkpoints, | |
| average_checkpoints_with_averaged_model, | |
| find_checkpoints, | |
| load_checkpoint, | |
| ) | |
| from icefall.decode import ctc_greedy_search # noqa | |
| from icefall.utils import make_pad_mask, AttributeDict # noqa | |
| from lhotse import CutSet, Fbank, FbankConfig, load_manifest_lazy # noqa | |
| from lhotse.dataset import DynamicBucketingSampler, K2SpeechRecognitionDataset # noqa | |
| from lhotse.dataset.input_strategies import OnTheFlyFeatures # noqa | |
| LOG_EPS = math.log(1e-10) | |
| WB = "▁" | |
| def wer_counts(ref, hyp): | |
| n, m = len(ref), len(hyp) | |
| dp = list(range(m + 1)) | |
| for i in range(1, n + 1): | |
| prev = dp[0] | |
| dp[0] = i | |
| for j in range(1, m + 1): | |
| cur = dp[j] | |
| if ref[i - 1] == hyp[j - 1]: | |
| dp[j] = prev | |
| else: | |
| dp[j] = 1 + min(prev, dp[j], dp[j - 1]) | |
| prev = cur | |
| return dp[m], n | |
| def tokens_to_words(tokens): | |
| return "".join(tokens).replace(WB, " ").split() | |
| def get_dl(cuts_path, max_dur): | |
| cuts = load_manifest_lazy(cuts_path) | |
| ds = K2SpeechRecognitionDataset( | |
| input_strategy=OnTheFlyFeatures(Fbank(FbankConfig(num_mel_bins=80))), | |
| return_cuts=True, | |
| ) | |
| sampler = DynamicBucketingSampler(cuts, max_duration=max_dur, shuffle=False) | |
| return DataLoader(ds, sampler=sampler, batch_size=None, num_workers=2) | |
| def decode_set(model, lexicon, dl, device, causal): | |
| tot_edits = tot_words = 0 | |
| n_utt = 0 | |
| samples = [] | |
| for batch in dl: | |
| feature = batch["inputs"].to(device) | |
| sup = batch["supervisions"] | |
| feature_lens = sup["num_frames"].to(device) | |
| if causal: | |
| pad_len = 30 | |
| feature_lens = feature_lens + pad_len | |
| feature = torch.nn.functional.pad(feature, (0, 0, 0, pad_len), value=LOG_EPS) | |
| x, x_lens = model.encoder_embed(feature, feature_lens) | |
| mask = make_pad_mask(x_lens) | |
| x = x.permute(1, 0, 2) | |
| enc, enc_lens = model.encoder(x, x_lens, mask) | |
| enc = enc.permute(1, 0, 2) | |
| ctc_out = model.ctc_output(enc) | |
| hyp_tokens = ctc_greedy_search(ctc_out, enc_lens) | |
| refs = sup["text"] | |
| for i, ids in enumerate(hyp_tokens): | |
| hyp = tokens_to_words([lexicon.token_table[j] for j in ids]) | |
| ref = refs[i].replace(WB, " ").split() | |
| e, n = wer_counts(ref, hyp) | |
| tot_edits += e | |
| tot_words += n | |
| n_utt += 1 | |
| if len(samples) < 5: | |
| samples.append((" ".join(ref), " ".join(hyp))) | |
| wer = 100.0 * tot_edits / max(1, tot_words) | |
| return wer, tot_edits, tot_words, n_utt, samples | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--exp-dir", required=True) | |
| ap.add_argument("--lang-dir", required=True) | |
| ap.add_argument("--manifest-dir", required=True) | |
| ap.add_argument("--epoch", type=int, default=30) | |
| ap.add_argument("--avg", type=int, default=5) | |
| ap.add_argument("--use-averaged-model", type=int, default=1) | |
| ap.add_argument("--max-duration", type=int, default=200) | |
| ap.add_argument("--eval-sets", nargs="+", default=["iv_hi", "svarah", "call_test"]) | |
| ap.add_argument("--out", default=None) | |
| add_model_arguments(ap) | |
| args = ap.parse_args() | |
| params = get_params() | |
| params.update(vars(args)) | |
| device = torch.device("cuda", 0) | |
| lexicon = Lexicon(params.lang_dir) | |
| params.blank_id = lexicon.token_table["<blk>"] | |
| params.vocab_size = max(lexicon.tokens) + 1 | |
| params.decoding_method = "ctc-greedy-search" | |
| model = get_model(params) | |
| exp = params.exp_dir | |
| if params.use_averaged_model: | |
| start = params.epoch - params.avg + 1 | |
| fns = [f"{exp}/epoch-{e}.pt" for e in range(start, params.epoch + 1)] | |
| fns = [f for f in fns if os.path.exists(f)] | |
| assert fns, f"no ckpts for epoch {start}..{params.epoch} in {exp}" | |
| # simple average (avoids needing model_avg); fine for reporting | |
| model.load_state_dict(average_checkpoints(fns, device=device), strict=False) | |
| print(f"averaged {len(fns)} ckpts: {[os.path.basename(f) for f in fns]}") | |
| else: | |
| load_checkpoint(f"{exp}/epoch-{params.epoch}.pt", model) | |
| model.to(device).eval() | |
| nparam = sum(p.numel() for p in model.parameters()) | |
| print(f"model params: {nparam/1e6:.1f}M vocab={params.vocab_size}") | |
| results = {"epoch": params.epoch, "avg": params.avg, "params_M": round(nparam / 1e6, 1), "sets": {}} | |
| for name in params.eval_sets: | |
| cp = os.path.join(params.manifest_dir, f"cuts_eval_{name}.jsonl.gz") | |
| if not os.path.exists(cp): | |
| print(f"skip {name}: {cp} missing"); continue | |
| dl = get_dl(cp, params.max_duration) | |
| wer, e, n, nu, samples = decode_set(model, lexicon, dl, device, bool(params.causal)) | |
| results["sets"][name] = {"wer": round(wer, 2), "edits": e, "ref_words": n, "utts": nu} | |
| print(f"\n==== {name}: WER {wer:.2f}% ({e}/{n} words, {nu} utts) ====") | |
| for r, h in samples: | |
| print(f" REF: {r}") | |
| print(f" HYP: {h}") | |
| out = params.out or os.path.join(exp, f"wer_epoch{params.epoch}_avg{params.avg}.json") | |
| with open(out, "w") as f: | |
| json.dump(results, f, ensure_ascii=False, indent=2) | |
| print("\nSUMMARY:", json.dumps({k: v["wer"] for k, v in results["sets"].items()})) | |
| print("saved", out) | |
| if __name__ == "__main__": | |
| main() | |