From 2f422ddbcfc19bc080e70c54696cf7a5b438a212 Mon Sep 17 00:00:00 2001 From: Sai Ankith Averineni Date: Sat, 26 Sep 2026 16:10:14 -0700 Subject: [PATCH] [AMD] Honor an explicit triton moe_runner_backend for mxfp8 on ROCm The mxfp8 constraint pass adds 'triton' to the allowed runners only when is_gfx95_mxfp8, so on every other ROCm part an explicit --moe-runner-backend triton fell through to the ot in allowed branch and was replaced by flashinfer_trtllm: mxfp8 quantization supports only cutlass, deep_gemm, flashinfer_megamoe, flashinfer_trtllm, flashinfer_trtllm_routed backends. Overriding 'triton'. Serving then died at startup, because flashinfer_trtllm's MoE apply path imports flashinfer, which is absent on ROCm. Neither escape hatch could recover: the AMD early-return in fp8.py requires runner_backend.is_aiter(), and create_moe_runner's ROCm fallback only fires when the backend is still 'auto' -- both dead once the gate had pinned a non-auto value. Being a post-process hook, the pass also clobbers the model overrides that ask for triton themselves. Widen llowed for an explicit triton request on any ROCm part. Every other entry in MXFP8_MOE_RUNNER_BACKEND_CHOICES is CUDA-only, so triton is the only selectable runner there. mxfp8_default is deliberately untouched, so 'auto' still resolves exactly as before on every platform and nothing changes for anyone who does not ask for triton by name. The aiter branch stays gated on gfx95, since its MoE quant info is built only there. Verified on MI455x (gfx1250) serving MiniMax-M3-MXFP8 at TP=2 and TP=4: zero override warnings, GSM8K 0.972 / 0.973 over 1319 questions. --- python/sglang/srt/arg_groups/overrides.py | 6 +- .../unit/test_mxfp8_moe_runner_backend.py | 92 +++++++++++++++++++ 2 files changed, 97 insertions(+), 1 deletion(-) create mode 100644 test/registered/unit/test_mxfp8_moe_runner_backend.py diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index 81f8c4cbe115..306b001a985a 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -1412,8 +1412,12 @@ def _moe_runner_backend_quant_constraints(view: Any) -> dict: is_gfx95_mxfp8 = get_platform().is_hip and is_gfx95_supported() allowed = list(MXFP8_MOE_RUNNER_BACKEND_CHOICES) - if is_gfx95_mxfp8: + # Every other entry is CUDA-only. Honor an explicit triton request on ROCm + # instead of sending it back to flashinfer_trtllm, whose MoE apply path + # imports flashinfer. + if is_gfx95_mxfp8 or (get_platform().is_hip and moe_runner_backend == "triton"): allowed.append("triton") + if is_gfx95_mxfp8: # the aiter MXFP8 MoE quant info is built only when aiter is enabled if envs.SGLANG_USE_AITER.get(): allowed.append("aiter") diff --git a/test/registered/unit/test_mxfp8_moe_runner_backend.py b/test/registered/unit/test_mxfp8_moe_runner_backend.py new file mode 100644 index 000000000000..532a50034b5f --- /dev/null +++ b/test/registered/unit/test_mxfp8_moe_runner_backend.py @@ -0,0 +1,92 @@ +"""Unit tests for the mxfp8 moe_runner_backend constraints in +srt/arg_groups/overrides.py.""" + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + +import contextlib +import unittest +from types import SimpleNamespace +from unittest.mock import patch + +from sglang.srt.arg_groups.overrides import ( + ResolvedView, + _moe_runner_backend_quant_constraints, +) +from sglang.srt.runtime_context import override_platform +from sglang.test.test_utils import CustomTestCase + + +@contextlib.contextmanager +def _platform(*, is_hip: bool, is_gfx95: bool = False): + """gfx95 is read straight from utils.common, not through the platform + probes, so it is patched where overrides.py imported it.""" + with override_platform(is_hip=is_hip, is_npu=False): + with patch( + "sglang.srt.arg_groups.overrides.is_gfx95_supported", + return_value=is_gfx95, + ): + yield + + +def _resolve(backend: str, quantization: str = "mxfp8") -> str: + """Run the pass and return the backend it resolves to. The pass declares + only what it changes, so an empty declaration means the request survived. + + moe_a2a_backend is "none" throughout: the pass reads it to pick the default + (flashinfer_megamoe a2a forces a matching runner), and that interaction is + out of scope here. + """ + view = ResolvedView( + SimpleNamespace( + moe_runner_backend=backend, + quantization=quantization, + moe_a2a_backend="none", + ) + ) + declared = _moe_runner_backend_quant_constraints(view) + return declared.get("moe_runner_backend", backend) + + +class TestMxfp8MoeRunnerBackend(CustomTestCase): + def test_rocm_honors_explicit_triton(self): + # triton is the only mxfp8 runner ROCm can use: every other entry in + # MXFP8_MOE_RUNNER_BACKEND_CHOICES is CUDA-only, and flashinfer_trtllm's + # MoE apply path imports flashinfer, which is absent on ROCm. + for is_gfx95 in (True, False): + with self.subTest(is_gfx95=is_gfx95): + with _platform(is_hip=True, is_gfx95=is_gfx95): + self.assertEqual(_resolve("triton"), "triton") + + def test_cuda_still_rejects_triton(self): + with _platform(is_hip=False): + self.assertEqual(_resolve("triton"), "flashinfer_trtllm") + + def test_auto_default_is_unchanged(self): + # The fix widens `allowed`, not the default: "auto" must resolve exactly + # as it did before on every platform. + cases = ( + (dict(is_hip=True, is_gfx95=True), "triton"), + (dict(is_hip=True, is_gfx95=False), "flashinfer_trtllm"), + (dict(is_hip=False, is_gfx95=False), "flashinfer_trtllm"), + ) + for facts, expected in cases: + with self.subTest(**facts): + with _platform(**facts): + self.assertEqual(_resolve("auto"), expected) + + def test_unsupported_backend_is_still_overridden(self): + # The widened `allowed` must not turn the gate into a pass-through: a + # name outside MXFP8_MOE_RUNNER_BACKEND_CHOICES still falls to the + # platform default. + with _platform(is_hip=True, is_gfx95=False): + self.assertEqual(_resolve("marlin"), "flashinfer_trtllm") + + def test_non_mxfp8_quantization_is_untouched(self): + with _platform(is_hip=True, is_gfx95=False): + self.assertEqual(_resolve("triton", quantization="fp8"), "triton") + + +if __name__ == "__main__": + unittest.main()