| |
| """Analyse profonde : d'ou vient l'erreur des champions sur la validation ? |
| Distingue clips pathologiques FIXABLES (longs, mal decodes) vs NON FIXABLES |
| (refs cassees a la source : audio court / texte long, mots/sec implausible).""" |
| import json |
| import numpy as np |
| import soundfile as sf |
| import torch |
| import jiwer |
| from transformers import AutoModelForCTC, AutoProcessor |
|
|
| MODELS = {"lin": "lin_s4_best", "sna": "sna_ps_best", "lug": "lug_r_best"} |
|
|
|
|
| def norm(s): |
| return " ".join(str(s).replace("|", " ").split()) |
|
|
|
|
| for LG, name in MODELS.items(): |
| P = f"/root/models/{name}" |
| rows = [json.loads(l) for l in open(f"/scratch/prep/manifests/waxal_{LG}_validation.jsonl", encoding="utf-8")] |
| rows = [r for r in rows if r.get("text", "").strip()] |
| proc = AutoProcessor.from_pretrained(P) |
| m = AutoModelForCTC.from_pretrained(P, torch_dtype=torch.bfloat16).cuda().eval() |
| rows.sort(key=lambda r: -r["duration"]) |
| out = {} |
| cur, bud, bs = [], 0.0, [] |
| for r in rows: |
| if cur and bud + r["duration"] > 120: |
| bs.append(cur); cur, bud = [], 0.0 |
| cur.append(r); bud += r["duration"] |
| if cur: |
| bs.append(cur) |
| with torch.inference_mode(): |
| for b in bs: |
| au = [sf.read(r["audio"], dtype="float32")[0] for r in b] |
| f = proc(au, sampling_rate=16000, return_tensors="pt", padding=True) |
| f = {k: v.to("cuda", dtype=torch.bfloat16 if v.dtype == torch.float32 else v.dtype) for k, v in f.items()} |
| ids = m(**f).logits.float().argmax(-1).cpu().numpy() |
| for r, s in zip(b, proc.batch_decode(ids)): |
| out[r["id"]] = norm(s) |
|
|
| patho, wps_all = [], [] |
| for r in rows: |
| ref = norm(r["text"]); hyp = out[r["id"]] |
| rw, hw, dur = len(ref.split()), len(hyp.split()), r["duration"] |
| wps = rw / max(dur, 0.1) |
| wer = jiwer.wer([ref], [hyp]) if ref else 0 |
| wps_all.append(wps) |
| if hw <= 2 or wer >= 1.0: |
| patho.append((r["id"], round(dur, 1), rw, hw, round(wps, 1), round(wer, 2))) |
|
|
| print(f"=== {LG} validation : {len(rows)} clips | pathologiques (hyp<=2 mots OU WER>=100%): " |
| f"{len(patho)} ({100*len(patho)/len(rows):.1f}%)") |
| if patho: |
| durs = [p[1] for p in patho]; wpss = [p[4] for p in patho]; rws = [p[2] for p in patho] |
| n_short = sum(1 for d in durs if d < 3) |
| n_wps = sum(1 for w in wpss if w > 4) |
| n_longok = sum(1 for _, d, _, _, w, _ in patho if d > 25 and w < 4) |
| print(f" duree patho med={np.median(durs):.1f}s max={max(durs):.1f}s | <3s: {n_short} | >25s&wps<4 (long fixable?): {n_longok}") |
| print(f" mots/sec patho med={np.median(wpss):.1f} | wps>4 (REF CASSEE): {n_wps}/{len(patho)} ({100*n_wps/len(patho):.0f}%)") |
| print(f" ref_mots patho med={int(np.median(rws))}") |
| n_all_bad = sum(1 for w in wps_all if w > 4) |
| print(f" TOUS clips avec mots/sec>4 (refs artefacts, non transcriptibles): {n_all_bad} ({100*n_all_bad/len(wps_all):.1f}%)") |
| print(" 4 refs longues patho:", sorted(patho, key=lambda x: -x[2])[:4]) |
| print("ANALYSE_DONE") |
|
|