waxal2026-backup / scripts /route_devhard.py
Pricile's picture
compactage apres suppression luganda
6eed659
Raw
History Blame Contribute Delete
5.54 kB
#!/usr/bin/env python3
"""GATE du pipeline route sur le dev-difficile (locuteurs disjoints, refs+langue connues).
Mesure : (1) precision LID mms-lid-256 contrainte a lin/lug/sna, (2) score MMS-seul par langue,
(3) score ROUTE (lug-predit->Whisper sinon MMS) par VRAIE langue + MACRO, (4) Whisper pur sur lug.
=> decide si router lug->Whisper bat le MMS-seul AVANT de decoder la Phase 2.
"""
import argparse, json
from collections import Counter, defaultdict
import soundfile as sf, torch, jiwer
from transformers import (AutoModelForCTC, AutoProcessor,
WhisperForConditionalGeneration, WhisperProcessor,
AutoModelForAudioClassification, AutoFeatureExtractor)
SR = 16000
def norm(s):
return " ".join(str(s).replace("|", " ").split())
def is_halluc(hyp, dur):
w = hyp.split()
if not w:
return True
if len(w) >= 6 and Counter(w).most_common(1)[0][1] / len(w) > 0.5:
return True
if dur > 0 and len(w) / dur > 6.0:
return True
return False
def sc(refs, hyps):
P = [(r, h) for r, h in zip(refs, hyps) if r.strip()]
if not P:
return (0, 0, 0)
R = [r for r, _ in P]; H = [h for _, h in P]
wer = jiwer.wer(R, H); cer = jiwer.cer(R, H)
return wer, cer, 0.5 * wer + 0.5 * cer
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--whisper_model", default="Sunbird/asr-whisper-large-v3-salt")
ap.add_argument("--mms_model", required=True)
ap.add_argument("--lid_model", default="facebook/mms-lid-256")
ap.add_argument("--manifest", default="/root/devhard/devhard_all.jsonl")
ap.add_argument("--lang_token", type=int, default=50355)
ap.add_argument("--targets", default="lin,lug,sna")
ap.add_argument("--chunk_s", type=float, default=28.0)
a = ap.parse_args()
dev = "cuda"
lid = AutoModelForAudioClassification.from_pretrained(a.lid_model).to(dev).eval()
lidfe = AutoFeatureExtractor.from_pretrained(a.lid_model)
label2id = {v: k for k, v in lid.config.id2label.items()}
tgt = a.targets.split(",")
missing = [t for t in tgt if t not in label2id]
assert not missing, f"LID labels manquants {missing} ; ex={list(label2id)[:8]}"
tgt_ids = torch.tensor([label2id[t] for t in tgt])
mms = AutoModelForCTC.from_pretrained(a.mms_model, torch_dtype=torch.float16).to(dev).eval()
mmsp = AutoProcessor.from_pretrained(a.mms_model)
wp = WhisperProcessor.from_pretrained(a.whisper_model)
wm = WhisperForConditionalGeneration.from_pretrained(a.whisper_model, torch_dtype=torch.float16).to(dev).eval()
tr = wp.tokenizer.convert_tokens_to_ids("<|transcribe|>")
nt = wp.tokenizer.convert_tokens_to_ids("<|notimestamps|>")
forced = [[1, a.lang_token], [2, tr], [3, nt]]
def mms_decode(au):
x = mmsp(au, sampling_rate=SR, return_tensors="pt")
x = {k: (v.to(dev, torch.float16) if v.dtype == torch.float32 else v.to(dev)) for k, v in x.items()}
ids = mms(**x).logits.argmax(-1).cpu().numpy()
return norm(mmsp.batch_decode(ids)[0])
def whisper_decode(au):
win = int(a.chunk_s * SR)
pieces = [au[j:j + win] for j in range(0, len(au), win)] or [au]
o = []
for p in pieces:
f = wp(p, sampling_rate=SR, return_tensors="pt").input_features.to(dev, torch.float16)
ids = wm.generate(f, forced_decoder_ids=forced, max_new_tokens=220)
o.append(wp.batch_decode(ids, skip_special_tokens=True)[0])
return norm(" ".join(o))
rows = [json.loads(l) for l in open(a.manifest, encoding="utf-8")]
lid_ok = 0; conf = defaultdict(Counter)
by = defaultdict(lambda: {"ref": [], "mms": [], "route": [], "wpur": []})
with torch.inference_mode():
for i, r in enumerate(rows):
au = sf.read(r["audio"], dtype="float32")[0]; dur = len(au) / SR
true = r["lang"]; ref = norm(r["text"])
inp = lidfe(au[:SR * 20], sampling_rate=SR, return_tensors="pt")
inp = {k: v.to(dev) for k, v in inp.items()}
pred = tgt[int(torch.argmax(lid(**inp).logits[0][tgt_ids]).item())]
conf[true][pred] += 1
if pred == true:
lid_ok += 1
h_mms = mms_decode(au)
# route selon LID
if pred == "lug":
h_w = whisper_decode(au)
h_route = h_mms if is_halluc(h_w, dur) else h_w
else:
h_route = h_mms
by[true]["ref"].append(ref)
by[true]["mms"].append(h_mms)
by[true]["route"].append(h_route)
if true == "lug":
by[true]["wpur"].append(whisper_decode(au))
if (i + 1) % 100 == 0:
print(f"{i+1}/{len(rows)}", flush=True)
print(f"\nLID accuracy = {lid_ok/len(rows):.4f} ({lid_ok}/{len(rows)})")
for t in tgt:
print(f" LID vrai={t}: {dict(conf[t])}")
mms_macro = []; route_macro = []
for t in tgt:
d = by[t]
_, _, cm = sc(d["ref"], d["mms"]); _, _, cr = sc(d["ref"], d["route"])
mms_macro.append(1 - cm); route_macro.append(1 - cr)
line = f"{t}: MMS score={1-cm:.4f} | ROUTE score={1-cr:.4f}"
if t == "lug" and d["wpur"]:
_, _, cw = sc(d["ref"], d["wpur"])
line += f" | Whisper-pur(lug)={1-cw:.4f}"
print(line)
print(f"MACRO MMS-seul={sum(mms_macro)/len(mms_macro):.4f} | ROUTE={sum(route_macro)/len(route_macro):.4f}")
if __name__ == "__main__":
main()