""" HOARD — a Mac-first Transformer competitor, in MLX. "A dragon that hoards memory instead of burning compute." Design thesis: Apple Silicon has tens of TFLOPS but hundreds of GB of unified memory shared by CPU and GPU. So: hoard parameters and state (cheap), spend as few FLOPs per token as possible (scarce). Components ---------- GDNMixer Gated DeltaNet token mixer. Fixed-size fast-weight state (fp32), chunkwise-parallel training in pure MLX ops (autodiff for free), recurrent step for constant-memory decode. HoardMLP A very wide, sparsely-activated neuron space (BDH's n >> d idea) realised as product-key routed blocks computed with mx.gather_mm, i.e. the same machinery MLX uses for MoE. Only top-k blocks are touched per token. WindowAttn Small exact attention window (SDPA + band mask) for local precision. Cell [GDNMixer -> HoardMLP -> WindowAttn -> HoardMLP] with pre-norm residuals. A cell is weight-tied and looped n_loops times. HOARD Embedding -> looped cell(s) -> RMSNorm -> LM head. Baselines (for matched-scale comparisons, same file, same trainer): mixer="attn" full causal attention instead of GDN mlp="dense" ordinary ReLU^2 MLP instead of Hoard n_loops=1, n_cells=12 -> a plain Transformer Everything is plain MLX ops. No custom Metal kernels (they need hand-written VJPs). Recurrent state is always float32. """ from __future__ import annotations import math from dataclasses import dataclass, field, asdict from typing import Optional, List, Dict, Any import mlx.core as mx import mlx.nn as nn # --------------------------------------------------------------------------- # # Config # --------------------------------------------------------------------------- # @dataclass class HoardConfig: vocab_size: int = 50304 # GPT-2 BPE padded to a multiple of 64 d_model: int = 512 # depth n_cells: int = 1 # distinct (untied) cells n_loops: int = 4 # each cell is applied this many times loop_schedule: str = "fixed" # "fixed" | "random" (sample loops in [1, n_loops] each step) # token mixer mixer: str = "gdn" # "gdn" | "attn" n_heads: int = 4 head_dim_k: int = 128 head_dim_v: int = 128 conv_kernel: int = 4 chunk_size: int = 64 # window attention branch use_window_attn: bool = True window: int = 256 attn_heads: int = 8 # mlp mlp: str = "hoard" # "hoard" | "dense" hoard_n_sub: int = 32 # blocks = n_sub^2 hoard_block: int = 64 # neurons per block hoard_topk: int = 16 # active blocks per token hoard_router_dim: int = 64 hoard_balance_coef: float = 0.01 hoard_gate: bool = False # silu-gated hoard output (Meta memory-layer style) dense_mult: int = 4 # hidden = dense_mult * d_model for the dense baseline # misc tie_embeddings: bool = False rope_base: float = 10000.0 norm_eps: float = 1e-5 @property def n_blocks(self) -> int: return self.hoard_n_sub ** 2 def to_dict(self) -> Dict[str, Any]: return asdict(self) @classmethod def from_dict(cls, d: Dict[str, Any]) -> "HoardConfig": return cls(**{k: v for k, v in d.items() if k in cls.__dataclass_fields__}) # --------------------------------------------------------------------------- # # Helpers # --------------------------------------------------------------------------- # def l2norm(x: mx.array, eps: float = 1e-6) -> mx.array: return x * mx.rsqrt((x * x).sum(-1, keepdims=True) + eps) def relu2(x: mx.array) -> mx.array: return mx.square(nn.relu(x)) def band_mask(T: int, window: int, offset: int = 0) -> mx.array: """Boolean (T_q, T_k) mask allowing key j for query i iff 0 <= i - j < window. offset: number of cached keys preceding the queries (decode).""" Tk = T + offset i = mx.arange(T)[:, None] + offset j = mx.arange(Tk)[None, :] diff = i - j return (diff >= 0) & (diff < window) def _inv_unit_lower_small(A: mx.array) -> mx.array: """(I + A)^-1 for strictly-lower A by row-wise forward substitution. Backward stable: never forms powers of A.""" b = A.shape[-1] I = mx.eye(b, dtype=A.dtype) rows = [mx.broadcast_to(I[0], A.shape[:-2] + (1, b))] for i in range(1, b): Xprev = mx.concatenate(rows, axis=-2) # (..., i, b) xi = I[i] - A[..., i:i + 1, :i] @ Xprev # (..., 1, b) rows.append(xi) return mx.concatenate(rows, axis=-2) def inv_unit_lower(A: mx.array, block: int = 16) -> mx.array: """(I + A)^-1 for strictly-lower-triangular A of size C. Blocked forward substitution: row substitution inside `block`-sized diagonal blocks, block substitution across. The nilpotent-product identity (I-A)(I+A^2)(I+A^4)... is exact algebra but explicitly forms A^(2^k) intermediates; with correlated keys their magnitude explodes and fp32 rounding of them corrupts the inverse — measured on shakespeare as a ~500x-per-chunk state blowup ending in inf/nan. Substitution never forms powers of A, so the error stays proportional to the true conditioning. """ C = A.shape[-1] if C <= block: return _inv_unit_lower_small(A) assert C % block == 0, f"chunk {C} not divisible by block {block}" nb = C // block Ab = [[A[..., i * block:(i + 1) * block, j * block:(j + 1) * block] for j in range(nb)] for i in range(nb)] X = [[None] * nb for _ in range(nb)] for i in range(nb): X[i][i] = _inv_unit_lower_small(Ab[i][i]) zero = mx.zeros_like(Ab[0][0]) for i in range(1, nb): for j in range(i - 1, -1, -1): s = Ab[i][j] @ X[j][j] for k in range(j + 1, i): s = s + Ab[i][k] @ X[k][j] X[i][j] = -(X[i][i] @ s) rows = [mx.concatenate([X[i][j] if j <= i else zero for j in range(nb)], axis=-1) for i in range(nb)] return mx.concatenate(rows, axis=-2) # --------------------------------------------------------------------------- # # Gated Delta Rule — chunkwise parallel (training / prefill) and recurrent step # --------------------------------------------------------------------------- # def chunk_gated_delta_rule( q: mx.array, k: mx.array, v: mx.array, g: mx.array, beta: mx.array, chunk_size: int = 64, initial_state: Optional[mx.array] = None, ): """ q, k : (B, H, T, dk) -- expected L2-normalised v : (B, H, T, dv) g : (B, H, T) -- log decay per token, <= 0 beta : (B, H, T) -- write strength in (0, 1) Returns o: (B, H, T, dv) and final state S: (B, H, dk, dv) (float32). Recurrence being computed (per head, S is dk x dv): S_t = exp(g_t) * (I - beta_t k_t k_t^T) S_{t-1} + beta_t k_t v_t^T o_t = S_t^T q_t """ B, H, T, dk = q.shape dv = v.shape[-1] C = chunk_size pad = (C - T % C) % C if pad: q = mx.pad(q, [(0, 0), (0, 0), (0, pad), (0, 0)]) k = mx.pad(k, [(0, 0), (0, 0), (0, pad), (0, 0)]) v = mx.pad(v, [(0, 0), (0, 0), (0, pad), (0, 0)]) g = mx.pad(g, [(0, 0), (0, 0), (0, pad)]) beta = mx.pad(beta, [(0, 0), (0, 0), (0, pad)]) Tp = T + pad nC = Tp // C f32 = mx.float32 q = q.astype(f32) * (dk ** -0.5) k = k.astype(f32) v = v.astype(f32) g = g.astype(f32) beta = beta.astype(f32) q = q.reshape(B, H, nC, C, dk) k = k.reshape(B, H, nC, C, dk) v = v.reshape(B, H, nC, C, dv) g = g.reshape(B, H, nC, C) beta = beta.reshape(B, H, nC, C) G = mx.cumsum(g, axis=-1) # (B,H,nC,C) tril = mx.tril(mx.ones((C, C), dtype=mx.bool_)) strict = mx.tril(mx.ones((C, C), dtype=mx.bool_), k=-1) diff = G[..., :, None] - G[..., None, :] # G_i - G_j D = mx.exp(mx.where(tril, diff, -1e30)) # decay mask incl. diagonal D_strict = mx.where(strict, D, 0.0) kb = k * beta[..., None] vb = v * beta[..., None] A = (kb @ k.swapaxes(-1, -2)) * D_strict # strictly lower Tinv = inv_unit_lower(A) # (I + A)^-1 W = Tinv @ (kb * mx.exp(G)[..., None]) # (B,H,nC,C,dk) U = Tinv @ vb # (B,H,nC,C,dv) if initial_state is None: S = mx.zeros((B, H, dk, dv), dtype=f32) else: S = initial_state.astype(f32) outs = [] for i in range(nC): qi, ki, Gi = q[:, :, i], k[:, :, i], G[:, :, i] attn = (qi @ ki.swapaxes(-1, -2)) * D[:, :, i] # (B,H,C,C) v_new = U[:, :, i] - W[:, :, i] @ S # (B,H,C,dv) o_i = (qi * mx.exp(Gi)[..., None]) @ S + attn @ v_new outs.append(o_i) g_last = Gi[..., -1] # (B,H) kdec = ki * mx.exp(g_last[..., None] - Gi)[..., None] S = S * mx.exp(g_last)[..., None, None] + kdec.swapaxes(-1, -2) @ v_new o = mx.stack(outs, axis=2).reshape(B, H, Tp, dv) if pad: o = o[:, :, :T] return o, S def step_gated_delta_rule(q, k, v, g, beta, S): """Single-token recurrent update. q,k:(B,H,dk) v:(B,H,dv) g,beta:(B,H) S:(B,H,dk,dv).""" dk = q.shape[-1] f32 = mx.float32 q = q.astype(f32) * (dk ** -0.5) k = k.astype(f32); v = v.astype(f32) S = S * mx.exp(g.astype(f32))[..., None, None] pred = (k[..., None, :] @ S)[..., 0, :] # (B,H,dv) delta = beta.astype(f32)[..., None] * (v - pred) S = S + k[..., :, None] * delta[..., None, :] o = (q[..., None, :] @ S)[..., 0, :] return o, S # --------------------------------------------------------------------------- # # Modules # --------------------------------------------------------------------------- # class GDNMixer(nn.Module): """Gated DeltaNet token mixer (Qwen3-Next style: short conv, L2-norm q/k, data-dependent decay and beta, gated RMSNorm output).""" def __init__(self, cfg: HoardConfig): super().__init__() d, H, dk, dv = cfg.d_model, cfg.n_heads, cfg.head_dim_k, cfg.head_dim_v self.H, self.dk, self.dv, self.K, self.C = H, dk, dv, cfg.conv_kernel, cfg.chunk_size self.qkv_dim = H * (2 * dk + dv) self.qkv = nn.Linear(d, self.qkv_dim, bias=False) self.z = nn.Linear(d, H * dv, bias=False) self.a = nn.Linear(d, H, bias=False) self.b = nn.Linear(d, H, bias=False) self.conv = nn.Conv1d(self.qkv_dim, self.qkv_dim, self.K, groups=self.qkv_dim, bias=False) # Mamba2 / GDN style init: decay rate A in [1,16], dt in [1e-3, 1e-1] self.A_log = mx.log(mx.random.uniform(1.0, 16.0, (H,))) dt = mx.exp(mx.random.uniform(math.log(1e-3), math.log(1e-1), (H,))) self.dt_bias = dt + mx.log(-mx.expm1(-dt)) # inverse softplus self.norm = nn.RMSNorm(dv, eps=cfg.norm_eps) self.o = nn.Linear(H * dv, d, bias=False) def _log_decay(self, x): return -mx.exp(self.A_log) * nn.softplus(self.a(x) + self.dt_bias) # (B,T,H) <= 0 def _split(self, qkv, B, T): H, dk, dv = self.H, self.dk, self.dv q = qkv[..., : H * dk].reshape(B, T, H, dk) k = qkv[..., H * dk: 2 * H * dk].reshape(B, T, H, dk) v = qkv[..., 2 * H * dk:].reshape(B, T, H, dv) return l2norm(q), l2norm(k), v def __call__(self, x: mx.array, cache: Optional[dict] = None): """x: (B,T,d). If cache is given, runs the recurrent path token-by-token (T may be 1 or more) and updates cache in place. The whole mixer computes in fp32 regardless of parameter dtype: with bf16 activations feeding the delta-rule core, training goes non-finite within a few steps (measured on shakespeare). Cast up here, back at the end; the residual stream stays in the model dtype.""" in_dtype = x.dtype x = x.astype(mx.float32) B, T, _ = x.shape qkv = self.qkv(x) g = self._log_decay(x) beta = mx.sigmoid(self.b(x)) z = self.z(x).reshape(B, T, self.H, self.dv) if cache is None: xpad = mx.pad(qkv, [(0, 0), (self.K - 1, 0), (0, 0)]) qkv = nn.silu(self.conv(xpad)) q, k, v = self._split(qkv, B, T) o, _ = chunk_gated_delta_rule( q.transpose(0, 2, 1, 3), k.transpose(0, 2, 1, 3), v.transpose(0, 2, 1, 3), g.transpose(0, 2, 1), beta.transpose(0, 2, 1), self.C) o = o.transpose(0, 2, 1, 3) # (B,T,H,dv) else: if "conv" not in cache: cache["conv"] = mx.zeros((B, self.K - 1, self.qkv_dim), dtype=qkv.dtype) cache["S"] = mx.zeros((B, self.H, self.dk, self.dv), dtype=mx.float32) xpad = mx.concatenate([cache["conv"], qkv], axis=1) cache["conv"] = xpad[:, -(self.K - 1):] qkv = nn.silu(self.conv(xpad)) q, k, v = self._split(qkv, B, T) if T > 1: # prefill: chunkwise path seeded with the cached state o, S = chunk_gated_delta_rule( q.transpose(0, 2, 1, 3), k.transpose(0, 2, 1, 3), v.transpose(0, 2, 1, 3), g.transpose(0, 2, 1), beta.transpose(0, 2, 1), self.C, initial_state=cache["S"]) o = o.transpose(0, 2, 1, 3) else: # decode: one recurrent step, O(1) memory o, S = step_gated_delta_rule(q[:, 0], k[:, 0], v[:, 0], g[:, 0], beta[:, 0], cache["S"]) o = o[:, None] cache["S"] = S o = self.norm(o.astype(x.dtype)) * nn.silu(z) return self.o(o.reshape(B, T, self.H * self.dv)).astype(in_dtype) class WindowAttn(nn.Module): """Sliding-window (or full causal) softmax attention with RoPE.""" def __init__(self, cfg: HoardConfig, full: bool = False): super().__init__() d, H = cfg.d_model, cfg.attn_heads assert d % H == 0 self.H, self.hd, self.window, self.full = H, d // H, cfg.window, full self.qkv = nn.Linear(d, 3 * d, bias=False) self.o = nn.Linear(d, d, bias=False) self.rope = nn.RoPE(self.hd, base=cfg.rope_base) self.qn = nn.RMSNorm(self.hd, eps=cfg.norm_eps) self.kn = nn.RMSNorm(self.hd, eps=cfg.norm_eps) def __call__(self, x: mx.array, cache: Optional[dict] = None): B, T, d = x.shape q, k, v = mx.split(self.qkv(x), 3, axis=-1) q = self.qn(q.reshape(B, T, self.H, self.hd)).transpose(0, 2, 1, 3) k = self.kn(k.reshape(B, T, self.H, self.hd)).transpose(0, 2, 1, 3) v = v.reshape(B, T, self.H, self.hd).transpose(0, 2, 1, 3) offset = 0 if cache is not None and "k" in cache: offset = cache["pos"] q = self.rope(q, offset=offset) k = self.rope(k, offset=offset) if cache is not None: if "k" in cache: k = mx.concatenate([cache["k"], k], axis=2) v = mx.concatenate([cache["v"], v], axis=2) keep = k.shape[2] if self.full else min(k.shape[2], self.window) cache["k"], cache["v"] = k[:, :, -keep:], v[:, :, -keep:] cache["pos"] = offset + T n_cached = k.shape[2] - T mask = band_mask(T, 10**9 if self.full else self.window, n_cached) else: mask = "causal" if self.full else band_mask(T, self.window) o = mx.fast.scaled_dot_product_attention(q, k, v, scale=self.hd ** -0.5, mask=mask) return self.o(o.transpose(0, 2, 1, 3).reshape(B, T, d)) class DenseMLP(nn.Module): def __init__(self, cfg: HoardConfig): super().__init__() h = cfg.dense_mult * cfg.d_model self.up = nn.Linear(cfg.d_model, h, bias=False) self.down = nn.Linear(h, cfg.d_model, bias=False) def __call__(self, x, cache=None): return self.down(relu2(self.up(x))) class HoardMLP(nn.Module): """Product-key routed sparse neuron space. n_blocks = n_sub^2 blocks of `block` neurons each. A token's router query is split in two halves; each half scores n_sub sub-keys; the top-k of the k*k candidate sums select the blocks. Selected blocks are applied with mx.gather_mm (sorted by block during training for cache-friendly access). Parameters: 2 * n_blocks * block * d_model. Active per token: 2 * topk * block * d_model. """ def __init__(self, cfg: HoardConfig): super().__init__() d, ns, blk, r = cfg.d_model, cfg.hoard_n_sub, cfg.hoard_block, cfg.hoard_router_dim self.d, self.ns, self.blk, self.k = d, ns, blk, cfg.hoard_topk nb = ns * ns self.W_up = mx.random.normal((nb, blk, d)) * (1.0 / math.sqrt(d)) self.W_down = mx.random.normal((nb, blk, d)) * (1.0 / math.sqrt(blk * self.k)) self.router = nn.Linear(d, 2 * r, bias=False) self.rn1 = nn.RMSNorm(r, eps=cfg.norm_eps) self.rn2 = nn.RMSNorm(r, eps=cfg.norm_eps) self.K1 = mx.random.normal((ns, r)) * (1.0 / math.sqrt(r)) self.K2 = mx.random.normal((ns, r)) * (1.0 / math.sqrt(r)) self.balance_coef = cfg.hoard_balance_coef if cfg.hoard_gate: self.gate = nn.Linear(d, d, bias=False) self._bal = 0.0 def route(self, x: mx.array): """x: (N,d) -> block idx (N,k) int32, gates (N,k), balance loss scalar.""" N, k, ns = x.shape[0], self.k, self.ns q1, q2 = mx.split(self.router(x), 2, axis=-1) s1 = self.rn1(q1) @ self.K1.T # (N, ns) s2 = self.rn2(q2) @ self.K2.T kk = min(k, ns) i1 = mx.stop_gradient(mx.argpartition(-s1, kth=kk - 1, axis=-1)[:, :kk]) # (N,kk) i2 = mx.stop_gradient(mx.argpartition(-s2, kth=kk - 1, axis=-1)[:, :kk]) v1 = mx.take_along_axis(s1, i1, axis=-1) v2 = mx.take_along_axis(s2, i2, axis=-1) cand = (v1[:, :, None] + v2[:, None, :]).reshape(N, kk * kk) sel = mx.stop_gradient(mx.argpartition(-cand, kth=k - 1, axis=-1)[:, :k]) # (N,k) into kk*kk scores = mx.take_along_axis(cand, sel, axis=-1) b1 = mx.take_along_axis(i1, sel // kk, axis=-1) b2 = mx.take_along_axis(i2, sel % kk, axis=-1) idx = mx.stop_gradient((b1 * ns + b2).astype(mx.int32)) gates = mx.softmax(scores, axis=-1) # Switch-style balance loss, factorised over the two sub-key halves. bal = mx.array(0.0) if self.balance_coef > 0: for s, b in ((s1, b1), (s2, b2)): P = mx.softmax(s, axis=-1).mean(0) # (ns,) f = mx.zeros((ns,)).at[b.reshape(-1)].add(1.0) / (N * k) bal = bal + ns * (mx.stop_gradient(f) * P).sum() bal = bal * self.balance_coef return idx, gates, bal def __call__(self, x: mx.array, cache=None): B, T, d = x.shape N = B * T xf = x.reshape(N, d) idx, gates, bal = self.route(xf) self._bal = bal k = self.k if N * k > 64: # sort by block so gather_mm touches each block's weights contiguously flat = idx.reshape(-1) order = mx.stop_gradient(mx.argsort(flat)) inv = mx.stop_gradient(mx.argsort(order)) xs = xf[order // k][:, None, :] # (N*k,1,d) sidx = flat[order] h = mx.gather_mm(xs, self.W_up.swapaxes(-1, -2), rhs_indices=sidx, sorted_indices=True) h = relu2(h) # (N*k,1,blk) y = mx.gather_mm(h, self.W_down, rhs_indices=sidx, sorted_indices=True) # (N*k,1,d) y = y[inv].reshape(N, k, d) else: h = mx.gather_mm(xf[:, None, None, :], self.W_up.swapaxes(-1, -2), rhs_indices=idx) h = relu2(h) # (N,k,1,blk) y = mx.gather_mm(h, self.W_down, rhs_indices=idx)[:, :, 0] # (N,k,d) y = (y * gates[..., None]).sum(1) if hasattr(self, "gate"): y = y * nn.silu(self.gate(xf)) return y.reshape(B, T, d) class Cell(nn.Module): """One (tied) cell: mixer -> mlp -> [window attn -> mlp]. Pre-norm residual.""" def __init__(self, cfg: HoardConfig): super().__init__() self.cfg = cfg # mlp="mixed": hoard for mlp1, dense for mlp2 — hoard capacity added to # the hybrid backbone one slot per cell instead of everywhere. mk1 = HoardMLP if cfg.mlp in ("hoard", "mixed") else DenseMLP mk2 = HoardMLP if cfg.mlp == "hoard" else DenseMLP self.n1 = nn.RMSNorm(cfg.d_model, eps=cfg.norm_eps) self.mixer = GDNMixer(cfg) if cfg.mixer == "gdn" else WindowAttn(cfg, full=True) self.n2 = nn.RMSNorm(cfg.d_model, eps=cfg.norm_eps) self.mlp1 = mk1(cfg) if cfg.use_window_attn: self.n3 = nn.RMSNorm(cfg.d_model, eps=cfg.norm_eps) self.attn = WindowAttn(cfg, full=False) self.n4 = nn.RMSNorm(cfg.d_model, eps=cfg.norm_eps) self.mlp2 = mk2(cfg) def __call__(self, h: mx.array, cache: Optional[dict] = None): c = cache if cache is not None else {} h = h + self.mixer(self.n1(h), c.setdefault("mixer", {}) if cache is not None else None) h = h + self.mlp1(self.n2(h)) if self.cfg.use_window_attn: h = h + self.attn(self.n3(h), c.setdefault("attn", {}) if cache is not None else None) h = h + self.mlp2(self.n4(h)) return h def balance_loss(self): b = mx.array(0.0) for m in (self.mlp1, getattr(self, "mlp2", None)): if m is not None and hasattr(m, "_bal"): b = b + m._bal return b class HOARD(nn.Module): def __init__(self, cfg: HoardConfig): super().__init__() self.cfg = cfg self.embed = nn.Embedding(cfg.vocab_size, cfg.d_model) self.cells = [Cell(cfg) for _ in range(cfg.n_cells)] # learned per-loop offset so a tied cell knows which iteration it is in self.loop_embed = mx.zeros((cfg.n_loops, cfg.d_model)) self.nf = nn.RMSNorm(cfg.d_model, eps=cfg.norm_eps) if not cfg.tie_embeddings: self.head = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False) # ---- forward ---------------------------------------------------------- # def forward_hidden(self, ids: mx.array, n_loops: Optional[int] = None, cache: Optional[List[List[dict]]] = None, return_all_loops: bool = False): n_loops = n_loops or self.cfg.n_loops e = self.embed(ids) h = e per_loop = [] # gradient checkpointing: train.py installs _ckpt_cells; only valid cache-free ckpt_fns = getattr(self, "_ckpt_cells", None) if cache is None else None for l in range(n_loops): h = h + e + self.loop_embed[l] # input re-injection + loop id for u, cell in enumerate(self.cells): if ckpt_fns is not None: h = ckpt_fns[u](h) else: c = cache[l][u] if cache is not None else None h = cell(h, c) if return_all_loops: per_loop.append(h) return (h, per_loop) if return_all_loops else h def logits(self, h: mx.array): h = self.nf(h) if self.cfg.tie_embeddings: return self.embed.as_linear(h) return self.head(h) def __call__(self, ids: mx.array, n_loops: Optional[int] = None): return self.logits(self.forward_hidden(ids, n_loops)) def balance_loss(self): b = mx.array(0.0) for c in self.cells: b = b + c.balance_loss() return b # ---- decode ----------------------------------------------------------- # def new_cache(self): return [[{} for _ in self.cells] for _ in range(self.cfg.n_loops)] def generate(self, prompt: mx.array, max_new_tokens: int = 64, temperature: float = 1.0, top_k: int = 0, n_loops: Optional[int] = None): """prompt: (B,T) int. Constant-memory decode: GDN state + window KV per (loop, cell).""" cache = self.new_cache() n_loops = n_loops or self.cfg.n_loops h = self.forward_hidden(prompt, n_loops, cache) out = [prompt] logits = self.logits(h[:, -1:]) for _ in range(max_new_tokens): nxt = self._sample(logits[:, -1], temperature, top_k) out.append(nxt) h = self.forward_hidden(nxt, n_loops, cache) logits = self.logits(h) return mx.concatenate(out, axis=1) @staticmethod def _sample(logits, temperature, top_k): if temperature <= 0: return mx.argmax(logits, axis=-1)[:, None] logits = logits / temperature if top_k > 0: kth = mx.sort(logits, axis=-1)[:, -top_k][:, None] logits = mx.where(logits < kth, -1e9, logits) return mx.random.categorical(logits)[:, None] # ---- accounting ------------------------------------------------------- # def param_counts(self) -> Dict[str, int]: from mlx.utils import tree_flatten total = sum(v.size for _, v in tree_flatten(self.parameters())) hoard = sum(v.size for n, v in tree_flatten(self.parameters()) if "W_up" in n or "W_down" in n) cfg = self.cfg active_hoard = 0 if cfg.mlp in ("hoard", "mixed"): per_cell = 2 if (cfg.use_window_attn and cfg.mlp == "hoard") else 1 n_mlps = cfg.n_cells * per_cell active_hoard = n_mlps * 2 * cfg.hoard_topk * cfg.hoard_block * cfg.d_model active = total - hoard + active_hoard return {"total": total, "hoard": hoard, "active_per_token_per_loop": active, "active_x_loops": (active - self.embed.weight.size) * cfg.n_loops + self.embed.weight.size} def build_model(cfg: HoardConfig) -> HOARD: return HOARD(cfg)