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