File size: 7,758 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 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 | #!/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()
|