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")