Download source/scripts/prepare_expanded_data.py from andyshu/opensysone: direct link, hf CLI and curl.
- Browser
- Download file 18.4 kB
-
https://huggingface.co/andyshu/opensysone/resolve/main/source/scripts/prepare_expanded_data.py
- Command line
-
hf download hf://andyshu/opensysone/source/scripts/prepare_expanded_data.py
-
curl -L -o prepare_expanded_data.py https://huggingface.co/andyshu/opensysone/resolve/main/source/scripts/prepare_expanded_data.py
18.4 kB
| """Build an immutable train-only expansion while preserving frozen evaluation. | |
| Only labeled official TRAIN rows become training or diagnostic decisions. Source | |
| validation/test files are projected onto grouping metadata for exclusion only. | |
| The existing evaluation files are copied byte-for-byte; no predictions are read. | |
| """ | |
| import argparse | |
| from collections import Counter | |
| from datetime import datetime, timezone | |
| import hashlib | |
| import json | |
| import os | |
| from pathlib import Path | |
| import random | |
| import shutil | |
| import subprocess | |
| import tempfile | |
| import unicodedata | |
| import urllib.request | |
| VERSION = "public-decisions-v2" | |
| PROTECTED = ("validation", "calibration", "test", "holdout") | |
| SPLITS = ("train", *PROTECTED) | |
| SOURCES = { | |
| "hellaswag": { | |
| "repo": "Rowan/hellaswag", "provider": "huggingface", | |
| "revision": "218ec52e09a7e7462a5400043bb9a69a41d06b76", | |
| "license": "MIT", "license_reference": "README.md", | |
| "expected_train_rows": 39905, "group_columns": ["ctx", "source_id"], | |
| "transform": "ctx -> state; four endings -> choices; label -> target; source_id groups", | |
| "question": "Which continuation is most plausible given this context?", | |
| }, | |
| "piqa": { | |
| "repo": "ybisk/ybisk.github.io", "provider": "github", | |
| "revision": "21edab439af693b961be2f069e8690a88d3b4e37", | |
| "license": "AFL-3.0", "license_reference": "piqa/README.md", | |
| "expected_train_rows": 16113, "group_columns": ["goal"], | |
| "transform": "goal -> state; sol1/sol2 -> choices; train-labels.lst -> target; goal groups", | |
| "question": "Which solution best achieves the goal?", | |
| }, | |
| "commonsenseqa": { | |
| "repo": "tau/commonsense_qa", "provider": "huggingface", | |
| "revision": "94630fe30dad47192a8546eb75f094926d47e155", | |
| "license": "MIT", "license_reference": "README.md", | |
| "expected_train_rows": 9741, "group_columns": ["question"], | |
| "transform": "question -> state; choices.text -> choices; answerKey mapped through choices.label; question groups", | |
| "question": "Which candidate correctly answers the question in the state?", | |
| }, | |
| } | |
| def sha256(path): | |
| digest = hashlib.sha256() | |
| with Path(path).open("rb") as handle: | |
| for chunk in iter(lambda: handle.read(1024 * 1024), b""): | |
| digest.update(chunk) | |
| return digest.hexdigest() | |
| def normalized(text): | |
| return " ".join(unicodedata.normalize("NFKC", text).casefold().split()) | |
| def fingerprint(text): | |
| return hashlib.sha256(normalized(text).encode()).hexdigest() | |
| def json_write(path, value): | |
| Path(path).write_text(json.dumps(value, indent=2, sort_keys=True) + "\n") | |
| def read_rows(path): | |
| with Path(path).open() as handle: | |
| for line in handle: | |
| if line.strip(): | |
| yield json.loads(line) | |
| def source_keys(family, row): | |
| """Return only group/text identities; never examine evaluation labels.""" | |
| state = row[{"hellaswag": "ctx", "piqa": "goal", "commonsenseqa": "question"}[family]] | |
| if not isinstance(state, str) or not state.strip(): | |
| raise ValueError(f"{family}: missing nonempty state") | |
| source_group = row.get("source_id") if family == "hellaswag" else state | |
| if not isinstance(source_group, str) or not source_group.strip(): | |
| raise ValueError(f"{family}: missing source group") | |
| return family + ":" + fingerprint(source_group), fingerprint(state) | |
| def convert_row(family, index, raw): | |
| source = SOURCES[family] | |
| group, _ = source_keys(family, raw) | |
| if family == "hellaswag": | |
| state, choices, target = raw["ctx"], raw["endings"], int(raw["label"]) | |
| original_id = str(raw["ind"]) | |
| elif family == "piqa": | |
| state, choices, target = raw["goal"], [raw["sol1"], raw["sol2"]], int(raw["label"]) | |
| original_id = str(index) | |
| elif family == "commonsenseqa": | |
| state, choices = raw["question"], raw["choices"]["text"] | |
| target = raw["choices"]["label"].index(raw["answerKey"]) | |
| original_id = str(raw["id"]) | |
| else: | |
| raise ValueError("Unknown source family") | |
| expected_choices = {"hellaswag": 4, "piqa": 2, "commonsenseqa": 5}[family] | |
| if (len(choices) != expected_choices or not 0 <= target < len(choices) | |
| or any(not isinstance(choice, str) or not choice.strip() for choice in choices) | |
| or len({normalized(choice) for choice in choices}) != len(choices)): | |
| raise ValueError(f"{family}: invalid choices or target at row {index}") | |
| row_id = f"{family}:train:{original_id}" | |
| order = list(range(len(choices))) | |
| random.Random(fingerprint("expanded-choices:" + row_id)).shuffle(order) | |
| return {"id": row_id, "group": group, "family": family, "source_split": "train", | |
| "source_record_id": original_id, "source_row_index": index, | |
| "source_revision": source["revision"], "state": state, | |
| "question": source["question"], "choices": [choices[i] for i in order], | |
| "target": order.index(target)} | |
| def download_sources(raw_directory): | |
| """Download pinned artifacts and project nontraining files to group metadata.""" | |
| os.environ.setdefault("HF_HUB_DISABLE_XET", "1") | |
| from huggingface_hub import hf_hub_download | |
| import pyarrow.parquet as pq | |
| raw_directory = Path(raw_directory) | |
| raw_directory.mkdir(parents=True, exist_ok=False) | |
| train, reserved, provenance = {}, {}, {} | |
| for family, source in SOURCES.items(): | |
| directory = raw_directory / family | |
| directory.mkdir() | |
| files = {} | |
| def acquire(filename): | |
| if source["provider"] == "huggingface": | |
| path = Path(hf_hub_download(source["repo"], filename, repo_type="dataset", | |
| revision=source["revision"], local_dir=directory)) | |
| url = f"https://huggingface.co/datasets/{source['repo']}/resolve/{source['revision']}/{filename}" | |
| else: | |
| path = directory / filename | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| url = f"https://raw.githubusercontent.com/{source['repo']}/{source['revision']}/{filename}" | |
| with urllib.request.urlopen(url, timeout=120) as response: | |
| with path.open("xb") as handle: | |
| shutil.copyfileobj(response, handle) | |
| files[filename] = {"sha256": sha256(path), "bytes": path.stat().st_size, | |
| "url": url, "local_path": str(path.relative_to(raw_directory.parent))} | |
| return path | |
| acquire(source["license_reference"]) | |
| reserved[family] = [] | |
| split_counts = {} | |
| if source["provider"] == "huggingface": | |
| for split in ("train", "validation", "test"): | |
| path = acquire(f"data/{split}-00000-of-00001.parquet") | |
| columns = None if split == "train" else source["group_columns"] | |
| rows = pq.read_table(path, columns=columns).to_pylist() | |
| split_counts[split] = len(rows) | |
| if split == "train": | |
| train[family] = rows | |
| else: | |
| reserved[family].extend(rows) | |
| else: | |
| rows = list(read_rows(acquire("piqa/data/train.jsonl"))) | |
| labels = acquire("piqa/data/train-labels.lst").read_text().splitlines() | |
| if len(rows) != len(labels): | |
| raise ValueError("PIQA train/label count mismatch") | |
| train[family] = [{**row, "label": int(label)} for row, label in zip(rows, labels)] | |
| split_counts["train"] = len(rows) | |
| for split, name in (("validation", "valid.jsonl"), ("test", "tests.jsonl")): | |
| # Only the goal participates in exclusions; no evaluation label file is fetched. | |
| rows = [{"goal": row["goal"]} for row in read_rows(acquire("piqa/data/" + name))] | |
| split_counts[split] = len(rows) | |
| reserved[family].extend(rows) | |
| if len(train[family]) != source["expected_train_rows"]: | |
| raise ValueError(f"{family}: pinned train row count changed") | |
| provenance[family] = {**source, "files": files, "source_counts": split_counts, | |
| "evaluation_access": "group/text metadata only; no evaluation labels used"} | |
| print(json.dumps({"event": "source_ready", "family": family, "counts": split_counts}), flush=True) | |
| return train, reserved, provenance | |
| def prepare_dataset(base, output, source_train, source_reserved, source_provenance, | |
| max_new_per_family=16000, diagnostic_per_family=128, seed=917): | |
| """Write a complete new dataset. All source rows passed here are in-memory data.""" | |
| base, output = Path(base).resolve(), Path(output).resolve() | |
| if output == base or output.is_relative_to(base): | |
| raise ValueError("Expanded output must be separate from its immutable base") | |
| if max_new_per_family < 0 or diagnostic_per_family < 0: | |
| raise ValueError("Sampling limits must be nonnegative") | |
| base_manifest = json.loads((base / "manifest.json").read_text()) | |
| if set(base_manifest["split_sha256"]) != set(SPLITS): | |
| raise ValueError("Base dataset must have exactly the five frozen decision splits") | |
| for split, checksum in base_manifest["split_sha256"].items(): | |
| if sha256(base / f"{split}.jsonl") != checksum: | |
| raise ValueError("Base dataset hash mismatch: " + split) | |
| original = list(read_rows(base / "train.jsonl")) | |
| if any(row["family"] in {"social", "socialiqa", "social_i_qa"} for row in original): | |
| raise ValueError("Social IQA must never enter training") | |
| original_ids = {row["id"] for row in original} | |
| if len(original_ids) != len(original): | |
| raise ValueError("Duplicate original training ID") | |
| old_train_text = {fingerprint(row["state"]) for row in original} | |
| reserved_groups, reserved_text = set(), set() | |
| for split in PROTECTED: | |
| for row in read_rows(base / f"{split}.jsonl"): | |
| # Label/choice content does not affect exclusions or sampling. | |
| reserved_groups.add(row["group"]) | |
| reserved_text.add(fingerprint(row["state"])) | |
| if {row["group"] for row in original} & reserved_groups: | |
| raise ValueError("Base dataset has train/evaluation group overlap") | |
| additions, diagnostics, audit = [], [], {} | |
| seen_ids, seen_text = set(original_ids), set(old_train_text) | |
| for family in SOURCES: | |
| official_groups, official_text = set(), set() | |
| for row in source_reserved[family]: | |
| group, text = source_keys(family, row) | |
| official_groups.add(group) | |
| official_text.add(text) | |
| removed, candidates = Counter(), [] | |
| for index, raw in enumerate(source_train[family]): | |
| try: | |
| row = convert_row(family, index, raw) | |
| except (ValueError, TypeError, KeyError, IndexError): | |
| removed["invalid_source_row"] += 1 | |
| continue | |
| text = fingerprint(row["state"]) | |
| if row["group"] in official_groups or text in official_text: | |
| removed["official_evaluation_overlap"] += 1 | |
| elif row["group"] in reserved_groups or text in reserved_text: | |
| removed["original_evaluation_overlap"] += 1 | |
| elif row["id"] in seen_ids or text in seen_text: | |
| removed["duplicate_training_identity_or_state"] += 1 | |
| else: | |
| seen_ids.add(row["id"]) | |
| seen_text.add(text) | |
| candidates.append(row) | |
| groups = sorted({row["group"] for row in candidates}, | |
| key=lambda group: fingerprint(f"diagnostics:{seed}:{group}")) | |
| if len(groups) <= diagnostic_per_family: | |
| raise ValueError(f"{family}: insufficient groups after exclusions") | |
| diagnostic_groups = set(groups[:diagnostic_per_family]) | |
| diagnostic_rows = {} | |
| train_rows = [] | |
| for row in candidates: | |
| if row["group"] in diagnostic_groups: | |
| diagnostic_rows.setdefault(row["group"], row) | |
| else: | |
| train_rows.append(row) | |
| # Sampling is deterministic and independent of input/parquet row ordering. | |
| train_rows.sort(key=lambda row: (fingerprint(f"train:{seed}:{row['id']}"), row["id"])) | |
| selected = train_rows[:max_new_per_family] if max_new_per_family else train_rows | |
| additions.extend(selected) | |
| diagnostics.extend(diagnostic_rows[group] for group in groups[:diagnostic_per_family]) | |
| audit[family] = {"raw_train_rows": len(source_train[family]), "removed": dict(removed), | |
| "official_reserved_groups": len(official_groups), | |
| "diagnostic_groups": len(diagnostic_groups), | |
| "diagnostic_group_rows_excluded_from_training": len(candidates) - len(train_rows), | |
| "available_training_rows": len(train_rows), "retained_training_rows": len(selected), | |
| "cap_excluded_rows": len(train_rows) - len(selected)} | |
| if {row["group"] for row in additions} & {row["group"] for row in diagnostics}: | |
| raise ValueError("Expanded train/diagnostic group overlap") | |
| if {fingerprint(row["state"]) for row in additions} & reserved_text: | |
| raise ValueError("Expanded training state overlaps frozen evaluation") | |
| output.mkdir(parents=True, exist_ok=False) | |
| for split in PROTECTED: | |
| shutil.copyfile(base / f"{split}.jsonl", output / f"{split}.jsonl") | |
| original_bytes = (base / "train.jsonl").read_bytes() | |
| if original_bytes and not original_bytes.endswith(b"\n"): | |
| raise ValueError("Original training JSONL must end in a newline for byte-exact replay") | |
| # Original replay is a byte-exact prefix. The trainer shuffles all rows per epoch. | |
| with (output / "train.jsonl").open("wb") as handle: | |
| handle.write(original_bytes) | |
| for row in additions: | |
| handle.write((json.dumps(row, sort_keys=True) + "\n").encode()) | |
| diagnostic_path = output / "diagnostics" / "new_sources.jsonl" | |
| diagnostic_path.parent.mkdir() | |
| diagnostic_path.write_text("".join(json.dumps(row, sort_keys=True) + "\n" for row in diagnostics)) | |
| split_hashes = {split: sha256(output / f"{split}.jsonl") for split in SPLITS} | |
| protected_hashes = {split: base_manifest["split_sha256"][split] for split in PROTECTED} | |
| if any(split_hashes[split] != expected for split, expected in protected_hashes.items()): | |
| raise ValueError("Frozen evaluation changed during copy") | |
| counts = {**base_manifest["split_family_counts"], | |
| "train": dict(Counter(row["family"] for row in [*original, *additions]))} | |
| manifest = {"version": VERSION, "created_utc": datetime.now(timezone.utc).isoformat(), | |
| "base_dataset": {"path": str(base), "manifest_sha256": sha256(base / "manifest.json"), | |
| "split_sha256": base_manifest["split_sha256"]}, | |
| "sources": {**base_manifest["sources"], **source_provenance}, | |
| "source_script_sha256": sha256(__file__), "split_sha256": split_hashes, | |
| "protected_split_sha256": protected_hashes, "split_family_counts": counts, | |
| "diagnostics": {"path": str(diagnostic_path.relative_to(output)), "sha256": sha256(diagnostic_path), | |
| "family_counts": dict(Counter(row["family"] for row in diagnostics)), | |
| "selection_eligible": False, "source_split": "train"}, | |
| "sampling": {"seed": seed, "max_new_per_family": max_new_per_family, | |
| "diagnostic_groups_per_family": diagnostic_per_family, | |
| "trainer_sampler": "unchanged uniform rows without replacement per epoch", | |
| "original_replay_rows": len(original), "original_replay_bytes": len(original_bytes), | |
| "original_replay_sha256": base_manifest["split_sha256"]["train"], | |
| "new_training_rows": len(additions)}, | |
| "audit": audit, "group_leakage": False, | |
| "holdout": "Original Social IQA holdout preserved byte-for-byte; no Social IQA training", | |
| "normalization": "Unicode NFKC, casefold, collapsed whitespace, SHA256", | |
| "limitations": "Exact normalized states/groups only; semantic duplicates and pretraining contamination not excluded. New-source diagnostics are outside fixed checkpoint selection."} | |
| json_write(output / "manifest.json", manifest) | |
| return manifest | |
| def main(): | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument("--base", required=True) | |
| parser.add_argument("--output", required=True) | |
| parser.add_argument("--max-new-per-family", type=int, default=16000) | |
| parser.add_argument("--diagnostic-per-family", type=int, default=128) | |
| parser.add_argument("--seed", type=int, default=917) | |
| args = parser.parse_args() | |
| output = Path(args.output).expanduser().resolve() | |
| repository = Path(__file__).resolve().parents[1] | |
| if output.exists() or output.is_relative_to(repository): | |
| parser.error("Output must be a new directory outside the source repository") | |
| output.parent.mkdir(parents=True, exist_ok=True) | |
| staging = Path(tempfile.mkdtemp(prefix=f".{output.name}-preparing-", dir=output.parent)) | |
| train, reserved, provenance = download_sources(staging / "raw") | |
| prepared = staging / "dataset" | |
| manifest = prepare_dataset(args.base, prepared, train, reserved, provenance, | |
| args.max_new_per_family, args.diagnostic_per_family, args.seed) | |
| shutil.move(staging / "raw", prepared / "raw") | |
| manifest["source_git_commit"] = subprocess.check_output(["git", "rev-parse", "HEAD"], cwd=repository, text=True).strip() | |
| manifest["source_git_status"] = subprocess.check_output(["git", "status", "--porcelain"], cwd=repository, text=True) | |
| json_write(prepared / "manifest.json", manifest) | |
| if output.exists(): | |
| raise ValueError("Output appeared during preparation; refusing overwrite") | |
| prepared.rename(output) | |
| staging.rmdir() | |
| print(json.dumps({"event": "expanded_dataset_ready", "path": str(output), | |
| "manifest_sha256": sha256(output / "manifest.json"), | |
| "counts": manifest["split_family_counts"], "audit": manifest["audit"]}, indent=2), flush=True) | |
| if __name__ == "__main__": | |
| main() | |