pactbench / pact /build_grouped_train_val_split.py
BBoran's picture
Publish current portable PACTBench release
f1fc3a0 verified
Raw
History Blame Contribute Delete
4.89 kB
#!/usr/bin/env python3
"""Build a deterministic boss-stratified, fight-disjoint train/validation split."""
from __future__ import annotations
import argparse
import hashlib
import json
import os
import random
from collections import defaultdict
from typing import Any, Dict, List, Sequence, Tuple
import torch
def row_key(row: Dict[str, Any]) -> List[Any]:
return [str(row.get("boss")), int(row.get("fight", -1)), int(row.get("index", -1))]
def keys_sha256(keys: Sequence[Sequence[Any]]) -> str:
return hashlib.sha256(
json.dumps(keys, ensure_ascii=False, separators=(",", ":")).encode("utf-8")
).hexdigest()
def build_grouped_split(
rows: Sequence[Dict[str, Any]], val_fraction: float, seed: int
) -> Tuple[List[List[Any]], List[List[Any]], Dict[str, Any]]:
if not 0.0 < val_fraction < 0.5:
raise ValueError("val_fraction must be in (0, 0.5)")
grouped: Dict[str, Dict[int, List[int]]] = defaultdict(lambda: defaultdict(list))
for index, row in enumerate(rows):
grouped[str(row.get("boss"))][int(row.get("fight", -1))].append(index)
validation_groups = set()
by_boss: Dict[str, Any] = {}
for boss, fights in sorted(grouped.items()):
if len(fights) < 2:
raise ValueError(f"boss {boss} has fewer than two fights")
ordered = sorted(fights)
boss_seed = int.from_bytes(
hashlib.sha256(f"{seed}:{boss}".encode("utf-8")).digest()[:8], "big"
)
random.Random(boss_seed).shuffle(ordered)
target = sum(len(fights[fight]) for fight in ordered) * val_fraction
selected: List[int] = []
selected_rows = 0
for fight in ordered:
candidate = selected_rows + len(fights[fight])
if not selected or abs(candidate - target) <= abs(selected_rows - target):
selected.append(fight)
selected_rows = candidate
else:
break
if len(selected) == len(fights):
removed = selected.pop()
selected_rows -= len(fights[removed])
validation_groups.update((boss, fight) for fight in selected)
total_rows = sum(len(indices) for indices in fights.values())
by_boss[boss] = {
"total_rows": total_rows,
"train_rows": total_rows - selected_rows,
"validation_rows": selected_rows,
"total_fights": len(fights),
"validation_fights": len(selected),
}
train_keys: List[List[Any]] = []
validation_keys: List[List[Any]] = []
for row in rows:
key = row_key(row)
target = validation_keys if (key[0], key[1]) in validation_groups else train_keys
target.append(key)
if len(train_keys) + len(validation_keys) != len(rows):
raise AssertionError("split does not partition all rows")
audit = {
"schema_revision": "fight-disjoint-validation-split-v1",
"seed": seed,
"validation_fraction_requested": val_fraction,
"n_rows": len(rows),
"n_train": len(train_keys),
"n_validation": len(validation_keys),
"validation_fraction_actual": len(validation_keys) / max(1, len(rows)),
"by_boss": by_boss,
}
return train_keys, validation_keys, audit
def write_manifest(path: str, split: str, keys: List[List[Any]], audit: Dict[str, Any]) -> None:
value = {
**audit,
"split": split,
"n_rows": len(keys),
"row_keys_sha256": keys_sha256(keys),
"row_keys": keys,
}
os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
tmp = f"{path}.tmp.{os.getpid()}"
with open(tmp, "w", encoding="utf-8") as handle:
json.dump(value, handle, ensure_ascii=False, indent=2, allow_nan=False)
handle.flush()
os.fsync(handle.fileno())
os.replace(tmp, path)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--cache", required=True)
parser.add_argument("--train_out", required=True)
parser.add_argument("--validation_out", required=True)
parser.add_argument("--validation_fraction", type=float, default=0.1)
parser.add_argument("--seed", type=int, default=20260718)
args = parser.parse_args()
payload = torch.load(args.cache, map_location="cpu", weights_only=False, mmap=True)
rows = payload.get("samples") or []
if not rows:
raise ValueError(f"cache has no samples: {args.cache}")
train_keys, validation_keys, audit = build_grouped_split(
rows, args.validation_fraction, args.seed
)
write_manifest(args.train_out, "train", train_keys, audit)
write_manifest(args.validation_out, "validation", validation_keys, audit)
print(json.dumps(audit, ensure_ascii=False, indent=2))
print(f"wrote {args.train_out}")
print(f"wrote {args.validation_out}")
if __name__ == "__main__":
main()