[FlashInfer v0.6.4] [RL] Integrate FlashInfer mxfp8 gemm, MoE, and routed MoE (#19537)
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user