"""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] # digits 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}), # A softmax never emits an exact zero, so keep a little mass on S here; # with p(d) == 0 for every candidate the CE is +inf by definition. ("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"] # the identity, recomputed independently of the code 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")