waxal2026-backup / code /apply_punct.py
Pricile's picture
compactage apres suppression luganda
6eed659
Raw
History Blame Contribute Delete
7.76 kB
#!/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()