Avra98's picture
add data-audit and metric-interpretation scripts
96b4304 verified
Raw
History Blame Contribute Delete
3.48 kB
"""How often has each candidate digit actually been shown, so far and by the
end of the epoch.
Counts are a property of the stored pool: if digit d appears in c of puzzle p's
instances at that cell, then after P passes over the pair list the model has
seen (cell -> d) exactly c*P times. So the whole exposure picture follows from
the pool's per-(cell, candidate) count histogram times the pass count.
"""
import argparse
import numpy as np
D = "/tmp/sudoku_superposition"
CAND = "/tmp/sudoku_s12/train_cand_masks.npy"
STAGE = 0
BS = 64
def digits(m):
return [d for d in range(1, 10) if int(m) & (1 << (d - 1))]
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--step", type=int, required=True)
ap.add_argument("--puzzles", type=int, default=300)
args = ap.parse_args()
starts = np.load(f"{D}/train_starts.npy")
counts = np.load(f"{D}/train_counts.npy")
A = np.load(f"{D}/train_assignments.npy", mmap_mode="r")
masks = np.load(CAND, mmap_mode="r")
n_puz = counts.shape[0]
pairs = int(np.asarray(counts[:, STAGE]).astype(np.int64).sum())
per_pass = pairs / BS
P = args.step / per_pass
print(f"TRAIN {n_puz:,} puzzles stage-1 (puzzle, instance) pairs "
f"{pairs:,} {pairs / n_puz:.2f} instances/puzzle")
print(f"one pass = {per_pass:,.0f} steps at batch {BS}; "
f"one epoch = 5 passes = {5 * per_pass:,.0f} steps")
print(f"step {args.step:,} -> {P:.2f} passes done "
f"({P / 5 * 100:.1f}% of the epoch)")
print(f" each puzzle shown {P * pairs / n_puz:.1f} times so far, "
f"{5 * pairs / n_puz:.1f} by the end")
print(f" each specific instance seen {P:.2f} times so far, 5 at the end")
# per-(cell, candidate) counts inside the pool
hist = np.zeros(8, dtype=np.int64)
tot = 0
ncand_sum = 0
ncell = 0
for q in range(args.puzzles):
m = np.asarray(masks[q, STAGE])
n = int(counts[q, STAGE])
blk = np.asarray(A[starts[q, STAGE]:starts[q, STAGE] + n])
for c in range(81):
S = digits(m[c])
if len(S) < 2:
continue
ncell += 1
ncand_sum += len(S)
v, k = np.unique(blk[:, c], return_counts=True)
cnt = {int(a): int(b) for a, b in zip(v, k)}
for d in S:
hist[min(cnt.get(d, 0), 7)] += 1
tot += 1
print(f"\nover {args.puzzles} puzzles, {ncell:,} multi-candidate cells, "
f"mean |S| {ncand_sum / ncell:.2f}")
print(f"how many of the puzzle's ~6 instances carry each candidate:")
print(f" {'times in pool':>14} {'share of (cell,cand)':>21} "
f"{'seen so far':>12} {'seen by epoch end':>18}")
for c in range(8):
lab = f"{c}" if c < 7 else "7+"
print(f" {lab:>14} {hist[c] / tot * 100:>20.1f}% "
f"{c * P:>12.1f} {c * 5:>18}")
mean_c = float((np.arange(8) * hist).sum() / tot)
print(f"\n mean over candidates: {mean_c:.2f} instances carry it, so "
f"~{mean_c * P:.1f} sightings so far, ~{mean_c * 5:.1f} by epoch end")
print(f" never shown (0 in pool): {hist[0] / tot * 100:.1f}% of "
f"(cell, candidate) pairs -- stays 0 forever")
tst = np.load(f"{D}/test_counts.npy")
print(f"\nVALIDATION {tst.shape[0]:,} test puzzles available")
print(f" each eval uses eval_epochs=5 x batch {BS} = {5 * BS} puzzles")
if __name__ == "__main__":
main()