| |
| """ASR-judge a folder of synthesized wavs: faster-whisper large-v3 -> WER/CER. |
| |
| Compares ASR transcript vs reference text (both normalized: tashkeel stripped, |
| punctuation removed, alef/yaa variants unified) using jiwer. |
| |
| Usage: python asr_eval.py --wavdir /opt/work/eval/baseline |
| Writes <wavdir>/asr_report.json and prints a summary table. |
| """ |
|
|
| import argparse |
| import json |
| import re |
| from pathlib import Path |
|
|
| import jiwer |
|
|
| TASHKEEL_RE = re.compile("[\u0610-\u061a\u064b-\u065f\u0670\u06d6-\u06dc\u06df-\u06e8\u06ea-\u06ed\u0640]") |
| PUNCT_RE = re.compile(r"[^\w\s]|[_]", re.UNICODE) |
| WS_RE = re.compile(r"\s+") |
|
|
|
|
| def norm(t: str) -> str: |
| t = TASHKEEL_RE.sub("", t) |
| t = t.replace("أ", "ا").replace("إ", "ا").replace("آ", "ا") |
| t = t.replace("ى", "ي").replace("ة", "ه") |
| t = PUNCT_RE.sub(" ", t) |
| return WS_RE.sub(" ", t).strip() |
|
|
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--wavdir", required=True) |
| ap.add_argument("--model", default="large-v3") |
| args = ap.parse_args() |
| wavdir = Path(args.wavdir) |
|
|
| refs = {} |
| for line in (wavdir / "timing.jsonl").read_text(encoding="utf-8").splitlines(): |
| r = json.loads(line) |
| if "text" in r: |
| refs[r["id"]] = r["text"] |
|
|
| from faster_whisper import WhisperModel |
|
|
| m = WhisperModel(args.model, device="cuda", compute_type="float16") |
|
|
| rows = [] |
| for sid, ref in refs.items(): |
| wav = wavdir / f"{sid}.wav" |
| if not wav.exists(): |
| rows.append({"id": sid, "error": "missing_wav"}) |
| continue |
| segs, _ = m.transcribe(str(wav), language="ar", beam_size=5, vad_filter=False) |
| hyp = " ".join(s.text for s in segs).strip() |
| r_n, h_n = norm(ref), norm(hyp) |
| wer = jiwer.wer(r_n, h_n) if r_n else 1.0 |
| cer = jiwer.cer(r_n, h_n) if r_n else 1.0 |
| rows.append( |
| {"id": sid, "ref": ref, "hyp": hyp, "wer": round(wer, 3), "cer": round(cer, 3)} |
| ) |
| print(f"{sid:12s} WER={wer:.2f} CER={cer:.2f} | {hyp[:70]}") |
|
|
| ok = [r for r in rows if "wer" in r] |
| summary = { |
| "n": len(rows), |
| "n_ok": len(ok), |
| "mean_wer": round(sum(r["wer"] for r in ok) / max(len(ok), 1), 4), |
| "mean_cer": round(sum(r["cer"] for r in ok) / max(len(ok), 1), 4), |
| } |
| print("SUMMARY", json.dumps(summary)) |
| (wavdir / "asr_report.json").write_text( |
| json.dumps({"summary": summary, "rows": rows}, ensure_ascii=False, indent=1), |
| encoding="utf-8", |
| ) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|