Sudoku_superposition / code /build_instance_offsets.py
Avra98's picture
Add README and training/generation code
bb23b91 verified
Raw
History Blame Contribute Delete
1.76 kB
"""Precompute (puzzle, stage) -> row range into the instance assignments.
The index rows are grouped by puzzle then stage, so a start offset plus a count
per (puzzle, stage) is enough for O(1) sampling in the data loader.
Writes {split}_starts.npy (N, S) int32 and {split}_counts.npy (N, S) uint8.
"""
import argparse
import os
import numpy as np
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--dir", default="datasets_superposition")
ap.add_argument("--splits", default="train,test")
ap.add_argument("--stages", type=int, default=12)
args = ap.parse_args()
S = args.stages
for split in args.splits.split(","):
split = split.strip()
ix = np.load(os.path.join(args.dir, f"{split}_index.npy"), mmap_mode="r")
n_puzzles = int(np.array(ix[-1, 0])) + 1
key = np.array(ix[:, 0]).astype(np.int64) * S + np.array(ix[:, 1])
counts = np.bincount(key, minlength=n_puzzles * S).astype(np.int64)
if counts.max() > 255:
raise ValueError(f"count {counts.max()} exceeds uint8")
starts = np.concatenate([[0], np.cumsum(counts)[:-1]]).astype(np.int64)
starts = starts.reshape(n_puzzles, S).astype(np.int32)
counts = counts.reshape(n_puzzles, S).astype(np.uint8)
np.save(os.path.join(args.dir, f"{split}_starts.npy"), starts)
np.save(os.path.join(args.dir, f"{split}_counts.npy"), counts)
empty = int((counts == 0).sum())
print(f"[{split}] puzzles={n_puzzles:,} rows={len(ix):,} "
f"(puzzle,stage) cells with zero instances: {empty:,}")
print(f" mean instances per stage: "
+ " ".join(f"{counts[:, s].mean():.2f}" for s in range(S)))
if __name__ == "__main__":
main()