waxal2026-backup / archive /scripts /audit_STATSLING6.py
Pricile's picture
Upload folder using huggingface_hub
0e11779 verified
Raw
History Blame Contribute Delete
5.54 kB
# -*- coding: utf-8 -*-
"""Audit 6: test decisif de sur-regularisation par le LM.
Compare la distribution de frequence des tokens produits (lin=avec LM, sna=sans LM)
a celle des references. Verifie aussi si le ratio joint/separe est deja calibre."""
import csv, json
from collections import Counter
OUT = []
def P(*a):
s = " ".join(str(x) for x in a); OUT.append(s); print(s)
sub = {}
with open("/root/sub_LMA06.csv", encoding="utf-8") as f:
for row in csv.DictReader(f): sub[row["ID"]] = row["Target"]
lang = json.load(open("/root/test_lang.json", encoding="utf-8"))
hyp = {"lin": [], "sna": []}
for k, v in sub.items():
if lang.get(k) in hyp: hyp[lang[k]].append(v)
train, val = {"lin": [], "sna": []}, {"lin": [], "sna": []}
for L in ("lin", "sna"):
for sp, st in (("train", train), ("validation", val)):
with open("/scratch/prep/manifests/waxal_%s_%s.jsonl" % (L, sp), encoding="utf-8") as f:
for line in f:
if line.strip(): st[L].append(json.loads(line)["text"])
tot_w = sum(len(t.split()) for L in hyp for t in hyp[L])
W_ERR = 0.5/tot_w
PUNCT = ".,;:!?\"()«»"
def toks(t): return [x for x in (w.strip(PUNCT) for w in t.split()) if x]
# ===== A. profil de frequence : sur-regularisation ? =====
P("=== A. SUR-REGULARISATION PAR LE LM : profil de frequence des tokens produits ===")
P("lin = decode AVEC KenLM ; sna = decode SANS LM (temoin)")
P("")
BUCKETS = [(1, 10), (11, 50), (51, 200), (201, 1000), (1001, 5000), (5001, 10**9)]
for L in ("lin", "sna"):
tr = Counter()
for t in train[L]: tr.update(w.lower() for w in toks(t))
rank = {w: i+1 for i, (w, _) in enumerate(tr.most_common())}
def profile(texts):
c = Counter(); n = 0
for t in texts:
for w in toks(t):
wl = w.lower(); n += 1
r = rank.get(wl)
if r is None: c["OOV"] += 1; continue
for lo, hi in BUCKETS:
if lo <= r <= hi: c["%d-%d" % (lo, hi)] += 1; break
return c, n
cv, nv = profile(val[L]) # reference held-out = ce que la verite ressemble
ch, nh = profile(hyp[L])
P("--- %s --- (ref=VAL held-out, %d tokens ; hyp=test, %d tokens)" % (L, nv, nh))
P(" %-12s %9s %9s %8s %10s" % ("rang", "ref%", "hyp%", "ratio", "delta_tok"))
for lo, hi in BUCKETS:
k = "%d-%d" % (lo, hi)
a = 100.0*cv.get(k,0)/nv; b = 100.0*ch.get(k,0)/nh
d = (b-a)/100.0*nh
P(" %-12s %8.2f%% %8.2f%% %8.2f %+10.0f" % (k.replace("1000000000","inf"), a, b, (b/a if a else float('nan')), d))
a = 100.0*cv.get("OOV",0)/nv; b = 100.0*ch.get("OOV",0)/nh
P(" %-12s %8.2f%% %8.2f%% %8.2f %+10.0f" % ("OOV", a, b, (b/a if a else float('nan')), (b-a)/100.0*nh))
P(" -> deficit total en tokens rares (rang>1000 + OOV) = %+.0f tokens" % (
sum((100.0*ch.get("%d-%d"%(lo,hi),0)/nh - 100.0*cv.get("%d-%d"%(lo,hi),0)/nv)/100.0*nh
for lo, hi in BUCKETS if lo > 1000) + (100.0*ch.get("OOV",0)/nh - 100.0*cv.get("OOV",0)/nv)/100.0*nh))
# ===== B. le ratio joint/separe est-il deja calibre ? =====
P("")
P("=== B. RATIO JOINT/SEPARE : sommes-nous deja calibres ? (pas de biais => ne rien faire) ===")
for L in ("lin", "sna"):
def cnt(texts):
u = Counter(); b = Counter()
for t in texts:
ws = [w.lower() for w in toks(t)]; u.update(ws); b.update(zip(ws, ws[1:]))
return u, b
u_tr, b_tr = cnt(train[L]); u_va, b_va = cnt(val[L]); u_hy, b_hy = cnt(hyp[L])
P("--- %s ---" % L)
P(" %-24s %14s %14s %14s" % ("paire", "train j/(j+s)", "val j/(j+s)", "HYP j/(j+s)"))
cands = []
for (a, b), c in b_hy.items():
j = a + b
if u_tr.get(j, 0) < 20: continue
cands.append((c + u_hy.get(j, 0), a, b))
cands.sort(reverse=True)
tot_hj = tot_hs = tot_tj = tot_ts = 0
for _, a, b in cands[:14]:
tj, ts = u_tr.get(a+b,0), b_tr.get((a,b),0)
vj, vs = u_va.get(a+b,0), b_va.get((a,b),0)
hj, hs = u_hy.get(a+b,0), b_hy.get((a,b),0)
P(" %-24s %6d/%-6d %.2f %6d/%-6d %.2f %6d/%-6d %.2f"
% (a+"|"+b, tj, ts, tj/max(tj+ts,1), vj, vs, vj/max(vj+vs,1), hj, hs, hj/max(hj+hs,1)))
tot_hj += hj; tot_hs += hs; tot_tj += tj; tot_ts += ts
if tot_hj + tot_hs:
P(" AGREGE: train joint=%.3f | HYP joint=%.3f | ecart=%+.3f"
% (tot_tj/max(tot_tj+tot_ts,1), tot_hj/max(tot_hj+tot_hs,1),
tot_hj/max(tot_hj+tot_hs,1) - tot_tj/max(tot_tj+tot_ts,1)))
# nombre de tokens a deplacer pour recalibrer
target = tot_tj/max(tot_tj+tot_ts,1)
tot = tot_hj + tot_hs
move = target*tot - tot_hj
P(" -> pour recaler: deplacer %+.0f occurrences vers la forme jointe (sur %d)" % (move, tot))
# ===== C. verdict OOV chiffre pour le lingala =====
P("")
P("=== C. RECAP CHIFFRE DU BIAIS PRINCIPAL ===")
for L in ("lin", "sna"):
tr = Counter()
for t in train[L]: tr.update(w.lower() for w in toks(t))
vw = [w.lower() for t in val[L] for w in toks(t)]
hw = [w.lower() for t in hyp[L] for w in toks(t)]
nat = sum(1 for w in vw if w not in tr)/len(vw)
ours = sum(1 for w in hw if w not in tr)/len(hw)
P("%s | OOV ref(held-out)=%.2f%% OOV hyp=%.2f%% ratio=%.2f deficit=%.0f tokens = %.5f score bloque"
% (L, 100*nat, 100*ours, ours/nat, (nat-ours)*len(hw), max((nat-ours)*len(hw),0)*W_ERR))
with open("/root/audit_STATSLING6_out.txt","w",encoding="utf-8") as f: f.write("\n".join(OUT))
print("\n[OK] /root/audit_STATSLING6_out.txt")