Sudoku_superposition / code /wavecurriculum_run /verify_uniform_loader.py
Avra98's picture
sync training code: stage-1 instance-epoch sampler, multi-stage run, superposition metrics
6a1771b verified
Raw
History Blame Contribute Delete
4.83 kB
"""Verify the on-the-fly uniform instance sampler on the real stage-1 masks.
Replicates SudokuDataset.uniform_instance_values / apply_instance with numpy
only, then checks the three things the experiment depends on:
1. the clue block is untouched, so the input prompt is fixed per puzzle
2. every drawn digit lies in that cell's candidate set
3. over N draws each candidate of each cell turns up a near-equal number of
times, and at least 5 times
"""
import numpy as np
PUZ = ("/scratch/users/gatmiry/llm-reasoning-logic-puzzles/sudoku-code/"
"datasets/train_sudoku_puzzles.npy")
CAND = "/tmp/sudoku_s12/train_cand_masks.npy"
STAGE = 0
N_DRAWS = 64
K, LATENT_ID = 12, 10
def load_inputs(n):
raw = np.load(PUZ, mmap_mode="r")
rows = np.asarray(raw[:n]).astype(np.int64)
return np.delete(rows[:, 1:], np.arange(81) * 4 + 3, axis=1), rows[:, 0]
def uniform_instance_values(mask, rng):
"""Same argmax-of-random-keys draw as the loader."""
m = mask.astype(np.int64)
bits = ((m[:, None] >> np.arange(9)) & 1).astype(np.float64)
keys = rng.random_sample((81, 9)) * bits
vals = (keys.argmax(1) + 1).astype(np.int8)
return np.where(m > 0, vals, 0)
def apply_instance(seq, vals):
seq = seq.copy()
cells = seq[0::3] * 9 + seq[1::3]
new = vals[cells]
seq[2::3] = np.where(new > 0, new, seq[2::3])
return seq
def digits_of(m):
return [d for d in range(1, 10) if int(m) & (1 << (d - 1))]
def main():
n_show = 500
inputs, si_all = load_inputs(n_show)
masks = np.load(CAND, mmap_mode="r")
rng = np.random.RandomState(0)
p = 3
si, seq = int(si_all[p]), inputs[p]
mask = np.asarray(masks[p, STAGE])
tr = seq.reshape(-1, 3)
empty = (tr[si:, 0] * 9 + tr[si:, 1])
drawn = [apply_instance(seq, uniform_instance_values(mask, rng))
for _ in range(N_DRAWS)]
print(f"puzzle p={p}: {si} clues, {len(empty)} empty cells, "
f"{N_DRAWS} uniform draws")
clue_ok = all(np.array_equal(d[:3 * si], seq[:3 * si]) for d in drawn)
print(f" clue block identical in all {N_DRAWS} draws: {clue_ok}")
bad = sum(int(d[2::3][j + si]) not in digits_of(mask[empty[j]])
for d in drawn for j in range(len(empty)))
print(f" drawn digits outside the candidate set: {bad}")
blk = np.stack([d[2::3] for d in drawn])[:, si:] # (N, n_empty)
print(f"\ninstance = one digit for every empty cell. First 10 cells:")
print(" cells: " + " ".join(f"({c // 9},{c % 9})" for c in empty[:10]))
for i in range(4):
print(f" i{i + 1}: " + " ".join(
f"{int(x)}" for x in blk[i, :10]))
print(" S: " + " ".join(f"{digits_of(mask[c])}" for c in empty[:4])
+ " ...")
print(f"\nreading DOWN a cell's column over {N_DRAWS} draws:")
print(f" {'cell':>8} {'|S|':>3} {'S':<16}{'counts':<34}min")
shown = 0
worst = 10 ** 9
for j, c in enumerate(empty):
S = digits_of(mask[c])
if len(S) < 2:
continue
v, k = np.unique(blk[:, j], return_counts=True)
cnt = {int(a): int(b) for a, b in zip(v, k)}
worst = min(worst, min(cnt.get(d, 0) for d in S))
shown += 1
if shown > 8:
continue
cs = " ".join(f"{d}:{cnt.get(d, 0)}" for d in S)
print(f" ({c // 9},{c % 9}) {len(S):>3} {str(S):<16}{cs:<34}"
f"{min(cnt.get(d, 0) for d in S)}")
print(f" ... {shown} multi-candidate cells; rarest candidate anywhere in "
f"this puzzle appeared {worst}x (need >=5)")
print(f"\nover {n_show} puzzles:")
ge5 = tot = 0
spreads = []
for q in range(n_show):
m = np.asarray(masks[q, STAGE])
b = np.stack([uniform_instance_values(m, rng) for _ in range(N_DRAWS)])
for c in range(81):
S = digits_of(m[c])
if len(S) < 2:
continue
v, k = np.unique(b[:, c], return_counts=True)
cnt = {int(a): int(b_) for a, b_ in zip(v, k)}
tot += len(S)
ge5 += sum(1 for d in S if cnt.get(d, 0) >= 5)
pr = np.array([cnt.get(d, 0) for d in S], dtype=float)
pr /= pr.sum()
nz = pr[pr > 0]
spreads.append(float(-(nz * np.log(nz)).sum()) / np.log(len(S)))
print(f" (cell, candidate) pairs seen >=5 times: {ge5 / tot:.4f}")
print(f" mean spread H(p)/log|S|: {np.mean(spreads):.4f} "
f"(1.0 = uniform superposition)")
full = np.concatenate([drawn[0][:3 * si],
np.full(K, LATENT_ID, dtype=drawn[0].dtype),
drawn[0][3 * si:]])
print(f"\ntoken sequence: {len(full)} = {3 * si} clue + {K} latent + "
f"{3 * (81 - si)} output")
if __name__ == "__main__":
main()