File size: 1,555 Bytes
c6be157
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
"""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()