Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion vllm/config/kernel.py
Original file line number Diff line number Diff line change
Expand Up @@ -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+)
Expand Down
4 changes: 4 additions & 0 deletions vllm/model_executor/kernels/linear/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -212,6 +213,7 @@ def _get_linear_backend() -> str:
},
"flashinfer_cutedsl": {
FlashInferCuteDslNvFp4LinearKernel,
FlashInferCutedslMxfp8LinearKernel,
},
"flashinfer_trtllm": {
FlashInferTrtllmNvFp4LinearKernel,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -1036,6 +1039,7 @@ def register_linear_kernel(
"MxFp4LinearLayerConfig",
"FlashInferMxFp4LinearKernel",
"MarlinMxFp4LinearKernel",
"FlashInferCutedslMxfp8LinearKernel",
"FlashInferCutlassMxfp8LinearKernel",
"MarlinMxfp8LinearKernel",
"XPUMxFp8LinearKernel",
Expand Down
82 changes: 82 additions & 0 deletions vllm/model_executor/kernels/linear/mxfp8/flashinfer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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)
Loading