diff --git a/python/sglang/srt/models/qwen3_5_mtp.py b/python/sglang/srt/models/qwen3_5_mtp.py index c6b21f0c8acf..b90d6f728c57 100644 --- a/python/sglang/srt/models/qwen3_5_mtp.py +++ b/python/sglang/srt/models/qwen3_5_mtp.py @@ -69,21 +69,35 @@ def _mtp_quant_config(quant_config): return None if is_npu() and get_spec().speculative_draft_model_quantization is None: return None - # Quark-quantized Qwen3.5 MXFP4 checkpoints ship the MTP module in bf16; - # every `mtp.*` layer appears under the quantization exclude list. Detect - # that and skip quantization here so linear/MoE weight loaders allocate - # bf16 shapes (see sgl-project/sglang#23113). + # Some Quark-quantized Qwen3.5 MXFP4 checkpoints ship the MTP module + # entirely in bf16, listing every `mtp.*` layer under the quantization + # exclude list. Skip quantization for those so linear/MoE weight loaders + # allocate bf16 shapes (see sgl-project/sglang#23146). + # + # Others are mixed: the routed experts stay MXFP4 while attention, the + # shared expert and fc are excluded. Skipping there would make the MoE + # loader allocate bf16 experts that the MXFP4 checkpoint shards no longer + # fit. The routed experts are the bulk of the draft, so use them as the + # signal and skip only when they are excluded too; the per-layer + # exclusions keep the remaining bf16 modules bf16 on their own. if quant_config and quant_config.get_name() == "quark": - exclude_layers = getattr(quant_config, "exclude_layers", []) - if any( - isinstance(layer, str) and layer.startswith("mtp.") - for layer in exclude_layers - ): + mtp_excludes = [ + layer + for layer in getattr(quant_config, "exclude_layers", []) + if isinstance(layer, str) and layer.startswith("mtp.") + ] + if mtp_excludes and any("mlp.experts" in layer for layer in mtp_excludes): return None return quant_config class Qwen3_5ForCausalLMMTP(nn.Module): + # The loader reads this off the model class and hands it to the quant + # config, which needs it to expand fused module names (qkv_proj -> + # q/k/v_proj) before matching them against an exclude list. Without it an + # excluded attention projection is not recognised as excluded. + packed_modules_mapping = Qwen3_5ForCausalLM.packed_modules_mapping + @staticmethod def shared_experts_fusion_disable_reason(hf_config, quant_config): return Qwen3_5ForCausalLM.shared_experts_fusion_disable_reason( diff --git a/test/registered/unit/models/test_qwen3_5_mtp_quant_config.py b/test/registered/unit/models/test_qwen3_5_mtp_quant_config.py new file mode 100644 index 000000000000..98b88ecbd5bc --- /dev/null +++ b/test/registered/unit/models/test_qwen3_5_mtp_quant_config.py @@ -0,0 +1,107 @@ +import unittest + +from sglang.srt.layers.quantization.quark.utils import should_ignore_layer +from sglang.srt.models.qwen3_5 import Qwen3_5ForCausalLM +from sglang.srt.models.qwen3_5_mtp import Qwen3_5ForCausalLMMTP, _mtp_quant_config +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + +# The real `mtp.*` exclude entries of the two AMD Quark MXFP4 checkpoints. +# Quark names layers in the checkpoint namespace, where the draft is prefixed +# `mtp.`. The two lists are identical apart from the routed experts: +# Qwen3.5-397B excludes all 512x3 draft expert projections, so its whole MTP +# module is bf16; Qwen3.8-2.4T excludes none of them, so its draft experts stay +# MXFP4 while attention, the shared expert and fc are bf16. +_MIXED_EXCLUDES = [ # amd/Qwen3.8-2.4T-A95B-Quark-MXFP4 (all 10 mtp.* entries) + "mtp.fc", + "mtp.layers.0.mlp.gate", + "mtp.layers.0.mlp.shared_expert.down_proj", + "mtp.layers.0.mlp.shared_expert.gate_proj", + "mtp.layers.0.mlp.shared_expert.up_proj", + "mtp.layers.0.mlp.shared_expert_gate", + "mtp.layers.0.self_attn.k_proj", + "mtp.layers.0.self_attn.o_proj", + "mtp.layers.0.self_attn.q_proj", + "mtp.layers.0.self_attn.v_proj", +] +_ALL_BF16_EXCLUDES = _MIXED_EXCLUDES + [ # amd/Qwen3.5-397B-A17B-MXFP4 + f"mtp.layers.0.mlp.experts.{expert}.{proj}" + for expert in range(512) + for proj in ("gate_proj", "up_proj", "down_proj") +] + + +class _FakeQuantConfig: + def __init__(self, name, exclude_layers): + self._name = name + self.exclude_layers = list(exclude_layers) + + def get_name(self): + return self._name + + +class TestQwen3_5MTPQuantConfig(CustomTestCase): + def test_mixed_quark_checkpoint_keeps_quantization(self): + """Routed experts stay MXFP4, so the draft must stay quantized.""" + quant_config = _FakeQuantConfig("quark", _MIXED_EXCLUDES) + + self.assertIs(_mtp_quant_config(quant_config), quant_config) + + def test_fully_bf16_quark_checkpoint_skips_quantization(self): + """Regression guard for #23146: a bf16 draft must stay unquantized.""" + quant_config = _FakeQuantConfig("quark", _ALL_BF16_EXCLUDES) + + self.assertIsNone(_mtp_quant_config(quant_config)) + + def test_the_two_checkpoints_differ_only_in_the_routed_experts(self): + """Guards the premise of the fix, not the implementation.""" + extra = set(_ALL_BF16_EXCLUDES) - set(_MIXED_EXCLUDES) + + self.assertTrue(all("mlp.experts" in layer for layer in extra)) + self.assertEqual(len(extra), 512 * 3) + + def test_quark_checkpoint_without_mtp_excludes_keeps_quantization(self): + quant_config = _FakeQuantConfig("quark", ["model.layers.0.self_attn.q_proj"]) + + self.assertIs(_mtp_quant_config(quant_config), quant_config) + + def test_non_quark_quant_config_is_untouched(self): + quant_config = _FakeQuantConfig("fp8", _MIXED_EXCLUDES) + + self.assertIs(_mtp_quant_config(quant_config), quant_config) + + def test_mtp_reuses_target_packed_modules_mapping(self): + self.assertEqual( + Qwen3_5ForCausalLMMTP.packed_modules_mapping, + Qwen3_5ForCausalLM.packed_modules_mapping, + ) + + def test_excluded_fused_attention_projection_is_ignored(self): + """The mapping has to expand qkv_proj before the exclude list matches. + + Without it `should_ignore_layer` compares the fused name against + `q_proj`/`k_proj`/`v_proj` entries, finds nothing, and the layer is + quantized against bf16 checkpoint shards. + """ + ignored = should_ignore_layer( + "mtp.layers.0.self_attn.qkv_proj", + ignore=_MIXED_EXCLUDES, + fused_mapping=Qwen3_5ForCausalLMMTP.packed_modules_mapping, + ) + + self.assertTrue(ignored) + + def test_unexcluded_fused_attention_projection_is_not_ignored(self): + ignored = should_ignore_layer( + "model.layers.0.self_attn.qkv_proj", + ignore=_MIXED_EXCLUDES, + fused_mapping=Qwen3_5ForCausalLMMTP.packed_modules_mapping, + ) + + self.assertFalse(ignored) + + +if __name__ == "__main__": + unittest.main()