fix(minimax): use routed TRT-LLM for NVFP4 MoE auto on SM100 (#32229)
Co-authored-by: Mohammad Miadh Angkad <176301910+mmangkad@users.noreply.github.com>
This commit is contained in:
co-authored by
Mohammad Miadh Angkad
parent
410088c91e
commit
5d85f25f75
@@ -731,10 +731,21 @@ def _mimo_v2_overrides(server_args: Any, hf_config: Any) -> dict:
|
|||||||
|
|
||||||
@_register_for("MiniMaxM2ForCausalLM")
|
@_register_for("MiniMaxM2ForCausalLM")
|
||||||
def _minimax_m2_overrides(server_args: Any, hf_config: Any) -> dict:
|
def _minimax_m2_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||||
|
overrides = {"enable_tf32_matmul": True}
|
||||||
logger.info(
|
logger.info(
|
||||||
"Enable TF32 matmul for MiniMaxM2ForCausalLM model to improve gate gemm performance."
|
"Enable TF32 matmul for MiniMaxM2ForCausalLM model to improve gate gemm performance."
|
||||||
)
|
)
|
||||||
return {"enable_tf32_matmul": True}
|
if (
|
||||||
|
is_sm100_supported()
|
||||||
|
and server_args.moe_runner_backend == "auto"
|
||||||
|
and server_args.get_model_config().quantization == "modelopt_fp4"
|
||||||
|
):
|
||||||
|
overrides["moe_runner_backend"] = "flashinfer_trtllm_routed"
|
||||||
|
logger.info(
|
||||||
|
"Use flashinfer_trtllm_routed as MoE runner backend on SM10X "
|
||||||
|
"for MiniMaxM2ForCausalLM with modelopt_fp4."
|
||||||
|
)
|
||||||
|
return overrides
|
||||||
|
|
||||||
|
|
||||||
@_register_for("MiniMaxM3SparseForCausalLM", "MiniMaxM3SparseForConditionalGeneration")
|
@_register_for("MiniMaxM3SparseForCausalLM", "MiniMaxM3SparseForConditionalGeneration")
|
||||||
|
|||||||
@@ -25,6 +25,7 @@ from sglang.srt.arg_groups.overrides import (
|
|||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
get_context,
|
get_context,
|
||||||
|
get_exec,
|
||||||
get_server_args,
|
get_server_args,
|
||||||
reset_context,
|
reset_context,
|
||||||
)
|
)
|
||||||
@@ -346,6 +347,51 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
|||||||
self.assertTrue(flags.enable_tf32_matmul)
|
self.assertTrue(flags.enable_tf32_matmul)
|
||||||
self.assertFalse(flags.enable_multi_layer_eagle) # pristine materialize
|
self.assertFalse(flags.enable_multi_layer_eagle) # pristine materialize
|
||||||
|
|
||||||
|
def test_minimax_m2_sm10x_nvfp4_uses_routed_trtllm(self):
|
||||||
|
"""MiniMax-M2 NVFP4 auto must avoid the unsupported plain TRT-LLM path."""
|
||||||
|
with patch.object(overrides_module, "is_sm100_supported", return_value=True):
|
||||||
|
explicit = self._construct(
|
||||||
|
"MiniMaxM2ForCausalLM",
|
||||||
|
"llama",
|
||||||
|
quantization="modelopt_fp4",
|
||||||
|
moe_runner_backend="flashinfer_cutlass",
|
||||||
|
)
|
||||||
|
non_nvfp4 = self._construct(
|
||||||
|
"MiniMaxM2ForCausalLM", "llama", quantization="fp8"
|
||||||
|
)
|
||||||
|
nvfp4 = self._construct(
|
||||||
|
"MiniMaxM2ForCausalLM", "llama", quantization="modelopt_fp4"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(explicit.moe_runner_backend, "flashinfer_cutlass")
|
||||||
|
self.assertEqual(non_nvfp4.moe_runner_backend, "auto")
|
||||||
|
self.assertEqual(nvfp4.moe_runner_backend, "flashinfer_trtllm_routed")
|
||||||
|
self.assertTrue(nvfp4.disable_shared_experts_fusion)
|
||||||
|
self.assertIn(
|
||||||
|
(
|
||||||
|
"_minimax_m2_overrides",
|
||||||
|
{
|
||||||
|
"enable_tf32_matmul": True,
|
||||||
|
"moe_runner_backend": "flashinfer_trtllm_routed",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
nvfp4._resolved_overrides,
|
||||||
|
)
|
||||||
|
self.assertIn(
|
||||||
|
("_moe_runner_fusion_disable", {"disable_shared_experts_fusion": True}),
|
||||||
|
nvfp4._resolved_overrides,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Thor (SM110) and other architectures keep the existing auto behavior.
|
||||||
|
with patch.object(overrides_module, "is_sm100_supported", return_value=False):
|
||||||
|
non_sm10x = self._construct(
|
||||||
|
"MiniMaxM2ForCausalLM", "llama", quantization="modelopt_fp4"
|
||||||
|
)
|
||||||
|
self.assertEqual(non_sm10x.moe_runner_backend, "auto")
|
||||||
|
|
||||||
|
self._publish(nvfp4)
|
||||||
|
self.assertEqual(get_exec().moe.moe_runner_backend, "flashinfer_trtllm_routed")
|
||||||
|
|
||||||
def test_mimo_v2_declarations(self):
|
def test_mimo_v2_declarations(self):
|
||||||
# Callable-level golden: MiMoV2 archs are hybrid (config-shape heavy),
|
# Callable-level golden: MiMoV2 archs are hybrid (config-shape heavy),
|
||||||
# so the declaration is pinned directly for both provider inputs.
|
# so the declaration is pinned directly for both provider inputs.
|
||||||
|
|||||||
Reference in New Issue
Block a user