Avra98's picture
sync training code: stage-1 instance-epoch sampler, multi-stage run, superposition metrics
6a1771b verified
Raw
History Blame Contribute Delete
2.72 kB
"""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")