From dfc37a06c5819467daadac3ca16b500b062c87f1 Mon Sep 17 00:00:00 2001 From: Thump604 Date: Fri, 24 Apr 2026 00:17:44 -0500 Subject: [PATCH 1/9] Fix DeepSeek V4 MLX conversion support --- mlx_lm/convert.py | 16 ++- mlx_lm/models/deepseek_v4.py | 168 ++++++++++++++++++++++--- tests/test_models.py | 231 +++++++++++++++++++++++++++++++++++ 3 files changed, 395 insertions(+), 20 deletions(-) diff --git a/mlx_lm/convert.py b/mlx_lm/convert.py index ab3fc62ac..6e4f51f0c 100644 --- a/mlx_lm/convert.py +++ b/mlx_lm/convert.py @@ -63,14 +63,26 @@ def mixed_quant_predicate( or index >= 7 * num_layers // 8 or (index - num_layers // 8) % 3 == 2 ) + always_more_bits = ( + "lm_head" in path + or "embed_tokens" in path + or "wq_a" in path + or "wq_b" in path + or "wkv" in path + or "wo_a" in path + or "wo_b" in path + or "compressor" in path + or "indexer" in path + or "shared_experts" in path + ) + if always_more_bits: + return {"group_size": group_size, "bits": high_bits, "mode": mode} if ( "v_proj" in path or "v_a_proj" in path or "v_b_proj" in path ) and use_more_bits: return {"group_size": group_size, "bits": high_bits, "mode": mode} if "down_proj" in path and use_more_bits: return {"group_size": group_size, "bits": high_bits, "mode": mode} - if "lm_head" in path: - return {"group_size": group_size, "bits": high_bits, "mode": mode} return {"group_size": group_size, "bits": low_bits, "mode": mode} diff --git a/mlx_lm/models/deepseek_v4.py b/mlx_lm/models/deepseek_v4.py index f71520fbe..aae0ac912 100644 --- a/mlx_lm/models/deepseek_v4.py +++ b/mlx_lm/models/deepseek_v4.py @@ -8,6 +8,7 @@ # # Reference: deepseek-ai/DeepSeek-V4 (Apr 2026). mHC: arXiv:2512.24880. +import math from dataclasses import dataclass, field from typing import Dict, List, Optional @@ -18,7 +19,6 @@ from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .cache import KVCache, RotatingKVCache from .pipeline import PipelineMixin -from .rope_utils import initialize_rope from .switch_layers import SwitchGLU @@ -81,6 +81,84 @@ class ModelArgs(BaseModelArgs): quantization_config: Optional[Dict] = None +class DeepseekV4RoPE(nn.Module): + """DeepSeek-V4 rotary embedding. + + The reference implementation applies RoPE to the KV tensor before attention + and applies the conjugate rotation to the attention output. The generic MLX + RoPE layers do not expose an inverse path, so keep the small DeepSeek-specific + implementation here. + """ + + def __init__( + self, + dims: int, + base: float, + scaling_config: Optional[Dict] = None, + ): + super().__init__() + self.dims = dims + + inv_freq = 1.0 / (base ** (mx.arange(0, dims, 2, dtype=mx.float32) / dims)) + rope_type = None + if scaling_config is not None: + rope_type = scaling_config.get("type") or scaling_config.get("rope_type") + + if rope_type in ("yarn", "deepseek_yarn"): + factor = scaling_config["factor"] + original_max_position_embeddings = scaling_config[ + "original_max_position_embeddings" + ] + beta_fast = scaling_config.get("beta_fast", 32) + beta_slow = scaling_config.get("beta_slow", 1) + + def correction_dim(num_rotations): + return ( + dims + * math.log( + original_max_position_embeddings + / (num_rotations * 2 * math.pi) + ) + / (2 * math.log(base)) + ) + + low = math.floor(correction_dim(beta_fast)) + high = math.ceil(correction_dim(beta_slow)) + low = max(low, 0) + high = min(high, dims - 1) + if low == high: + high += 0.001 + + ramp = (mx.arange(dims // 2, dtype=mx.float32) - low) / (high - low) + smooth = 1 - mx.clip(ramp, 0, 1) + inv_freq = inv_freq / factor * (1 - smooth) + inv_freq * smooth + elif rope_type not in (None, "default", "linear"): + raise ValueError(f"Unsupported DeepSeek-V4 RoPE type {rope_type}") + + self.inv_freq = inv_freq + + def __call__(self, x: mx.array, offset: int = 0, inverse: bool = False): + dtype = x.dtype + T = x.shape[-2] + pos = mx.arange(offset, offset + T, dtype=mx.float32) + theta = pos[:, None] * self.inv_freq[None, :] + if inverse: + theta = -theta + + broadcast_shape = (1,) * (x.ndim - 2) + theta.shape + cos = mx.cos(theta).reshape(broadcast_shape).astype(dtype) + sin = mx.sin(theta).reshape(broadcast_shape).astype(dtype) + + rot = x[..., : self.dims].reshape(*x.shape[:-1], self.dims // 2, 2) + x0 = rot[..., 0] + x1 = rot[..., 1] + y = mx.stack((x0 * cos - x1 * sin, x0 * sin + x1 * cos), axis=-1) + y = y.reshape(*x.shape[:-1], self.dims) + if x.shape[-1] == self.dims: + return y + return mx.concatenate([y, x[..., self.dims :]], axis=-1) + + # --------------------------------------------------------------------------- # # mHC (Manifold-constrained Hyper-Connections) # # --------------------------------------------------------------------------- # @@ -349,13 +427,24 @@ def __init__(self, args: ModelArgs, compress_ratio: int, head_dim: int): self.head_dim = head_dim self.rope_head_dim = args.qk_rope_head_dim self.ratio = compress_ratio - self.wkv = nn.Linear(self.dim, head_dim, bias=False) - self.wgate = nn.Linear(self.dim, head_dim, bias=False) - self.ape = mx.zeros((compress_ratio, head_dim), dtype=mx.float32) + self.overlap = compress_ratio == 4 + out_dim = head_dim * (2 if self.overlap else 1) + self.wkv = nn.Linear(self.dim, out_dim, bias=False) + self.wgate = nn.Linear(self.dim, out_dim, bias=False) + self.ape = mx.zeros((compress_ratio, out_dim), dtype=mx.float32) self.norm = nn.RMSNorm(head_dim, eps=args.rms_norm_eps) + def _overlap_transform(self, tensor: mx.array, value: float) -> mx.array: + B, S, R, _ = tensor.shape + D = self.head_dim + out = mx.full((B, S, 2 * R, D), value, dtype=tensor.dtype) + out[:, :, R:] = tensor[:, :, :, D:] + out[:, 1:, :R] = tensor[:, :-1, :, :D] + return out + def __call__(self, x: mx.array) -> mx.array: - # Prefill-only MVP: chunk x into non-overlapping windows of `ratio` tokens. + # Prefill-only MVP: chunk x into windows of `ratio` tokens. Ratio-4 + # layers use the overlapping layout from the reference implementation. # Returns compressed KV: [B, S//ratio, head_dim] (bf16). B, S, _ = x.shape r = self.ratio @@ -363,8 +452,11 @@ def __call__(self, x: mx.array) -> mx.array: if keep == 0: return mx.zeros((B, 0, self.head_dim), dtype=x.dtype) xf = x[:, :keep].astype(mx.float32) - kv = self.wkv(xf).reshape(B, keep // r, r, self.head_dim) - score = self.wgate(xf).reshape(B, keep // r, r, self.head_dim) + self.ape + kv = self.wkv(xf).reshape(B, keep // r, r, -1) + score = self.wgate(xf).reshape(B, keep // r, r, -1) + self.ape + if self.overlap: + kv = self._overlap_transform(kv, 0.0) + score = self._overlap_transform(score, float("-inf")) weights = mx.softmax(score, axis=2, precise=True) kv = (kv * weights).sum(axis=2) return self.norm(kv.astype(x.dtype)) @@ -427,20 +519,16 @@ def __init__(self, args: ModelArgs, layer_idx: int): self.wo_a = nn.Linear(group_feat, self.n_groups * self.o_lora_rank, bias=False) self.wo_b = nn.Linear(self.n_groups * self.o_lora_rank, self.dim, bias=args.attention_bias) - # rope (sliding layers use base theta; compressed layers use YaRN + compress_rope_theta) + # RoPE: sliding layers use base theta; compressed layers use YaRN with + # compress_rope_theta. DeepSeek-V4 also inverse-rotates the attention + # output rope dims after sparse attention. if self.compress_ratio: base = args.compress_rope_theta scaling = args.rope_scaling else: base = args.rope_theta scaling = None - self.rope = initialize_rope( - dims=self.rope_head_dim, - base=base, - traditional=True, - max_position_embeddings=args.max_position_embeddings, - scaling_config=scaling, - ) + self.rope = DeepseekV4RoPE(self.rope_head_dim, base, scaling) # Compressor / Indexer — present only when compress_ratio > 0 if self.compress_ratio: @@ -477,9 +565,19 @@ def __call__(self, x: mx.array, mask=None, cache=None): # Standard SDPA (compressed KV + topk deferred to v0.2) out = scaled_dot_product_attention( - q, k, v, cache=cache, scale=self.scale, mask=mask, + q, + k, + v, + cache=cache, + scale=self.scale, + mask=mask, + sinks=self.attn_sink, ) + out_nope, out_pe = mx.split(out, [self.nope_head_dim], axis=-1) + out_pe = self.rope(out_pe, offset=offset, inverse=True) + out = mx.concatenate([out_nope, out_pe], axis=-1) + # Grouped low-rank projection: [B, n_heads, S, head_dim] -> [B, S, n_heads*head_dim] out = out.transpose(0, 2, 1, 3).reshape(B, S, self.n_heads * self.head_dim) # Split into o_groups along the head_dim*n_heads axis; apply per-group wo_a via einsum. @@ -691,9 +789,18 @@ def sanitize(self, weights: Dict[str, mx.array]) -> Dict[str, mx.array]: new[k] = v weights = new - # 2) FP8 block dequant: `X.weight` + `X.scale` -> dequantized bf16 `X.weight` + def _scale_to_float(scale: mx.array) -> mx.array: + if scale.dtype == mx.uint8: + return mx.exp2(scale.astype(mx.float32) - 127.0) + return scale.astype(mx.float32) + + # 2) FP8/FP4 block dequant: + # `X.weight` + `X.scale` -> dequantized bf16 `X.weight` + # Routed experts in Flash are FP4-packed int8; other scaled matrices + # are FP8 e4m3 with 128x128 block scales. def _dequant_fp8_block(weight: mx.array, scale: mx.array, bs: int = 128) -> mx.array: weight = mx.from_fp8(weight, dtype=mx.bfloat16) + scale = _scale_to_float(scale) m, n = weight.shape pad_b = (-m) % bs pad_s = (-n) % bs @@ -702,11 +809,36 @@ def _dequant_fp8_block(weight: mx.array, scale: mx.array, bs: int = 128) -> mx.a weight = (weight * scale[:, None, :, None]).reshape(m + pad_b, n + pad_s) return weight[:m, :n].astype(mx.bfloat16) + def _dequant_fp4_block(weight: mx.array, scale: mx.array, bs: int = 32) -> mx.array: + table = mx.array( + [ + 0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0, + 0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0, + ], + dtype=mx.float32, + ) + packed = weight.astype(mx.uint8) + low = packed & 0x0F + high = (packed >> 4) & 0x0F + unpacked = mx.stack([mx.take(table, low), mx.take(table, high)], axis=-1) + unpacked = unpacked.reshape(weight.shape[0], weight.shape[1] * 2) + scale = mx.repeat(_scale_to_float(scale), bs, axis=-1) + return (unpacked * scale).astype(mx.bfloat16) + new = {} for k, v in weights.items(): if k.endswith(".scale"): wk = k[:-len(".scale")] + ".weight" - if wk in weights and weights[wk].dtype in (mx.uint8,): + weight = weights.get(wk) + if ( + weight is not None + and ".ffn.experts." in wk + and "shared_experts" not in wk + and weight.dtype in (mx.int8, mx.uint8) + and v.shape[-1] * 16 == weight.shape[-1] + ): + new[wk] = _dequant_fp4_block(weight, v) + elif weight is not None and weight.dtype in (mx.uint8,): new[wk] = _dequant_fp8_block(weights[wk], v) else: new[k] = v diff --git a/tests/test_models.py b/tests/test_models.py index 6e1fcd96e..51df81f84 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -1,6 +1,7 @@ # Copyright © 2024 Apple Inc. import copy import importlib +import math import unittest import mlx.core as mx @@ -1422,6 +1423,236 @@ def test_deepseek_v3(self): model, args.model_type, args.vocab_size, args.num_hidden_layers ) + def test_deepseek_v4_rope_inverse(self): + from mlx_lm.models.deepseek_v4 import DeepseekV4RoPE + + scaling = { + "type": "yarn", + "factor": 16, + "original_max_position_embeddings": 65536, + "beta_fast": 32, + "beta_slow": 1, + } + rope = DeepseekV4RoPE(8, 160000, scaling) + x = mx.random.uniform(shape=(1, 2, 4, 8)) + + y = rope(x, offset=3) + z = rope(y, offset=3, inverse=True) + self.assertTrue(mx.allclose(x, z, rtol=1e-5, atol=1e-5)) + + inv_freq = 1.0 / (160000 ** (mx.arange(0, 8, 2, dtype=mx.float32) / 8)) + + def correction_dim(num_rotations): + return ( + 8 + * math.log(65536 / (num_rotations * 2 * math.pi)) + / (2 * math.log(160000)) + ) + + low = max(math.floor(correction_dim(32)), 0) + high = min(math.ceil(correction_dim(1)), 7) + if low == high: + high += 0.001 + ramp = (mx.arange(4, dtype=mx.float32) - low) / (high - low) + smooth = 1 - mx.clip(ramp, 0, 1) + inv_freq = inv_freq / 16 * (1 - smooth) + inv_freq * smooth + + theta = mx.arange(3, 7, dtype=mx.float32)[:, None] * inv_freq[None, :] + cos = mx.cos(theta).reshape(1, 1, 4, 4) + sin = mx.sin(theta).reshape(1, 1, 4, 4) + rot = x.reshape(1, 2, 4, 4, 2) + expected = mx.stack( + ( + rot[..., 0] * cos - rot[..., 1] * sin, + rot[..., 0] * sin + rot[..., 1] * cos, + ), + axis=-1, + ).reshape(1, 2, 4, 8) + self.assertTrue(mx.allclose(y, expected, rtol=1e-5, atol=1e-5)) + + def test_deepseek_v4(self): + from mlx_lm.models import deepseek_v4 + + args = deepseek_v4.ModelArgs( + model_type="deepseek_v4", + vocab_size=1024, + hidden_size=128, + num_hidden_layers=4, + num_attention_heads=4, + num_key_value_heads=1, + q_lora_rank=32, + o_lora_rank=16, + o_groups=2, + head_dim=32, + qk_rope_head_dim=8, + sliding_window=16, + compress_ratios=[0, 0, 4, 0], + index_n_heads=4, + index_head_dim=16, + index_topk=8, + moe_intermediate_size=32, + n_routed_experts=4, + n_shared_experts=1, + num_experts_per_tok=2, + num_hash_layers=1, + hc_mult=2, + hc_sinkhorn_iters=2, + max_position_embeddings=256, + rope_scaling={ + "beta_fast": 32, + "beta_slow": 1, + "factor": 2, + "original_max_position_embeddings": 128, + "type": "yarn", + }, + ) + model = deepseek_v4.Model(args) + self.assertEqual(len(model.layers), args.num_hidden_layers) + self.assertEqual(model.model_type, args.model_type) + self.assertEqual( + model.layers[2].attn.compressor.wkv.weight.shape, + (2 * args.head_dim, args.hidden_size), + ) + self.assertEqual( + model.layers[2].attn.indexer.compressor.wkv.weight.shape, + (2 * args.index_head_dim, args.hidden_size), + ) + + for dtype in [mx.float32, mx.float16]: + model.update( + tree_map( + lambda p: p.astype(dtype) + if mx.issubdtype(p.dtype, mx.floating) + else p, + model.parameters(), + ) + ) + + inputs = mx.array([[0, 1, 2, 3, 4]], dtype=mx.int32) + outputs = model(inputs) + self.assertEqual(outputs.shape, (1, 5, args.vocab_size)) + self.assertEqual(outputs.dtype, dtype) + + cache = model.make_cache() + self.assertIsInstance(cache[0], RotatingKVCache) + self.assertIsInstance(cache[2], KVCache) + outputs = model(inputs[:, :3], cache=cache) + self.assertEqual(outputs.shape, (1, 3, args.vocab_size)) + self.assertEqual(outputs.dtype, dtype) + outputs = model(inputs[:, 3:4], cache=cache) + self.assertEqual(outputs.shape, (1, 1, args.vocab_size)) + self.assertEqual(outputs.dtype, dtype) + + def test_mixed_quant_preserves_deepseek_v4_attention_paths(self): + from mlx_lm.convert import mixed_quant_predicate_builder + from mlx_lm.models import deepseek_v4 + + args = deepseek_v4.ModelArgs( + model_type="deepseek_v4", + vocab_size=128, + hidden_size=64, + num_hidden_layers=4, + num_attention_heads=4, + q_lora_rank=16, + o_lora_rank=8, + o_groups=2, + head_dim=16, + qk_rope_head_dim=4, + sliding_window=16, + compress_ratios=[0, 0, 4, 0], + index_n_heads=4, + index_head_dim=8, + index_topk=4, + moe_intermediate_size=16, + n_routed_experts=4, + n_shared_experts=1, + num_experts_per_tok=2, + num_hash_layers=1, + hc_mult=2, + hc_sinkhorn_iters=2, + ) + model = deepseek_v4.Model(args) + modules = dict(model.named_modules()) + predicate = mixed_quant_predicate_builder("mixed_3_6", model, group_size=32) + + high = {"group_size": 32, "bits": 6, "mode": "affine"} + low = {"group_size": 32, "bits": 3, "mode": "affine"} + for path in [ + "model.layers.0.attn.wq_a", + "model.layers.0.attn.wq_b", + "model.layers.0.attn.wkv", + "model.layers.0.attn.wo_a", + "model.layers.0.attn.wo_b", + "model.layers.2.attn.compressor.wkv", + "model.layers.2.attn.indexer.wq_b", + "model.layers.0.ffn.shared_experts.down_proj", + "model.embed_tokens", + "lm_head", + ]: + self.assertEqual(predicate(path, modules[path]), high) + + self.assertEqual( + predicate( + "model.layers.0.ffn.switch_mlp.gate_proj", + modules["model.layers.0.ffn.switch_mlp.gate_proj"], + ), + low, + ) + + def test_deepseek_v4_sanitize_unpacks_fp4_experts(self): + from mlx_lm.models import deepseek_v4 + + args = deepseek_v4.ModelArgs( + model_type="deepseek_v4", + vocab_size=128, + hidden_size=32, + num_hidden_layers=1, + num_attention_heads=4, + q_lora_rank=16, + o_lora_rank=8, + o_groups=2, + head_dim=16, + qk_rope_head_dim=4, + moe_intermediate_size=2, + n_routed_experts=2, + n_shared_experts=1, + num_experts_per_tok=1, + hc_mult=2, + hc_sinkhorn_iters=2, + ) + model = deepseek_v4.Model(args) + + packed = mx.array( + [ + [0x21] * 16, + [0xFE] * 16, + ], + dtype=mx.int8, + ) + weights = { + "layers.0.ffn.experts.0.w1.weight": packed, + "layers.0.ffn.experts.0.w1.scale": mx.ones((2, 1), dtype=mx.float32), + "layers.0.ffn.experts.1.w1.weight": packed, + "layers.0.ffn.experts.1.w1.scale": mx.ones((2, 1), dtype=mx.float32), + } + + converted = model.sanitize(weights) + key = "model.layers.0.ffn.switch_mlp.gate_proj.weight" + self.assertIn(key, converted) + self.assertEqual(converted[key].shape, (2, 2, 32)) + self.assertTrue( + mx.array_equal( + converted[key][0, 0, :4].astype(mx.float32), + mx.array([0.5, 1.0, 0.5, 1.0], dtype=mx.float32), + ) + ) + self.assertTrue( + mx.array_equal( + converted[key][0, 1, :4].astype(mx.float32), + mx.array([-4.0, -6.0, -4.0, -6.0], dtype=mx.float32), + ) + ) + def test_gemma2(self): from mlx_lm.models import gemma2 From 874af5595f33833b9afd5341678e4713d0fa54f0 Mon Sep 17 00:00:00 2001 From: Thump604 Date: Fri, 24 Apr 2026 00:20:11 -0500 Subject: [PATCH 2/9] Handle DeepSeek V4 FP4 scale decoding --- mlx_lm/models/deepseek_v4.py | 2 +- tests/test_models.py | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/mlx_lm/models/deepseek_v4.py b/mlx_lm/models/deepseek_v4.py index aae0ac912..faaeb4ea0 100644 --- a/mlx_lm/models/deepseek_v4.py +++ b/mlx_lm/models/deepseek_v4.py @@ -791,7 +791,7 @@ def sanitize(self, weights: Dict[str, mx.array]) -> Dict[str, mx.array]: def _scale_to_float(scale: mx.array) -> mx.array: if scale.dtype == mx.uint8: - return mx.exp2(scale.astype(mx.float32) - 127.0) + return mx.exp((scale.astype(mx.float32) - 127.0) * math.log(2.0)) return scale.astype(mx.float32) # 2) FP8/FP4 block dequant: diff --git a/tests/test_models.py b/tests/test_models.py index 51df81f84..545ed700d 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -1631,9 +1631,9 @@ def test_deepseek_v4_sanitize_unpacks_fp4_experts(self): ) weights = { "layers.0.ffn.experts.0.w1.weight": packed, - "layers.0.ffn.experts.0.w1.scale": mx.ones((2, 1), dtype=mx.float32), + "layers.0.ffn.experts.0.w1.scale": mx.full((2, 1), 127, dtype=mx.uint8), "layers.0.ffn.experts.1.w1.weight": packed, - "layers.0.ffn.experts.1.w1.scale": mx.ones((2, 1), dtype=mx.float32), + "layers.0.ffn.experts.1.w1.scale": mx.full((2, 1), 127, dtype=mx.uint8), } converted = model.sanitize(weights) From 38520a172172dd8dd32b999b1f702ebe50395c38 Mon Sep 17 00:00:00 2001 From: Thump604 Date: Fri, 24 Apr 2026 00:21:02 -0500 Subject: [PATCH 3/9] Cover DeepSeek V4 FP8 block dequantization --- tests/test_models.py | 40 ++++++++++++++++++++++++++++++++++++++++ 1 file changed, 40 insertions(+) diff --git a/tests/test_models.py b/tests/test_models.py index 545ed700d..c17c100e7 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -1653,6 +1653,46 @@ def test_deepseek_v4_sanitize_unpacks_fp4_experts(self): ) ) + def test_deepseek_v4_sanitize_dequantizes_fp8_blocks(self): + from mlx_lm.models import deepseek_v4 + + args = deepseek_v4.ModelArgs( + model_type="deepseek_v4", + vocab_size=128, + hidden_size=32, + num_hidden_layers=1, + num_attention_heads=4, + q_lora_rank=16, + o_lora_rank=8, + o_groups=2, + head_dim=16, + qk_rope_head_dim=4, + moe_intermediate_size=2, + n_routed_experts=2, + n_shared_experts=1, + num_experts_per_tok=1, + hc_mult=2, + hc_sinkhorn_iters=2, + ) + model = deepseek_v4.Model(args) + weight = mx.to_fp8(mx.ones((128, 128), dtype=mx.float32)) + converted = model.sanitize( + { + "layers.0.attn.wkv.weight": weight, + "layers.0.attn.wkv.scale": mx.full((1, 1), 127, dtype=mx.uint8), + } + ) + key = "model.layers.0.attn.wkv.weight" + self.assertIn(key, converted) + self.assertTrue( + mx.allclose( + converted[key].astype(mx.float32), + mx.ones((128, 128), dtype=mx.float32), + rtol=1e-5, + atol=1e-5, + ) + ) + def test_gemma2(self): from mlx_lm.models import gemma2 From c9dc7325d7182bad971f8ca2e2e823f42990e2d2 Mon Sep 17 00:00:00 2001 From: Thump604 Date: Fri, 24 Apr 2026 01:00:16 -0500 Subject: [PATCH 4/9] Load DeepSeek V4 E8M0 scale metadata --- mlx_lm/utils.py | 61 +++++++++++++++++++++++++++++++++++++++++++- tests/test_models.py | 48 ++++++++++++++++++++++++++++++++++ 2 files changed, 108 insertions(+), 1 deletion(-) diff --git a/mlx_lm/utils.py b/mlx_lm/utils.py index ef3d266b9..4d1760183 100644 --- a/mlx_lm/utils.py +++ b/mlx_lm/utils.py @@ -8,6 +8,8 @@ import os import resource import shutil +import struct +import warnings from pathlib import Path from textwrap import dedent from typing import ( @@ -279,6 +281,62 @@ def load_config(model_path: Path) -> dict: return config +def _reinterpret_safetensor_e8m0_scales_as_uint8(path: str) -> bool: + """Rewrite safetensors E8M0 scale metadata to U8 in-place. + + DeepSeek-V4 stores FP8/FP4 block scales as float8_e8m0fnu. The payload is + one byte per element, and the model sanitizer decodes those exponent bytes. + MLX currently rejects the safetensors dtype before the sanitizer can run, so + reinterpret the header as uint8 while leaving tensor bytes untouched. + """ + with open(path, "r+b") as f: + header_len = struct.unpack(" header_len: + raise RuntimeError( + f"Cannot reinterpret F8_E8M0 safetensors header in {path}: " + "rewritten header is larger than original header." + ) + + f.seek(8) + f.write(new_header) + f.write(b" " * (header_len - len(new_header))) + + return True + + +def _load_safetensors(path: str, *, allow_e8m0_uint8: bool = False) -> dict: + try: + return mx.load(path) + except RuntimeError as e: + if "F8_E8M0" not in str(e) or not allow_e8m0_uint8: + raise + + if _reinterpret_safetensor_e8m0_scales_as_uint8(path): + warnings.warn( + f"Reinterpreted F8_E8M0 scale metadata as uint8 in {path}. " + "Tensor bytes were not changed.", + RuntimeWarning, + stacklevel=2, + ) + return mx.load(path) + + def load_model( model_path: Path, lazy: bool = False, @@ -319,8 +377,9 @@ def load_model( raise FileNotFoundError(f"No safetensors found in {model_path}") weights = {} + allow_e8m0_uint8 = config.get("model_type") == "deepseek_v4" for wf in weight_files: - weights.update(mx.load(wf)) + weights.update(_load_safetensors(wf, allow_e8m0_uint8=allow_e8m0_uint8)) if (model_file := config.get("model_file")) is not None: spec = importlib.util.spec_from_file_location( diff --git a/tests/test_models.py b/tests/test_models.py index c17c100e7..fcb2cf2e1 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -1693,6 +1693,54 @@ def test_deepseek_v4_sanitize_dequantizes_fp8_blocks(self): ) ) + def test_deepseek_v4_loads_e8m0_scales_as_uint8(self): + try: + import tempfile + from pathlib import Path + + import torch + from safetensors.torch import save_file + except ImportError: + self.skipTest("torch and safetensors are required for this test") + + if not hasattr(torch, "float8_e4m3fn") or not hasattr(torch, "float8_e8m0fnu"): + self.skipTest("torch build does not expose required float8 dtypes") + + from mlx_lm.utils import _load_safetensors + + with tempfile.TemporaryDirectory() as tmpdir: + path = Path(tmpdir) / "model.safetensors" + save_file( + { + "weight": torch.tensor([1.0, -2.0], dtype=torch.float32).to( + torch.float8_e4m3fn + ), + "scale": torch.tensor([[1.0, 2.0]], dtype=torch.float32).to( + torch.float8_e8m0fnu + ), + }, + str(path), + ) + + with self.assertRaisesRegex(RuntimeError, "F8_E8M0"): + mx.load(str(path)) + + loaded = _load_safetensors(str(path), allow_e8m0_uint8=True) + self.assertEqual(loaded["scale"].dtype, mx.uint8) + self.assertEqual(loaded["weight"].dtype, mx.uint8) + self.assertTrue( + mx.array_equal( + loaded["scale"], + mx.array([[127, 128]], dtype=mx.uint8), + ) + ) + self.assertTrue( + mx.allclose( + mx.from_fp8(loaded["weight"], dtype=mx.float32), + mx.array([1.0, -2.0], dtype=mx.float32), + ) + ) + def test_gemma2(self): from mlx_lm.models import gemma2 From 6091a3764621f2ff90ec790a2193c338f6fe88c0 Mon Sep 17 00:00:00 2001 From: Thump604 Date: Fri, 24 Apr 2026 01:01:39 -0500 Subject: [PATCH 5/9] Map DeepSeek V4 final norm weight --- mlx_lm/models/deepseek_v4.py | 1 + tests/test_models.py | 2 ++ 2 files changed, 3 insertions(+) diff --git a/mlx_lm/models/deepseek_v4.py b/mlx_lm/models/deepseek_v4.py index faaeb4ea0..9be1cba64 100644 --- a/mlx_lm/models/deepseek_v4.py +++ b/mlx_lm/models/deepseek_v4.py @@ -849,6 +849,7 @@ def _dequant_fp4_block(weight: mx.array, scale: mx.array, bs: int = 32) -> mx.ar # 3) Remap top-level names to our module structure top_remap = { "embed.weight": "model.embed_tokens.weight", + "norm.weight": "model.norm.weight", "head.weight": "lm_head.weight", "hc_head_fn": "model.hc_head.fn", "hc_head_base": "model.hc_head.base", diff --git a/tests/test_models.py b/tests/test_models.py index fcb2cf2e1..ff46b960b 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -1680,10 +1680,12 @@ def test_deepseek_v4_sanitize_dequantizes_fp8_blocks(self): { "layers.0.attn.wkv.weight": weight, "layers.0.attn.wkv.scale": mx.full((1, 1), 127, dtype=mx.uint8), + "norm.weight": mx.ones((32,), dtype=mx.float32), } ) key = "model.layers.0.attn.wkv.weight" self.assertIn(key, converted) + self.assertIn("model.norm.weight", converted) self.assertTrue( mx.allclose( converted[key].astype(mx.float32), From 3ce16ddbf0130ea438496499366459eab53fe558 Mon Sep 17 00:00:00 2001 From: Thump604 Date: Fri, 24 Apr 2026 01:02:59 -0500 Subject: [PATCH 6/9] Keep DeepSeek V4 RoPE frequencies derived --- mlx_lm/models/deepseek_v4.py | 7 ++++++- tests/test_models.py | 2 ++ 2 files changed, 8 insertions(+), 1 deletion(-) diff --git a/mlx_lm/models/deepseek_v4.py b/mlx_lm/models/deepseek_v4.py index 9be1cba64..28a583b21 100644 --- a/mlx_lm/models/deepseek_v4.py +++ b/mlx_lm/models/deepseek_v4.py @@ -135,7 +135,12 @@ def correction_dim(num_rotations): elif rope_type not in (None, "default", "linear"): raise ValueError(f"Unsupported DeepSeek-V4 RoPE type {rope_type}") - self.inv_freq = inv_freq + # This is derived from config, not a checkpoint parameter. + self._inv_freq = (inv_freq,) + + @property + def inv_freq(self): + return self._inv_freq[0] def __call__(self, x: mx.array, offset: int = 0, inverse: bool = False): dtype = x.dtype diff --git a/tests/test_models.py b/tests/test_models.py index ff46b960b..3df027b8d 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -1509,6 +1509,8 @@ def test_deepseek_v4(self): model = deepseek_v4.Model(args) self.assertEqual(len(model.layers), args.num_hidden_layers) self.assertEqual(model.model_type, args.model_type) + parameter_names = {name for name, _ in tree_flatten(model.parameters())} + self.assertNotIn("model.layers.0.attn.rope.inv_freq", parameter_names) self.assertEqual( model.layers[2].attn.compressor.wkv.weight.shape, (2 * args.head_dim, args.hidden_size), From ab196c4225e6b87ca78d723773dccbec71d0cb47 Mon Sep 17 00:00:00 2001 From: Thump604 Date: Fri, 24 Apr 2026 01:05:20 -0500 Subject: [PATCH 7/9] Load tokenizers with unknown model configs --- mlx_lm/tokenizer_utils.py | 29 +++++++++++++++++++++++++---- tests/test_tokenizers.py | 38 ++++++++++++++++++++++++++++++++++++++ 2 files changed, 63 insertions(+), 4 deletions(-) diff --git a/mlx_lm/tokenizer_utils.py b/mlx_lm/tokenizer_utils.py index c7e50fbe7..f4ba58706 100644 --- a/mlx_lm/tokenizer_utils.py +++ b/mlx_lm/tokenizer_utils.py @@ -5,7 +5,7 @@ from json import JSONDecodeError from typing import Any, Dict, List, Optional -from transformers import AutoTokenizer, PreTrainedTokenizerFast +from transformers import AutoTokenizer, PreTrainedConfig, PreTrainedTokenizerFast class StreamingDetokenizer: @@ -611,9 +611,30 @@ def load( tokenizer_config_file = model_path / "tokenizer_config.json" chat_template = None - tokenizer = AutoTokenizer.from_pretrained( - model_path, **(tokenizer_config_extra or {}) - ) + tokenizer_config_extra = tokenizer_config_extra or {} + try: + tokenizer = AutoTokenizer.from_pretrained(model_path, **tokenizer_config_extra) + except (AttributeError, ValueError) as e: + message = str(e) + if ( + "config" in tokenizer_config_extra + or ( + "deepseek_v4" not in message + and "max_position_embeddings" not in message + ) + ): + raise + warnings.warn( + "Falling back to generic tokenizer config because Transformers does " + f"not recognize this model config: {e}", + RuntimeWarning, + stacklevel=2, + ) + tokenizer = AutoTokenizer.from_pretrained( + model_path, + config=PreTrainedConfig(), + **tokenizer_config_extra, + ) tokenizer_config = tokenizer.init_kwargs diff --git a/tests/test_tokenizers.py b/tests/test_tokenizers.py index 54906af1c..e18be957d 100644 --- a/tests/test_tokenizers.py +++ b/tests/test_tokenizers.py @@ -2,13 +2,17 @@ import unittest from pathlib import Path +from tempfile import TemporaryDirectory +from unittest.mock import patch from huggingface_hub import snapshot_download +from transformers import PreTrainedConfig from mlx_lm.tokenizer_utils import ( BPEStreamingDetokenizer, NaiveStreamingDetokenizer, SPMStreamingDetokenizer, + load as load_tokenizer_impl, ) from mlx_lm.utils import load_tokenizer @@ -109,6 +113,40 @@ def test_thinking(self): self.assertIsNone(tokenizer.think_start_id) self.assertIsNone(tokenizer.think_end_id) + def test_unknown_model_config_tokenizer_fallback(self): + class MockTokenizer: + eos_token_id = 1 + chat_template = None + init_kwargs = {} + + def get_vocab(self): + return {} + + calls = [] + + def from_pretrained(*args, **kwargs): + calls.append(kwargs) + if len(calls) == 1: + raise AttributeError( + "'PreTrainedConfig' object has no attribute " + "'max_position_embeddings'" + ) + return MockTokenizer() + + with TemporaryDirectory() as tmpdir: + tokenizer_json = Path(tmpdir) / "tokenizer.json" + tokenizer_json.write_text("{}", encoding="utf-8") + + with patch( + "mlx_lm.tokenizer_utils.AutoTokenizer.from_pretrained", + side_effect=from_pretrained, + ): + tokenizer = load_tokenizer_impl(Path(tmpdir)) + + self.assertEqual(tokenizer.eos_token_id, 1) + self.assertEqual(len(calls), 2) + self.assertIsInstance(calls[1]["config"], PreTrainedConfig) + if __name__ == "__main__": unittest.main() From 69087368a335a2a03cec83069f365d98ee86cea1 Mon Sep 17 00:00:00 2001 From: Thump604 Date: Fri, 24 Apr 2026 01:53:29 -0500 Subject: [PATCH 8/9] Cast DeepSeek V4 attention sinks for SDPA --- mlx_lm/models/deepseek_v4.py | 2 +- tests/test_models.py | 4 +++- 2 files changed, 4 insertions(+), 2 deletions(-) diff --git a/mlx_lm/models/deepseek_v4.py b/mlx_lm/models/deepseek_v4.py index 28a583b21..46f597332 100644 --- a/mlx_lm/models/deepseek_v4.py +++ b/mlx_lm/models/deepseek_v4.py @@ -576,7 +576,7 @@ def __call__(self, x: mx.array, mask=None, cache=None): cache=cache, scale=self.scale, mask=mask, - sinks=self.attn_sink, + sinks=self.attn_sink.astype(q.dtype), ) out_nope, out_pe = mx.split(out, [self.nope_head_dim], axis=-1) diff --git a/tests/test_models.py b/tests/test_models.py index 3df027b8d..e8675ba11 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -1520,7 +1520,7 @@ def test_deepseek_v4(self): (2 * args.index_head_dim, args.hidden_size), ) - for dtype in [mx.float32, mx.float16]: + for dtype in [mx.float32, mx.float16, mx.bfloat16]: model.update( tree_map( lambda p: p.astype(dtype) @@ -1529,6 +1529,8 @@ def test_deepseek_v4(self): model.parameters(), ) ) + for layer in model.model.layers: + layer.attn.attn_sink = layer.attn.attn_sink.astype(mx.float32) inputs = mx.array([[0, 1, 2, 3, 4]], dtype=mx.int32) outputs = model(inputs) From 9c990f4e0039f144bf7a3a70150c9e5e7634e3d3 Mon Sep 17 00:00:00 2001 From: Thump604 Date: Fri, 24 Apr 2026 02:00:16 -0500 Subject: [PATCH 9/9] Support quantized DeepSeek V4 output projection --- mlx_lm/models/deepseek_v4.py | 49 +++++++++++++++++++++++++++++++----- tests/test_models.py | 37 +++++++++++++++++++++++++++ 2 files changed, 80 insertions(+), 6 deletions(-) diff --git a/mlx_lm/models/deepseek_v4.py b/mlx_lm/models/deepseek_v4.py index 46f597332..f4b69bb59 100644 --- a/mlx_lm/models/deepseek_v4.py +++ b/mlx_lm/models/deepseek_v4.py @@ -541,6 +541,48 @@ def __init__(self, args: ModelArgs, layer_idx: int): if self.compress_ratio == 4: self.indexer = Indexer(args, self.compress_ratio) + def _grouped_output_projection(self, out: mx.array) -> mx.array: + # DeepSeek-V4 stores wo_a as grouped low-rank blocks. QuantizedLinear + # packs the per-group input dimension, so grouped slicing happens on + # output rows while each group uses the full packed input row. + B, S = out.shape[:2] + group_feat = (self.n_heads * self.head_dim) // self.n_groups + out = out.reshape(B, S, self.n_groups, group_feat) + + if isinstance(self.wo_a, nn.QuantizedLinear): + pieces = [] + for group_idx in range(self.n_groups): + rows = slice( + group_idx * self.o_lora_rank, + (group_idx + 1) * self.o_lora_rank, + ) + biases = ( + self.wo_a.biases[rows] + if self.wo_a.biases is not None + else None + ) + y = mx.quantized_matmul( + out[:, :, group_idx, :], + self.wo_a.weight[rows], + scales=self.wo_a.scales[rows], + biases=biases, + transpose=True, + group_size=self.wo_a.group_size, + bits=self.wo_a.bits, + mode=self.wo_a.mode, + ) + if "bias" in self.wo_a: + y = y + self.wo_a.bias[rows] + pieces.append(y) + return mx.concatenate(pieces, axis=-1) + + wa = self.wo_a.weight.reshape(self.n_groups, self.o_lora_rank, group_feat) + out = mx.einsum("bsgd,grd->bsgr", out, wa) + out = out.reshape(B, S, self.n_groups * self.o_lora_rank) + if "bias" in self.wo_a: + out = out + self.wo_a.bias + return out + def __call__(self, x: mx.array, mask=None, cache=None): B, S, _ = x.shape @@ -585,12 +627,7 @@ def __call__(self, x: mx.array, mask=None, cache=None): # Grouped low-rank projection: [B, n_heads, S, head_dim] -> [B, S, n_heads*head_dim] out = out.transpose(0, 2, 1, 3).reshape(B, S, self.n_heads * self.head_dim) - # Split into o_groups along the head_dim*n_heads axis; apply per-group wo_a via einsum. - out = out.reshape(B, S, self.n_groups, -1) - # wo_a.weight shape is [n_groups * o_lora_rank, group_feat]; reshape to [n_groups, o_lora_rank, group_feat] - wa = self.wo_a.weight.reshape(self.n_groups, self.o_lora_rank, -1) - out = mx.einsum("bsgd,grd->bsgr", out, wa) # [B,S,n_groups,o_lora_rank] - out = out.reshape(B, S, self.n_groups * self.o_lora_rank) + out = self._grouped_output_projection(out) return self.wo_b(out) diff --git a/tests/test_models.py b/tests/test_models.py index e8675ba11..d52178360 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -1603,6 +1603,43 @@ def test_mixed_quant_preserves_deepseek_v4_attention_paths(self): low, ) + def test_deepseek_v4_quantized_grouped_output_projection(self): + from mlx_lm.models import deepseek_v4 + + args = deepseek_v4.ModelArgs( + model_type="deepseek_v4", + vocab_size=128, + hidden_size=64, + num_hidden_layers=1, + num_attention_heads=4, + q_lora_rank=16, + o_lora_rank=8, + o_groups=2, + head_dim=16, + qk_rope_head_dim=4, + sliding_window=16, + compress_ratios=[0], + moe_intermediate_size=16, + n_routed_experts=4, + n_shared_experts=1, + num_experts_per_tok=2, + num_hash_layers=1, + hc_mult=2, + hc_sinkhorn_iters=2, + ) + attn = deepseek_v4.V4Attention(args, layer_idx=0) + attn.wo_a = nn.QuantizedLinear.from_linear( + attn.wo_a, + group_size=32, + bits=6, + mode="affine", + ) + + out = mx.random.uniform(shape=(1, 3, args.num_attention_heads * args.head_dim)) + y = attn._grouped_output_projection(out) + mx.eval(y) + self.assertEqual(y.shape, (1, 3, args.o_groups * args.o_lora_rank)) + def test_deepseek_v4_sanitize_unpacks_fp4_experts(self): from mlx_lm.models import deepseek_v4