From 2bd8a0c628fdba9c3b73f2b189fe701ccbd8a983 Mon Sep 17 00:00:00 2001 From: Yongye Zhu Date: Fri, 19 Jun 2026 06:09:35 +0000 Subject: [PATCH] [Kernel] Add FlashInferCutedslMxfp8LinearKernel (cute-dsl mm_mxfp8) Add an MXFP8 W8A8 linear GEMM that drives FlashInfer's mm_mxfp8(..., backend="cute-dsl"), sibling to the existing CUTLASS kernel. The cute-dsl backend consumes the same 1D swizzled F8_128x4 scales the CUTLASS path already produces, so weight/activation prep is identical and output is bit-identical; only the backend string and support gate differ. Gate to sm_100/sm_103 (matching FlashInfer's supported_compute_capability [100, 103]) plus has_flashinfer_cutedsl(); auto-selects on those archs and falls through to CUTLASS elsewhere. Reachable explicitly via --linear-backend flashinfer_cutedsl. Verified on SM100: bit-identical parity vs the CUTLASS kernel, correct backend dispatch, and tests/models/quantization/test_mxfp8.py test_mxfp8_generation[dense] passing (cute-dsl by default). AI assistance (Claude Code) was used for this change. Co-Authored-By: Claude Opus 4.8 Signed-off-by: Yongye Zhu Signed-off-by: Yongye Zhu --- vllm/config/kernel.py | 2 +- .../model_executor/kernels/linear/__init__.py | 4 + .../kernels/linear/mxfp8/flashinfer.py | 82 +++++++++++++++++++ 3 files changed, 87 insertions(+), 1 deletion(-) diff --git a/vllm/config/kernel.py b/vllm/config/kernel.py index cd1408c3e7fa..e9f41c538ad9 100644 --- a/vllm/config/kernel.py +++ b/vllm/config/kernel.py @@ -199,7 +199,7 @@ class KernelConfig: - "auto": Automatically select the best backend based on model and hardware - "cutlass": Use CUTLASS-based kernels - "flashinfer_cutlass": Use FlashInfer with CUTLASS kernels - - "flashinfer_cutedsl": Use FlashInfer with CuteDSL kernels + - "flashinfer_cutedsl": Use FlashInfer with CuTe-DSL kernels (NVFP4, MXFP8) - "flashinfer_trtllm": Use FlashInfer with TensorRT-LLM kernels - "flashinfer_cudnn": Use FlashInfer with cuDNN kernels - "flashinfer_b12x": Use FlashInfer b12x CuteDSL NVFP4 GEMM (SM120+) diff --git a/vllm/model_executor/kernels/linear/__init__.py b/vllm/model_executor/kernels/linear/__init__.py index 58ba7c8cb403..4ac8d49cd58e 100644 --- a/vllm/model_executor/kernels/linear/__init__.py +++ b/vllm/model_executor/kernels/linear/__init__.py @@ -88,6 +88,7 @@ EmulationMxfp8LinearKernel, ) from vllm.model_executor.kernels.linear.mxfp8.flashinfer import ( + FlashInferCutedslMxfp8LinearKernel, FlashInferCutlassMxfp8LinearKernel, ) from vllm.model_executor.kernels.linear.mxfp8.marlin import ( @@ -212,6 +213,7 @@ def _get_linear_backend() -> str: }, "flashinfer_cutedsl": { FlashInferCuteDslNvFp4LinearKernel, + FlashInferCutedslMxfp8LinearKernel, }, "flashinfer_trtllm": { FlashInferTrtllmNvFp4LinearKernel, @@ -385,6 +387,7 @@ def _filter_kernels_by_backend( # in priority/performance order (when available) _POSSIBLE_MXFP8_KERNELS: dict[PlatformEnum, list[type[Mxfp8LinearKernel]]] = { PlatformEnum.CUDA: [ + FlashInferCutedslMxfp8LinearKernel, FlashInferCutlassMxfp8LinearKernel, MarlinMxfp8LinearKernel, EmulationMxfp8LinearKernel, @@ -1036,6 +1039,7 @@ def register_linear_kernel( "MxFp4LinearLayerConfig", "FlashInferMxFp4LinearKernel", "MarlinMxFp4LinearKernel", + "FlashInferCutedslMxfp8LinearKernel", "FlashInferCutlassMxfp8LinearKernel", "MarlinMxfp8LinearKernel", "XPUMxFp8LinearKernel", diff --git a/vllm/model_executor/kernels/linear/mxfp8/flashinfer.py b/vllm/model_executor/kernels/linear/mxfp8/flashinfer.py index 8188fd596096..d26e5579edbb 100644 --- a/vllm/model_executor/kernels/linear/mxfp8/flashinfer.py +++ b/vllm/model_executor/kernels/linear/mxfp8/flashinfer.py @@ -11,6 +11,7 @@ ) from vllm.platforms import current_platform from vllm.utils import flashinfer as vllm_flashinfer +from vllm.utils.flashinfer import has_flashinfer_cutedsl from .Mxfp8LinearKernel import Mxfp8LinearKernel, Mxfp8LinearLayerConfig @@ -91,3 +92,84 @@ def apply_weights( output_shape = (*input_shape[:-1], N) return output.view(output_shape) + + +class FlashInferCutedslMxfp8LinearKernel(Mxfp8LinearKernel): + """MXFP8 W8A8 GEMM via FlashInfer CuTe-DSL (SM100/SM103).""" + + @classmethod + def is_supported( + cls, compute_capability: int | None = None + ) -> tuple[bool, str | None]: + if not ( + current_platform.is_cuda() + and current_platform.is_device_capability_family(100) + ): + return False, "requires sm_100/sm_103 (Blackwell)" + if not has_flashinfer_cutedsl(): + return False, "requires FlashInfer CuTe-DSL module" + return True, None + + @classmethod + def can_implement(cls, c: Mxfp8LinearLayerConfig) -> tuple[bool, str | None]: + return True, None + + def process_weights_after_loading(self, layer: torch.nn.Module) -> None: + weight = layer.weight.data # [N, K] + N, K = weight.shape + + scale_k = K // MXFP8_BLOCK_SIZE + weight_scale_2d = layer.weight_scale.data[:N, :scale_k].contiguous() + weight_scale_swizzled = swizzle_mxfp8_scale(weight_scale_2d, M=N, K=K) + + # Store weight column-major [K, N] as mm_mxfp8 expects for operand B. + layer.weight = Parameter(weight.contiguous().t(), requires_grad=False) + layer.weight_scale = Parameter( + weight_scale_swizzled.contiguous(), requires_grad=False + ) + + def apply_weights( + self, + layer: torch.nn.Module, + x: torch.Tensor, + bias: torch.Tensor | None = None, + ) -> torch.Tensor: + weight = layer.weight # [K, N], column-major + weight_scale = layer.weight_scale + out_dtype = x.dtype + K, N = weight.shape + + input_shape = x.shape + input_2d = x.view(-1, K) + min_dim = 128 + + assert min_dim <= K, ( + f"mm_mxfp8 requires K >= {min_dim}, got K={K}. " + f"in_features is too small for mm_mxfp8." + ) + assert K % MXFP8_BLOCK_SIZE == 0, ( + f"mm_mxfp8 requires K to be divisible by {MXFP8_BLOCK_SIZE}, got K={K}." + ) + assert min_dim <= N, ( + f"mm_mxfp8 requires N >= {min_dim}, got N={N}. " + f"out_features is too small for mm_mxfp8." + ) + + input_mxfp8, input_scale = mxfp8_e4m3_quantize( + input_2d, is_sf_swizzled_layout=True + ) + + output = vllm_flashinfer.mm_mxfp8( + input_mxfp8, + weight, + input_scale, + weight_scale, + out_dtype=out_dtype, + backend="cute-dsl", + ) + + if bias is not None: + output = output + bias + + output_shape = (*input_shape[:-1], N) + return output.view(output_shape)