add flashinfer cute-dsl backend for mxfp8 gemm (#34042)

Co-authored-by: Brayden Zhong <brayden@radixark.ai>
This commit is contained in:
Carrie Chen
2026-08-13 08:50:01 +08:00
committed by GitHub
co-authored by Brayden Zhong
parent 40eaf34428
commit 6a5a9eccaa
6 changed files with 97 additions and 38 deletions
@@ -66,8 +66,13 @@ def _fp8_block_backends():
def _mxfp8_backends():
# MXFP8 linear is validated on SM100/103 only.
if 100 <= get_device_sm() < 110:
return ["triton", "flashinfer_trtllm", "flashinfer_cutlass"]
if get_device_sm() in (100, 103):
return [
"auto",
"flashinfer_trtllm",
"flashinfer_cutlass",
"flashinfer_cutedsl",
]
return []
@@ -188,15 +193,50 @@ class TestMxfp8LinearBackends(_LinearBackendCheck):
def _run(self, backend: str):
self._check_backend(backend, _mxfp8_backends(), MXFP8_SHAPES, self._build_layer)
def test_triton(self):
self._run("triton")
def test_flashinfer_trtllm(self):
self._run("flashinfer_trtllm")
def test_flashinfer_cutlass(self):
self._run("flashinfer_cutlass")
def test_flashinfer_cutedsl(self):
self._run("flashinfer_cutedsl")
def test_auto(self):
if "auto" not in _mxfp8_backends():
self.skipTest(f"auto not in SM{get_device_sm()} MXFP8 backend set")
with mock.patch.object(
fp8_utils,
"FP8_GEMM_RUNNER_BACKEND",
Fp8GemmRunnerBackend.AUTO,
):
self.assertEqual(
fp8_utils.resolve_mxfp8_dense_gemm_backend(),
fp8_utils.Mxfp8DenseGemmBackend.FLASHINFER_CUTEDSL,
)
self._run("auto")
@unittest.skipUnless(get_device_sm() >= 100, "Requires Blackwell FlashInfer")
def test_auto_falls_back_when_cutedsl_is_unsupported(self):
with (
mock.patch.object(
fp8_utils,
"FP8_GEMM_RUNNER_BACKEND",
Fp8GemmRunnerBackend.AUTO,
),
mock.patch.object(fp8_utils, "get_device_sm", return_value=107),
mock.patch.object(
fp8_utils._raw_flashinfer_mm_mxfp8,
"is_backend_supported",
return_value=False,
) as is_backend_supported,
):
self.assertEqual(
fp8_utils.resolve_mxfp8_dense_gemm_backend(),
fp8_utils.Mxfp8DenseGemmBackend.FLASHINFER_CUTLASS,
)
is_backend_supported.assert_called_once_with("cute-dsl", 107)
@unittest.skipIf(get_device_sm() < 90, "FP8 GEMM backends require SM90+")
class TestModeloptFp8PerTensorLinear(_LinearBackendCheck):