File size: 5,149 Bytes
c6be157
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
"""Train GNOMON on a bigram task.



Task: tokens generated from a sharp row-stochastic transition matrix.

With temperature=0.3, the transition has significant structure — Bayes

CE is well below uniform, so the model has room to show learning.



The coupling KL should fall from ~0.04 to <0.005 during training.

"""
import os
import sys
import time

sys.path.insert(0, os.path.join(os.path.dirname(__file__), ".."))

import jax
import jax.numpy as jnp
import numpy as np

from gnomon import init_gnomon, init_mode_state, make_forward_jit
from gnomon.train import train_step
from gnomon.params import adam_init


def make_transition(rng, V, temperature=0.3):
    """Sharper than the previous version: temperature scales the noise."""
    logits = rng.standard_normal((V, V)) * temperature
    logits += 1.5 * np.eye(V)
    p = np.exp(logits - logits.max(axis=1, keepdims=True))
    p /= p.sum(axis=1, keepdims=True)
    return p.astype(np.float32)


def row_entropy(p):
    """Average per-row entropy of the transition matrix."""
    return float(-(p * np.log(p + 1e-12)).sum(axis=1).mean())


def sample_batch(rng, B, S, V, transition):
    ids = np.zeros((B, S), dtype=np.int32)
    ids[:, 0] = rng.integers(0, V, B)
    for t in range(1, S):
        prev = ids[:, t - 1]
        probs = transition[prev]
        cum = probs.cumsum(axis=1)
        u = rng.random((B, 1))
        ids[:, t] = (u > cum).sum(axis=1)
    return jnp.asarray(ids)


def coupling_kl(params, ids, mode, mode_sig, key, n_heads):
    raw_fn = make_forward_jit(n_heads=n_heads, use_reconstructed_kv=False)
    rec_fn = make_forward_jit(n_heads=n_heads, use_reconstructed_kv=True)
    out_raw = raw_fn(params, ids, mode, mode_sig, key)
    out_rec = rec_fn(params, ids, mode, mode_sig, key)
    jax.block_until_ready(out_raw["logits"])
    p = jax.nn.softmax(out_raw["logits"], axis=-1)
    q = jax.nn.log_softmax(out_rec["logits"], axis=-1)
    return float(jnp.mean(jnp.sum(p * (jnp.log(p + 1e-9) - q), axis=-1)))


def main():
    V, D, L, H, IB_K, S, B = 32, 64, 2, 4, 16, 24, 32
    STEPS = 2000
    LR = 2e-3
    D_SIGNAL = 16

    rng = np.random.default_rng(0)
    transition = make_transition(rng, V, temperature=0.3)
    bayes_ce = row_entropy(transition)
    uniform_ce = float(np.log(V))

    key = jax.random.PRNGKey(0)
    key, sub = jax.random.split(key)
    params = init_gnomon(
        sub, vocab_size=V, d_model=D, n_layers=L, n_heads=H,
        ib_k=IB_K, seq_len=S, d_signal=D_SIGNAL,
    )
    opt = adam_init(params)
    mode = init_mode_state(d_mod=4)

    step_fn = jax.jit(
        lambda p, o, b, m, r: train_step(p, o, b, m, r, n_heads=H, lr=LR)
    )

    print(f"Task:  bigram Markov chain (V={V}, temperature=0.3)")
    print(f"Bayes CE (row entropy): {bayes_ce:.4f}")
    print(f"Uniform CE:             {uniform_ce:.4f}")
    print(f"Model: {L} layers, d={D}, heads={H}, IB k={IB_K}")
    print(f"Train: {STEPS} steps, batch {B}×{S}, Adam lr={LR}")
    print()
    print(f"{'step':>5}  {'loss':>8}  {'ce':>8}  {'ib':>8}  {'coupling KL':>12}")
    print("-" * 52)

    key, sub = jax.random.split(key)
    ids_probe = sample_batch(rng, 4, S, V, transition)
    mode_sig_probe = jnp.zeros((4, S, D_SIGNAL))
    kl_0 = coupling_kl(params, ids_probe, mode, mode_sig_probe, sub, H)
    print(f"{'init':>5}  {'—':>8}  {'—':>8}  {'—':>8}  {kl_0:>12.4f}")

    t0 = time.perf_counter()
    final_ce = None
    final_kl = kl_0
    for step in range(STEPS):
        key, sub = jax.random.split(key)
        ids = sample_batch(rng, B, S, V, transition)
        mode_sig = jnp.zeros((B, S, D_SIGNAL))
        batch = {"input_ids": ids, "labels": ids, "mode_signal": mode_sig}
        params, opt, loss, aux = step_fn(params, opt, batch, mode, sub)
        final_ce = float(aux["ce"])

        if step % 200 == 0 or step == STEPS - 1:
            key, sub = jax.random.split(key)
            final_kl = coupling_kl(params, ids_probe, mode, mode_sig_probe, sub, H)
            print(f"{step:>5}  {float(loss):>8.4f}  {float(aux['ce']):>8.4f}  "
                  f"{float(aux['ib']):>8.4f}  {final_kl:>12.4f}")

    total = time.perf_counter() - t0
    print("-" * 52)
    print(f"{STEPS} steps in {total:.1f}s ({STEPS / total:.1f} steps/s)")
    print()

    print(f"Uniform CE:             {uniform_ce:.4f}")
    print(f"Bayes CE:               {bayes_ce:.4f}")
    print(f"Final model CE:         {final_ce:.4f}")
    print(f"Signal recovered:       "
          f"{100 * (uniform_ce - final_ce) / (uniform_ce - bayes_ce):.1f}% "
          f"of the uniform→Bayes gap")
    print(f"Coupling KL:            {kl_0:.4f} → {final_kl:.4f} "
          f"({kl_0 / max(final_kl, 1e-6):.0f}x tighter)")
    print()
    if final_ce <= bayes_ce + 0.3:
        print("  → Model captured the transition structure.")
    elif final_ce <= uniform_ce - 0.3:
        print("  → Model learned partial structure.")
    else:
        print("  → No learning. Check hyperparameters.")


if __name__ == "__main__":
    main()