| |
| """TRANSFERT DE PONCTUATION Whisper -> CTC (le levier le plus prometteur restant). |
| Constat : notre CTC a les BONS MOTS mais ~0 virgule (0.29‰ vs 5.22‰ dans les réfs, coût |
| mesuré +0.0144). Whisper fine-tuné a des mots moins bons mais produit la ponctuation. |
| => On aligne les deux hypothèses mot à mot et on ne transfère QUE les virgules, uniquement |
| là où les mots CONCORDENT. Le contenu du CTC n'est jamais modifié : risque borné. |
| Gate : devhard-lin, cible = battre 0.3457 (config du record). |
| """ |
| import difflib, json, os, pickle, re |
| import jiwer, numpy as np, soundfile as sf, torch |
| from multiprocessing import Pool |
| from pyctcdecode import build_ctcdecoder |
| from transformers import (AutoProcessor, WhisperForConditionalGeneration, WhisperProcessor) |
|
|
| M1 = "/root/models/joint_cont_best" |
| ARPA = "/scratch/lm/lin_5g.arpa" |
| WM = os.environ.get("WMODEL", "/scratch/runs/whisper_lin/final") |
| AUD = "/root/devhard_audio" |
| SR = 16000 |
| CACHE_W = "/root/whisper_hyps_devhard_lin.json" |
|
|
|
|
| def comb(refs, hyps): |
| pr = [(r, h) for r, h in zip(refs, hyps) if r.strip()] |
| a = [x for x, _ in pr]; b = [y for _, y in pr] |
| w = jiwer.wer(a, b); c = jiwer.cer(a, b) |
| return w, c, 0.5 * w + 0.5 * c |
|
|
|
|
| def strip_p(w): |
| return re.sub(r"[^\w'ɛɔ]", "", w.lower()) |
|
|
|
|
| def transfer_commas(ctc, whi, need_match=True): |
| """Ajoute une virgule au mot i du CTC si le mot aligné de Whisper en porte une.""" |
| cw = ctc.split(); ww = whi.split() |
| if not cw or not ww: |
| return ctc |
| a = [strip_p(x) for x in cw]; b = [strip_p(x) for x in ww] |
| sm = difflib.SequenceMatcher(None, a, b, autojunk=False) |
| out = list(cw) |
| for i1, i2, j1, j2 in [(o[1], o[2], o[3], o[4]) for o in sm.get_opcodes() if o[0] == "equal"]: |
| for k in range(i2 - i1): |
| if ww[j1 + k].endswith(","): |
| t = out[i1 + k] |
| if not t.endswith((",", ".", "!", "?")): |
| out[i1 + k] = t + "," |
| return " ".join(out) |
|
|
|
|
| def main(): |
| rows = [json.loads(l) for l in open("/root/devhard/devhard_linsna.jsonl", encoding="utf-8")] |
| sub = [r for r in rows if r["lang"] == "lin"] |
| for r in sub: |
| r["audio"] = os.path.join(AUD, os.path.basename(r["audio"])) |
| sub = [r for r in sub if os.path.exists(r["audio"])] |
| refs = [r["text"] for r in sub] |
| print("devhard-lin %d clips" % len(sub), flush=True) |
|
|
| |
| L1 = pickle.load(open("/scratch/lm/logits_lin.pkl", "rb")) |
| tok = AutoProcessor.from_pretrained(M1).tokenizer |
| v = tok.get_vocab() |
| lab = [None] * len(v) |
| for t, i in v.items(): |
| lab[i] = t |
| lab[tok.word_delimiter_token_id] = " " |
| lab[tok.unk_token_id] = "⁇" |
| lab[tok.pad_token_id] = "" |
| greedy = [" ".join(tok.decode(l.argmax(-1)).replace("|", " ").split()) for l in L1] |
| dec = build_ctcdecoder(lab, kenlm_model_path=ARPA, alpha=0.5, beta=1.0, lm_score_boundary=True) |
| with Pool(8) as p: |
| db = [" ".join(x.split()) for x in dec.decode_batch(p, L1, beam_width=64)] |
| cc = lambda h, g: (g[:1] + h[1:]) if (h and g) else h |
| ctc = [cc(h, g) for h, g in zip(db, greedy)] |
| base = comb(refs, ctc) |
| print("CTC config record : WER %.4f CER %.4f COMBINE %.4f" % base, flush=True) |
|
|
| |
| if os.path.exists(CACHE_W): |
| whi = json.load(open(CACHE_W, encoding="utf-8")) |
| else: |
| proc = WhisperProcessor.from_pretrained(WM, language="ln", task="transcribe") |
| m = WhisperForConditionalGeneration.from_pretrained(WM, dtype=torch.bfloat16).cuda().eval() |
| m.generation_config.language = "ln"; m.generation_config.task = "transcribe" |
| m.generation_config.forced_decoder_ids = None |
| whi = [] |
| with torch.inference_mode(): |
| for i in range(0, len(sub), 8): |
| b = sub[i:i + 8] |
| au = [sf.read(r["audio"], dtype="float32")[0] for r in b] |
| au = [x.mean(1) if x.ndim > 1 else x for x in au] |
| x = proc(au, sampling_rate=SR, return_tensors="pt").input_features.cuda().to(torch.bfloat16) |
| g = m.generate(x, num_beams=5, max_new_tokens=200, language="ln", task="transcribe", |
| no_repeat_ngram_size=4, repetition_penalty=1.1) |
| whi += [" ".join(t.strip().split()) for t in proc.batch_decode(g, skip_special_tokens=True)] |
| json.dump(whi, open(CACHE_W, "w", encoding="utf-8"), ensure_ascii=False) |
| del m; torch.cuda.empty_cache() |
| print("Whisper : COMBINE %.4f" % comb(refs, whi)[2], flush=True) |
|
|
| |
| out = [transfer_commas(c, w) for c, w in zip(ctc, whi)] |
| m2 = comb(refs, out) |
| nref = sum(t.count(",") for t in refs) |
| print("\nvirgules : refs %d | CTC %d | Whisper %d | apres transfert %d" |
| % (nref, sum(t.count(",") for t in ctc), sum(t.count(",") for t in whi), |
| sum(t.count(",") for t in out)), flush=True) |
| print("CTC + virgules Whisper : WER %.4f CER %.4f COMBINE %.4f (%+.4f)" |
| % (m2[0], m2[1], m2[2], m2[2] - base[2]), flush=True) |
| print("%s" % ("✅ GAIN — a deployer" if m2[2] < base[2] - 0.002 else |
| ("~ neutre" if m2[2] < base[2] else "❌ degrade")), flush=True) |
| for i in range(3): |
| if ctc[i] != out[i]: |
| print("\n CTC : %s" % ctc[i][:100]) |
| print(" WHIS : %s" % whi[i][:100]) |
| print(" FUSION: %s" % out[i][:100]) |
| print(" REF : %s" % refs[i][:100]) |
| json.dump({"base": base[2], "fused": m2[2], "delta": m2[2] - base[2]}, |
| open("/root/comma_transfer.json", "w")) |
| print("\nCOMMA_TRANSFER_DONE", flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|