"""Train a tiny GNOMON on random data for a few hundred steps.""" import os import sys sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..")) import jax import jax.numpy as jnp from gnomon import init_gnomon, gnomon_forward, train_step from gnomon.params import adam_init from gnomon.neuromod import init_mode_state def make_batch(key, B=4, S=32, V=64, d_signal=32): k1, k2, k3 = jax.random.split(key, 3) input_ids = jax.random.randint(k1, (B, S), 0, V) labels = jax.random.randint(k2, (B, S), 0, V) mode_signal = jax.random.normal(k3, (B, S, d_signal)) return {"input_ids": input_ids, "labels": labels, "mode_signal": mode_signal} def main(): V, D, L, H, IB_K, S = 64, 64, 2, 4, 16, 32 key = jax.random.PRNGKey(0) k1, k2 = jax.random.split(key) params = init_gnomon(k1, vocab_size=V, d_model=D, n_layers=L, n_heads=H, ib_k=IB_K, seq_len=S, d_signal=32) 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) ) for step in range(50): key, sub = jax.random.split(key) batch = make_batch(sub, V=V, d_signal=32) params, opt, loss, aux = step_fn(params, opt, batch, mode, sub) if step % 10 == 0: print(f"step {step:03d} loss={float(loss):.4f} " f"ce={float(aux['ce']):.4f} ib={float(aux['ib']):.4f}") print("done") if __name__ == "__main__": main()