waxal2026-backup / phase2_corrected /code /punct_restore.py
Pricile's picture
compactage apres suppression luganda
6eed659
Raw
History Blame Contribute Delete
7.58 kB
#!/usr/bin/env python3
"""Restaurateur PONCTUATION + CASSE entraine sur le TEXTE WAXAL train (aucune donnee externe).
Token-classification (mBERT) : pour chaque mot -> (punct in {O,COMMA,PERIOD}, case in {L,C,U}).
Entree = mots minuscules sans ponctuation. Reconstruit texte ponctue+capitalise.
Teste offline sur devhard (combine avant/apres) pour valider SANS soumission.
"""
import argparse, json, re, random, unicodedata
import numpy as np, torch
from datasets import Dataset
from transformers import (AutoTokenizer, AutoModelForTokenClassification,
TrainingArguments, Trainer, DataCollatorForTokenClassification)
PUNCT = ["O", "COMMA", "PERIOD"]
CASE = ["L", "C", "U"]
LABELS = [f"{c}|{p}" for c in CASE for p in PUNCT] # 9 classes combinees
L2I = {l: i for i, l in enumerate(LABELS)}
I2L = {i: l for l, i in L2I.items()}
PMAP = {"COMMA": ",", "PERIOD": ".", "O": ""}
def clean_word(w):
"""Retire ponctuation de fin, renvoie (mot_nu_minuscule, punct_label, case_label)."""
m = re.search(r"([.,])\s*$", w)
punct = "O"
if m:
punct = "COMMA" if m.group(1) == "," else "PERIOD"
core = re.sub(r"^[\"'(\-]+|[\"').,!?;:\-]+$", "", w)
if not core:
return None
if core.isupper() and len(core) > 1:
case = "U"
elif core[:1].isupper():
case = "C"
else:
case = "L"
return core.lower(), punct, case
def sent_to_example(text):
words, labels = [], []
for w in text.split():
r = clean_word(w)
if r is None:
continue
core, punct, case = r
words.append(core)
labels.append(L2I[f"{case}|{punct}"])
return words, labels
def restore(words, preds):
out = []
for w, p in zip(words, preds):
case, punct = I2L[p].split("|")
ww = w.upper() if case == "U" else (w.capitalize() if case == "C" else w)
out.append(ww + PMAP[punct])
return " ".join(out)
def read_jsonl(p):
return [json.loads(l) for l in open(p, encoding="utf-8")]
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--train", nargs="+", required=True)
ap.add_argument("--hyps", required=True) # devhard_joint_hyps.json (id,lang,ref,hyp)
ap.add_argument("--out", default="/scratch/ftruns/punct_mbert")
ap.add_argument("--base", default="bert-base-multilingual-cased")
ap.add_argument("--epochs", type=float, default=3)
ap.add_argument("--bs", type=int, default=32)
ap.add_argument("--lr", type=float, default=3e-5)
ap.add_argument("--maxlen", type=int, default=128)
ap.add_argument("--build_only", action="store_true")
ap.add_argument("--seed", type=int, default=42)
a = ap.parse_args()
random.seed(a.seed); np.random.seed(a.seed); torch.manual_seed(a.seed)
rows = []
for m in a.train:
for r in read_jsonl(m):
t = unicodedata.normalize("NFC", r["text"]).strip()
if not t:
continue
w, l = sent_to_example(t)
if 1 <= len(w) <= a.maxlen:
rows.append({"words": w, "labels": l})
random.shuffle(rows)
nval = max(200, len(rows) // 20)
val_rows, train_rows = rows[:nval], rows[nval:]
# distribution des labels
from collections import Counter
cnt = Counter(l for r in train_rows for l in r["labels"])
print(f"exemples train={len(train_rows)} val={len(val_rows)}")
print("distribution labels:", {I2L[i]: cnt.get(i, 0) for i in range(len(LABELS))})
if a.build_only:
print("BUILD_ONLY_DONE"); return
tok = AutoTokenizer.from_pretrained(a.base)
def encode(batch):
enc = tok(batch["words"], is_split_into_words=True, truncation=True,
max_length=a.maxlen, padding=False)
all_labels = []
for i, labs in enumerate(batch["labels"]):
wids = enc.word_ids(batch_index=i)
prev, lab = None, []
for wid in wids:
if wid is None:
lab.append(-100)
elif wid != prev:
lab.append(labs[wid])
else:
lab.append(-100)
prev = wid
all_labels.append(lab)
enc["labels"] = all_labels
return enc
train_ds = Dataset.from_list(train_rows).map(encode, batched=True, remove_columns=["words"])
val_ds = Dataset.from_list(val_rows).map(encode, batched=True, remove_columns=["words"])
model = AutoModelForTokenClassification.from_pretrained(
a.base, num_labels=len(LABELS), id2label=I2L, label2id=L2I)
coll = DataCollatorForTokenClassification(tok)
def metrics(p):
preds = np.argmax(p.predictions, -1)
mask = p.label_ids != -100
acc = (preds[mask] == p.label_ids[mask]).mean()
# F1 sur les classes ponctuees (non-O)
yl, yp = p.label_ids[mask], preds[mask]
def punct_of(i): return I2L[i].split("|")[1]
tp = sum(1 for t, h in zip(yl, yp) if punct_of(t) != "O" and t == h)
fp = sum(1 for t, h in zip(yl, yp) if punct_of(h) != "O" and t != h)
fn = sum(1 for t, h in zip(yl, yp) if punct_of(t) != "O" and t != h)
prec = tp / max(tp + fp, 1); rec = tp / max(tp + fn, 1)
f1 = 2 * prec * rec / max(prec + rec, 1e-9)
return {"acc": acc, "punct_f1": f1}
targs = TrainingArguments(
output_dir=a.out, per_device_train_batch_size=a.bs, per_device_eval_batch_size=64,
num_train_epochs=a.epochs, learning_rate=a.lr, warmup_ratio=0.1, bf16=True,
eval_strategy="epoch", save_strategy="epoch", save_total_limit=1,
load_best_model_at_end=True, metric_for_best_model="punct_f1", greater_is_better=True,
logging_steps=100, report_to=[], seed=a.seed)
trainer = Trainer(model=model, args=targs, train_dataset=train_ds, eval_dataset=val_ds,
data_collator=coll, compute_metrics=metrics, processing_class=tok)
trainer.train()
trainer.save_model(a.out); tok.save_pretrained(a.out)
print("PUNCT_TRAIN_DONE", json.dumps(trainer.evaluate(), default=float))
# ---- APPLIQUE au devhard + mesure combine avant/apres ----
import jiwer
D = json.load(open(a.hyps, encoding="utf-8"))
model.eval().cuda()
def restore_text(text):
words = [clean_word(w)[0] for w in text.split() if clean_word(w)]
if not words:
return text
enc = tok([words], is_split_into_words=True, truncation=True, max_length=a.maxlen,
padding=True, return_tensors="pt").to("cuda")
with torch.inference_mode():
logits = model(**enc).logits[0]
wids = enc.word_ids(batch_index=0)
preds, seen = [], set()
for j, wid in enumerate(wids):
if wid is not None and wid not in seen:
preds.append(int(logits[j].argmax())); seen.add(wid)
preds += [L2I["L|O"]] * (len(words) - len(preds))
return restore(words, preds[:len(words)])
def comb(refs, hyps):
pr = [(r, h) for r, h in zip(refs, hyps) if r.strip()]
r = [x for x, _ in pr]; h = [x for _, x in pr]
return 0.5 * jiwer.wer(r, h) + 0.5 * jiwer.cer(r, h)
for lang in ["ALL", "lin", "sna"]:
sub = [d for d in D if lang == "ALL" or d["lang"] == lang]
R = [d["ref"] for d in sub]; H = [d["hyp"] for d in sub]
Hr = [restore_text(h) for h in H]
print(f"{lang}: combine RAW={comb(R,H):.4f} -> RESTORE={comb(R,Hr):.4f}", flush=True)
print("PUNCT_APPLY_DONE", flush=True)
if __name__ == "__main__":
main()