File size: 2,555 Bytes
6eed659
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
#!/usr/bin/env python3
"""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()  # [1, T, V]
                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()