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)