From df06e0366220bcabeaa0312e4192b81fa7a3a2bb Mon Sep 17 00:00:00 2001 From: Siming Deng Date: Tue, 7 Jul 2026 11:31:59 +0800 Subject: [PATCH] [MLX] Size the attention KV pool at the compute dtype for quantized models (#30097) Co-authored-by: siming-deng --- .../srt/hardware_backend/mlx/model_runner.py | 14 +- .../mlx/test_attention_patching.py | 1 + .../mlx/test_mlx_pool_dtype.py | 121 ++++++++++++++++++ 3 files changed, 133 insertions(+), 3 deletions(-) create mode 100644 test/registered/unit/hardware_backend/mlx/test_mlx_pool_dtype.py diff --git a/python/sglang/srt/hardware_backend/mlx/model_runner.py b/python/sglang/srt/hardware_backend/mlx/model_runner.py index 6711252cc..d76671c94 100644 --- a/python/sglang/srt/hardware_backend/mlx/model_runner.py +++ b/python/sglang/srt/hardware_backend/mlx/model_runner.py @@ -512,9 +512,17 @@ class MlxModelRunner: if hasattr(sample_attn, "k_proj") and hasattr(sample_attn.k_proj, "weight"): dtype = sample_attn.k_proj.weight.dtype if dtype not in _MLX_KV_FLOAT_DTYPES: - # QuantizedLinear stores packed weights as integers, while the KV - # cache stores dequantized projection outputs. - dtype = mx.float32 + # QuantizedLinear packs weights as integers, but the KV cache + # stores dequantized projection outputs, which are produced in + # the compute dtype carried by the quantization scales. Storing + # at that dtype instead of float32 halves pool bytes per slot + # and keeps prefix-hit forwards in the same dtype as the no-hit + # path (a float32 pool promoted every post-hit concat). + scales = getattr(sample_attn.k_proj, "scales", None) + if scales is not None and scales.dtype in _MLX_KV_FLOAT_DTYPES: + dtype = scales.dtype + else: + dtype = mx.float32 return n_kv_heads, head_dim, dtype def _get_attn_config(self) -> tuple[int, int, mx.Dtype]: diff --git a/test/registered/unit/hardware_backend/mlx/test_attention_patching.py b/test/registered/unit/hardware_backend/mlx/test_attention_patching.py index 6cfd25d37..8c84b9069 100644 --- a/test/registered/unit/hardware_backend/mlx/test_attention_patching.py +++ b/test/registered/unit/hardware_backend/mlx/test_attention_patching.py @@ -1149,6 +1149,7 @@ class TestMlxOverlapScheduler(unittest.TestCase): self.assertTrue(torch.equal(schedule_batch.input_ids, token_ids)) self.assertIs(scheduler.processed_batch, batch_copy) self.assertIs(scheduler.processed_result, scheduler.tp_worker.result) + self.assertEqual(scheduler.forward_ct, 1) def test_overlap_loop_materializes_prefill_input_ids(self): # Regression: the MLX overlap loop must materialize batch.input_ids diff --git a/test/registered/unit/hardware_backend/mlx/test_mlx_pool_dtype.py b/test/registered/unit/hardware_backend/mlx/test_mlx_pool_dtype.py new file mode 100644 index 000000000..821bd71b4 --- /dev/null +++ b/test/registered/unit/hardware_backend/mlx/test_mlx_pool_dtype.py @@ -0,0 +1,121 @@ +"""Unit tests for MLX attention-KV pool dtype inference. + +The shared pool stores dequantized projection outputs, so its dtype must +follow the model's compute dtype, not the packed integer dtype of +QuantizedLinear weights. Storing float32 for quantized models (the old +fallback) doubled pool bytes per slot for no precision gain — the source +values are bf16/fp16 — and made prefix-hit forwards run in float32 (the +pool gather promoted the concat) while no-hit forwards ran in the compute +dtype. These tests pin the inference rules: + +- unquantized projections: weight dtype (unchanged behavior); +- quantized projections: the quantization ``scales`` dtype (== compute + dtype); +- quantized without usable scales: float32 (the conservative fallback). +""" + +from __future__ import annotations + +import importlib.util +import unittest + +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=4, suite="base-a-test-cpu") + +_HAS_MLX = ( + importlib.util.find_spec("mlx") is not None + and importlib.util.find_spec("mlx_lm") is not None +) +_SKIP_REASON = "requires mlx + mlx_lm" + +if _HAS_MLX: + import mlx.core as mx + import mlx.nn as nn + from mlx_lm.models import qwen2 + + from sglang.srt.hardware_backend.mlx.kv_cache import ( + MlxModelCacheLayout, + find_attention_layers, + ) + from sglang.srt.hardware_backend.mlx.model_runner import MlxModelRunner + + +def _tiny_qwen2_model(): + """Randomly initialized 2-layer dense qwen2 (plain full attention).""" + args = qwen2.ModelArgs( + model_type="qwen2", + hidden_size=64, + num_hidden_layers=2, + intermediate_size=128, + num_attention_heads=4, + num_key_value_heads=2, + rms_norm_eps=1e-6, + vocab_size=128, + rope_theta=10000.0, + ) + return qwen2.Model(args) + + +def _runner_for(model): + layers, attrs = find_attention_layers(model) + runner = MlxModelRunner.__new__(MlxModelRunner) + runner._cache_layout = MlxModelCacheLayout.from_attention_discovery(layers, attrs) + return runner + + +@unittest.skipUnless(_HAS_MLX, _SKIP_REASON) +class TestPoolDtypeInference(CustomTestCase): + def test_unquantized_model_uses_weight_dtype(self): + model = _tiny_qwen2_model() + model.set_dtype(mx.float16) + _, _, dtype = _runner_for(model)._get_attn_config() + self.assertEqual(dtype, mx.float16) + + def test_quantized_model_uses_scales_dtype(self): + model = _tiny_qwen2_model() + model.set_dtype(mx.bfloat16) + nn.quantize(model, group_size=64, bits=4) + attn = model.model.layers[0].self_attn + self.assertNotIn(attn.k_proj.weight.dtype, {mx.float16, mx.bfloat16}) + self.assertEqual(attn.k_proj.scales.dtype, mx.bfloat16) + _, _, dtype = _runner_for(model)._get_attn_config() + self.assertEqual(dtype, mx.bfloat16) + + def test_quantized_fp16_model_uses_scales_dtype(self): + model = _tiny_qwen2_model() + model.set_dtype(mx.float16) + nn.quantize(model, group_size=64, bits=4) + _, _, dtype = _runner_for(model)._get_attn_config() + self.assertEqual(dtype, mx.float16) + + def test_quantized_without_usable_scales_falls_back_to_float32(self): + model = _tiny_qwen2_model() + model.set_dtype(mx.bfloat16) + nn.quantize(model, group_size=64, bits=4) + for layer in model.model.layers: + # Simulate a packed layer whose scales are not a float array + # (e.g. an exotic quant format): the conservative fallback must + # hold. + layer.self_attn.k_proj.scales = layer.self_attn.k_proj.scales.astype( + mx.uint32 + ) + _, _, dtype = _runner_for(model)._get_attn_config() + self.assertEqual(dtype, mx.float32) + + def test_pool_bytes_per_slot_halves_for_bf16_quantized_model(self): + # The practical effect: bytes/slot uses dtype.size, so bf16 halves + # the fp32 fallback and the auto-sized pool fits ~2x the tokens. + model = _tiny_qwen2_model() + model.set_dtype(mx.bfloat16) + nn.quantize(model, group_size=64, bits=4) + n_kv_heads, head_dim, dtype = _runner_for(model)._get_attn_config() + num_layers = 2 + bytes_per_slot = 2 * num_layers * n_kv_heads * head_dim * dtype.size + fp32_bytes_per_slot = 2 * num_layers * n_kv_heads * head_dim * 4 + self.assertEqual(bytes_per_slot * 2, fp32_bytes_per_slot) + + +if __name__ == "__main__": + unittest.main()