diff --git a/python/sglang/kernels/ops/gemm/__init__.py b/python/sglang/kernels/ops/gemm/__init__.py index b6670d8bd..936f0987d 100644 --- a/python/sglang/kernels/ops/gemm/__init__.py +++ b/python/sglang/kernels/ops/gemm/__init__.py @@ -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); diff --git a/python/sglang/srt/layers/quantization/fp8_utils.py b/python/sglang/srt/layers/quantization/fp8_utils.py index bbd70a438..74453d3e5 100755 --- a/python/sglang/srt/layers/quantization/fp8_utils.py +++ b/python/sglang/srt/layers/quantization/fp8_utils.py @@ -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 diff --git a/test/registered/kernels/ops/layernorm/test_kernels_namespace.py b/test/registered/kernels/ops/layernorm/test_kernels_namespace.py index 1dec217dc..d8483c4bf 100644 --- a/test/registered/kernels/ops/layernorm/test_kernels_namespace.py +++ b/test/registered/kernels/ops/layernorm/test_kernels_namespace.py @@ -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():