Fix Muse Glimmer ModelOpt mixed weight mapping (#37510)
This commit is contained in:
@@ -40,6 +40,7 @@ from sglang.srt.model_loader.weight_utils import (
|
||||
get_quant_config,
|
||||
)
|
||||
from sglang.srt.models.minimax_m3 import MiniMaxM3SparseForCausalLM
|
||||
from sglang.srt.models.muse_glimmer import MuseGlimmerForConditionalGeneration
|
||||
from sglang.srt.models.utils import WeightsMapper
|
||||
from sglang.srt.utils import get_device
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
@@ -889,6 +890,69 @@ class TestModelOptMixedPrecisionConfig(CustomTestCase):
|
||||
["language_model.lm_head", "lm_head"],
|
||||
)
|
||||
|
||||
def test_muse_glimmer_mixed_precision_resolves_runtime_names(self):
|
||||
"""The vendor keys quant metadata under ``model.language_model.*``;
|
||||
it must resolve for the ``model.*`` modules the runtime builds.
|
||||
"""
|
||||
quant_config = ModelOptMixedPrecisionConfig.from_config(
|
||||
{
|
||||
"quant_algo": "MIXED_PRECISION",
|
||||
"quantized_layers": {
|
||||
"model.language_model.layers.0.mlp.gate_proj": {
|
||||
"quant_algo": "W4A16_NVFP4",
|
||||
"group_size": 16,
|
||||
},
|
||||
"model.language_model.layers.0.mlp.up_proj": {
|
||||
"quant_algo": "W4A16_NVFP4",
|
||||
"group_size": 16,
|
||||
},
|
||||
"model.language_model.layers.0.self_attn.q_proj": {
|
||||
"quant_algo": "FP8"
|
||||
},
|
||||
"model.language_model.layers.0.self_attn.k_proj": {
|
||||
"quant_algo": "FP8"
|
||||
},
|
||||
"model.language_model.layers.0.self_attn.v_proj": {
|
||||
"quant_algo": "FP8"
|
||||
},
|
||||
"model.language_model.layers.0.self_attn.gate_proj": {
|
||||
"quant_algo": "FP8"
|
||||
},
|
||||
"lm_head": {"quant_algo": "W4A16_NVFP4", "group_size": 16},
|
||||
"model.vision_tower.layers.0.attn.q_proj": {"quant_algo": "FP8"},
|
||||
},
|
||||
"packed_modules_mapping": (
|
||||
MuseGlimmerForConditionalGeneration.packed_modules_mapping
|
||||
),
|
||||
}
|
||||
)
|
||||
quant_config.apply_weight_name_mapper(
|
||||
MuseGlimmerForConditionalGeneration.hf_to_sglang_mapper
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
quant_config._resolve_quant_algo("model.layers.0.mlp.gate_up_proj"),
|
||||
"W4A16_NVFP4",
|
||||
)
|
||||
# Attention stays unfused whenever a quant_config is present, so q/k/v
|
||||
# resolve per shard; only the MLP goes through packed_modules_mapping.
|
||||
self.assertEqual(
|
||||
quant_config._resolve_quant_algo("model.layers.0.self_attn.q_proj"),
|
||||
"FP8",
|
||||
)
|
||||
self.assertEqual(
|
||||
quant_config._resolve_quant_algo(
|
||||
"model.layers.0.self_attn.output_gate_proj"
|
||||
),
|
||||
"FP8",
|
||||
)
|
||||
self.assertEqual(quant_config._resolve_quant_algo("lm_head"), "W4A16_NVFP4")
|
||||
# The vision tower hangs off the entry class, not off ``model``.
|
||||
self.assertEqual(
|
||||
quant_config._resolve_quant_algo("vision_tower.layers.0.attn.q_proj"),
|
||||
"FP8",
|
||||
)
|
||||
|
||||
def test_nemotron_mixed_precision_with_nvfp4_layers_uses_modelopt_mixed(self):
|
||||
model_config = ModelConfig.__new__(ModelConfig)
|
||||
model_config.hf_config = MagicMock()
|
||||
|
||||
Reference in New Issue
Block a user