openjev-e4b / eval /analysis /contamination_audit.py
bambamdevs's picture
Publish OpenJEV E4B 1.0
03223d7
Raw History Blame Contribute Delete
3.04 kB
"""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()