File size: 4,891 Bytes
f1fc3a0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 | #!/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()
|