[MLX] Size the attention KV pool at the compute dtype for quantized models (#30097)
Co-authored-by: siming-deng <deng_siming@apple.com>
This commit is contained in:
co-authored by
siming-deng
parent
c3da0a2582
commit
df06e03662
@@ -512,8 +512,16 @@ class MlxModelRunner:
|
|||||||
if hasattr(sample_attn, "k_proj") and hasattr(sample_attn.k_proj, "weight"):
|
if hasattr(sample_attn, "k_proj") and hasattr(sample_attn.k_proj, "weight"):
|
||||||
dtype = sample_attn.k_proj.weight.dtype
|
dtype = sample_attn.k_proj.weight.dtype
|
||||||
if dtype not in _MLX_KV_FLOAT_DTYPES:
|
if dtype not in _MLX_KV_FLOAT_DTYPES:
|
||||||
# QuantizedLinear stores packed weights as integers, while the KV
|
# QuantizedLinear packs weights as integers, but the KV cache
|
||||||
# cache stores dequantized projection outputs.
|
# 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
|
dtype = mx.float32
|
||||||
return n_kv_heads, head_dim, dtype
|
return n_kv_heads, head_dim, dtype
|
||||||
|
|
||||||
|
|||||||
@@ -1149,6 +1149,7 @@ class TestMlxOverlapScheduler(unittest.TestCase):
|
|||||||
self.assertTrue(torch.equal(schedule_batch.input_ids, token_ids))
|
self.assertTrue(torch.equal(schedule_batch.input_ids, token_ids))
|
||||||
self.assertIs(scheduler.processed_batch, batch_copy)
|
self.assertIs(scheduler.processed_batch, batch_copy)
|
||||||
self.assertIs(scheduler.processed_result, scheduler.tp_worker.result)
|
self.assertIs(scheduler.processed_result, scheduler.tp_worker.result)
|
||||||
|
self.assertEqual(scheduler.forward_ct, 1)
|
||||||
|
|
||||||
def test_overlap_loop_materializes_prefill_input_ids(self):
|
def test_overlap_loop_materializes_prefill_input_ids(self):
|
||||||
# Regression: the MLX overlap loop must materialize batch.input_ids
|
# Regression: the MLX overlap loop must materialize batch.input_ids
|
||||||
|
|||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user