[Kernel] Route large SM90 row/column-scaled FP8 GEMMs to Torch (#34318)
This commit is contained in:
@@ -4,6 +4,7 @@ from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
from sglang.kernels.fused_op import BaseFusedOp, register_fused_op
|
||||
from sglang.kernels.registry import register_kernel
|
||||
from sglang.kernels.selector import get_kernel
|
||||
from sglang.kernels.spec import (
|
||||
@@ -17,19 +18,150 @@ if TYPE_CHECKING:
|
||||
import torch
|
||||
|
||||
_CUDA = frozenset({CapabilityRequirement.CUDA})
|
||||
_SM90 = frozenset({CapabilityRequirement.cuda(min_sm=(9, 0), max_sm=(9, 0))})
|
||||
|
||||
register_kernel(
|
||||
KernelSpec(
|
||||
op="gemm.fp8_scaled_mm",
|
||||
backend=KernelBackend.AOT,
|
||||
target="sgl_kernel:fp8_scaled_mm",
|
||||
format_signature=FormatSignature(
|
||||
supported_dtypes=("float8_e4m3fn",),
|
||||
description="C = (A_fp8 @ B_fp8) * scales_a * scales_b (+ bias)",
|
||||
|
||||
def _prefer_torch_rowwise_fp8(
|
||||
mat_a: torch.Tensor,
|
||||
mat_b: torch.Tensor,
|
||||
scales_a: torch.Tensor,
|
||||
scales_b: torch.Tensor,
|
||||
out_dtype: torch.dtype,
|
||||
bias: Optional[torch.Tensor],
|
||||
) -> bool:
|
||||
"""Whether Torch's SM90 NVJet kernel wins for this row/column-scaled shape."""
|
||||
import torch
|
||||
|
||||
if (
|
||||
mat_a.device.type != "cuda"
|
||||
or mat_b.device != mat_a.device
|
||||
or not hasattr(torch, "_scaled_mm")
|
||||
or out_dtype != torch.bfloat16
|
||||
or bias is not None
|
||||
or mat_a.dtype != torch.float8_e4m3fn
|
||||
or mat_b.dtype != torch.float8_e4m3fn
|
||||
or mat_a.ndim != 2
|
||||
or mat_b.ndim != 2
|
||||
or mat_a.stride(1) != 1
|
||||
or mat_b.stride(0) != 1
|
||||
):
|
||||
return False
|
||||
|
||||
m, k = mat_a.shape
|
||||
n = mat_b.shape[1]
|
||||
# This path is intentionally row/column scaled, never tensorwise: A has
|
||||
# one independent FP32 scale per input row and B one per output column.
|
||||
if (
|
||||
scales_a.dtype != torch.float32
|
||||
or scales_b.dtype != torch.float32
|
||||
or scales_a.device != mat_a.device
|
||||
or scales_b.device != mat_a.device
|
||||
or not scales_a.is_contiguous()
|
||||
or not scales_b.is_contiguous()
|
||||
or scales_a.numel() != m
|
||||
or scales_b.numel() != n
|
||||
):
|
||||
return False
|
||||
|
||||
# Tuned on H100 over MiniMax-H3's complete dense shape set: four
|
||||
# production sequence lengths and TP1/2/4/8 (64 shapes). This selector
|
||||
# chose the measured winner for every shape while retaining the AOT kernel
|
||||
# for the smaller-K projections where NVJet loses.
|
||||
return (k >= 5376 and n >= 3584) or (k >= 3584 and m >= 8192)
|
||||
|
||||
|
||||
class Fp8ScaledMMOp(BaseFusedOp):
|
||||
"""FP8 GEMM with independent per-row A and per-column B scales."""
|
||||
|
||||
op = "gemm.fp8_scaled_mm"
|
||||
priority = (KernelBackend.AOT, KernelBackend.TORCH)
|
||||
capabilities = {
|
||||
KernelBackend.AOT: _CUDA,
|
||||
KernelBackend.TORCH: _SM90,
|
||||
}
|
||||
format_signature = FormatSignature(
|
||||
supported_dtypes=("float8_e4m3fn",),
|
||||
description=(
|
||||
"C[M,N] = (A_fp8[M,K] @ B_fp8[K,N]) * scale_a[M,1] * scale_b[1,N] (+ bias)"
|
||||
),
|
||||
description="FP8 scaled matmul (sgl_kernel wheel).",
|
||||
)
|
||||
)
|
||||
descriptions = {
|
||||
KernelBackend.AOT: "Row/column-scaled FP8 matmul (sgl_kernel wheel).",
|
||||
KernelBackend.TORCH: "Row/column-scaled FP8 matmul (Torch NVJet on SM90).",
|
||||
}
|
||||
|
||||
def backend_eligible(
|
||||
self,
|
||||
backend: KernelBackend,
|
||||
mat_a: torch.Tensor,
|
||||
mat_b: torch.Tensor,
|
||||
scales_a: torch.Tensor,
|
||||
scales_b: torch.Tensor,
|
||||
out_dtype: torch.dtype,
|
||||
bias: Optional[torch.Tensor] = None,
|
||||
) -> bool:
|
||||
if (
|
||||
backend is KernelBackend.AOT
|
||||
and super().backend_eligible(
|
||||
KernelBackend.TORCH,
|
||||
mat_a,
|
||||
mat_b,
|
||||
scales_a,
|
||||
scales_b,
|
||||
out_dtype,
|
||||
bias,
|
||||
)
|
||||
and _prefer_torch_rowwise_fp8(
|
||||
mat_a, mat_b, scales_a, scales_b, out_dtype, bias
|
||||
)
|
||||
):
|
||||
return False
|
||||
return super().backend_eligible(
|
||||
backend, mat_a, mat_b, scales_a, scales_b, out_dtype, bias
|
||||
)
|
||||
|
||||
def forward_native(
|
||||
self,
|
||||
mat_a: torch.Tensor,
|
||||
mat_b: torch.Tensor,
|
||||
scales_a: torch.Tensor,
|
||||
scales_b: torch.Tensor,
|
||||
out_dtype: torch.dtype,
|
||||
bias: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
import torch
|
||||
|
||||
m, n = mat_a.shape[0], mat_b.shape[1]
|
||||
scale_a = scales_a.reshape(m, 1) if scales_a.numel() == m else scales_a
|
||||
scale_b = scales_b.reshape(1, n) if scales_b.numel() == n else scales_b
|
||||
return torch._scaled_mm(
|
||||
mat_a,
|
||||
mat_b,
|
||||
scale_a=scale_a,
|
||||
scale_b=scale_b,
|
||||
out_dtype=out_dtype,
|
||||
bias=bias,
|
||||
use_fast_accum=True,
|
||||
)
|
||||
|
||||
def forward_aot(
|
||||
self,
|
||||
mat_a: torch.Tensor,
|
||||
mat_b: torch.Tensor,
|
||||
scales_a: torch.Tensor,
|
||||
scales_b: torch.Tensor,
|
||||
out_dtype: torch.dtype,
|
||||
bias: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
import sgl_kernel
|
||||
|
||||
return sgl_kernel.fp8_scaled_mm(
|
||||
mat_a, mat_b, scales_a, scales_b, out_dtype, bias
|
||||
)
|
||||
|
||||
|
||||
_FP8_SCALED_MM = register_fused_op(Fp8ScaledMMOp(), __name__, "_FP8_SCALED_MM")
|
||||
|
||||
register_kernel(
|
||||
KernelSpec(
|
||||
op="gemm.bmm_fp8",
|
||||
@@ -91,10 +223,8 @@ def fp8_scaled_mm(
|
||||
out_dtype: torch.dtype,
|
||||
bias: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
"""FP8 scaled matmul: ``(mat_a @ mat_b) * scales_a * scales_b (+ bias)``."""
|
||||
return get_kernel("gemm.fp8_scaled_mm", KernelBackend.AOT)(
|
||||
mat_a, mat_b, scales_a, scales_b, out_dtype, bias
|
||||
)
|
||||
"""FP8 matmul with per-row A and per-column B scales."""
|
||||
return _FP8_SCALED_MM(mat_a, mat_b, scales_a, scales_b, out_dtype, bias)
|
||||
|
||||
|
||||
def bmm_fp8(
|
||||
@@ -133,7 +263,13 @@ def dsv3_router_gemm(
|
||||
return impl(hidden_states, router_weights, out_dtype, output)
|
||||
|
||||
|
||||
__all__ = ["fp8_scaled_mm", "bmm_fp8", "dsv3_fused_a_gemm", "dsv3_router_gemm"]
|
||||
__all__ = [
|
||||
"Fp8ScaledMMOp",
|
||||
"fp8_scaled_mm",
|
||||
"bmm_fp8",
|
||||
"dsv3_fused_a_gemm",
|
||||
"dsv3_router_gemm",
|
||||
]
|
||||
|
||||
|
||||
# LoRA SGMV Triton kernels migrated into this group (from lora/triton_ops);
|
||||
|
||||
@@ -202,8 +202,7 @@ if _use_aiter:
|
||||
|
||||
|
||||
if _is_cuda:
|
||||
from sgl_kernel import fp8_scaled_mm
|
||||
|
||||
from sglang.kernels.ops.gemm import fp8_scaled_mm
|
||||
from sglang.kernels.ops.gemm.fp8_blockwise_gemm import fp8_blockwise_scaled_mm
|
||||
from sglang.srt.utils.patch_torch import register_fake_if_exists
|
||||
|
||||
|
||||
@@ -25,7 +25,7 @@ EXPECTED = {
|
||||
"activation.relu2": {"jit", "torch", "torch_compile"},
|
||||
"layernorm.rmsnorm": {"aot", "jit", "aiter", "torch_npu", "torch", "torch_compile"},
|
||||
"layernorm.gemma_rmsnorm": {"aot", "jit", "torch_npu", "torch", "torch_compile"},
|
||||
"gemm.fp8_scaled_mm": {"aot"},
|
||||
"gemm.fp8_scaled_mm": {"aot", "torch", "torch_compile"},
|
||||
"moe.moe_align_block_size": {"aot", "jit"},
|
||||
"quantization.nvfp4_gemm_swiglu_nvfp4_quant": {"cute_dsl"},
|
||||
"kvcache.reshape_and_cache_flash": {"triton"},
|
||||
@@ -84,7 +84,20 @@ def test_sparse_linear_attention_registry_targets_forward_kernel():
|
||||
|
||||
|
||||
def test_single_backend_resolves_without_backend():
|
||||
assert K.select_kernel("gemm.fp8_scaled_mm").backend is KernelBackend.AOT
|
||||
assert (
|
||||
K.select_kernel("kvcache.reshape_and_cache_flash").backend
|
||||
is KernelBackend.TRITON
|
||||
)
|
||||
|
||||
|
||||
def test_fp8_scaled_mm_requires_explicit_registry_backend(monkeypatch):
|
||||
monkeypatch.setattr(sel, "_platform", lambda: _SM90)
|
||||
with pytest.raises(ValueError, match="multiple backends"):
|
||||
K.select_kernel("gemm.fp8_scaled_mm")
|
||||
assert (
|
||||
K.select_kernel("gemm.fp8_scaled_mm", backend=KernelBackend.AOT).backend
|
||||
is KernelBackend.AOT
|
||||
)
|
||||
|
||||
|
||||
def test_unknown_op_or_backend_raises():
|
||||
|
||||
Reference in New Issue
Block a user