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()