Sudoku_superposition / code /smoke_instance_loader.py
Avra98's picture
Add README and training/generation code
bb23b91 verified
Raw
History Blame Contribute Delete
2.8 kB
"""Sanity-check: instance rewrite keeps clues, varies empties, stays in-set."""
import os
import sys
import numpy as np
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "wavecurriculum_run"))
os.environ.setdefault("CUDA_VISIBLE_DEVICES", "")
from train.data import CurriculumState, SudokuDataset # noqa: E402
class Cfg:
seed = 7
seq_order = "solver-order"
num_latent_slots = 12
latent_token_id = 10
curriculum_max_stage = 12
data_curriculum = "none"
cand_slot_mode = "depth"
passes_per_stage = 1
level_balanced_sampling = 0
train_puzzle_path = "datasets/train_sudoku_puzzles.npy"
test_puzzle_path = "datasets/test_sudoku_puzzles.npy"
train_cand_masks_path = "datasets_multicandidate_s12/train_cand_masks.npy"
test_cand_masks_path = "datasets_multicandidate_s12/test_cand_masks.npy"
instance_dir = "datasets_superposition"
train_meta_path = None
test_meta_path = None
def main():
cur = CurriculumState(stage=1, max_stage=12)
ds = SudokuDataset(Cfg(), train=True, curriculum=cur)
assert ds.instances is not None
assert ds.cand_masks is not None
n_check = 8
for stage in (0, 5, 11):
cur.stage = stage + 1
clue_ok = clue_tot = 0
empty_in = empty_tot = empty_eq = 0
for idx in range(n_check):
base = ds.train_inputs[idx].copy()
si = int(ds.train_start_index[idx, 0])
rewritten = ds.apply_instance(base, idx, stage)
sol = ds.train_puzzles[idx]
mask = np.array(ds.cand_masks[idx, stage])
orig = {(int(t[0]), int(t[1])): int(t[2])
for t in base.reshape(81, 3)}
for t in rewritten[: 3 * si].reshape(-1, 3):
r, c, v = int(t[0]), int(t[1]), int(t[2])
clue_tot += 1
clue_ok += int(v == orig[(r, c)] == int(sol[r * 9 + c]))
for t in rewritten[3 * si:].reshape(-1, 3):
r, c, v = int(t[0]), int(t[1]), int(t[2])
bits = int(mask[r * 9 + c])
empty_tot += 1
empty_in += int(1 <= v <= 9 and (bits >> (v - 1)) & 1)
empty_eq += int(v == int(sol[r * 9 + c]))
print(f"stage {stage}: clues {clue_ok}/{clue_tot} "
f"in-set {empty_in}/{empty_tot} "
f"==sol {empty_eq}/{empty_tot}")
assert clue_ok == clue_tot, "clue values must stay the original clues"
assert empty_in == empty_tot, "every empty must stay in the stage set"
if stage == 11:
assert empty_eq == empty_tot, "stage 11 must be the unique solution"
else:
assert empty_eq < empty_tot, "early stages must still be ambiguous"
print("SMOKE OK")
if __name__ == "__main__":
main()