Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 23 additions & 9 deletions python/sglang/srt/models/qwen3_5_mtp.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
107 changes: 107 additions & 0 deletions test/registered/unit/models/test_qwen3_5_mtp_quant_config.py
Original file line number Diff line number Diff line change
@@ -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()
Loading