"""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()