From 7c195b9151627320b2a4688251adc12975712a86 Mon Sep 17 00:00:00 2001 From: Tuan Nguyen Gia Date: Sat, 12 Sep 2026 10:50:41 +0300 Subject: [PATCH] [AMD] Fix Quark load of MiniMax-M3 MXFP4 index_qkv_proj (#37254) --- python/sglang/srt/models/minimax_m3.py | 3 +- python/sglang/srt/models/minimax_m3_vl.py | 3 +- .../layers/quantization/test_quark_utils.py | 51 ++++++++++++++++++- 3 files changed, 54 insertions(+), 3 deletions(-) diff --git a/python/sglang/srt/models/minimax_m3.py b/python/sglang/srt/models/minimax_m3.py index fa2c30723..443de34b6 100644 --- a/python/sglang/srt/models/minimax_m3.py +++ b/python/sglang/srt/models/minimax_m3.py @@ -1559,7 +1559,8 @@ class MiniMaxM3SparseForCausalLM(nn.Module): ) packed_modules_mapping = { "qkv_proj": ["q_proj", "k_proj", "v_proj"], - "index_qkv_proj": ["index_q_proj", "index_k_proj", "index_v_proj"], + # no index_v_proj in the M3 checkpoint + "index_qkv_proj": ["index_q_proj", "index_k_proj"], "gate_up_proj": ["gate_proj", "up_proj"], } diff --git a/python/sglang/srt/models/minimax_m3_vl.py b/python/sglang/srt/models/minimax_m3_vl.py index 5eed7654c..68a6e39b2 100644 --- a/python/sglang/srt/models/minimax_m3_vl.py +++ b/python/sglang/srt/models/minimax_m3_vl.py @@ -63,7 +63,8 @@ class MiniMaxM3SparseForConditionalGeneration(nn.Module): ) packed_modules_mapping = { "qkv_proj": ["q_proj", "k_proj", "v_proj"], - "index_qkv_proj": ["index_q_proj", "index_k_proj", "index_v_proj"], + # no index_v_proj in the M3 checkpoint + "index_qkv_proj": ["index_q_proj", "index_k_proj"], "gate_up_proj": ["gate_proj", "up_proj"], } diff --git a/test/registered/unit/layers/quantization/test_quark_utils.py b/test/registered/unit/layers/quantization/test_quark_utils.py index 6ae108bd4..5ed8240b2 100644 --- a/test/registered/unit/layers/quantization/test_quark_utils.py +++ b/test/registered/unit/layers/quantization/test_quark_utils.py @@ -8,10 +8,59 @@ import unittest import torch -from sglang.srt.layers.quantization.quark.utils import e8m0_to_f32 +from sglang.srt.layers.quantization.quark.utils import ( + e8m0_to_f32, + should_ignore_layer, +) from sglang.test.test_utils import CustomTestCase +class TestShouldIgnoreLayer(CustomTestCase): + """MiniMax-M3 MXFP4: sparse index_qkv_proj packs only q/k (the DSA value + projection is disabled, so index_v_proj is absent on disk).""" + + _LAYER = "language_model.model.layers.3.self_attn.index_qkv_proj" + _IGNORE = ( + "language_model.model.layers.3.self_attn.index_q_proj", + "language_model.model.layers.3.self_attn.index_k_proj", + ) + # The fix lives in the model: index_qkv_proj maps to only q/k (no v). + _MAPPING = { + "index_qkv_proj": ["index_q_proj", "index_k_proj"], + } + + def test_minimax_dsa_index_qkv_ignored(self): + # Both present shards are excluded -> fused module stays bf16, no raise. + self.assertTrue(should_ignore_layer(self._LAYER, self._IGNORE, self._MAPPING)) + + def test_all_shards_agree_still_works(self): + layer = "model.layers.0.self_attn.qkv_proj" + ignore = ( + "model.layers.0.self_attn.q_proj", + "model.layers.0.self_attn.k_proj", + "model.layers.0.self_attn.v_proj", + ) + mapping = {"qkv_proj": ["q_proj", "k_proj", "v_proj"]} + self.assertTrue(should_ignore_layer(layer, ignore, mapping)) + + def test_no_shards_ignored(self): + layer = "model.layers.0.self_attn.qkv_proj" + mapping = {"qkv_proj": ["q_proj", "k_proj", "v_proj"]} + self.assertFalse(should_ignore_layer(layer, (), mapping)) + + def test_mixed_schemes_raise(self): + # Safety net preserved: if a fused module genuinely mixes excluded and + # quantized shards, the loader must fail loudly rather than guess. + layer = "model.layers.0.self_attn.qkv_proj" + ignore = ( + "model.layers.0.self_attn.q_proj", + "model.layers.0.self_attn.k_proj", + ) # v_proj NOT excluded -> inconsistent with q/k + mapping = {"qkv_proj": ["q_proj", "k_proj", "v_proj"]} + with self.assertRaises(ValueError): + should_ignore_layer(layer, ignore, mapping) + + class TestE8M0ToF32(CustomTestCase): """Cover OCP MX-format v1.0 e8m0 decoding: encoded 0..254 -> 2^(x-127); encoded 255 -> NaN.