Gala-598M-MLX / model.py
junafinity's picture
Gala-598M: weights, code, logs, report
76bbe95 verified
Raw History Blame Contribute Delete
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
# --------------------------------------------------------------------------- #
@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)