txgsync/Maple-Preview-oQ4e
0243
1# Copyright © 2026 DeepGrove AI.2 3from dataclasses import dataclass4from functools import partial5from typing import Any, List, Optional6 7import mlx.core as mx8import mlx.nn as nn9 10# Absolute imports so this file also works standalone when shipped inside a11# checkpoint and loaded via the config's `model_file` (trust_remote_code).12from mlx_lm.models.activations import swiglu13from mlx_lm.models.base import (14 BaseModelArgs,15 create_attention_mask,16 scaled_dot_product_attention,17)18from mlx_lm.models.cache import KVCache, RotatingKVCache19from mlx_lm.models.rope_utils import initialize_rope20from mlx_lm.models.switch_layers import SwitchLinear21 22# SwiGLU clamp for the MoE experts only (the dense MapleMLP is unclamped);23# part of the trained forward pass, not an optional guard.24MLP_CLAMP = 7.025 26 27@partial(mx.compile, shapeless=True)28def clamped_swiglu(gate, x):29 # Python floats, not 0-d arrays, so bf16 activations stay bf16.30 return nn.silu(mx.minimum(gate, MLP_CLAMP)) * mx.clip(x, -MLP_CLAMP, MLP_CLAMP)31 32 33def _matches(fast, reference, tol=2e-2):34 """One-time self-check for a hand-written Metal kernel.35 36 Every fast path below has a portable equivalent, and each is used only37 after its outputs have been compared against that equivalent once, on the38 live weights. This file ships inside checkpoints and runs on whatever mlx39 and GPU the user has, so a kernel that fails to compile, silently mismatches40 the config it was templated for, or drifts from a future mlx must degrade to41 the portable path rather than corrupt the token stream.42 43 Both callables return a tuple of arrays. The kernels stay in bounds for any44 config (loop counts are integer-divided from the templated dims), so a45 config they cannot handle shows up here as wrong values, not as a fault.46 """47 try:48 got, want = fast(), reference()49 mx.eval(got, want)50 except Exception:51 return False52 return len(got) == len(want) and all(53 g.shape == w.shape54 and bool(55 mx.allclose(g.astype(mx.float32), w.astype(mx.float32), rtol=tol, atol=tol)56 )57 for g, w in zip(got, want)58 )59 60 61class MapleRMSNorm(nn.Module):62 """RMSNorm with the weight multiply in float32.63 64 The reference rounds only the finished product; mx.fast.rms_norm rounds65 the normalized activation first (~1% per element). Float32 inputs to the66 same kernel reproduce the reference bit-for-bit.67 """68 69 def __init__(self, dims: int, eps: float = 1e-6):70 super().__init__()71 self.weight = mx.ones((dims,))72 self.eps = eps73 74 def __call__(self, x: mx.array) -> mx.array:75 return mx.fast.rms_norm(76 x.astype(mx.float32), self.weight.astype(mx.float32), self.eps77 ).astype(x.dtype)78 79 80def _make_add_rms_norm_kernel(eps):81 """Residual add + RMSNorm in ONE dispatch for single-token decode.82 83 Emits both h = x + r (the residual stream, rounded once like a bf16 add)84 and hn = rmsnorm(h) with the weight multiply in fp32 (reference85 semantics, identical to MapleRMSNorm). Folding the add into the norm and86 skipping the astype round-trips replaces ~4 dispatches with 1, and the87 decode step is bounded by its serial dispatch chain, not by this math.88 """89 source = """90 uint tid = thread_position_in_threadgroup.x;91 constexpr uint N = DIM;92 constexpr uint PT = N / 256u;93 float hb[PT];94 float ss = 0.0f;95 for (uint i = 0; i < PT; ++i) {96 uint j = tid * PT + i;97 float v = (float)x[j] + (float)r[j];98 T_ vb = (T_)v; // one rounding, same as a bf16 add99 h_out[j] = vb;100 hb[i] = (float)vb; // norm sees the rounded stream101 ss += hb[i] * hb[i];102 }103 ss = simd_sum(ss);104 threadgroup float sums[8];105 uint sg = tid / 32u;106 uint lane = tid % 32u;107 if (lane == 0u) sums[sg] = ss;108 threadgroup_barrier(mem_flags::mem_threadgroup);109 float tot = 0.0f;110 for (uint i = 0; i < 8u; ++i) tot += sums[i];111 float scale = metal::rsqrt(tot / (float)N + EPS_);112 for (uint i = 0; i < PT; ++i) {113 uint j = tid * PT + i;114 hn_out[j] = (T_)(hb[i] * scale * (float)w[j]);115 }116 """.replace("EPS_", f"{eps:.10e}f")117 tag = f"{eps:.3e}".replace(".", "_").replace("-", "m").replace("+", "p")118 return mx.fast.metal_kernel(119 name=f"maple_add_rms_norm_{tag}",120 input_names=["x", "r", "w"],121 output_names=["h_out", "hn_out"],122 source=source,123 )124 125 126_add_rms_kernels = {}127 128 129def _add_rms_norm(h, r, w, eps):130 kernel = _add_rms_kernels.get(eps)131 if kernel is None:132 kernel = _add_rms_kernels[eps] = _make_add_rms_norm_kernel(eps)133 return kernel(134 inputs=[h.reshape(-1), r.reshape(-1), w],135 template=[("T_", h.dtype), ("DIM", h.shape[-1])],136 grid=(256, 1, 1),137 threadgroup=(256, 1, 1),138 output_shapes=[h.shape, h.shape],139 output_dtypes=[h.dtype, h.dtype],140 )141 142 143def _add_rms_norm_ok(dim, dtype, w, eps):144 x = mx.random.normal((1, 1, dim), key=mx.random.key(0)).astype(dtype)145 r = mx.random.normal((1, 1, dim), key=mx.random.key(1)).astype(dtype)146 return _matches(147 lambda: _add_rms_norm(x, r, w, eps),148 lambda: (149 x + r,150 mx.fast.rms_norm(151 (x + r).astype(mx.float32), w.astype(mx.float32), eps152 ).astype(dtype),153 ),154 )155 156 157# Inlined rather than imported from switch_layers: those helpers are private158# (underscore-prefixed), and this file must keep loading against whatever159# mlx-lm a user has installed when it ships inside a checkpoint.160def _gather_sort(x, indices):161 *_, M = indices.shape162 indices = indices.flatten()163 order = mx.argsort(indices)164 inv_order = mx.argsort(order)165 return x.flatten(0, -3)[order // M], indices[order], inv_order166 167 168def _scatter_unsort(x, inv_order, shape=None):169 x = x[inv_order]170 if shape is not None:171 x = mx.unflatten(x, 0, shape)172 return x173 174 175@dataclass176class ModelArgs(BaseModelArgs):177 model_type: str = "maple"178 hidden_size: int = 2048179 intermediate_size: int = 5120180 moe_intermediate_size: int = 512181 num_hidden_layers: int = 24182 num_attention_heads: int = 16183 num_key_value_heads: int = 4184 head_dim: int = 128185 num_experts: int = 256186 num_experts_per_tok: int = 8187 first_k_dense_replace: int = 0188 rms_norm_eps: float = 1e-6189 rope_theta: float = 10000.0190 rope_scaling: Optional[dict] = None191 partial_rotary_factor: float = 0.5192 max_position_embeddings: int = 140000193 vocab_size: int = 151936194 sliding_window: int = 512195 layer_types: Optional[List[str]] = None196 use_qk_norm: bool = True197 use_bias: bool = False198 tie_word_embeddings: bool = False199 # FlashHead metadata written by `mlx_lm.ternary --flash-head`. The exact200 # lm_head is the default; opt in to the approximate fast head with201 # mlx_lm.load(..., model_config={"use_flash_head": True}).202 flash_head: Optional[dict] = None203 use_flash_head: bool = False204 # Populated from the checkpoint's config; sanitize() reads group_size from205 # it to expand row-scale (`row_alpha`) ternary tensors.206 quantization: Optional[dict] = None207 208 def __post_init__(self):209 # Single source of truth for per-layer attention types: attention210 # (RoPE/NoPE), masks, and caches all read this resolved list.211 if not self.layer_types:212 self.layer_types = ["full_attention"] * self.num_hidden_layers213 214 215def _make_qk_norm_rope_kernel():216 """Fused per-head RMSNorm + partial RoPE for single-token decode.217 218 One dispatch replaces q_norm, k_norm and two rope calls. One simdgroup per219 head: normalize head_dim values, scale by the head's norm weight, and220 rotate the first ROPE_DIM dims (non-traditional pairing i, i+R/2) at the221 given position. NoPE layers pass ROPE_DIM=0.222 """223 source = """224 uint head = thread_position_in_grid.y;225 uint lane = thread_position_in_grid.x;226 227 constexpr int per_lane = HEAD_DIM / 32;228 const device T_* xh = x + head * HEAD_DIM;229 const device T_* wh = w + head * HEAD_DIM;230 device T_* oh = out + head * HEAD_DIM;231 232 float ss = 0.0f;233 for (int i = 0; i < per_lane; ++i) {234 float v = (float)xh[lane * per_lane + i];235 ss += v * v;236 }237 ss = simd_sum(ss);238 float pos = pos_eps[0];239 float eps = pos_eps[1];240 float scale = metal::rsqrt(ss / HEAD_DIM + eps);241 242 for (int i = 0; i < per_lane; ++i) {243 int j = lane * per_lane + i;244 float v = (float)xh[j] * scale * (float)wh[j];245 if (ROPE_DIM > 0 && j < ROPE_DIM) {246 constexpr int rhalf = ROPE_DIM > 0 ? ROPE_DIM / 2 : 1;247 int p = j < rhalf ? j : j - rhalf;248 float theta = pos * inv_freq[p];249 float c = metal::cos(theta);250 float s = metal::sin(theta);251 int j2 = j < rhalf ? j + rhalf : j - rhalf;252 float u = (float)xh[j2] * scale * (float)wh[j2];253 v = j < rhalf ? (v * c - u * s) : (v * c + u * s);254 }255 oh[j] = (T_)v;256 }257 """258 return mx.fast.metal_kernel(259 name="maple_qk_norm_rope",260 input_names=["x", "w", "inv_freq", "pos_eps"],261 output_names=["out"],262 source=source,263 )264 265 266_qk_norm_rope_kernel = _make_qk_norm_rope_kernel()267 268 269class MapleAttention(nn.Module):270 def __init__(self, args: ModelArgs, layer_idx: int):271 super().__init__()272 self.num_attention_heads = args.num_attention_heads273 self.num_key_value_heads = args.num_key_value_heads274 self.head_dim = args.head_dim or args.hidden_size // args.num_attention_heads275 self.scale = self.head_dim**-0.5276 self.use_qk_norm = args.use_qk_norm277 278 # q/k/v are stored fused (one matmul per step); sanitize() concatenates279 # the checkpoint's split projections.280 self.qkv_proj = nn.Linear(281 args.hidden_size,282 (args.num_attention_heads + 2 * args.num_key_value_heads) * self.head_dim,283 bias=args.use_bias,284 )285 self.o_proj = nn.Linear(286 args.num_attention_heads * self.head_dim,287 args.hidden_size,288 bias=args.use_bias,289 )290 291 if args.use_qk_norm:292 self.q_norm = MapleRMSNorm(self.head_dim, eps=args.rms_norm_eps)293 self.k_norm = MapleRMSNorm(self.head_dim, eps=args.rms_norm_eps)294 self._eps = args.rms_norm_eps295 self._rope_base = args.rope_theta296 self._qk_w = None297 self._inv_freq = None298 self._fused_qk = None # None = unprobed, then True/False299 300 # Maple applies RoPE only on sliding-window layers; full-attention301 # layers use no positional encoding (NoPE).302 self.use_rope = args.layer_types[layer_idx] == "sliding_attention"303 if self.use_rope:304 rope_dim = int(self.head_dim * args.partial_rotary_factor)305 self.rope = initialize_rope(306 rope_dim,307 args.rope_theta,308 traditional=False,309 scaling_config=args.rope_scaling,310 max_position_embeddings=args.max_position_embeddings,311 )312 313 def _qk_fused(self, qk, offset):314 """Both norms and both rope applications in one dispatch."""315 if self._qk_w is None:316 n_q = self.num_attention_heads317 n_kv = self.num_key_value_heads318 self._qk_w = mx.contiguous(319 mx.concatenate(320 [321 mx.broadcast_to(self.q_norm.weight[None], (n_q, self.head_dim)),322 mx.broadcast_to(323 self.k_norm.weight[None], (n_kv, self.head_dim)324 ),325 ]326 )327 )328 if self.use_rope:329 half = self.rope.dims // 2330 self._inv_freq = self._rope_base ** (331 -mx.arange(half, dtype=mx.float32) / half332 )333 else:334 self._inv_freq = mx.ones((1,), dtype=mx.float32)335 mx.eval(self._qk_w, self._inv_freq)336 337 # cache.offset is a Python int for a plain cache but an mx.array for338 # the batched caches; coerce so the pos/eps pair is always uniform.339 pos_eps = mx.array([float(offset), self._eps], dtype=mx.float32)340 return _qk_norm_rope_kernel(341 inputs=[qk, self._qk_w, self._inv_freq, pos_eps],342 template=[343 ("T_", qk.dtype),344 ("HEAD_DIM", self.head_dim),345 ("ROPE_DIM", self.rope.dims if self.use_rope else 0),346 ],347 grid=(32, qk.shape[0], 1),348 threadgroup=(32, 1, 1),349 output_shapes=[qk.shape],350 output_dtypes=[qk.dtype],351 )[0]352 353 def _qk_reference(self, qk, offset):354 """The same result from stock ops: fallback, and the yardstick the355 fused kernel is checked against."""356 n_q = self.num_attention_heads357 q = self.q_norm(qk[None, :n_q, None, :])358 k = self.k_norm(qk[None, n_q:, None, :])359 if self.use_rope:360 q = self.rope(q, offset=offset)361 k = self.rope(k, offset=offset)362 return mx.concatenate([q, k], axis=1).reshape(qk.shape)363 364 def __call__(365 self,366 x: mx.array,367 mask: Optional[mx.array] = None,368 cache: Optional[Any] = None,369 ) -> mx.array:370 B, L, _ = x.shape371 372 qkv = self.qkv_proj(x)373 374 if B == 1 and L == 1 and self.use_qk_norm:375 n_q = self.num_attention_heads376 n_kv = self.num_key_value_heads377 qk_size = (n_q + n_kv) * self.head_dim378 qk = qkv.reshape(-1)[:qk_size].reshape(n_q + n_kv, self.head_dim)379 if self._fused_qk is None:380 # A nonzero position, so a broken rotation cannot pass.381 self._fused_qk = _matches(382 lambda: (self._qk_fused(qk, 7),),383 lambda: (self._qk_reference(qk, 7),),384 )385 offset = cache.offset if cache is not None else 0386 out = (self._qk_fused if self._fused_qk else self._qk_reference)(qk, offset)387 queries = out[:n_q].reshape(1, n_q, 1, self.head_dim)388 keys = out[n_q:].reshape(1, n_kv, 1, self.head_dim)389 values = qkv.reshape(-1)[qk_size:].reshape(1, n_kv, 1, self.head_dim)390 else:391 q_size = self.num_attention_heads * self.head_dim392 kv_size = self.num_key_value_heads * self.head_dim393 q, k, v = mx.split(qkv, [q_size, q_size + kv_size], axis=-1)394 395 queries = q.reshape(B, L, self.num_attention_heads, self.head_dim)396 keys = k.reshape(B, L, self.num_key_value_heads, self.head_dim)397 values = v.reshape(B, L, self.num_key_value_heads, self.head_dim)398 399 if self.use_qk_norm:400 queries = self.q_norm(queries)401 keys = self.k_norm(keys)402 403 queries = queries.transpose(0, 2, 1, 3)404 keys = keys.transpose(0, 2, 1, 3)405 values = values.transpose(0, 2, 1, 3)406 407 if self.use_rope:408 offset = cache.offset if cache is not None else 0409 queries = self.rope(queries, offset=offset)410 keys = self.rope(keys, offset=offset)411 412 if cache is not None:413 keys, values = cache.update_and_fetch(keys, values)414 415 output = scaled_dot_product_attention(416 queries, keys, values, cache=cache, scale=self.scale, mask=mask417 )418 419 output = output.transpose(0, 2, 1, 3).reshape(B, L, -1)420 return self.o_proj(output)421 422 423class MapleMLP(nn.Module):424 def __init__(self, args: ModelArgs, intermediate_size: Optional[int] = None):425 super().__init__()426 intermediate_size = intermediate_size or args.intermediate_size427 self.gate_proj = nn.Linear(428 args.hidden_size, intermediate_size, bias=args.use_bias429 )430 self.up_proj = nn.Linear(431 args.hidden_size, intermediate_size, bias=args.use_bias432 )433 self.down_proj = nn.Linear(434 intermediate_size, args.hidden_size, bias=args.use_bias435 )436 437 def __call__(self, x) -> mx.array:438 # Dense / shared-expert MLP: no clamp; only the MoE experts clamp.439 # Unused at first_k_dense_replace=0 with no shared experts, but keep440 # it faithful.441 return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x)))442 443 444@mx.compile445def group_expert_select(gates, top_k):446 # Maple routes with a plain softmax over all experts followed by top-k447 # selection and renormalization, computed in float32.448 scores = mx.softmax(gates.astype(mx.float32), axis=-1)449 inds = mx.argpartition(scores, kth=-top_k, axis=-1)[..., -top_k:]450 scores = mx.take_along_axis(scores, inds, axis=-1)451 scores = scores / (scores.sum(axis=-1, keepdims=True) + 1e-20)452 return inds, scores453 454 455def _make_fused_router_kernel():456 """Router gemv + softmax + top-8 + renormalize in ONE dispatch (+18%).457 458 Replaces ~6 kernels per layer. NE/32 threadgroups each compute 32 logits,459 keep them in float32 (`router_dtype: fp32`), and publish through an460 atomic-float scratch (plain device stores are not reliably visible across461 threadgroups on Apple GPUs); the last threadgroup to arrive does the462 softmax + top-8 + renorm.463 464 `ctr_in` is a persistent arrival counter, not an input: every dispatch465 must see it at zero, so the electing threadgroup resets it on its way out466 and each MapleGate keeps its own. Election on a stale counter would read467 unwritten scratch, so nothing else may share the buffer.468 """469 source = """470 constexpr uint NE = NEXP;471 constexpr uint D = DIM;472 constexpr uint NTG = NE / 32u;473 constexpr uint TM = 4u;474 constexpr uint TN = 4u;475 constexpr uint BLOCKN = 32u * TN;476 constexpr uint NITER = D / BLOCKN;477 478 uint tid = thread_position_in_threadgroup.x;479 uint tgid = threadgroup_position_in_grid.x;480 uint n_threads = 256u;481 uint sg_id = tid / 32u;482 uint lane = tid % 32u;483 uint n_sg = n_threads / 32u;484 485 uint row0 = tgid * (n_sg * TM) + sg_id * TM;486 float result[TM] = {0.0f, 0.0f, 0.0f, 0.0f};487 uint bn = lane * TN;488 for (uint i = 0u; i < NITER; ++i) {489 float v[TN];490 for (uint tn = 0u; tn < TN; ++tn) v[tn] = float(x[bn + tn]);491 for (uint tm = 0u; tm < TM; ++tm) {492 const device T_* wrow = w + (ulong)(row0 + tm) * D;493 T_ inter[TN];494 for (uint tn = 0u; tn < TN; ++tn) inter[tn] = wrow[bn + tn];495 for (uint tn = 0u; tn < TN; ++tn) result[tm] += inter[tn] * v[tn];496 }497 bn += BLOCKN;498 }499 for (uint tm = 0u; tm < TM; ++tm) {500 for (ushort sn = 16; sn >= 1; sn >>= 1) {501 result[tm] += simd_shuffle_down(result[tm], sn);502 }503 }504 device atomic_float* ls = (device atomic_float*)logits_scratch;505 if (lane == 0u) {506 for (uint tm = 0u; tm < TM; ++tm) {507 atomic_store_explicit(&ls[row0 + tm], result[tm],508 memory_order_relaxed);509 }510 }511 512 threadgroup_barrier(mem_flags::mem_device);513 threadgroup uint last_flag;514 if (tid == 0u) {515 device atomic_uint* ctr = (device atomic_uint*)ctr_in;516 uint prev = atomic_fetch_add_explicit(ctr, 1u, memory_order_relaxed);517 uint last = (prev == NTG - 1u) ? 1u : 0u;518 if (last == 1u) atomic_store_explicit(ctr, 0u, memory_order_relaxed);519 last_flag = last;520 }521 threadgroup_barrier(mem_flags::mem_threadgroup);522 if (last_flag == 0u) return;523 threadgroup_barrier(mem_flags::mem_device);524 525 float my_max = -1e30f;526 for (uint e = tid; e < NE; e += n_threads) {527 float v = atomic_load_explicit(&ls[e], memory_order_relaxed);528 if (v > my_max) my_max = v;529 }530 for (int off = 16; off > 0; off >>= 1) {531 float other = simd_shuffle_down(my_max, off);532 if (other > my_max) my_max = other;533 }534 threadgroup float sg_red[16];535 if (lane == 0u) sg_red[sg_id] = my_max;536 threadgroup_barrier(mem_flags::mem_threadgroup);537 if (tid == 0u) {538 float m = sg_red[0];539 for (uint s = 1u; s < n_sg; s++) if (sg_red[s] > m) m = sg_red[s];540 sg_red[0] = m;541 }542 threadgroup_barrier(mem_flags::mem_threadgroup);543 float lmax = sg_red[0];544 545 threadgroup float scores[NE];546 float my_sum = 0.0f;547 for (uint e = tid; e < NE; e += n_threads) {548 float lv = atomic_load_explicit(&ls[e], memory_order_relaxed);549 float v = metal::exp(lv - lmax);550 scores[e] = v;551 my_sum += v;552 }553 for (int off = 16; off > 0; off >>= 1) {554 my_sum += simd_shuffle_down(my_sum, off);555 }556 threadgroup_barrier(mem_flags::mem_threadgroup);557 if (lane == 0u) sg_red[sg_id] = my_sum;558 threadgroup_barrier(mem_flags::mem_threadgroup);559 if (tid == 0u) {560 float ssum = sg_red[0];561 for (uint i = 1u; i < n_sg; i++) ssum += sg_red[i];562 sg_red[0] = ssum;563 }564 threadgroup_barrier(mem_flags::mem_threadgroup);565 float inv_total = 1.0f / (sg_red[0] + 1e-20f);566 for (uint e = tid; e < NE; e += n_threads) {567 scores[e] = scores[e] * inv_total;568 }569 threadgroup_barrier(mem_flags::mem_threadgroup);570 571 threadgroup int topk_idx[8];572 threadgroup float topk_val[8];573 threadgroup uint8_t used[NE];574 for (uint e = tid; e < NE; e += n_threads) used[e] = 0;575 threadgroup_barrier(mem_flags::mem_threadgroup);576 577 for (int k = 0; k < 8; k++) {578 float my_best = -1e30f;579 int my_idx = 0;580 for (int e = int(tid); e < int(NE); e += int(n_threads)) {581 if (!used[e] && scores[e] > my_best) {582 my_best = scores[e];583 my_idx = e;584 }585 }586 for (int off = 16; off > 0; off >>= 1) {587 float other_v = simd_shuffle_down(my_best, off);588 int other_i = simd_shuffle_down(my_idx, off);589 if (other_v > my_best) { my_best = other_v; my_idx = other_i; }590 }591 threadgroup float sg_vals[16];592 threadgroup int sg_idxs[16];593 if (lane == 0u) { sg_vals[sg_id] = my_best; sg_idxs[sg_id] = my_idx; }594 threadgroup_barrier(mem_flags::mem_threadgroup);595 if (tid == 0u) {596 float bv = sg_vals[0]; int bi = sg_idxs[0];597 for (uint s = 1u; s < n_sg; s++) {598 if (sg_vals[s] > bv) { bv = sg_vals[s]; bi = sg_idxs[s]; }599 }600 topk_val[k] = bv; topk_idx[k] = bi;601 used[bi] = 1;602 }603 threadgroup_barrier(mem_flags::mem_threadgroup);604 }605 606 if (tid < 8u) {607 float sel_sum = 0.0f;608 for (int i = 0; i < 8; i++) sel_sum += topk_val[i];609 out_indices[tid] = topk_idx[tid];610 out_scores[tid] = float(topk_val[tid] / (sel_sum + 1e-20f));611 }612"""613 return mx.fast.metal_kernel(614 name="maple_fused_router",615 input_names=["x", "w", "ctr_in"],616 output_names=["out_indices", "out_scores", "logits_scratch"],617 source=source,618 )619 620 621_fused_router_kernel = _make_fused_router_kernel()622 623 624class MapleGate(nn.Module):625 def __init__(self, args: ModelArgs):626 super().__init__()627 self.top_k = args.num_experts_per_tok628 self.num_experts = args.num_experts629 self.hidden_size = args.hidden_size630 # Kept as a raw parameter (not nn.Linear) so quantization never631 # touches it. The matmul accumulates in float32 and selection runs on632 # float32 scores.633 self.weight = mx.zeros((args.num_experts, args.hidden_size))634 self._router_ctr = None635 self._fused = None # None = unprobed, then True/False636 637 def _fused_call(self, x):638 if self._router_ctr is None:639 self._router_ctr = mx.zeros((8,), dtype=mx.uint32)640 mx.eval(self._router_ctr)641 inds, scores, _ = _fused_router_kernel(642 inputs=[x.reshape(-1), self.weight, self._router_ctr],643 template=[644 ("T_", self.weight.dtype),645 ("NEXP", self.num_experts),646 ("DIM", self.hidden_size),647 ],648 grid=((self.num_experts // 32) * 256, 1, 1),649 threadgroup=(256, 1, 1),650 output_shapes=[(8,), (8,), (self.num_experts,)],651 output_dtypes=[mx.int32, mx.float32, mx.float32],652 )653 shape = x.shape[:-1] + (self.top_k,)654 return inds.reshape(shape), scores.reshape(shape)655 656 def _reference(self, x):657 # `router_dtype: fp32`. In bf16 the near-tied top-8 boundary flips a658 # few percent of picks per layer, which compounds over 24 layers.659 gates = x.astype(mx.float32) @ self.weight.astype(mx.float32).T660 return group_expert_select(gates, self.top_k)661 662 def _probe(self, x):663 # Not _matches(): the two paths may order the selected experts664 # differently, and an exact tie at the top-k boundary may legitimately665 # pick either of the tied experts. Compare the sorted score vectors,666 # and bound-check the ids since a bad one indexes the expert gather.667 try:668 inds, scores = self._fused_call(x)669 ref_inds, ref_scores = self._reference(x)670 mx.eval(inds, scores, ref_inds, ref_scores)671 except Exception:672 return False673 return (674 inds.shape == ref_inds.shape675 and bool(mx.all((inds >= 0) & (inds < self.num_experts)))676 and bool(mx.allclose(mx.sort(scores), mx.sort(ref_scores), atol=1e-5))677 )678 679 def __call__(self, x):680 if self._fused is not False and x.size == self.hidden_size:681 if self._fused is None:682 self._fused = self._probe(x)683 if self._fused:684 return self._fused_call(x)685 return self._reference(x)686 687 688@partial(mx.compile, shapeless=True)689def aggregate_expert_outputs(expert_outputs, scores):690 # Combined in float32, rounded once at the end (reference `moe_infer`).691 return (692 (expert_outputs.astype(mx.float32) * scores[..., None])693 .sum(axis=-2)694 .astype(expert_outputs.dtype)695 )696 697 698class MapleSwitchGLU(nn.Module):699 """SwitchGLU with the up and gate projections fused into one gather700 matmul; sanitize() concatenates the checkpoint's split tensors."""701 702 def __init__(self, input_dims, hidden_dims, num_experts, bias=False):703 super().__init__()704 self.up_gate_proj = SwitchLinear(705 input_dims, 2 * hidden_dims, num_experts, bias=bias706 )707 self.down_proj = SwitchLinear(hidden_dims, input_dims, num_experts, bias=bias)708 709 def __call__(self, x, indices):710 x = mx.expand_dims(x, (-2, -3))711 712 do_sort = indices.size >= 64713 idx = indices714 inv_order = None715 if do_sort:716 x, idx, inv_order = _gather_sort(x, indices)717 718 x_up, x_gate = mx.split(719 self.up_gate_proj(x, idx, sorted_indices=do_sort), 2, axis=-1720 )721 x = self.down_proj(clamped_swiglu(x_gate, x_up), idx, sorted_indices=do_sort)722 723 if do_sort:724 x = _scatter_unsort(x, inv_order, indices.shape)725 726 return x.squeeze(-2)727 728 729class MapleSparseMoeBlock(nn.Module):730 def __init__(self, args: ModelArgs):731 super().__init__()732 self.gate = MapleGate(args)733 self.switch_mlp = MapleSwitchGLU(734 args.hidden_size,735 args.moe_intermediate_size,736 args.num_experts,737 bias=args.use_bias,738 )739 740 def __call__(self, x):741 inds, scores = self.gate(x)742 y = self.switch_mlp(x, inds)743 return aggregate_expert_outputs(y, scores)744 745 746class MapleDecoderLayer(nn.Module):747 def __init__(self, args: ModelArgs, layer_idx: int):748 super().__init__()749 self.self_attn = MapleAttention(args, layer_idx)750 self.mlp = (751 MapleSparseMoeBlock(args)752 if layer_idx >= args.first_k_dense_replace753 else MapleMLP(args)754 )755 self.input_layernorm = MapleRMSNorm(args.hidden_size, eps=args.rms_norm_eps)756 self.post_attention_layernorm = MapleRMSNorm(757 args.hidden_size, eps=args.rms_norm_eps758 )759 760 def __call__(761 self,762 x: mx.array,763 mask: Optional[mx.array] = None,764 cache: Optional[Any] = None,765 ) -> mx.array:766 r = self.self_attn(self.input_layernorm(x), mask, cache)767 h = x + r768 r = self.mlp(self.post_attention_layernorm(h))769 return h + r770 771 772class MapleModel(nn.Module):773 def __init__(self, args: ModelArgs):774 super().__init__()775 self.args = args776 self.word_embeddings = nn.Embedding(args.vocab_size, args.hidden_size)777 self.layers = [778 MapleDecoderLayer(args, layer_idx=i) for i in range(args.num_hidden_layers)779 ]780 self.norm = MapleRMSNorm(args.hidden_size, eps=args.rms_norm_eps)781 782 self.layer_types = args.layer_types783 self.window_size = args.sliding_window784 self.swa_idx = (785 self.layer_types.index("sliding_attention")786 if "sliding_attention" in self.layer_types787 else None788 )789 self.ga_idx = (790 self.layer_types.index("full_attention")791 if "full_attention" in self.layer_types792 else None793 )794 self._fused_add_norm = None # None = unprobed, then True/False795 self._zero = None796 797 def _decode_fused(self, h, cache, full_mask, swa_mask):798 """Decode loop with residual adds folded into the norms.799 800 Carries (h, r) instead of adding r back each step, so every801 add+norm pair is one dispatch. Identical arithmetic: the kernel802 rounds the sum once (as the bf16 add did) and norms the rounded803 stream with an fp32 weight multiply.804 """805 if self._zero is None:806 self._zero = mx.zeros(h.shape, h.dtype)807 mx.eval(self._zero)808 r = self._zero # x + 0 is exact in bf16809 for layer, c, layer_type in zip(self.layers, cache, self.layer_types):810 mask = full_mask if layer_type == "full_attention" else swa_mask811 ln = layer.input_layernorm812 h, hn = _add_rms_norm(h, r, ln.weight, ln.eps)813 r = layer.self_attn(hn, mask, c)814 ln = layer.post_attention_layernorm815 h, hn = _add_rms_norm(h, r, ln.weight, ln.eps)816 r = layer.mlp(hn)817 return _add_rms_norm(h, r, self.norm.weight, self.norm.eps)[1]818 819 def __call__(820 self,821 inputs: mx.array,822 cache: Optional[Any] = None,823 ):824 h = self.word_embeddings(inputs)825 826 if cache is None:827 cache = [None] * len(self.layers)828 829 full_mask = None830 swa_mask = None831 if self.ga_idx is not None:832 full_mask = create_attention_mask(h, cache[self.ga_idx])833 if self.swa_idx is not None:834 swa_mask = create_attention_mask(835 h, cache[self.swa_idx], window_size=self.window_size836 )837 838 if h.size == h.shape[-1]:839 if self._fused_add_norm is None:840 self._fused_add_norm = _add_rms_norm_ok(841 h.shape[-1], h.dtype, self.norm.weight, self.norm.eps842 )843 if self._fused_add_norm:844 return self._decode_fused(h, cache, full_mask, swa_mask)845 846 for layer, c, layer_type in zip(self.layers, cache, self.layer_types):847 mask = full_mask if layer_type == "full_attention" else swa_mask848 h = layer(h, mask, c)849 850 return self.norm(h)851 852 853class FlashHead(nn.Module):854 """Two-phase approximate lm_head for single-stream decode.855 856 Phase one scores quantized cluster centroids of the vocabulary; phase two857 computes exact logits only for the tokens of the top ``n_probes`` clusters858 (plus a fixed set of forced control tokens such as EOS). All other logits859 are -inf, so greedy decoding is exact whenever the true argmax lies in the860 probed clusters. Prefill and batched calls use the exact lm_head.861 862 Reference: FlashHead — Efficient Drop-in Replacement for the863 Classification Head in Language Model Inference.864 """865 866 def __init__(self, args: ModelArgs):867 super().__init__()868 meta = args.flash_head869 if not meta.get("scaled_centroids"):870 raise ValueError(871 "FlashHead metadata predates scaled centroids; regenerate with "872 "`python -m mlx_lm.ternary <checkpoint> --flash-head-only`."873 )874 n_clusters = meta["n_clusters"]875 cluster_size = meta["cluster_size"]876 # Default matches the converter's `--probes` default; every generated877 # checkpoint records the value explicitly.878 self.n_probes = min(meta.get("n_probes", 512), n_clusters)879 self.head_group_size = meta.get("head_group_size", 64)880 self.head_bits = meta.get("head_bits", 4)881 # Centroids are directions, pre-scaled at generation time by the882 # largest lm_head row norm in their cluster: that upper-bounds the883 # cluster's best logit, so high-frequency small-norm tokens are still884 # probed, and scoring stays a single matmul.885 self.centroids = nn.QuantizedLinear(886 args.hidden_size,887 n_clusters,888 bias=False,889 group_size=meta.get("group_size", 64),890 bits=meta.get("bits", 4),891 )892 self.token_map = mx.zeros((n_clusters, cluster_size), dtype=mx.int32)893 # Cluster-ordered copy of the quantized lm_head: subset logits are one894 # gather_qmm over the probed 32-row blocks, with no per-step gather.895 # It is a row-permutation of lm_head by token_map and nothing more, so896 # it is derived rather than stored: Model.sanitize rebuilds it at load.897 hidden = args.hidden_size898 self.head = {899 "weight": mx.zeros(900 (n_clusters, cluster_size, hidden * self.head_bits // 32),901 dtype=mx.uint32,902 ),903 "scales": mx.zeros(904 (n_clusters, cluster_size, hidden // self.head_group_size),905 dtype=mx.bfloat16,906 ),907 "biases": mx.zeros(908 (n_clusters, cluster_size, hidden // self.head_group_size),909 dtype=mx.bfloat16,910 ),911 }912 self._force_ids = mx.array(meta.get("force_tokens", []), dtype=mx.int32)913 self._force_rows = None914 915 def __call__(self, h: mx.array, lm_head: nn.Module) -> mx.array:916 hv = h[:, -1, :]917 top = mx.argpartition(self.centroids(hv), kth=-self.n_probes, axis=-1)[918 ..., -self.n_probes :919 ] # [1, n_probes]920 oids = self.token_map[top[0]].reshape(-1)921 922 logits = mx.gather_qmm(923 hv.reshape(1, 1, 1, 1, -1),924 self.head["weight"],925 self.head["scales"],926 self.head["biases"],927 rhs_indices=top[:, None, :],928 transpose=True,929 group_size=self.head_group_size,930 bits=self.head_bits,931 ).reshape(-1)932 933 if self._force_ids.size:934 if self._force_rows is None:935 self._force_rows = (936 lm_head.weight[self._force_ids],937 lm_head.scales[self._force_ids],938 lm_head.biases[self._force_ids],939 )940 mx.eval(*self._force_rows)941 fw, fs, fb = self._force_rows942 force_logits = mx.quantized_matmul(943 hv,944 fw,945 scales=fs,946 biases=fb,947 transpose=True,948 group_size=lm_head.group_size,949 bits=lm_head.bits,950 mode=getattr(lm_head, "mode", "affine"),951 )[0]952 oids = mx.concatenate([oids, self._force_ids])953 logits = mx.concatenate([logits, force_logits])954 955 vocab_size = lm_head.weight.shape[0]956 full = mx.full((1, 1, vocab_size), float("-inf"), dtype=logits.dtype)957 full[0, 0, oids] = logits958 return full959 960 961class Model(nn.Module):962 def __init__(self, args: ModelArgs):963 super().__init__()964 self.args = args965 self.model_type = args.model_type966 self.model = MapleModel(args)967 if not args.tie_word_embeddings:968 self.lm_head = nn.Linear(args.hidden_size, args.vocab_size, bias=False)969 if args.flash_head and args.use_flash_head and not args.tie_word_embeddings:970 self.lm_head_flash = FlashHead(args)971 else:972 self.lm_head_flash = None973 974 def __call__(975 self,976 inputs: mx.array,977 cache=None,978 ):979 out = self.model(inputs, cache)980 if self.args.tie_word_embeddings:981 return self.model.word_embeddings.as_linear(out)982 if (983 self.lm_head_flash is not None984 and out.shape[0] == 1985 and out.shape[1] == 1986 and isinstance(self.lm_head, nn.QuantizedLinear)987 and getattr(self.lm_head, "mode", "affine") == "affine"988 ):989 return self.lm_head_flash(out, self.lm_head)990 return self.lm_head(out)991 992 def sanitize(self, weights):993 if self.args.tie_word_embeddings:994 # Drop the head entirely (weight + quantization scales/biases).995 weights = {k: v for k, v in weights.items() if not k.startswith("lm_head.")}996 997 # FlashHead disabled (e.g. model_config={"flash_head": None}): drop its998 # tensors so checkpoints that carry them still load.999 if self.lm_head_flash is None:1000 weights = {1001 k: v for k, v in weights.items() if not k.startswith("lm_head_flash.")1002 }1003 else:1004 # Folded into the centroid rows at generation time; older shards1005 # still carry the tensor.1006 weights.pop("lm_head_flash.cluster_scale", None)1007 # `lm_head_flash.head.*` is lm_head permuted by token_map (see1008 # mlx_lm.ternary.generate_flash_head), so it is pure redundancy on1009 # disk. Checkpoints may ship it or omit it; reconcile both here.1010 if "lm_head_flash.head.weight" not in weights:1011 token_map = weights["lm_head_flash.token_map"]1012 order = token_map.reshape(-1)1013 for k in ("weight", "scales", "biases"):1014 weights[f"lm_head_flash.head.{k}"] = weights[f"lm_head.{k}"][1015 order1016 ].reshape(*token_map.shape, -1)1017 1018 # Ternary tensors carry one scale per output row, so checkpoints store1019 # it once as `row_alpha` and omit biases entirely (bias == -scale).1020 # Expand here so everything downstream — fusion below, and mlx's own1021 # quantized kernels — sees the per-group layout. Checkpoints written1022 # with `--group-scales` have no row_alpha and pass straight through.1023 row_alpha_keys = [k for k in weights if k.endswith(".row_alpha")]1024 if row_alpha_keys:1025 group_size = (self.args.quantization or {}).get("group_size", 128)1026 for key in row_alpha_keys:1027 alpha = weights.pop(key)1028 prefix = key[: -len(".row_alpha")]1029 packed = weights.get(f"{prefix}.weight")1030 if packed is None:1031 continue1032 # 2-bit packing stores 16 codes per uint32 word.1033 n_groups = (packed.shape[-1] * 16) // group_size1034 scales = mx.contiguous(1035 mx.broadcast_to(alpha[..., None], (*alpha.shape, n_groups))1036 )1037 weights[f"{prefix}.scales"] = scales1038 weights[f"{prefix}.biases"] = -scales1039 1040 # Stack per-expert weights from the Hugging Face layout into the1041 # SwitchGLU layout. Already-converted checkpoints pass through.1042 for l in range(self.args.num_hidden_layers):1043 prefix = f"model.layers.{l}"1044 for m in ["gate_proj", "down_proj", "up_proj"]:1045 for k in ["weight", "scales", "biases", "bias"]:1046 if f"{prefix}.mlp.experts.0.{m}.{k}" in weights:1047 to_join = [1048 weights.pop(f"{prefix}.mlp.experts.{e}.{m}.{k}")1049 for e in range(self.args.num_experts)1050 ]1051 weights[f"{prefix}.mlp.switch_mlp.{m}.{k}"] = mx.stack(to_join)1052 1053 # Fuse split projections: q/k/v -> qkv_proj (rows), MoE up/gate ->1054 # up_gate_proj (per-expert rows). Row-wise quantized tensors1055 # (weight/scales/biases) concatenate losslessly along the output1056 # axis.1057 for suffix in ["weight", "scales", "biases", "bias"]:1058 qkv = [1059 f"{prefix}.self_attn.{p}.{suffix}"1060 for p in ("q_proj", "k_proj", "v_proj")1061 ]1062 if qkv[0] in weights:1063 weights[f"{prefix}.self_attn.qkv_proj.{suffix}"] = mx.concatenate(1064 [weights.pop(k) for k in qkv], axis=01065 )1066 up = f"{prefix}.mlp.switch_mlp.up_proj.{suffix}"1067 gate = f"{prefix}.mlp.switch_mlp.gate_proj.{suffix}"1068 if up in weights:1069 weights[f"{prefix}.mlp.switch_mlp.up_gate_proj.{suffix}"] = (1070 mx.concatenate([weights.pop(up), weights.pop(gate)], axis=1)1071 )1072 1073 return weights1074 1075 def make_cache(self):1076 caches = []1077 for layer_type in self.model.layer_types:1078 if layer_type == "sliding_attention":1079 caches.append(RotatingKVCache(max_size=self.args.sliding_window))1080 else:1081 caches.append(KVCache())1082 return caches1083 1084 @property1085 def layers(self):1086 return self.model.layers1087 1088 @property1089 def quant_predicate(self):1090 def predicate(path, _):1091 if path.endswith("lm_head") or "word_embeddings" in path:1092 return {"group_size": 64, "bits": 4}1093 return True1094 1095 return predicate1096 