Vengadanathan's picture
Phase-1: from-scratch Zipformer-M CTC streaming (Hindi/Hinglish) + full training scripts
e146811 verified
Raw
History Blame Contribute Delete
6.29 kB
#!/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)
@torch.no_grad()
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()