| |
| """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] |
| 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) |
| 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:] |
| |
| 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() |
| |
| 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)) |
|
|
| |
| 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() |
|
|