Text Generation
MLX
English
apple-silicon
pretrained-from-scratch
gated-deltanet
linear-attention
product-key-memory
long-context
Instructions to use junafinity/Gala-598M-MLX with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use junafinity/Gala-598M-MLX with MLX:
# Make sure mlx-lm is installed # pip install --upgrade mlx-lm # if on a CUDA device, also pip install mlx[cuda] # Generate text with mlx-lm from mlx_lm import load, generate model, tokenizer = load("junafinity/Gala-598M-MLX") prompt = "Once upon a time in" text = generate(model, tokenizer, prompt=prompt, verbose=True) - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
- MLX LM
How to use junafinity/Gala-598M-MLX with MLX LM:
Generate or start a chat session
# Install MLX LM uv tool install mlx-lm # Generate some text mlx_lm.generate --model "junafinity/Gala-598M-MLX" --prompt "Once upon a time"
- Atomic Chat
Download model.py from junafinity/Gala-598M-MLX: direct link, hf CLI and curl.
- Browser
- Download file 26.2 kB
-
https://huggingface.co/junafinity/Gala-598M-MLX/resolve/main/model.py
- Command line
-
hf download hf://junafinity/Gala-598M-MLX/model.py
-
curl -L -o model.py https://huggingface.co/junafinity/Gala-598M-MLX/resolve/main/model.py
26.2 kB
| """ | |
| 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 | |
| # --------------------------------------------------------------------------- # | |
| 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 | |
| def n_blocks(self) -> int: | |
| return self.hoard_n_sub ** 2 | |
| def to_dict(self) -> Dict[str, Any]: | |
| return asdict(self) | |
| 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) | |
| 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) | |