IncidentResponseDetective / scripts /verify_splits.py
90shikhar08's picture
feat(splits): immutable train/val/test seed splits v1.0.0
16607b8
Raw
History Blame Contribute Delete
2.86 kB
"""Verify train/val/test seed split JSON files (disjointness, counts, checksums)."""
from __future__ import annotations
import json
import sys
from pathlib import Path
ROOT = Path(__file__).resolve().parent.parent
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
import splits # noqa: E402
_EXPECTED_TOTAL = {"train": 2100, "val": 450, "test": 450}
_EXPECTED_PER_ARCH = {"train": 700, "val": 150, "test": 150}
def _fail(msg: str) -> None:
print(msg, file=sys.stderr)
sys.exit(1)
def main() -> None:
docs: dict[str, dict] = {}
sets_by_split: dict[str, set[int]] = {}
for name in splits._SPLIT_NAMES:
path = splits._split_path(name)
if not path.is_file():
_fail(f"missing split file: {path}")
data = json.loads(path.read_text(encoding="utf-8"))
docs[name] = data
if data.get("split") != name:
_fail(f"{path}: expected split={name!r}, got {data.get('split')!r}")
if data.get("version") != splits.SPLIT_VERSION:
_fail(
f"{path}: expected version={splits.SPLIT_VERSION!r}, got {data.get('version')!r}"
)
if data.get("master_seed") != splits.MASTER_SEED:
_fail(
f"{path}: expected master_seed={splits.MASTER_SEED!r}, got {data.get('master_seed')!r}"
)
sb = data["seeds_by_archetype"]
actual = splits._checksum_seeds_only(sb)
expected = data["checksum"]
if actual != expected:
_fail(
f"{path}: checksum mismatch: expected {expected!r}, got {actual!r}"
)
for arch in ("easy", "medium", "hard"):
n = len(sb[arch])
want = _EXPECTED_PER_ARCH[name]
if n != want:
_fail(
f"{path}: {arch}: expected {want} seeds, got {n}"
)
flat = [*sb["easy"], *sb["medium"], *sb["hard"]]
if len(flat) != _EXPECTED_TOTAL[name]:
_fail(
f"{path}: expected {_EXPECTED_TOTAL[name]} total seeds, got {len(flat)}"
)
if len(set(flat)) != len(flat):
_fail(f"{path}: duplicate seed within split")
sets_by_split[name] = set(flat)
if (
sets_by_split["train"] & sets_by_split["val"]
or sets_by_split["train"] & sets_by_split["test"]
or sets_by_split["val"] & sets_by_split["test"]
):
_fail("splits overlap: same seed appears in more than one split")
union_size = len(
sets_by_split["train"] | sets_by_split["val"] | sets_by_split["test"]
)
if union_size != 3000:
_fail(f"expected 3000 globally unique seeds, got {union_size}")
print(
"verify_splits: OK (checksums, per-archetype counts, global disjointness)"
)
if __name__ == "__main__":
main()