"""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)