Text Generation
MLX
Safetensors
English
maple
mixture-of-experts
quantized
experimental
openmed
conversational
custom_code
4-bit precision
Instructions to use OpenMed/maple-preview-4bit-mlx with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use OpenMed/maple-preview-4bit-mlx with MLX:
# Make sure mlx-lm is installed # pip install --upgrade mlx-lm # Generate text with mlx-lm from mlx_lm import load, generate model, tokenizer = load("OpenMed/maple-preview-4bit-mlx") prompt = "Write a story about Einstein" messages = [{"role": "user", "content": prompt}] prompt = tokenizer.apply_chat_template( messages, add_generation_prompt=True ) text = generate(model, tokenizer, prompt=prompt, verbose=True) - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
- Pi
How to use OpenMed/maple-preview-4bit-mlx with Pi:
Start the MLX server
# Install MLX LM: uv tool install mlx-lm # Start a local OpenAI-compatible server: mlx_lm.server --model "OpenMed/maple-preview-4bit-mlx"
Configure the model in Pi
# Install Pi: npm install -g @mariozechner/pi-coding-agent # Add to ~/.pi/agent/models.json: { "providers": { "mlx-lm": { "baseUrl": "http://localhost:8080/v1", "api": "openai-completions", "apiKey": "none", "models": [ { "id": "OpenMed/maple-preview-4bit-mlx" } ] } } }Run Pi
# Start Pi in your project directory: pi
- OpenClaw new
How to use OpenMed/maple-preview-4bit-mlx with OpenClaw:
Start the MLX server
# Install MLX LM: uv tool install mlx-lm # Start a local OpenAI-compatible server: mlx_lm.server --model "OpenMed/maple-preview-4bit-mlx"
Configure OpenClaw
# Install OpenClaw: npm install -g openclaw@latest # Register the local server and set it as the default model: openclaw onboard --non-interactive --mode local \ --auth-choice custom-api-key \ --custom-base-url http://127.0.0.1:8080/v1 \ --custom-model-id "OpenMed/maple-preview-4bit-mlx" \ --custom-provider-id mlx-lm \ --custom-compatibility openai \ --custom-text-input \ --accept-risk \ --skip-health
Run OpenClaw
openclaw agent --local --agent main --message "Hello from Hugging Face"
- MLX LM
How to use OpenMed/maple-preview-4bit-mlx with MLX LM:
Generate or start a chat session
# Install MLX LM uv tool install mlx-lm # Interactive chat REPL mlx_lm.chat --model "OpenMed/maple-preview-4bit-mlx"
Run an OpenAI-compatible server
# Install MLX LM uv tool install mlx-lm # Start the server mlx_lm.server --model "OpenMed/maple-preview-4bit-mlx" # Calling the OpenAI-compatible server with curl curl -X POST "http://localhost:8000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "OpenMed/maple-preview-4bit-mlx", "messages": [ {"role": "user", "content": "Hello"} ] }' - Hermes Agent
How to use OpenMed/maple-preview-4bit-mlx with Hermes Agent:
Start the MLX server
# Install MLX LM: uv tool install mlx-lm # Start a local OpenAI-compatible server: mlx_lm.server --model "OpenMed/maple-preview-4bit-mlx"
Configure Hermes
# Install Hermes: curl -fsSL https://hermes-agent.nousresearch.com/install.sh | bash hermes setup # Point Hermes at the local server: hermes config set model.provider custom hermes config set model.base_url http://127.0.0.1:8080/v1 hermes config set model.default OpenMed/maple-preview-4bit-mlx
Run Hermes
hermes
| # Copyright © 2026 DeepGrove AI. | |
| from dataclasses import dataclass | |
| from functools import partial | |
| from typing import Any, List, Optional | |
| import mlx.core as mx | |
| import mlx.nn as nn | |
| # Absolute imports so this file also works standalone when shipped inside a | |
| # checkpoint and loaded via the config's `model_file` (trust_remote_code). | |
| from mlx_lm.models.activations import swiglu | |
| from mlx_lm.models.base import ( | |
| BaseModelArgs, | |
| create_attention_mask, | |
| scaled_dot_product_attention, | |
| ) | |
| from mlx_lm.models.cache import KVCache, RotatingKVCache | |
| from mlx_lm.models.rope_utils import initialize_rope | |
| from mlx_lm.models.switch_layers import SwitchLinear | |
| # SwiGLU clamp for the MoE experts only (the dense MapleMLP is unclamped); | |
| # part of the trained forward pass, not an optional guard. | |
| MLP_CLAMP = 7.0 | |
| def clamped_swiglu(gate, x): | |
| # Python floats, not 0-d arrays, so bf16 activations stay bf16. | |
| return nn.silu(mx.minimum(gate, MLP_CLAMP)) * mx.clip(x, -MLP_CLAMP, MLP_CLAMP) | |
| class MapleRMSNorm(nn.Module): | |
| """RMSNorm with the weight multiply in float32. | |
| The reference rounds only the finished product; mx.fast.rms_norm rounds | |
| the normalized activation first (~1% per element). Float32 inputs to the | |
| same kernel reproduce the reference bit-for-bit. | |
| """ | |
| def __init__(self, dims: int, eps: float = 1e-6): | |
| super().__init__() | |
| self.weight = mx.ones((dims,)) | |
| self.eps = eps | |
| def __call__(self, x: mx.array) -> mx.array: | |
| return mx.fast.rms_norm( | |
| x.astype(mx.float32), self.weight.astype(mx.float32), self.eps | |
| ).astype(x.dtype) | |
| def _make_add_rms_norm_kernel(eps): | |
| """Residual add + RMSNorm in ONE dispatch for single-token decode. | |
| Emits both h = x + r (the residual stream, rounded once like a bf16 add) | |
| and hn = rmsnorm(h) with the weight multiply in fp32 (reference | |
| semantics, identical to MapleRMSNorm). Folding the add into the norm and | |
| skipping the astype round-trips replaces ~4 dispatches with 1, and the | |
| decode step is bounded by its serial dispatch chain, not by this math. | |
| """ | |
| source = """ | |
| uint tid = thread_position_in_threadgroup.x; | |
| constexpr uint N = DIM; | |
| constexpr uint PT = N / 256u; | |
| float hb[PT]; | |
| float ss = 0.0f; | |
| for (uint i = 0; i < PT; ++i) { | |
| uint j = tid * PT + i; | |
| float v = (float)x[j] + (float)r[j]; | |
| T_ vb = (T_)v; // one rounding, same as a bf16 add | |
| h_out[j] = vb; | |
| hb[i] = (float)vb; // norm sees the rounded stream | |
| ss += hb[i] * hb[i]; | |
| } | |
| ss = simd_sum(ss); | |
| threadgroup float sums[8]; | |
| uint sg = tid / 32u; | |
| uint lane = tid % 32u; | |
| if (lane == 0u) sums[sg] = ss; | |
| threadgroup_barrier(mem_flags::mem_threadgroup); | |
| float tot = 0.0f; | |
| for (uint i = 0; i < 8u; ++i) tot += sums[i]; | |
| float scale = metal::rsqrt(tot / (float)N + EPS_); | |
| for (uint i = 0; i < PT; ++i) { | |
| uint j = tid * PT + i; | |
| hn_out[j] = (T_)(hb[i] * scale * (float)w[j]); | |
| } | |
| """.replace("EPS_", f"{eps:.10e}f") | |
| tag = f"{eps:.3e}".replace(".", "_").replace("-", "m").replace("+", "p") | |
| return mx.fast.metal_kernel( | |
| name=f"maple_add_rms_norm_{tag}", | |
| input_names=["x", "r", "w"], | |
| output_names=["h_out", "hn_out"], | |
| source=source, | |
| ) | |
| _add_rms_kernels = {} | |
| def _add_rms_norm(h, r, w, eps): | |
| kernel = _add_rms_kernels.get(eps) | |
| if kernel is None: | |
| kernel = _add_rms_kernels[eps] = _make_add_rms_norm_kernel(eps) | |
| return kernel( | |
| inputs=[h.reshape(-1), r.reshape(-1), w], | |
| template=[("T_", h.dtype), ("DIM", h.shape[-1])], | |
| grid=(256, 1, 1), | |
| threadgroup=(256, 1, 1), | |
| output_shapes=[h.shape, h.shape], | |
| output_dtypes=[h.dtype, h.dtype], | |
| ) | |
| # Inlined rather than imported from switch_layers: those helpers are private | |
| # (underscore-prefixed), and this file must keep loading against whatever | |
| # mlx-lm a user has installed when it ships inside a checkpoint. | |
| def _gather_sort(x, indices): | |
| *_, M = indices.shape | |
| indices = indices.flatten() | |
| order = mx.argsort(indices) | |
| inv_order = mx.argsort(order) | |
| return x.flatten(0, -3)[order // M], indices[order], inv_order | |
| def _scatter_unsort(x, inv_order, shape=None): | |
| x = x[inv_order] | |
| if shape is not None: | |
| x = mx.unflatten(x, 0, shape) | |
| return x | |
| class ModelArgs(BaseModelArgs): | |
| model_type: str = "maple" | |
| hidden_size: int = 2048 | |
| intermediate_size: int = 5120 | |
| moe_intermediate_size: int = 512 | |
| num_hidden_layers: int = 24 | |
| num_attention_heads: int = 16 | |
| num_key_value_heads: int = 4 | |
| head_dim: int = 128 | |
| num_experts: int = 256 | |
| num_experts_per_tok: int = 8 | |
| first_k_dense_replace: int = 0 | |
| rms_norm_eps: float = 1e-6 | |
| rope_theta: float = 10000.0 | |
| rope_scaling: Optional[dict] = None | |
| partial_rotary_factor: float = 0.5 | |
| max_position_embeddings: int = 140000 | |
| vocab_size: int = 151936 | |
| sliding_window: int = 512 | |
| layer_types: Optional[List[str]] = None | |
| use_qk_norm: bool = True | |
| use_bias: bool = False | |
| tie_word_embeddings: bool = False | |
| # FlashHead metadata written by `mlx_lm.ternary --flash-head`. The exact | |
| # lm_head is the default; opt in to the approximate fast head with | |
| # mlx_lm.load(..., model_config={"use_flash_head": True}). | |
| flash_head: Optional[dict] = None | |
| use_flash_head: bool = False | |
| # Populated from the checkpoint's config; sanitize() reads group_size from | |
| # it to expand row-scale (`row_alpha`) ternary tensors. | |
| quantization: Optional[dict] = None | |
| def __post_init__(self): | |
| # Single source of truth for per-layer attention types: attention | |
| # (RoPE/NoPE), masks, and caches all read this resolved list. | |
| if not self.layer_types: | |
| self.layer_types = ["full_attention"] * self.num_hidden_layers | |
| def _make_qk_norm_rope_kernel(): | |
| """Fused per-head RMSNorm + partial RoPE for single-token decode. | |
| One dispatch replaces q_norm, k_norm and two rope calls. One simdgroup per | |
| head: normalize head_dim values, scale by the head's norm weight, and | |
| rotate the first ROPE_DIM dims (non-traditional pairing i, i+R/2) at the | |
| given position. NoPE layers pass ROPE_DIM=0. | |
| """ | |
| source = """ | |
| uint head = thread_position_in_grid.y; | |
| uint lane = thread_position_in_grid.x; | |
| constexpr int per_lane = HEAD_DIM / 32; | |
| const device T_* xh = x + head * HEAD_DIM; | |
| const device T_* wh = w + head * HEAD_DIM; | |
| device T_* oh = out + head * HEAD_DIM; | |
| float ss = 0.0f; | |
| for (int i = 0; i < per_lane; ++i) { | |
| float v = (float)xh[lane * per_lane + i]; | |
| ss += v * v; | |
| } | |
| ss = simd_sum(ss); | |
| float pos = pos_eps[0]; | |
| float eps = pos_eps[1]; | |
| float scale = metal::rsqrt(ss / HEAD_DIM + eps); | |
| for (int i = 0; i < per_lane; ++i) { | |
| int j = lane * per_lane + i; | |
| float v = (float)xh[j] * scale * (float)wh[j]; | |
| if (ROPE_DIM > 0 && j < ROPE_DIM) { | |
| constexpr int rhalf = ROPE_DIM > 0 ? ROPE_DIM / 2 : 1; | |
| int p = j < rhalf ? j : j - rhalf; | |
| float theta = pos * inv_freq[p]; | |
| float c = metal::cos(theta); | |
| float s = metal::sin(theta); | |
| int j2 = j < rhalf ? j + rhalf : j - rhalf; | |
| float u = (float)xh[j2] * scale * (float)wh[j2]; | |
| v = j < rhalf ? (v * c - u * s) : (v * c + u * s); | |
| } | |
| oh[j] = (T_)v; | |
| } | |
| """ | |
| return mx.fast.metal_kernel( | |
| name="maple_qk_norm_rope", | |
| input_names=["x", "w", "inv_freq", "pos_eps"], | |
| output_names=["out"], | |
| source=source, | |
| ) | |
| _qk_norm_rope_kernel = _make_qk_norm_rope_kernel() | |
| class MapleAttention(nn.Module): | |
| def __init__(self, args: ModelArgs, layer_idx: int): | |
| super().__init__() | |
| self.num_attention_heads = args.num_attention_heads | |
| self.num_key_value_heads = args.num_key_value_heads | |
| self.head_dim = args.head_dim or args.hidden_size // args.num_attention_heads | |
| self.scale = self.head_dim**-0.5 | |
| self.use_qk_norm = args.use_qk_norm | |
| # q/k/v are stored fused (one matmul per step); sanitize() concatenates | |
| # the checkpoint's split projections. | |
| self.qkv_proj = nn.Linear( | |
| args.hidden_size, | |
| (args.num_attention_heads + 2 * args.num_key_value_heads) | |
| * self.head_dim, | |
| bias=args.use_bias, | |
| ) | |
| self.o_proj = nn.Linear( | |
| args.num_attention_heads * self.head_dim, | |
| args.hidden_size, | |
| bias=args.use_bias, | |
| ) | |
| if args.use_qk_norm: | |
| self.q_norm = MapleRMSNorm(self.head_dim, eps=args.rms_norm_eps) | |
| self.k_norm = MapleRMSNorm(self.head_dim, eps=args.rms_norm_eps) | |
| self._eps = args.rms_norm_eps | |
| self._rope_base = args.rope_theta | |
| self._qk_w = None | |
| self._inv_freq = None | |
| # Maple applies RoPE only on sliding-window layers; full-attention | |
| # layers use no positional encoding (NoPE). | |
| self.use_rope = args.layer_types[layer_idx] == "sliding_attention" | |
| if self.use_rope: | |
| rope_dim = int(self.head_dim * args.partial_rotary_factor) | |
| self.rope = initialize_rope( | |
| rope_dim, | |
| args.rope_theta, | |
| traditional=False, | |
| scaling_config=args.rope_scaling, | |
| max_position_embeddings=args.max_position_embeddings, | |
| ) | |
| def __call__( | |
| self, | |
| x: mx.array, | |
| mask: Optional[mx.array] = None, | |
| cache: Optional[Any] = None, | |
| ) -> mx.array: | |
| B, L, _ = x.shape | |
| qkv = self.qkv_proj(x) | |
| q_size = self.num_attention_heads * self.head_dim | |
| kv_size = self.num_key_value_heads * self.head_dim | |
| if B == 1 and L == 1 and self.use_qk_norm: | |
| # Single-token decode: one fused dispatch for both norms and both | |
| # rope applications. | |
| n_q = self.num_attention_heads | |
| n_kv = self.num_key_value_heads | |
| if self._qk_w is None: | |
| self._qk_w = mx.contiguous( | |
| mx.concatenate( | |
| [ | |
| mx.broadcast_to( | |
| self.q_norm.weight[None], (n_q, self.head_dim) | |
| ), | |
| mx.broadcast_to( | |
| self.k_norm.weight[None], (n_kv, self.head_dim) | |
| ), | |
| ] | |
| ) | |
| ) | |
| if self.use_rope: | |
| half = self.rope.dims // 2 | |
| self._inv_freq = self._rope_base ** ( | |
| -mx.arange(half, dtype=mx.float32) / half | |
| ) | |
| else: | |
| self._inv_freq = mx.ones((1,), dtype=mx.float32) | |
| mx.eval(self._qk_w, self._inv_freq) | |
| # cache.offset is a Python int for a plain cache but an mx.array | |
| # for the server's mergeable prompt cache; coerce to a scalar so | |
| # the pos/eps pair is always uniform. | |
| offset = cache.offset if cache is not None else 0 | |
| pos_eps = mx.array([float(offset), self._eps], dtype=mx.float32) | |
| qk = qkv.reshape(-1)[: (n_q + n_kv) * self.head_dim].reshape( | |
| n_q + n_kv, self.head_dim | |
| ) | |
| out = _qk_norm_rope_kernel( | |
| inputs=[qk, self._qk_w, self._inv_freq, pos_eps], | |
| template=[ | |
| ("T_", qkv.dtype), | |
| ("HEAD_DIM", self.head_dim), | |
| ("ROPE_DIM", self.rope.dims if self.use_rope else 0), | |
| ], | |
| grid=(32, n_q + n_kv, 1), | |
| threadgroup=(32, 1, 1), | |
| output_shapes=[qk.shape], | |
| output_dtypes=[qkv.dtype], | |
| )[0] | |
| queries = out[:n_q].reshape(1, n_q, 1, self.head_dim) | |
| keys = out[n_q:].reshape(1, n_kv, 1, self.head_dim) | |
| values = qkv.reshape(-1)[(n_q + n_kv) * self.head_dim :].reshape( | |
| 1, n_kv, 1, self.head_dim | |
| ) | |
| else: | |
| q, k, v = mx.split(qkv, [q_size, q_size + kv_size], axis=-1) | |
| queries = q.reshape(B, L, self.num_attention_heads, self.head_dim) | |
| keys = k.reshape(B, L, self.num_key_value_heads, self.head_dim) | |
| values = v.reshape(B, L, self.num_key_value_heads, self.head_dim) | |
| if self.use_qk_norm: | |
| queries = self.q_norm(queries) | |
| keys = self.k_norm(keys) | |
| queries = queries.transpose(0, 2, 1, 3) | |
| keys = keys.transpose(0, 2, 1, 3) | |
| values = values.transpose(0, 2, 1, 3) | |
| if self.use_rope: | |
| offset = cache.offset if cache is not None else 0 | |
| queries = self.rope(queries, offset=offset) | |
| keys = self.rope(keys, offset=offset) | |
| if cache is not None: | |
| keys, values = cache.update_and_fetch(keys, values) | |
| output = scaled_dot_product_attention( | |
| queries, keys, values, cache=cache, scale=self.scale, mask=mask | |
| ) | |
| output = output.transpose(0, 2, 1, 3).reshape(B, L, -1) | |
| return self.o_proj(output) | |
| class MapleMLP(nn.Module): | |
| def __init__(self, args: ModelArgs, intermediate_size: Optional[int] = None): | |
| super().__init__() | |
| intermediate_size = intermediate_size or args.intermediate_size | |
| self.gate_proj = nn.Linear(args.hidden_size, intermediate_size, bias=args.use_bias) | |
| self.up_proj = nn.Linear(args.hidden_size, intermediate_size, bias=args.use_bias) | |
| self.down_proj = nn.Linear(intermediate_size, args.hidden_size, bias=args.use_bias) | |
| def __call__(self, x) -> mx.array: | |
| # Dense / shared-expert MLP: no clamp; only the MoE experts clamp. | |
| # Unused at first_k_dense_replace=0 with no shared experts, but keep | |
| # it faithful. | |
| return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) | |
| def group_expert_select(gates, top_k): | |
| # Maple routes with a plain softmax over all experts followed by top-k | |
| # selection and renormalization, computed in float32. | |
| scores = mx.softmax(gates.astype(mx.float32), axis=-1) | |
| inds = mx.argpartition(scores, kth=-top_k, axis=-1)[..., -top_k:] | |
| scores = mx.take_along_axis(scores, inds, axis=-1) | |
| scores = scores / (scores.sum(axis=-1, keepdims=True) + 1e-20) | |
| return inds, scores | |
| def _make_fused_router_kernel(): | |
| """Router gemv + softmax + top-8 + renormalize in ONE dispatch (+18%). | |
| Replaces ~6 kernels per layer. NE/32 threadgroups each compute 32 logits, | |
| keep them in float32 (`router_dtype: fp32`), and publish through an | |
| atomic-float scratch (plain device stores are not reliably visible across | |
| threadgroups on Apple GPUs); the last threadgroup to arrive does the | |
| softmax + top-8 + renorm. | |
| """ | |
| source = """ | |
| constexpr uint NE = NEXP; | |
| constexpr uint D = DIM; | |
| constexpr uint NTG = NE / 32u; | |
| constexpr uint TM = 4u; | |
| constexpr uint TN = 4u; | |
| constexpr uint BLOCKN = 32u * TN; | |
| constexpr uint NITER = D / BLOCKN; | |
| uint tid = thread_position_in_threadgroup.x; | |
| uint tgid = threadgroup_position_in_grid.x; | |
| uint n_threads = 256u; | |
| uint sg_id = tid / 32u; | |
| uint lane = tid % 32u; | |
| uint n_sg = n_threads / 32u; | |
| uint row0 = tgid * (n_sg * TM) + sg_id * TM; | |
| float result[TM] = {0.0f, 0.0f, 0.0f, 0.0f}; | |
| uint bn = lane * TN; | |
| for (uint i = 0u; i < NITER; ++i) { | |
| float v[TN]; | |
| for (uint tn = 0u; tn < TN; ++tn) v[tn] = float(x[bn + tn]); | |
| for (uint tm = 0u; tm < TM; ++tm) { | |
| const device T_* wrow = w + (ulong)(row0 + tm) * D; | |
| T_ inter[TN]; | |
| for (uint tn = 0u; tn < TN; ++tn) inter[tn] = wrow[bn + tn]; | |
| for (uint tn = 0u; tn < TN; ++tn) result[tm] += inter[tn] * v[tn]; | |
| } | |
| bn += BLOCKN; | |
| } | |
| for (uint tm = 0u; tm < TM; ++tm) { | |
| for (ushort sn = 16; sn >= 1; sn >>= 1) { | |
| result[tm] += simd_shuffle_down(result[tm], sn); | |
| } | |
| } | |
| device atomic_float* ls = (device atomic_float*)logits_scratch; | |
| if (lane == 0u) { | |
| for (uint tm = 0u; tm < TM; ++tm) { | |
| atomic_store_explicit(&ls[row0 + tm], result[tm], | |
| memory_order_relaxed); | |
| } | |
| } | |
| threadgroup_barrier(mem_flags::mem_device); | |
| threadgroup uint last_flag; | |
| if (tid == 0u) { | |
| device atomic_uint* ctr = (device atomic_uint*)ctr_in; | |
| uint prev = atomic_fetch_add_explicit(ctr, 1u, memory_order_relaxed); | |
| uint last = (prev == NTG - 1u) ? 1u : 0u; | |
| if (last == 1u) atomic_store_explicit(ctr, 0u, memory_order_relaxed); | |
| last_flag = last; | |
| } | |
| threadgroup_barrier(mem_flags::mem_threadgroup); | |
| if (last_flag == 0u) return; | |
| threadgroup_barrier(mem_flags::mem_device); | |
| float my_max = -1e30f; | |
| for (uint e = tid; e < NE; e += n_threads) { | |
| float v = atomic_load_explicit(&ls[e], memory_order_relaxed); | |
| if (v > my_max) my_max = v; | |
| } | |
| for (int off = 16; off > 0; off >>= 1) { | |
| float other = simd_shuffle_down(my_max, off); | |
| if (other > my_max) my_max = other; | |
| } | |
| threadgroup float sg_red[16]; | |
| if (lane == 0u) sg_red[sg_id] = my_max; | |
| threadgroup_barrier(mem_flags::mem_threadgroup); | |
| if (tid == 0u) { | |
| float m = sg_red[0]; | |
| for (uint s = 1u; s < n_sg; s++) if (sg_red[s] > m) m = sg_red[s]; | |
| sg_red[0] = m; | |
| } | |
| threadgroup_barrier(mem_flags::mem_threadgroup); | |
| float lmax = sg_red[0]; | |
| threadgroup float scores[NE]; | |
| float my_sum = 0.0f; | |
| for (uint e = tid; e < NE; e += n_threads) { | |
| float lv = atomic_load_explicit(&ls[e], memory_order_relaxed); | |
| float v = metal::exp(lv - lmax); | |
| scores[e] = v; | |
| my_sum += v; | |
| } | |
| for (int off = 16; off > 0; off >>= 1) { | |
| my_sum += simd_shuffle_down(my_sum, off); | |
| } | |
| threadgroup_barrier(mem_flags::mem_threadgroup); | |
| if (lane == 0u) sg_red[sg_id] = my_sum; | |
| threadgroup_barrier(mem_flags::mem_threadgroup); | |
| if (tid == 0u) { | |
| float ssum = sg_red[0]; | |
| for (uint i = 1u; i < n_sg; i++) ssum += sg_red[i]; | |
| sg_red[0] = ssum; | |
| } | |
| threadgroup_barrier(mem_flags::mem_threadgroup); | |
| float inv_total = 1.0f / (sg_red[0] + 1e-20f); | |
| for (uint e = tid; e < NE; e += n_threads) { | |
| scores[e] = scores[e] * inv_total; | |
| } | |
| threadgroup_barrier(mem_flags::mem_threadgroup); | |
| threadgroup int topk_idx[8]; | |
| threadgroup float topk_val[8]; | |
| threadgroup uint8_t used[NE]; | |
| for (uint e = tid; e < NE; e += n_threads) used[e] = 0; | |
| threadgroup_barrier(mem_flags::mem_threadgroup); | |
| for (int k = 0; k < 8; k++) { | |
| float my_best = -1e30f; | |
| int my_idx = 0; | |
| for (int e = int(tid); e < int(NE); e += int(n_threads)) { | |
| if (!used[e] && scores[e] > my_best) { | |
| my_best = scores[e]; | |
| my_idx = e; | |
| } | |
| } | |
| for (int off = 16; off > 0; off >>= 1) { | |
| float other_v = simd_shuffle_down(my_best, off); | |
| int other_i = simd_shuffle_down(my_idx, off); | |
| if (other_v > my_best) { my_best = other_v; my_idx = other_i; } | |
| } | |
| threadgroup float sg_vals[16]; | |
| threadgroup int sg_idxs[16]; | |
| if (lane == 0u) { sg_vals[sg_id] = my_best; sg_idxs[sg_id] = my_idx; } | |
| threadgroup_barrier(mem_flags::mem_threadgroup); | |
| if (tid == 0u) { | |
| float bv = sg_vals[0]; int bi = sg_idxs[0]; | |
| for (uint s = 1u; s < n_sg; s++) { | |
| if (sg_vals[s] > bv) { bv = sg_vals[s]; bi = sg_idxs[s]; } | |
| } | |
| topk_val[k] = bv; topk_idx[k] = bi; | |
| used[bi] = 1; | |
| } | |
| threadgroup_barrier(mem_flags::mem_threadgroup); | |
| } | |
| if (tid < 8u) { | |
| float sel_sum = 0.0f; | |
| for (int i = 0; i < 8; i++) sel_sum += topk_val[i]; | |
| out_indices[tid] = topk_idx[tid]; | |
| out_scores[tid] = float(topk_val[tid] / (sel_sum + 1e-20f)); | |
| } | |
| """ | |
| return mx.fast.metal_kernel( | |
| name="maple_fused_router", | |
| input_names=["x", "w", "ctr_in"], | |
| output_names=["out_indices", "out_scores", "logits_scratch"], | |
| source=source, | |
| ) | |
| _fused_router_kernel = _make_fused_router_kernel() | |
| class MapleGate(nn.Module): | |
| def __init__(self, args: ModelArgs): | |
| super().__init__() | |
| self.top_k = args.num_experts_per_tok | |
| self.num_experts = args.num_experts | |
| self.hidden_size = args.hidden_size | |
| # Kept as a raw parameter (not nn.Linear) so quantization never | |
| # touches it. The matmul accumulates in float32 and selection runs on | |
| # float32 scores. | |
| self.weight = mx.zeros((args.num_experts, args.hidden_size)) | |
| self._router_ctr = None | |
| self._router_probed = False | |
| self._fused_ok = ( | |
| args.num_experts % 32 == 0 | |
| and args.hidden_size % 128 == 0 | |
| and args.num_experts_per_tok == 8 | |
| ) | |
| def _fused(self, x): | |
| if self._router_ctr is None: | |
| self._router_ctr = mx.zeros((8,), dtype=mx.uint32) | |
| mx.eval(self._router_ctr) | |
| inds, scores, _ = _fused_router_kernel( | |
| inputs=[x.reshape(-1), self.weight, self._router_ctr], | |
| template=[ | |
| ("T_", self.weight.dtype), | |
| ("NEXP", self.num_experts), | |
| ("DIM", self.hidden_size), | |
| ], | |
| grid=((self.num_experts // 32) * 256, 1, 1), | |
| threadgroup=(256, 1, 1), | |
| output_shapes=[(8,), (8,), (self.num_experts,)], | |
| output_dtypes=[mx.int32, mx.float32, mx.float32], | |
| ) | |
| shape = x.shape[:-1] + (8,) | |
| return inds.reshape(shape), scores.reshape(shape) | |
| def __call__(self, x): | |
| if ( | |
| self._fused_ok | |
| and x.size == self.hidden_size | |
| and self.weight.dtype == mx.bfloat16 | |
| ): | |
| try: | |
| inds, scores = self._fused(x) | |
| if not self._router_probed: | |
| # mlx is lazy: force one eval so a kernel failure surfaces | |
| # here and latches the fallback. | |
| mx.eval(inds, scores) | |
| self._router_probed = True | |
| return inds, scores | |
| except Exception: | |
| self._fused_ok = False | |
| # `router_dtype: fp32`. In bf16 the near-tied top-8 boundary flips a | |
| # few percent of picks per layer, which compounds over 24 layers. | |
| gates = x.astype(mx.float32) @ self.weight.astype(mx.float32).T | |
| return group_expert_select(gates, self.top_k) | |
| def aggregate_expert_outputs(expert_outputs, scores): | |
| # Combined in float32, rounded once at the end (reference `moe_infer`). | |
| return ( | |
| (expert_outputs.astype(mx.float32) * scores[..., None]) | |
| .sum(axis=-2) | |
| .astype(expert_outputs.dtype) | |
| ) | |
| class MapleSwitchGLU(nn.Module): | |
| """SwitchGLU with the up and gate projections fused into one gather | |
| matmul; sanitize() concatenates the checkpoint's split tensors.""" | |
| def __init__(self, input_dims, hidden_dims, num_experts, bias=False): | |
| super().__init__() | |
| self.up_gate_proj = SwitchLinear( | |
| input_dims, 2 * hidden_dims, num_experts, bias=bias | |
| ) | |
| self.down_proj = SwitchLinear(hidden_dims, input_dims, num_experts, bias=bias) | |
| def __call__(self, x, indices): | |
| x = mx.expand_dims(x, (-2, -3)) | |
| do_sort = indices.size >= 64 | |
| idx = indices | |
| inv_order = None | |
| if do_sort: | |
| x, idx, inv_order = _gather_sort(x, indices) | |
| x_up, x_gate = mx.split( | |
| self.up_gate_proj(x, idx, sorted_indices=do_sort), 2, axis=-1 | |
| ) | |
| x = self.down_proj(clamped_swiglu(x_gate, x_up), idx, sorted_indices=do_sort) | |
| if do_sort: | |
| x = _scatter_unsort(x, inv_order, indices.shape) | |
| return x.squeeze(-2) | |
| class MapleSparseMoeBlock(nn.Module): | |
| def __init__(self, args: ModelArgs): | |
| super().__init__() | |
| self.gate = MapleGate(args) | |
| self.switch_mlp = MapleSwitchGLU( | |
| args.hidden_size, | |
| args.moe_intermediate_size, | |
| args.num_experts, | |
| bias=args.use_bias, | |
| ) | |
| def __call__(self, x): | |
| inds, scores = self.gate(x) | |
| y = self.switch_mlp(x, inds) | |
| return aggregate_expert_outputs(y, scores) | |
| class MapleDecoderLayer(nn.Module): | |
| def __init__(self, args: ModelArgs, layer_idx: int): | |
| super().__init__() | |
| self.self_attn = MapleAttention(args, layer_idx) | |
| self.mlp = ( | |
| MapleSparseMoeBlock(args) | |
| if layer_idx >= args.first_k_dense_replace | |
| else MapleMLP(args) | |
| ) | |
| self.input_layernorm = MapleRMSNorm(args.hidden_size, eps=args.rms_norm_eps) | |
| self.post_attention_layernorm = MapleRMSNorm( | |
| args.hidden_size, eps=args.rms_norm_eps | |
| ) | |
| def __call__( | |
| self, | |
| x: mx.array, | |
| mask: Optional[mx.array] = None, | |
| cache: Optional[Any] = None, | |
| ) -> mx.array: | |
| r = self.self_attn(self.input_layernorm(x), mask, cache) | |
| h = x + r | |
| r = self.mlp(self.post_attention_layernorm(h)) | |
| return h + r | |
| class MapleModel(nn.Module): | |
| def __init__(self, args: ModelArgs): | |
| super().__init__() | |
| self.args = args | |
| self.word_embeddings = nn.Embedding(args.vocab_size, args.hidden_size) | |
| self.layers = [ | |
| MapleDecoderLayer(args, layer_idx=i) | |
| for i in range(args.num_hidden_layers) | |
| ] | |
| self.norm = MapleRMSNorm(args.hidden_size, eps=args.rms_norm_eps) | |
| self.layer_types = args.layer_types | |
| self.window_size = args.sliding_window | |
| self.swa_idx = ( | |
| self.layer_types.index("sliding_attention") | |
| if "sliding_attention" in self.layer_types | |
| else None | |
| ) | |
| self.ga_idx = ( | |
| self.layer_types.index("full_attention") | |
| if "full_attention" in self.layer_types | |
| else None | |
| ) | |
| self._fused_add_norm = None # None = unprobed, then True/False | |
| self._zero = None | |
| def _decode_fused(self, h, cache, full_mask, swa_mask): | |
| """Decode loop with residual adds folded into the norms. | |
| Carries (h, r) instead of adding r back each step, so every | |
| add+norm pair is one dispatch. Identical arithmetic: the kernel | |
| rounds the sum once (as the bf16 add did) and norms the rounded | |
| stream with an fp32 weight multiply. | |
| """ | |
| if self._zero is None: | |
| self._zero = mx.zeros(h.shape, h.dtype) | |
| mx.eval(self._zero) | |
| r = self._zero # x + 0 is exact in bf16 | |
| for layer, c, layer_type in zip(self.layers, cache, self.layer_types): | |
| mask = full_mask if layer_type == "full_attention" else swa_mask | |
| ln = layer.input_layernorm | |
| h, hn = _add_rms_norm(h, r, ln.weight, ln.eps) | |
| r = layer.self_attn(hn, mask, c) | |
| ln = layer.post_attention_layernorm | |
| h, hn = _add_rms_norm(h, r, ln.weight, ln.eps) | |
| r = layer.mlp(hn) | |
| return _add_rms_norm(h, r, self.norm.weight, self.norm.eps)[1] | |
| def __call__( | |
| self, | |
| inputs: mx.array, | |
| cache: Optional[Any] = None, | |
| ): | |
| h = self.word_embeddings(inputs) | |
| if cache is None: | |
| cache = [None] * len(self.layers) | |
| full_mask = None | |
| swa_mask = None | |
| if self.ga_idx is not None: | |
| full_mask = create_attention_mask(h, cache[self.ga_idx]) | |
| if self.swa_idx is not None: | |
| swa_mask = create_attention_mask( | |
| h, cache[self.swa_idx], window_size=self.window_size | |
| ) | |
| if h.size == h.shape[-1] and h.shape[-1] % 256 == 0: | |
| if self._fused_add_norm is None: | |
| # Probe on dummy data, outside the real graph: a failure here | |
| # must not leave the caches half-updated. | |
| try: | |
| z = mx.zeros((1, 1, h.shape[-1]), h.dtype) | |
| mx.eval(_add_rms_norm(z, z, self.norm.weight, self.norm.eps)) | |
| self._fused_add_norm = True | |
| except Exception: | |
| self._fused_add_norm = False | |
| if self._fused_add_norm: | |
| return self._decode_fused(h, cache, full_mask, swa_mask) | |
| for layer, c, layer_type in zip(self.layers, cache, self.layer_types): | |
| mask = full_mask if layer_type == "full_attention" else swa_mask | |
| h = layer(h, mask, c) | |
| return self.norm(h) | |
| class FlashHead(nn.Module): | |
| """Two-phase approximate lm_head for single-stream decode. | |
| Phase one scores quantized cluster centroids of the vocabulary; phase two | |
| computes exact logits only for the tokens of the top ``n_probes`` clusters | |
| (plus a fixed set of forced control tokens such as EOS). All other logits | |
| are -inf, so greedy decoding is exact whenever the true argmax lies in the | |
| probed clusters. Prefill and batched calls use the exact lm_head. | |
| Reference: FlashHead — Efficient Drop-in Replacement for the | |
| Classification Head in Language Model Inference. | |
| """ | |
| def __init__(self, args: ModelArgs): | |
| super().__init__() | |
| meta = args.flash_head | |
| n_clusters = meta["n_clusters"] | |
| cluster_size = meta["cluster_size"] | |
| # Default matches the converter's `--probes` default; every generated | |
| # checkpoint records the value explicitly. | |
| self.n_probes = min(meta.get("n_probes", 512), n_clusters) | |
| self.head_group_size = meta.get("head_group_size", 64) | |
| self.head_bits = meta.get("head_bits", 4) | |
| self.centroids = nn.QuantizedLinear( | |
| args.hidden_size, | |
| n_clusters, | |
| bias=False, | |
| group_size=meta.get("group_size", 64), | |
| bits=meta.get("bits", 4), | |
| ) | |
| self.token_map = mx.zeros((n_clusters, cluster_size), dtype=mx.int32) | |
| # Per-cluster max lm_head row norm. Centroids are directions; scaling | |
| # by the largest member norm upper-bounds the cluster's best logit so | |
| # high-frequency small-norm tokens are still probed. Newer checkpoints | |
| # fold the scale into the centroid rows at generation time. | |
| self.cluster_scale = mx.ones((n_clusters,), dtype=mx.bfloat16) | |
| self._scaled_centroids = bool(meta.get("scaled_centroids", False)) | |
| # mlx >= the indexed_qmv release computes the subset logits in one | |
| # dispatch straight from the flat head; older mlx uses a gather over | |
| # the cluster-ordered head copy. | |
| self._has_indexed_qmv = hasattr(mx.fast, "indexed_qmv") | |
| # Cluster-ordered copy of the quantized lm_head: subset logits are one | |
| # gather_qmm over the probed 32-row blocks, with no per-step gather. | |
| # It is a row-permutation of lm_head by token_map and nothing more, so | |
| # it is derived rather than stored: Model.sanitize rebuilds it at load | |
| # when this path is live. The indexed_qmv path never reads it, so on | |
| # those builds it is not allocated at all (~175 MB of the head saved). | |
| hidden = args.hidden_size | |
| self.head = ( | |
| {} | |
| if self._has_indexed_qmv | |
| else { | |
| "weight": mx.zeros( | |
| (n_clusters, cluster_size, hidden * self.head_bits // 32), | |
| dtype=mx.uint32, | |
| ), | |
| "scales": mx.zeros( | |
| (n_clusters, cluster_size, hidden // self.head_group_size), | |
| dtype=mx.bfloat16, | |
| ), | |
| "biases": mx.zeros( | |
| (n_clusters, cluster_size, hidden // self.head_group_size), | |
| dtype=mx.bfloat16, | |
| ), | |
| } | |
| ) | |
| self._force_ids = mx.array(meta.get("force_tokens", []), dtype=mx.int32) | |
| self._force_rows = None | |
| def __call__(self, h: mx.array, lm_head: nn.Module) -> mx.array: | |
| hv = h[:, -1, :] | |
| sims = self.centroids(hv) | |
| if not self._scaled_centroids: | |
| sims = sims * self.cluster_scale | |
| top = mx.argpartition(sims, kth=-self.n_probes, axis=-1)[ | |
| ..., -self.n_probes : | |
| ] # [1, n_probes] | |
| oids = self.token_map[top[0]].reshape(-1) | |
| if self._has_indexed_qmv: | |
| if self._force_ids.size: | |
| oids = mx.concatenate([oids, self._force_ids]) | |
| logits = mx.fast.indexed_qmv( | |
| hv[0], | |
| lm_head.weight, | |
| lm_head.scales, | |
| lm_head.biases, | |
| oids, | |
| group_size=lm_head.group_size, | |
| bits=lm_head.bits, | |
| ) | |
| vocab_size = lm_head.weight.shape[0] | |
| full = mx.full((1, 1, vocab_size), float("-inf"), dtype=logits.dtype) | |
| full[0, 0, oids] = logits | |
| return full | |
| logits = mx.gather_qmm( | |
| hv.reshape(1, 1, 1, 1, -1), | |
| self.head["weight"], | |
| self.head["scales"], | |
| self.head["biases"], | |
| rhs_indices=top[:, None, :], | |
| transpose=True, | |
| group_size=self.head_group_size, | |
| bits=self.head_bits, | |
| ).reshape(-1) | |
| if self._force_ids.size: | |
| if self._force_rows is None: | |
| self._force_rows = ( | |
| lm_head.weight[self._force_ids], | |
| lm_head.scales[self._force_ids], | |
| lm_head.biases[self._force_ids], | |
| ) | |
| mx.eval(*self._force_rows) | |
| fw, fs, fb = self._force_rows | |
| force_logits = mx.quantized_matmul( | |
| hv, | |
| fw, | |
| scales=fs, | |
| biases=fb, | |
| transpose=True, | |
| group_size=lm_head.group_size, | |
| bits=lm_head.bits, | |
| mode=getattr(lm_head, "mode", "affine"), | |
| )[0] | |
| oids = mx.concatenate([oids, self._force_ids]) | |
| logits = mx.concatenate([logits, force_logits]) | |
| vocab_size = lm_head.weight.shape[0] | |
| full = mx.full((1, 1, vocab_size), float("-inf"), dtype=logits.dtype) | |
| full[0, 0, oids] = logits | |
| return full | |
| class Model(nn.Module): | |
| def __init__(self, args: ModelArgs): | |
| super().__init__() | |
| self.args = args | |
| self.model_type = args.model_type | |
| self.model = MapleModel(args) | |
| if not args.tie_word_embeddings: | |
| self.lm_head = nn.Linear(args.hidden_size, args.vocab_size, bias=False) | |
| if (args.flash_head and args.use_flash_head and not args.tie_word_embeddings): | |
| self.lm_head_flash = FlashHead(args) | |
| else: | |
| self.lm_head_flash = None | |
| def __call__( | |
| self, | |
| inputs: mx.array, | |
| cache=None, | |
| ): | |
| out = self.model(inputs, cache) | |
| if self.args.tie_word_embeddings: | |
| return self.model.word_embeddings.as_linear(out) | |
| if ( | |
| self.lm_head_flash is not None | |
| and out.shape[0] == 1 | |
| and out.shape[1] == 1 | |
| and isinstance(self.lm_head, nn.QuantizedLinear) | |
| and getattr(self.lm_head, "mode", "affine") == "affine" | |
| ): | |
| return self.lm_head_flash(out, self.lm_head) | |
| return self.lm_head(out) | |
| def sanitize(self, weights): | |
| if self.args.tie_word_embeddings: | |
| # Drop the head entirely (weight + quantization scales/biases). | |
| weights = { | |
| k: v for k, v in weights.items() if not k.startswith("lm_head.") | |
| } | |
| # FlashHead disabled (e.g. model_config={"flash_head": None}): drop its | |
| # tensors so checkpoints that carry them still load. | |
| if self.lm_head_flash is None: | |
| weights = { | |
| k: v for k, v in weights.items() if not k.startswith("lm_head_flash.") | |
| } | |
| else: | |
| # `lm_head_flash.head.*` is lm_head permuted by token_map (see | |
| # mlx_lm.ternary.generate_flash_head), so it is pure redundancy on | |
| # disk. Checkpoints may ship it or omit it; reconcile both here. | |
| if self.lm_head_flash._has_indexed_qmv: | |
| # Dead on this build: indexed_qmv reads the flat lm_head. | |
| weights = { | |
| k: v | |
| for k, v in weights.items() | |
| if not k.startswith("lm_head_flash.head.") | |
| } | |
| elif "lm_head_flash.head.weight" not in weights: | |
| token_map = weights["lm_head_flash.token_map"] | |
| order = token_map.reshape(-1) | |
| for k in ("weight", "scales", "biases"): | |
| weights[f"lm_head_flash.head.{k}"] = weights[f"lm_head.{k}"][ | |
| order | |
| ].reshape(*token_map.shape, -1) | |
| # Ternary tensors carry one scale per output row, so checkpoints store | |
| # it once as `row_alpha` and omit biases entirely (bias == -scale). | |
| # Expand here so everything downstream — fusion below, and mlx's own | |
| # quantized kernels — sees the per-group layout. Checkpoints written | |
| # with `--group-scales` have no row_alpha and pass straight through. | |
| row_alpha_keys = [k for k in weights if k.endswith(".row_alpha")] | |
| if row_alpha_keys: | |
| group_size = (self.args.quantization or {}).get("group_size", 128) | |
| for key in row_alpha_keys: | |
| alpha = weights.pop(key) | |
| prefix = key[: -len(".row_alpha")] | |
| packed = weights.get(f"{prefix}.weight") | |
| if packed is None: | |
| continue | |
| # 2-bit packing stores 16 codes per uint32 word. | |
| n_groups = (packed.shape[-1] * 16) // group_size | |
| scales = mx.contiguous( | |
| mx.broadcast_to(alpha[..., None], (*alpha.shape, n_groups)) | |
| ) | |
| weights[f"{prefix}.scales"] = scales | |
| weights[f"{prefix}.biases"] = -scales | |
| # Stack per-expert weights from the Hugging Face layout into the | |
| # SwitchGLU layout. Already-converted checkpoints pass through. | |
| for l in range(self.args.num_hidden_layers): | |
| prefix = f"model.layers.{l}" | |
| for m in ["gate_proj", "down_proj", "up_proj"]: | |
| for k in ["weight", "scales", "biases", "bias"]: | |
| if f"{prefix}.mlp.experts.0.{m}.{k}" in weights: | |
| to_join = [ | |
| weights.pop(f"{prefix}.mlp.experts.{e}.{m}.{k}") | |
| for e in range(self.args.num_experts) | |
| ] | |
| weights[f"{prefix}.mlp.switch_mlp.{m}.{k}"] = mx.stack(to_join) | |
| # Fuse split projections: q/k/v -> qkv_proj (rows), MoE up/gate -> | |
| # up_gate_proj (per-expert rows). Row-wise quantized tensors | |
| # (weight/scales/biases) concatenate losslessly along the output | |
| # axis. | |
| for suffix in ["weight", "scales", "biases", "bias"]: | |
| qkv = [f"{prefix}.self_attn.{p}.{suffix}" for p in ("q_proj", "k_proj", "v_proj")] | |
| if qkv[0] in weights: | |
| weights[f"{prefix}.self_attn.qkv_proj.{suffix}"] = mx.concatenate( | |
| [weights.pop(k) for k in qkv], axis=0 | |
| ) | |
| up = f"{prefix}.mlp.switch_mlp.up_proj.{suffix}" | |
| gate = f"{prefix}.mlp.switch_mlp.gate_proj.{suffix}" | |
| if up in weights: | |
| weights[f"{prefix}.mlp.switch_mlp.up_gate_proj.{suffix}"] = ( | |
| mx.concatenate([weights.pop(up), weights.pop(gate)], axis=1) | |
| ) | |
| return weights | |
| def make_cache(self): | |
| caches = [] | |
| for layer_type in self.model.layer_types: | |
| if layer_type == "sliding_attention": | |
| caches.append(RotatingKVCache(max_size=self.args.sliding_window)) | |
| else: | |
| caches.append(KVCache()) | |
| return caches | |
| def layers(self): | |
| return self.model.layers | |
| def quant_predicate(self): | |
| def predicate(path, _): | |
| if path.endswith("lm_head") or "word_embeddings" in path: | |
| return {"group_size": 64, "bits": 4} | |
| return True | |
| return predicate | |