| """Worked example: one cell with candidate set S = {1,2,4}, four model behaviours. |
| |
| Shows that excess = log(1/mass) + KL(Uniform(S) || p/mass) separates "support |
| leaks outside S" from "support inside S is not uniform". |
| """ |
| import sys |
| import os |
| import types |
|
|
| import numpy as np |
|
|
| for name in ("jax", "jax.numpy", "flax", "flax.training", |
| "flax.training.common_utils", "train.model"): |
| sys.modules.setdefault(name, types.ModuleType(name)) |
| sys.modules["jax"].numpy = sys.modules["jax.numpy"] |
| sys.modules["flax"].training = sys.modules["flax.training"] |
| sys.modules["flax.training"].common_utils = sys.modules[ |
| "flax.training.common_utils"] |
| sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) |
| from train.evaluater import _cand_value_stats |
|
|
| CAND = [1, 2, 4] |
| BITS = int(sum(1 << (d - 1) for d in CAND)) |
|
|
| cases = [ |
| ("A perfect superposition ", {1: 1 / 3, 2: 1 / 3, 4: 1 / 3}), |
| ("B uniform over all 9 ", {d: 1 / 9 for d in range(1, 10)}), |
| ("C in S but collapsed on 4 ", {1: 0.05, 2: 0.05, 4: 0.90}), |
| ("D leaks AND collapsed ", {1: 0.05, 2: 0.05, 4: 0.50, |
| 7: 0.20, 8: 0.20}), |
| |
| |
| ("E almost all mass outside S", {1: 0.005, 2: 0.005, 4: 0.005, 7: 0.985}), |
| ] |
|
|
| hdr = (f"{'case':<28}{'mass':>7}{'out':>8}{'leak':>9}{'kl':>9}" |
| f"{'excess':>9}{'spread':>8}{'score':>8}") |
| print(f"candidate set S = {{{','.join(map(str, CAND))}}}, |S| = {len(CAND)}, " |
| f"log|S| = {np.log(len(CAND)):.4f}") |
| print(hdr) |
| print("-" * len(hdr)) |
| for label, pmap in cases: |
| p = np.zeros(9) |
| for d, v in pmap.items(): |
| p[d - 1] = v |
| assert abs(p.sum() - 1.0) < 1e-9, (label, p.sum()) |
| logp = np.log(np.maximum(p, 1e-300))[None, :] |
| st = _cand_value_stats(logp, np.array([BITS])) |
| excess = st["ce"] - st["floor"] |
| mass = st["mass"] |
| leak = np.log(1.0 / max(mass, 1e-300)) |
| kl = st["kl_multi"] |
| |
| assert abs(excess - (leak + kl)) < 1e-8, (label, excess, leak + kl) |
| print(f"{label:<28}{mass:>7.3f}{1 - mass:>8.3f}{leak:>9.3f}{kl:>9.3f}" |
| f"{excess:>9.3f}{st['spread']:>8.3f}{np.exp(-excess):>8.3f}") |
|
|
| print() |
| print("out = val_acc = 1 - mass = P(digit outside S)") |
| print("leak = log(1/mass) -> 0 when no support outside S") |
| print("kl = KL(Unif(S)||p/m) -> 0 when support inside S is uniform") |
| print("excess = leak + kl -> 0 iff BOTH; this is the promotion signal") |
| print("score = exp(-excess) -> 1.0 iff both; compared to SUDOKU_PROMOTE_ACC") |
|
|