File size: 2,796 Bytes
bb23b91
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
"""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()