"""Pinned public sources -> frozen grouped decision splits; no customer data. Original source splits take priority over derived cuts. Exact normalized source groups are reserved before transformations or sampling. This does not detect all semantic duplicates or prove absence from the pretrained model's corpus. """ import argparse from collections import Counter, defaultdict import csv import hashlib import io import json import os from pathlib import Path import random import sys import urllib.request import zipfile os.environ["HF_HUB_DISABLE_XET"] = "1" from huggingface_hub import hf_hub_download import pyarrow.parquet as pq VERSION = "public-decisions-v1" SOURCES = { "snli": {"repo": "stanfordnlp/snli", "revision": "cdb5c3d5eed6ead6e5a341c8e56e669bb666725b", "license": "CC-BY-SA-4.0", "path": "plain_text"}, "boolq": {"repo": "google/boolq", "revision": "35b264d03638db9f4ce671b711558bf7ff0f80d5", "license": "CC-BY-SA-3.0", "path": "data"}, "arc": {"repo": "allenai/ai2_arc", "revision": "210d026faf9955653af8916fad021475a3f00453", "license": "CC-BY-SA-4.0"}, "banking": {"repo": "PolyAI-LDN/task-specific-datasets", "revision": "57ec275d8078af65b7731c2a98be812d844a6d6b", "license": "CC-BY-4.0"}, "social": {"repo": "allenai/social_i_qa", "revision": "8835ceb9141d7896d9d968634a9b21ae440e3ec5", "url": "https://storage.googleapis.com/ai2-mosaic/public/socialiqa/socialiqa-train-dev.zip", "license": "CC-BY-4.0", "holdout_only": True}, } def digest(value): return hashlib.sha256(value.encode()).hexdigest() def normalized(value): return " ".join(value.casefold().split()) def group_id(family, value): return family + ":" + digest(normalized(value)) def main(): parser = argparse.ArgumentParser() parser.add_argument("--output", required=True) args = parser.parse_args() out = Path(args.output).resolve() out.mkdir(parents=True, exist_ok=False) raw = out / "raw" raw.mkdir() hashes = {} pools = defaultdict(lambda: defaultdict(list)) def parquet(family, split, path): source = SOURCES[family] file = Path(hf_hub_download(source["repo"], f"{path}/{split}-00000-of-00001.parquet", repo_type="dataset", revision=source["revision"], local_dir=raw / family)) hashes[str(file.relative_to(out))] = hashlib.sha256(file.read_bytes()).hexdigest() return pq.read_table(file).to_pylist() def download(name, url): file = raw / name with urllib.request.urlopen(url, timeout=120) as response: file.write_bytes(response.read()) hashes[str(file.relative_to(out))] = hashlib.sha256(file.read_bytes()).hexdigest() return file.read_bytes() def add(family, split, index, state, question, choices, target, group): if not 0 <= target < len(choices) or len(set(choices)) != len(choices): return pools[family][split].append({"id": f"{family}:{split}:{index}", "group": group_id(family, group), "family": family, "source_split": split, "state": state, "question": question, "choices": choices, "target": target}) for split in ("train", "validation", "test"): for i, row in enumerate(parquet("snli", split, "plain_text")): if row["label"] not in (0, 1, 2): continue add("snli", split, i, f"PREMISE: {row['premise']}\nHYPOTHESIS: {row['hypothesis']}", "How does the premise relate to the hypothesis?", ["The premise entails the hypothesis.", "The hypothesis is neither established nor contradicted.", "The premise contradicts the hypothesis."], row["label"], row["premise"]) print("SNLI downloaded", flush=True) for split in ("train", "validation"): for i, row in enumerate(parquet("boolq", split, "data")): add("boolq", split, i, row["passage"], row["question"], ["no", "yes"], int(row["answer"]), row["passage"]) print("BoolQ downloaded", flush=True) for path in ("ARC-Challenge", "ARC-Easy"): for split in ("train", "validation", "test"): for i, row in enumerate(parquet("arc", split, path)): labels = row["choices"]["label"] if row["answerKey"] not in labels: continue add("arc", split, path + ":" + str(i), row["question"], "Which candidate is the correct answer to this science question?", row["choices"]["text"], labels.index(row["answerKey"]), row["question"]) print("ARC downloaded", flush=True) banking = SOURCES["banking"] prefix = f"https://raw.githubusercontent.com/{banking['repo']}/{banking['revision']}/banking_data/" categories = json.loads(download("banking-categories.json", prefix + "categories.json")) for split in ("train", "test"): content = download(f"banking-{split}.csv", prefix + split + ".csv").decode() for i, row in enumerate(csv.DictReader(io.StringIO(content))): target = categories.index(row["category"]) rng = random.Random(digest(f"banking:{split}:{i}")) others = rng.sample([j for j in range(len(categories)) if j != target], 3) choices = [categories[j].replace("_", " ") for j in [target, *others]] add("banking", split, i, row["text"], "Which banking support intent matches this request?", choices, 0, row["text"]) print("Banking77 downloaded (fixed four-choice subsets)", flush=True) archive = zipfile.ZipFile(io.BytesIO(download("socialiqa-train-dev.zip", SOURCES["social"]["url"]))) # Deliberately do not read Social IQA's training split: this family is untouched. dev_name = next(name for name in archive.namelist() if name.endswith("/dev.jsonl")) label_name = next(name for name in archive.namelist() if name.endswith("/dev-labels.lst")) rows = archive.read(dev_name).decode().splitlines() labels = archive.read(label_name).decode().splitlines() if len(rows) != len(labels): raise ValueError("Social IQA dev/label length mismatch") for i, (line, label) in enumerate(zip(rows, labels)): row = json.loads(line) add("social", "validation", i, row["context"], row["question"], [row["answerA"], row["answerB"], row["answerC"]], int(label) - 1, row["context"]) derived = defaultdict(list) audit = {} for family, splits in pools.items(): # Source evaluation groups win, including rows not retained in the small cut. evaluation = splits.get("test", []) + splits.get("validation", []) reserved = {row["group"] for row in evaluation} train = [row for row in splits.get("train", []) if row["group"] not in reserved] # Banking has no source validation: reserve stable training groups for dev/cal. if family == "banking": train_groups = sorted({row["group"] for row in train}, key=lambda g: digest("dev:" + g)) held = set(train_groups[:384]) splits["validation"] = [row for row in train if row["group"] in held] train = [row for row in train if row["group"] not in held] chosen_groups = {} for source_split, destinations in ( ("test", [("test", 512)]), ("validation", [("validation", 128), ("calibration", 128)])): if family == "social": destinations = [("holdout", 768)] elif family == "boolq": if source_split == "test": continue destinations = [("test", 512), ("validation", 128), ("calibration", 128)] candidates = splits.get(source_split, []) # Reserve complete groups; source test takes priority over source validation. if source_split == "validation" and splits.get("test"): test_groups = {row["group"] for row in splits["test"]} candidates = [row for row in candidates if row["group"] not in test_groups] groups = sorted({row["group"] for row in candidates}, key=lambda g: digest("cut:" + g)) offset = 0 for destination, count in destinations: selection = set(groups[offset:offset + count]) offset += count chosen_groups[destination] = selection subset = [row for row in candidates if row["group"] in selection] # Keep one source question per group to make the frozen cut small and independent. by_group = {} for row in subset: by_group.setdefault(row["group"], row) derived[destination].extend(by_group.values()) if family != "social": # Deduplicate exact input+question pairs, keeping premise groups intact. unique = {} for row in train: unique.setdefault(digest(normalized(row["state"]) + normalized(row["question"])), row) train = list(unique.values()) random.Random(917).shuffle(train) train = train[:20000] if family == "snli" else train derived["train"].extend(train) audit[family] = {"source_counts": {s: len(r) for s, r in splits.items()}, "retained_train": len(train), "group_definition": "normalized premise/passage/request/question/context"} for split, rows in derived.items(): for row in rows: # Seeded candidate permutations precede all training/evaluation. old_target = row["choices"][row["target"]] random.Random(digest("choices:" + row["id"])).shuffle(row["choices"]) row["target"] = row["choices"].index(old_target) if split in ("test", "holdout"): if row["family"] == "snli": row["question"] = "Which description best characterizes the evidential relationship between these two statements?" elif row["family"] == "arc": row["question"] = "Select the scientifically correct response to the question in the state." elif row["family"] == "banking": row["question"] = "Which issue should the support team address for this customer?" random.Random(1122).shuffle(rows) (out / f"{split}.jsonl").write_text("".join(json.dumps(row) + "\n" for row in rows)) split_groups = {s: {row["group"] for row in rows} for s, rows in derived.items()} for first, a in split_groups.items(): for second, b in split_groups.items(): if first != second and a & b: raise ValueError(f"Group leakage: {first}/{second}") manifest = {"version": VERSION, "sources": SOURCES, "raw_sha256": hashes, "source_script_sha256": hashlib.sha256(Path(__file__).read_bytes()).hexdigest(), "split_sha256": {s: hashlib.sha256((out / f"{s}.jsonl").read_bytes()).hexdigest() for s in derived}, "split_family_counts": {s: dict(Counter(r["family"] for r in rows)) for s, rows in derived.items()}, "audit": audit, "banking_protocol": "target plus three uniformly sampled negatives, seeded and shuffled; not 77-way accuracy", "holdout": "Social IQA family; no training or calibration rows", "group_leakage": False, "limitations": "Exact normalized groups only; pretrained contamination and semantic duplicates not excluded."} (out / "manifest.json").write_text(json.dumps(manifest, indent=2) + "\n") print(json.dumps(manifest["split_family_counts"], indent=2), flush=True) if __name__ == "__main__": main()