#!/usr/bin/env python3 """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)