Team Ai
Modelpublic

txgsync/Maple-Preview-oQ4e

sourceHugging Facemitupdated 2mo agoView on Hugging Face
0likes243downloads
maple.py1096 linesDownload Raw Back to root
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