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