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