#!/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()