gnomon / examples /train_task.py
zeechimp's picture
Upload 15 files
c6be157 verified
Raw History Blame Contribute Delete
5.15 kB
"""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()