[Kernel] Route large SM90 row/column-scaled FP8 GEMMs to Torch (#34318)

This commit is contained in:
Gregory Leleytner
2026-08-29 07:35:37 +08:00
committed by GitHub
parent 96a4dcdde8
commit e1b3bba3cc
3 changed files with 167 additions and 19 deletions
+151 -15
View File
@@ -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():