| |
| """Ensemble logit-averaging de N checkpoints multilingues (MEME vocab) sur Phase 2, clip-par-clip. |
| Usage: multi_ensemble.py --models /root/models/joint_cont2_best /root/models/joint_cont_best --out /root/sub_ens.csv |
| Meme vocab requis (famille joint) -> les logits s'additionnent index par index. Meme archi -> meme nb de frames.""" |
| import argparse, csv, glob, os |
| import soundfile as sf, torch |
| from transformers import AutoModelForCTC, AutoProcessor |
| SR = 16000 |
|
|
|
|
| def norm(s): |
| return " ".join(str(s).replace("|", " ").split()) |
|
|
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--models", nargs="+", required=True) |
| ap.add_argument("--audio_dir", default="/root/phase2_audio/audio") |
| ap.add_argument("--test_csv", default="/root/Test_phase2.csv") |
| ap.add_argument("--out", required=True) |
| a = ap.parse_args() |
| test_ids = [r["ID"] for r in csv.DictReader(open(a.test_csv, encoding="utf-8"))] |
| procs = [AutoProcessor.from_pretrained(m) for m in a.models] |
| models = [AutoModelForCTC.from_pretrained(m, torch_dtype=torch.bfloat16).cuda().eval() for m in a.models] |
| v0 = procs[0].tokenizer.get_vocab() |
| for p in procs[1:]: |
| assert p.tokenizer.get_vocab() == v0, "VOCABS DIFFERENTS -> ensemble logit invalide (arret)" |
| dec = procs[0] |
| files = sorted(glob.glob(os.path.join(a.audio_dir, "*.wav"))) |
| out = {} |
| with torch.inference_mode(): |
| for i, f in enumerate(files): |
| au = sf.read(f, dtype="float32")[0] |
| logits_sum = None |
| for pm, mdl in zip(procs, models): |
| x = pm(au, sampling_rate=SR, return_tensors="pt", padding=True) |
| x = {k: (v.to("cuda", dtype=torch.bfloat16) if v.dtype == torch.float32 else v.to("cuda")) for k, v in x.items()} |
| lg = mdl(**x).logits.float() |
| logits_sum = lg if logits_sum is None else logits_sum + lg |
| ids = logits_sum.argmax(-1).cpu().numpy() |
| out[os.path.splitext(os.path.basename(f))[0]] = norm(dec.batch_decode(ids)[0]) |
| if (i + 1) % 200 == 0: |
| print(f"{i+1}/{len(files)} clips", flush=True) |
| empt = sum(1 for t in test_ids if not out.get(t, "").strip()) |
| with open(a.out, "w", newline="", encoding="utf-8") as fo: |
| w = csv.writer(fo); w.writerow(["ID", "Target"]) |
| for t in test_ids: |
| w.writerow([t, out.get(t) or "a"]) |
| print(f"ENSEMBLE_DONE {a.out} | {len(test_ids)} IDs | vides={empt}", flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|