[FlashInfer v0.6.4] [RL] Integrate FlashInfer mxfp8 gemm, MoE, and routed MoE (#19537)

This commit is contained in:
Ziang Li
2026-03-10 15:37:57 -07:00
committed by GitHub
parent bd460e9565
commit 76ee4bb98c
14 changed files with 671 additions and 86 deletions
@@ -185,6 +185,8 @@ def _check_cutlass_block_fp8_hardware_support() -> bool:
if is_blackwell_supported() and is_flashinfer_available():
from flashinfer import mm_mxfp8 as _raw_flashinfer_mm_mxfp8
from flashinfer import mxfp8_quantize as _raw_flashinfer_mxfp8_quantize
from flashinfer.gemm import gemm_fp8_nt_groupwise as _raw_gemm_fp8_nt_groupwise
from sglang.srt.utils.custom_op import register_custom_op
@@ -242,6 +244,62 @@ if is_blackwell_supported() and is_flashinfer_available():
backend=backend,
)
# Wrap MXFP8 ops as custom ops so torch.compile does not trace into
# flashinfer's JIT compilation path (filesystem checks/cubin loader).
def _fake_flashinfer_mxfp8_quantize(
input: torch.Tensor,
_is_sf_swizzled_layout: bool = True,
alignment: int = 32,
) -> Tuple[torch.Tensor, torch.Tensor]:
# Fake mode only needs dtypes and output rank to propagate compile graph.
# The scale tensor shape is not consumed before the following fake mm op.
k_aligned = ((input.shape[1] + alignment - 1) // alignment) * alignment
q_input = input.new_empty(
(input.shape[0], k_aligned), dtype=torch.float8_e4m3fn
)
scale = input.new_empty((1,), dtype=torch.uint8)
return q_input, scale
@register_custom_op(
op_name="flashinfer_mxfp8_quantize",
mutates_args=[],
fake_impl=_fake_flashinfer_mxfp8_quantize,
)
def flashinfer_mxfp8_quantize(
input: torch.Tensor,
is_sf_swizzled_layout: bool = True,
alignment: int = 32,
) -> Tuple[torch.Tensor, torch.Tensor]:
return _raw_flashinfer_mxfp8_quantize(
input,
is_sf_swizzled_layout=is_sf_swizzled_layout,
alignment=alignment,
)
@register_custom_op(
op_name="flashinfer_mm_mxfp8",
mutates_args=[],
fake_impl=lambda q_input, weight_t, x_scale_u8, weight_scale_t, out_dtype, backend="auto": (
q_input.new_empty((q_input.shape[0], weight_t.shape[1]), dtype=out_dtype)
),
)
def flashinfer_mm_mxfp8(
q_input: torch.Tensor,
weight_t: torch.Tensor,
x_scale_u8: torch.Tensor,
weight_scale_t: torch.Tensor,
out_dtype: torch.dtype,
backend: str = "auto",
) -> torch.Tensor:
return _raw_flashinfer_mm_mxfp8(
q_input,
weight_t,
x_scale_u8,
weight_scale_t,
out_dtype=out_dtype,
backend=backend,
)
if is_sm90_supported() and is_flashinfer_available():
# FlashInfer SM90 DeepGEMM with automatic swapAB optimization for small M
@@ -266,6 +324,18 @@ def dispatch_w8a8_block_fp8_linear() -> Callable:
return _dispatch_auto_backend()
def dispatch_w8a8_mxfp8_linear() -> Callable:
"""Dispatch MXFP8 linear kernel by --fp8-gemm-backend.
For MXFP8, Triton remains the default path. We only route to FlashInfer
when backend is explicitly set to flashinfer_trtllm.
"""
backend = get_fp8_gemm_runner_backend()
if backend.is_flashinfer_trtllm():
return flashinfer_mxfp8_blockscaled_linear
return triton_mxfp8_blockscaled_linear
def _dispatch_explicit_backend(backend: Fp8GemmRunnerBackend) -> Callable:
"""Dispatch based on explicitly selected backend."""
if backend.is_flashinfer_trtllm():
@@ -843,6 +913,61 @@ def triton_mxfp8_blockscaled_linear(
return output.to(dtype=output_dtype).view(*output_shape)
def flashinfer_mxfp8_blockscaled_linear(
input: torch.Tensor,
weight: torch.Tensor,
weight_scale: torch.Tensor,
input_scale: Optional[torch.Tensor] = None,
bias: Optional[torch.Tensor] = None,
output_dtype: Optional[torch.dtype] = None,
) -> torch.Tensor:
"""MXFP8 dense linear via FlashInfer mm_mxfp8."""
input_2d = input.view(-1, input.shape[-1]).contiguous()
output_shape = [*input.shape[:-1], weight.shape[0]]
m, k = input_2d.shape
n, k_w = weight.shape
if k != k_w:
raise ValueError(f"Input K={k} does not match weight K={k_w}.")
if k % 32 != 0:
raise ValueError(f"K={k} must be divisible by 32 for MXFP8.")
if weight.dtype != torch.float8_e4m3fn:
raise TypeError("MXFP8 weight must be FP8 E4M3.")
if input_scale is None:
q_input, x_scale_u8 = flashinfer_mxfp8_quantize(
input_2d, is_sf_swizzled_layout=True, alignment=32
)
else:
q_input = input_2d
if output_dtype is None:
if input_2d.dtype in (torch.float16, torch.bfloat16, torch.float32):
output_dtype = input_2d.dtype
else:
output_dtype = torch.bfloat16
# Ensure transposed tensors are contiguous for FlashInfer's internal runner.
weight_t = weight.contiguous().t()
weight_scale_t = (
weight_scale.contiguous().t()
if weight_scale.ndim == 2
else weight_scale.contiguous()
)
output = flashinfer_mm_mxfp8(
q_input,
weight_t,
x_scale_u8,
weight_scale_t,
out_dtype=output_dtype,
backend="auto",
)
if bias is not None:
output += bias
return output.to(dtype=output_dtype).view(*output_shape)
def dequant_mxfp4(
w_block: torch.Tensor,
w_scale: torch.Tensor,