File size: 3,348 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
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
#!/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)