""" Multi-class Arabic dialect + MSA classifier. 14 labels: 13 dialects plus MSA. Input is the text alone - no prompt, no rule. Classes are imbalanced 14.3% (ma) to 0.9% (ly), so accuracy alone is misleading and macro-F1 selects the checkpoint. Two diagnostics beyond the headline: * per-class F1, since a rare dialect can vanish into a neighbour without moving accuracy; * MSA accuracy split by source. The MSA class mixes transcript rows (from the original `ar` label) with chunks of written articles. If the model is really learning "written register" rather than "MSA", it will score far better on the article slice than the transcript one. """ import json import os from collections import Counter, defaultdict from pathlib import Path import numpy as np import torch from torch.utils.data import Dataset from transformers import (AutoModelForSequenceClassification, AutoTokenizer, Trainer, TrainingArguments) os.environ.setdefault("CUDA_VISIBLE_DEVICES", "0") BASE_MODEL = os.environ.get("BASE_MODEL", "oddadmix/50M-2048-Emhotob") OUTPUT_DIR = os.environ.get("OUTPUT_DIR", "./Nawah-Guard-52M") DATA = os.environ.get("DATA_DIR", "data") MAX_LENGTH = int(os.environ.get("MAX_LENGTH", 192)) LEARNING_RATE = float(os.environ.get("LR", 3e-4)) EPOCHS = float(os.environ.get("EPOCHS", 3)) BATCH_SIZE = int(os.environ.get("BATCH_SIZE", 64)) WARMUP = int(os.environ.get("WARMUP", 500)) SEED = 42 META = json.load(open(os.path.join(os.environ.get("DATA_DIR","data"),"labels.json"), encoding="utf-8")) LABELS = META["labels"] UNSAFE = set(META["unsafe"]) def load_jsonl(p): with open(p, encoding="utf-8") as fh: return [json.loads(l) for l in fh] class TextDataset(Dataset): def __init__(self, rows, tok, max_length): self.rows, self.tok, self.max_length = rows, tok, max_length def __len__(self): return len(self.rows) def __getitem__(self, i): r = self.rows[i] ids = self.tok.encode(r["text"], add_special_tokens=False)[: self.max_length] if not ids: ids = [self.tok.pad_token_id] return {"input_ids": torch.tensor(ids, dtype=torch.long), "labels": torch.tensor(r["label_id"], dtype=torch.long)} class Collator: def __init__(self, pad_id): self.pad_id = pad_id def __call__(self, feats): n = max(len(f["input_ids"]) for f in feats) ids, att = [], [] for f in feats: pad = n - len(f["input_ids"]) ids.append(torch.cat([f["input_ids"], torch.full((pad,), self.pad_id, dtype=torch.long)])) att.append(torch.cat([torch.ones(len(f["input_ids"]), dtype=torch.long), torch.zeros(pad, dtype=torch.long)])) return {"input_ids": torch.stack(ids), "attention_mask": torch.stack(att), "labels": torch.stack([f["labels"] for f in feats])} def metrics_fn(p): logits, labels = p pred = np.asarray(logits).argmax(-1); labels = np.asarray(labels) acc = float((pred == labels).mean()) f1s = [] for c in range(len(LABELS)): tp = int(((pred == c) & (labels == c)).sum()) fp = int(((pred == c) & (labels != c)).sum()) fn = int(((pred != c) & (labels == c)).sum()) pr = tp / (tp + fp) if tp + fp else 0.0 rc = tp / (tp + fn) if tp + fn else 0.0 f1s.append(2 * pr * rc / (pr + rc) if pr + rc else 0.0) return {"accuracy": acc, "macro_f1": float(np.mean(f1s)), "majority_baseline": float(max(np.bincount(labels, minlength=len(LABELS))) / len(labels))} def main(): tok = AutoTokenizer.from_pretrained(BASE_MODEL) if tok.pad_token_id is None: tok.pad_token = tok.eos_token model = AutoModelForSequenceClassification.from_pretrained( BASE_MODEL, num_labels=len(LABELS), dtype=torch.float32) model.config.pad_token_id = tok.pad_token_id model.config.id2label = {i: l for i, l in enumerate(LABELS)} model.config.label2id = {l: i for i, l in enumerate(LABELS)} print(f"[*] {BASE_MODEL} | params {sum(p.numel() for p in model.parameters())/1e6:.2f}M " f"| {len(LABELS)} labels") train_rows = load_jsonl(os.path.join(DATA, "train.jsonl")) evals = {k: load_jsonl(os.path.join(DATA, f"eval_{k}.jsonl")) for k in ("held", "unseen_phrasing", "unseen_dialect", "hard")} evals = {k: v for k, v in evals.items() if v} print(f"[*] train {len(train_rows):,} | " + " | ".join(f"{k} {len(v):,}" for k, v in evals.items())) args = TrainingArguments( output_dir=OUTPUT_DIR, num_train_epochs=EPOCHS, per_device_train_batch_size=BATCH_SIZE, per_device_eval_batch_size=128, learning_rate=LEARNING_RATE, lr_scheduler_type="cosine", warmup_steps=WARMUP, max_grad_norm=1.0, bf16=True, logging_steps=200, eval_strategy="steps", eval_steps=1500, save_strategy="steps", save_steps=1500, save_total_limit=2, load_best_model_at_end=True, metric_for_best_model="eval_held_macro_f1", greater_is_better=True, report_to=[], seed=SEED, dataloader_num_workers=4, remove_unused_columns=False) trainer = Trainer(model=model, args=args, train_dataset=TextDataset(train_rows, tok, MAX_LENGTH), eval_dataset={k: TextDataset(v, tok, MAX_LENGTH) for k, v in evals.items()}, data_collator=Collator(tok.pad_token_id), compute_metrics=metrics_fn) trainer.train() trainer.save_model(OUTPUT_DIR); tok.save_pretrained(OUTPUT_DIR) final = trainer.evaluate() print("[*] final:", {k: round(v, 4) for k, v in final.items() if isinstance(v, float)}) # For a guardrail the headline is NOT accuracy. Over-blocking is what makes one unusable, # and it hides completely inside a macro-F1. Report it first, then the opposite failure. import torch as _t report = {} for split, rows in evals.items(): ds = TextDataset(rows, tok, MAX_LENGTH) pred = trainer.predict(ds).predictions.argmax(-1) per = defaultdict(lambda: [0, 0]); conf = Counter() over = [0, 0]; miss = [0, 0] for r, p_ in zip(rows, pred): got = LABELS[p_]; want = r["label"] per[want][1] += 1; per[want][0] += int(got == want) if got != want: conf[(want, got)] += 1 if want in ("safe", "safe_sensitive"): over[1] += 1; over[0] += int(got in UNSAFE) # benign flagged unsafe elif want in UNSAFE: miss[1] += 1; miss[0] += int(got not in UNSAFE) # unsafe let through acc = sum(v[0] for v in per.values()) / max(len(rows), 1) f1s = [] for c in LABELS: tp = sum(1 for r, p_ in zip(rows, pred) if LABELS[p_] == c and r["label"] == c) fp = sum(1 for r, p_ in zip(rows, pred) if LABELS[p_] == c and r["label"] != c) fn = sum(1 for r, p_ in zip(rows, pred) if LABELS[p_] != c and r["label"] == c) pr = tp / (tp + fp) if tp + fp else 0.0 rc = tp / (tp + fn) if tp + fn else 0.0 f1s.append(2 * pr * rc / (pr + rc) if pr + rc else 0.0) orr = over[0] / max(over[1], 1); mr = miss[0] / max(miss[1], 1) print(f"\n[{split}] n={len(rows):,}") print(f" OVER-REFUSAL (benign flagged unsafe) : {orr:.4f} n={over[1]:,}") print(f" MISS RATE (unsafe let through) : {mr:.4f} n={miss[1]:,}") print(f" accuracy {acc:.4f} macro-F1 {sum(f1s)/len(f1s):.4f}") print(" per-class acc (worst first): " + " ".join( f"{k}:{v[0]/v[1]:.2f}" for k, v in sorted(per.items(), key=lambda kv: kv[1][0]/kv[1][1])[:6])) print(" top confusions: " + " ".join(f"{a}->{b}:{n}" for (a, b), n in conf.most_common(5))) report[split] = {"over_refusal": orr, "miss_rate": mr, "accuracy": acc, "macro_f1": sum(f1s)/len(f1s), "n": len(rows), "per_class": {k: {"acc": v[0]/v[1], "n": v[1]} for k, v in per.items()}, "confusions": {f"{a}->{b}": n for (a, b), n in conf.most_common(20)}} Path(OUTPUT_DIR, "train_metrics.json").write_text(json.dumps( {"report": report, "base_model": BASE_MODEL, "log_history": trainer.state.log_history}, ensure_ascii=False, indent=2), encoding="utf-8") print(f"[+] done -> {OUTPUT_DIR}") if __name__ == "__main__": main()