Download source/scripts/prepare_public_data.py from andyshu/opensysone: direct link, hf CLI and curl.
- Browser
- Download file 11.9 kB
-
https://huggingface.co/andyshu/opensysone/resolve/main/source/scripts/prepare_public_data.py
- Command line
-
hf download hf://andyshu/opensysone/source/scripts/prepare_public_data.py
-
curl -L -o prepare_public_data.py https://huggingface.co/andyshu/opensysone/resolve/main/source/scripts/prepare_public_data.py
11.9 kB
| """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() | |