"""Decode-path comparison: tokens/sec and memory vs context length. HOARD carries constant state (GDN fast weights + a 256-token attention window) while the Transformer's KV cache grows with context. Measures batch-1 decode after prefills of increasing length. python bench_decode.py --runs runs/fw_hoard_small runs/fw_transformer_small \ --contexts 1024 4096 16384 32768 --tokens 96 """ import argparse, json, os, time import numpy as np import mlx.core as mx from model import HOARD, HoardConfig def load(run_dir): meta = json.load(open(os.path.join(run_dir, "config.json"))) cfg = HoardConfig.from_dict(meta["config"]) model = HOARD(cfg) model.load_weights(os.path.join(run_dir, "model.safetensors")) model.eval() return model, meta["args"]["config"] def main(): ap = argparse.ArgumentParser() ap.add_argument("--runs", nargs="+", required=True) ap.add_argument("--contexts", type=int, nargs="+", default=[1024, 4096, 16384, 32768]) ap.add_argument("--tokens", type=int, default=96) ap.add_argument("--data", default="data/fineweb/val.bin") a = ap.parse_args() toks = np.memmap(a.data, dtype=np.uint16, mode="r") print(f"{'model':<22}{'context':>9}{'prefill s':>11}{'decode tok/s':>14}{'peak GB':>9}") for rd in a.runs: model, name = load(rd) for L in a.contexts: mx.reset_peak_memory() prompt = mx.array(np.asarray(toks[:L]).astype(np.int32))[None] cache = model.new_cache() t0 = time.time() h = model.forward_hidden(prompt, None, cache) logits = model.logits(h[:, -1:]) mx.eval(logits) t_prefill = time.time() - t0 ids = mx.argmax(logits[:, -1], axis=-1)[:, None] t0 = time.time() for _ in range(a.tokens): h = model.forward_hidden(ids, None, cache) logits = model.logits(h) ids = mx.argmax(logits[:, -1], axis=-1)[:, None] mx.eval(ids) dt = time.time() - t0 print(f"{name:<22}{L:>9}{t_prefill:>11.2f}{a.tokens / dt:>14.1f}" f"{mx.get_peak_memory() / 2**30:>9.2f}") del model mx.clear_cache() if __name__ == "__main__": main()