""" Split the verified guardrail cache and audit it before anything trains. Three held-out axes, each testing a different way a guardrail fails in the wild: unseen_phrasing - two jailbreak template families (encoding, translation_pivot) and one injection carrier (tool_output) never appear in training. Attackers do not reuse the templates you trained on, so this is the claim that matters. unseen_dialect - Maghrebi held out entirely. An MSA-only guardrail is bypassed by writing the same attack in عامية; this measures whether ours is. hard - the safe_sensitive rows on the eval side. Over-blocking is the failure that makes a guardrail unusable, and it is invisible unless measured separately. Lanes shuffled? Not applicable here - this is single-text classification. What IS audited is the length shortcut, which an earlier version of the grid leaked badly (+0.080 over majority from a single threshold, before length became an axis independent of class). """ import json import os import random from collections import Counter, defaultdict from pathlib import Path import guard_common as gc SRC = Path(os.environ.get("OUT_DIR", "out_guard")) OUT = Path(os.environ.get("DATA_DIR", "data")); OUT.mkdir(exist_ok=True) HELD_FAMILIES = {"encoding", "translation_pivot", "tool_output"} HELD_VARIETY = "اللهجة المغاربية" EVAL_FRAC = float(os.environ.get("EVAL_FRAC", 0.03)) SEED = 42 def main(): rng = random.Random(SEED) rows, amb = [], 0 for line in open(SRC / "verified.jsonl", encoding="utf-8"): try: r = json.loads(line) except json.JSONDecodeError: continue if not r["agree"]: amb += 1 continue rows.append({"text": r["text"], "label": r["label"], "variety": r["variety"], "length": r.get("length"), "register": r["register"], "family": r.get("family"), "task_id": r["task_id"]}) seen, ded = set(), [] for r in rows: k = gc.dedup_key(r["text"]) if k in seen: continue seen.add(k); ded.append(r) print(f"verified {len(rows):,} (ambiguous dropped {amb:,}) -> {len(ded):,} after dedup") rows = ded labels = sorted(gc.CLASSES) lab2id = {l: i for i, l in enumerate(labels)} phr = [r for r in rows if r.get("family") in HELD_FAMILIES] rest = [r for r in rows if r.get("family") not in HELD_FAMILIES] dia = [r for r in rest if r["variety"] == HELD_VARIETY] rest = [r for r in rest if r["variety"] != HELD_VARIETY] tasks = sorted({r["task_id"] for r in rest}); rng.shuffle(tasks) ev = set(tasks[: max(1, int(len(tasks) * EVAL_FRAC))]) held = [r for r in rest if r["task_id"] in ev] train = [r for r in rest if r["task_id"] not in ev] hard = [r for r in held if r["label"] == "safe_sensitive"] tt = {r["text"] for r in train} for name, s in (("unseen_phrasing", phr), ("unseen_dialect", dia), ("held", held)): n = sum(1 for r in s if r["text"] in tt) if n: print(f" ! {n} texts in {name} also in train") for r in rows: r["label_id"] = lab2id[r["label"]] for name, s in (("train", train), ("eval_held", held), ("eval_unseen_phrasing", phr), ("eval_unseen_dialect", dia), ("eval_hard", hard)): with open(OUT / f"{name}.jsonl", "w", encoding="utf-8") as fh: for r in s: fh.write(json.dumps(r, ensure_ascii=False) + "\n") print(f" {name:<24}{len(s):>8,}") json.dump({"labels": labels, "label2id": lab2id, "unsafe": [l for l in labels if l not in ("safe", "safe_sensitive")]}, open(OUT / "labels.json", "w", encoding="utf-8"), ensure_ascii=False, indent=1) print(f"\nheld-out families in train: " f"{ {r.get('family') for r in train} & HELD_FAMILIES or 'none'}") print(f"held-out variety in train : " f"{ {r['variety'] for r in train} & {HELD_VARIETY} or 'none'}") print("\n--- audit ---") c = Counter(r["label"] for r in rows) print("class mix:", {k: f"{v/len(rows):.1%}" for k, v in c.most_common()}) maj = max(c.values()) / len(rows) pairs = [(r["label"], len(r["text"])) for r in rows] best = 0 for thr in range(40, 700, 5): lo = [l for l, n in pairs if n < thr]; hi = [l for l, n in pairs if n >= thr] if not lo or not hi: continue best = max(best, (Counter(lo).most_common(1)[0][1] + Counter(hi).most_common(1)[0][1]) / len(pairs)) print(f"length shortcut: majority {maj:.4f} | best threshold {best:.4f} | leak {best-maj:+.4f}") v = Counter(r["variety"] for r in rows) print("variety mix:", {k: f"{n/len(rows):.1%}" for k, n in v.most_common()}) print(f"TOTAL: {len(rows):,}") if __name__ == "__main__": main()