| |
| """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() |
|
|