"""Row-level text-overlap audit between training source splits and public panel rows. Training sources are the full upstream splits the compilers sample from (src/compile_datasets.py): cais/mmlu auxiliary_train, allenai/ai2_arc train, Rowan/hellaswag train. Using the full split is an upper bound on the sampled training rows. Public rows are regenerated with eval/harness/benchmark_public.py. A row matches when its normalized question text (lowercase, alphanumerics only, collapsed whitespace) hashes identically. "question+options" additionally requires the same normalized option multiset. """ from __future__ import annotations import hashlib import json import re import sys from pathlib import Path from datasets import load_dataset HARNESS = Path(__file__).resolve().parents[1] / "harness" sys.path.insert(0, str(HARNESS)) import benchmark_public as bp # noqa: E402 def norm(text: str) -> str: return " ".join(re.sub(r"[^a-z0-9]+", " ", str(text).lower()).split()) def h(text: str) -> str: return hashlib.sha256(norm(text).encode()).hexdigest() def hq(question: str, options) -> str: return h(question) + ":" + h("|".join(sorted(norm(o) for o in options))) def train_sources() -> dict[str, tuple[set[str], set[str]]]: out = {} aux = load_dataset("cais/mmlu", "all", split="auxiliary_train") out["mmlu_auxiliary_train"] = ( {h(x["question"]) for x in aux}, {hq(x["question"], x["choices"]) for x in aux}, ) q, qo = set(), set() for cfg in ("ARC-Challenge", "ARC-Easy"): for x in load_dataset("allenai/ai2_arc", cfg, split="train"): q.add(h(x["question"])) qo.add(hq(x["question"], x["choices"]["text"])) out["arc_train"] = (q, qo) hs = load_dataset("Rowan/hellaswag", split="train") out["hellaswag_train"] = ( {h(x["ctx"]) for x in hs}, {hq(x["ctx"], x["endings"]) for x in hs}, ) return out def main() -> None: sources = train_sources() panel = { "mmlu": bp.load_mmlu(), "arc_easy": bp.load_arc("ARC-Easy"), "arc_challenge": bp.load_arc("ARC-Challenge"), "hellaswag": bp.load_hellaswag(), } report = {"method": __doc__.strip().splitlines()[0], "sources": {k: len(v[0]) for k, v in sources.items()}, "overlap": {}} for task, items in panel.items(): report["overlap"][task] = {"n": len(items)} for src, (qset, qoset) in sources.items(): # HellaSwag panel rows carry the context in the question field. q_hits = sum(h(it["question"]) in qset or h(it["state"]) in qset for it in items) qo_hits = sum(hq(it["question"], it["options"]) in qoset or hq(it["state"], it["options"]) in qoset for it in items) report["overlap"][task][src] = {"question": q_hits, "question_and_options": qo_hits} print(json.dumps(report, indent=2)) if len(sys.argv) > 1: Path(sys.argv[1]).write_text(json.dumps(report, indent=2), encoding="utf-8") if __name__ == "__main__": main()