File size: 2,719 Bytes
6a1771b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 | """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")
|