#!/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()