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