Download examples/train_task.py from zeechimp/gnomon: direct link, hf CLI and curl.
- Browser
- Download file 5.15 kB
-
https://huggingface.co/zeechimp/gnomon/resolve/main/examples/train_task.py
- Command line
-
hf download hf://zeechimp/gnomon/examples/train_task.py
-
curl -L -o train_task.py https://huggingface.co/zeechimp/gnomon/resolve/main/examples/train_task.py
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() |