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