Gala-598M-MLX / tests /test_model.py
junafinity's picture
Gala-598M: weights, code, logs, report
76bbe95 verified
Raw History Blame Contribute Delete
3.97 kB
"""Run: python -m pytest tests/ -q (or just: python tests/test_model.py)"""
import sys, os
sys.path.insert(0, os.path.join(os.path.dirname(__file__), ".."))
import mlx.core as mx
import mlx.nn as nn
from mlx.utils import tree_flatten
from model import (HoardConfig, HOARD, HoardMLP, chunk_gated_delta_rule,
step_gated_delta_rule, inv_unit_lower, l2norm)
def tiny_cfg(**kw):
base = dict(vocab_size=256, d_model=64, n_cells=1, n_loops=2, n_heads=2, head_dim_k=16,
head_dim_v=16, chunk_size=8, window=8, attn_heads=4, hoard_n_sub=4,
hoard_block=8, hoard_topk=3, hoard_router_dim=8)
base.update(kw)
return HoardConfig(**base)
def test_inv_unit_lower():
mx.random.seed(0)
C = 16
A = mx.tril(mx.random.normal((3, C, C)), k=-1)
I = mx.eye(C)
prod = (I + A) @ inv_unit_lower(A)
assert mx.abs(prod - I).max().item() < 1e-4
def test_chunk_vs_recurrent():
mx.random.seed(1)
B, H, T, dk, dv = 2, 3, 37, 16, 12 # T not a multiple of chunk on purpose
q = l2norm(mx.random.normal((B, H, T, dk)))
k = l2norm(mx.random.normal((B, H, T, dk)))
v = mx.random.normal((B, H, T, dv))
g = -mx.exp(mx.random.normal((B, H, T))) * 0.1
beta = mx.sigmoid(mx.random.normal((B, H, T)))
o_chunk, S_chunk = chunk_gated_delta_rule(q, k, v, g, beta, chunk_size=8)
S = mx.zeros((B, H, dk, dv))
outs = []
for t in range(T):
o_t, S = step_gated_delta_rule(q[:, :, t], k[:, :, t], v[:, :, t], g[:, :, t], beta[:, :, t], S)
outs.append(o_t)
o_rec = mx.stack(outs, axis=2)
assert mx.abs(o_chunk - o_rec).max().item() < 1e-4, mx.abs(o_chunk - o_rec).max().item()
assert mx.abs(S_chunk - S).max().item() < 1e-4
def test_hoard_sorted_vs_unsorted():
mx.random.seed(2)
cfg = tiny_cfg()
m = HoardMLP(cfg)
x = mx.random.normal((2, 40, cfg.d_model)) # N*k > 64 -> sorted path
y_sorted = m(x)
m.k_backup = m.k
# force unsorted path by evaluating in small pieces
ys = mx.concatenate([m(x[:, i:i + 4]) for i in range(0, 40, 4)], axis=1)
assert mx.abs(y_sorted - ys).max().item() < 1e-4
def test_forward_backward_finite():
mx.random.seed(3)
for mixer, mlp in (("gdn", "hoard"), ("attn", "dense"), ("gdn", "dense")):
cfg = tiny_cfg(mixer=mixer, mlp=mlp)
model = HOARD(cfg)
ids = mx.random.randint(0, cfg.vocab_size, (2, 33))
def loss_fn(model, ids):
logits = model(ids[:, :-1])
ce = nn.losses.cross_entropy(logits, ids[:, 1:], reduction="mean")
return ce + model.balance_loss()
loss, grads = nn.value_and_grad(model, loss_fn)(model, ids)
mx.eval(loss, grads)
assert mx.isfinite(loss).item()
for n, g in tree_flatten(grads):
assert mx.isfinite(g).all().item(), n
assert loss.item() < 7.0 # ~ln(256)=5.5 + slack
def test_decode_matches_prefill():
"""Token-by-token decode with caches must equal one-shot logits."""
mx.random.seed(4)
cfg = tiny_cfg(window=6)
model = HOARD(cfg)
ids = mx.random.randint(0, cfg.vocab_size, (1, 21))
full = model(ids)
cache = model.new_cache()
outs = []
# prefill first 5, then decode one at a time
h = model.forward_hidden(ids[:, :5], cache=cache)
outs.append(model.logits(h))
for t in range(5, 21):
h = model.forward_hidden(ids[:, t:t + 1], cache=cache)
outs.append(model.logits(h))
inc = mx.concatenate(outs, axis=1)
err = mx.abs(full - inc).max().item()
assert err < 1e-3, err
def test_generate():
mx.random.seed(5)
cfg = tiny_cfg()
model = HOARD(cfg)
out = model.generate(mx.random.randint(0, cfg.vocab_size, (2, 7)), max_new_tokens=9, temperature=0.8, top_k=20)
assert out.shape == (2, 16)
if __name__ == "__main__":
for name, fn in list(globals().items()):
if name.startswith("test_"):
fn(); print("PASS", name)