#!/usr/bin/env python3 """Applique le modele de restauration de virgules a une soumission. Chaque virgule BIEN placee corrige 1 substitution de mot (+1 caractere manquant) ; chaque virgule MAL placee en CREE une. Le gain n est donc positif que si la precision depasse 50 %. => on n insere qu au-dessus d un SEUIL de probabilite, et on balaye ce seuil : seuil eleve = peu d insertions mais tres sures. SEUIL=1.01 => aucune insertion => reproduit BASE a l identique (CONTROLE). On traite les DEUX langues (references : 0.81 virgule/enonce en lingala, 0.51 en shona ; notre sortie : 0.05 et 0.13). Le rapport `taux atteint` permet de voir a quel seuil on se rapproche du taux de reference sans forcer. On ne touche QUE les virgules : ni la casse, ni le point final, ni l ordre des mots. """ import csv import json import os import numpy as np import torch from transformers import AutoModelForTokenClassification, AutoTokenizer BASE = os.environ.get("BASE", "/root/sub_XC010.csv") # VOTE : plusieurs modeles separes par des virgules -> on MOYENNE leurs probabilites. # Meme mecanisme que les juges decorreles (+0.000655) : deux modeles qui se trompent # DIFFEREMMENT relevent la precision ; deux modeles de la meme famille se trompent # ensemble. D ou mBERT (WordPiece, pre-entrainement Wikipedia) a cote d afro-xlmr # (SentencePiece, CC100 adapte aux langues africaines). MODELS = [x for x in os.environ.get("PUNCT_MODEL", "/scratch/runs/punct2/best").split(",") if x] LANGF = os.environ.get("LANGF", "/root/test_lang.json") THS = [float(x) for x in os.environ.get("THS", "1.01,0.9,0.8,0.7,0.6,0.5").split(",")] # SEUIL PAR LANGUE : a seuil global 0.60 le shona atteint 63 % de son taux de # reference (0.32/0.51) mais le lingala seulement 31 % (0.25/0.81). Le lingala est # deux fois plus en retard => un seuil unique le sous-sert. TH_LIN abaisse le sien # sans toucher au shona, deja proche de sa cible. TH_LIN = os.environ.get("TH_LIN", "") # ex "0.40" ; vide = meme seuil que le global # MESURE 8 aout : abaisser le seuil du LINGALA a 0.40 COUTE -0.000498 (RL060). # => le lingala n est pas bride, il est intrinsequement plus dur (predictions moins # fiables sur un texte plus bruite). Le test miroir sur le SHONA reste ouvert. TH_SNA = os.environ.get("TH_SNA", "") TAG = os.environ.get("TAG", "PC") MAXLEN = int(os.environ.get("MAXLEN", "192")) REF_RATE = {"lin": 0.81, "sna": 0.51} # virgules/enonce dans les references train def main(): base = {r["ID"]: r["Target"] for r in csv.DictReader(open(BASE, encoding="utf-8"))} lang = json.load(open(LANGF)) ids = list(base) print("base %d clips" % len(base), flush=True) # proba de virgule apres chaque mot, pour chaque clip PROB, WORDS, RAW = {}, {}, {} ACC = {} # somme des probabilites sur les modeles du vote B = 32 for MODEL in MODELS: tok = AutoTokenizer.from_pretrained(MODEL) model = AutoModelForTokenClassification.from_pretrained(MODEL).cuda().eval() with torch.inference_mode(): for i in range(0, len(ids), B): chunk = ids[i:i + B] wl = [] for k in chunk: # /!\ On ne RETIRE JAMAIS les virgules deja presentes : la base en a # 0.05/enonce (lin) et 0.13 (sna), et les effacer rendait le controle # rouge (52 clips modifies a seuil 1.01) tout en FAISANT BAISSER le # taux lingala sous celui de la base. On ne fait qu AJOUTER. # Le modele voit le texte sans virgule (comme a l entrainement), mais # la reconstruction repart du texte ORIGINAL. w = [x for x in str(base[k]).split() if x] or ["a"] wl.append([x.replace(",", "") or "a" for x in w]) RAW[k] = w x = tok(wl, is_split_into_words=True, truncation=True, max_length=MAXLEN, padding=True, return_tensors="pt") xx = {kk: vv.cuda() for kk, vv in x.items()} p = torch.softmax(model(**xx).logits.float(), -1)[:, :, 1].cpu().numpy() for j, k in enumerate(chunk): wid = x.word_ids(j) prev, pr = None, {} for t, w in enumerate(wid): if w is not None and w != prev: pr[w] = float(p[j, t]) prev = w WORDS[k] = wl[j] v = np.array([pr.get(n, 0.0) for n in range(len(wl[j]))]) ACC[k] = v if k not in ACC else ACC[k] + v if (i + B) % 320 == 0: print(" %d/%d" % (i + B, len(ids)), flush=True) del model torch.cuda.empty_cache() print("modele %s note" % os.path.basename(MODEL.rstrip("/")), flush=True) for k in ACC: PROB[k] = ACC[k] / len(MODELS) print("probabilites moyennees sur %d modele(s)" % len(MODELS), flush=True) from huggingface_hub import HfApi api = HfApi(token=open(os.path.expanduser("~/.cache/huggingface/token")).read().strip()) for th in THS: out, nins = dict(base), {"lin": 0, "sna": 0} ncl = {"lin": 0, "sna": 0} nchg = 0 for k in ids: lg = lang.get(k, "lin") ncl[lg] = ncl.get(lg, 0) + 1 # on repart du texte ORIGINAL : les virgules deja produites par le modele # acoustique sont CONSERVEES (on n en retire jamais aucune) w, p = list(RAW[k]), PROB[k] # jamais de virgule sur le DERNIER mot (il porte le point final) # /!\ le seuil NEUTRE (>1) doit desactiver TOUTE insertion, y compris # les seuils par langue : sinon le controle ponctue quand meme le shona # et ne controle plus rien (bug constate le 8 aout, 115 clips au lieu de 0). thl = th if th <= 1.0: if TH_LIN and lg == "lin": thl = float(TH_LIN) elif TH_SNA and lg == "sna": thl = float(TH_SNA) for n in range(min(len(w), len(p)) - 1): if p[n] >= thl and not w[n].endswith(","): w[n] = w[n] + "," nins[lg] = nins.get(lg, 0) + 1 h = " ".join(w) if h != base[k]: nchg += 1 out[k] = h empt = sum(1 for x in out.values() if not str(x).strip()) tag = "%s%03d" % (TAG, round(th * 100)) OUT = "/root/sub_%s.csv" % tag with open(OUT, "w", newline="", encoding="utf-8") as f: wr = csv.writer(f) wr.writerow(["ID", "Target"]) for k in base: wr.writerow([k, out[k] or "a"]) assert len(out) == 892 and empt == 0, "%s INVALIDE" % tag tl = sum(out[k].count(",") for k in ids if lang.get(k) == "lin") ts = sum(out[k].count(",") for k in ids if lang.get(k) == "sna") rl = tl / max(ncl["lin"], 1) rs = ts / max(ncl["sna"], 1) if th <= 1.0: api.upload_file(path_or_fileobj=OUT, path_in_repo="phase2_corrected/sub_%s.csv" % tag, repo_id="Pricile/waxal2026-backup", repo_type="model") flag = " <-- CONTROLE : doit etre 0" if th > 1.0 else "" print("%-8s seuil %.2f | clips modifies %3d/892 | virg/enonce lin %.2f (ref %.2f)" " sna %.2f (ref %.2f)%s" % (tag, th, nchg, rl, REF_RATE["lin"], rs, REF_RATE["sna"], flag), flush=True) print("APPLY_PUNCT_DONE", flush=True) if __name__ == "__main__": main()