| |
| """Ensembles de checkpoints par langue : moyenne des log-probs -> greedy -> score val complet.""" |
| import json |
| import jiwer |
| import numpy as np |
| import soundfile as sf |
| import torch |
| from transformers import Wav2Vec2BertForCTC, Wav2Vec2BertProcessor |
|
|
| SR = 16000 |
| COMBOS = { |
| "lin": {"single_s4": ["/root/models/lin_s4_best"], |
| "s2+s4": ["/root/models/lin_s2_best", "/root/models/lin_s4_best"], |
| "s2+s3+s4": ["/root/models/lin_s2_best", "/root/models/lin_s3_best", "/root/models/lin_s4_best"]}, |
| "lug": {"single_s2": ["/root/models/lug_s2_best"], |
| "s1+s2": ["/root/models/lug_s1_best", "/root/models/lug_s2_best"]}, |
| "sna": {"single_s2": ["/root/models/sna_s2_best"], |
| "s1+s2": ["/root/models/sna_s1_best", "/root/models/sna_s2_best"]}, |
| } |
|
|
| def read_jsonl(p): |
| rows = [] |
| for l in open(p, encoding="utf-8"): |
| r = json.loads(l) |
| r["text"] = " ".join(r.get("text", "").replace("|", " ").split()) |
| rows.append(r) |
| return [r for r in rows if r["text"]] |
|
|
| @torch.inference_mode() |
| def batch_logits(model, proc, audio): |
| feats = proc.feature_extractor(audio, sampling_rate=SR, return_tensors="pt", padding=True) |
| feats = {k: v.to("cuda", dtype=torch.bfloat16 if v.dtype == torch.float32 else v.dtype) for k, v in feats.items()} |
| return torch.log_softmax(model(**feats).logits.float(), dim=-1) |
|
|
| results = {} |
| for lang, combos in COMBOS.items(): |
| val = read_jsonl(f"/scratch/prep/manifests/waxal_{lang}_validation.jsonl") |
| val.sort(key=lambda r: -r["duration"]) |
| batches, cur, bud = [], [], 0.0 |
| for r in val: |
| if cur and bud + r["duration"] > 120.0: |
| batches.append(cur); cur, bud = [], 0.0 |
| cur.append(r); bud += r["duration"] |
| if cur: batches.append(cur) |
|
|
| model_dirs = sorted({d for dirs in combos.values() for d in dirs}) |
| loaded = {} |
| for d in model_dirs: |
| proc = Wav2Vec2BertProcessor.from_pretrained(d) |
| loaded[d] = (proc, Wav2Vec2BertForCTC.from_pretrained(d, torch_dtype=torch.bfloat16).cuda().eval()) |
|
|
| per_model_ids = {d: {} for d in model_dirs} |
| logits_cache = {} |
| for b in batches: |
| audio = [sf.read(r["audio"], dtype="float32")[0] for r in b] |
| lg = {d: batch_logits(m, p, audio) for d, (p, m) in loaded.items()} |
| for name, dirs in combos.items(): |
| avg = torch.stack([lg[d] for d in dirs]).mean(0) |
| ids = avg.argmax(-1).cpu().numpy() |
| tok = loaded[dirs[0]][0].tokenizer |
| for r, s in zip(b, tok.batch_decode(ids)): |
| logits_cache.setdefault(name, {})[r["id"]] = " ".join(s.replace("|", " ").split()) |
| refs = {r["id"]: r["text"] for r in val} |
| for name in combos: |
| order = [r["id"] for r in val] |
| hyps = [logits_cache[name][i] for i in order] |
| rr = [refs[i] for i in order] |
| wer, cer = jiwer.wer(rr, hyps), jiwer.cer(rr, hyps) |
| comb = 0.5 * wer + 0.5 * cer |
| results[f"{lang}/{name}"] = {"wer": round(wer, 4), "cer": round(cer, 4), "combine": round(comb, 4)} |
| print(f"{lang}/{name}: WER {wer:.4f} CER {cer:.4f} combine {comb:.4f}", flush=True) |
| for d in list(loaded): |
| del loaded[d] |
| torch.cuda.empty_cache() |
| json.dump(results, open("/root/models/ensemble_results.json", "w"), indent=1) |
| print("ENSEMBLE_DONE", flush=True) |
|
|