Delete sgl-kernel AOT bmm_fp8, use flashinfer.bmm_fp8 (#31202)

Co-authored-by: root <root@sgl-b300-inference.datacrunch.io>
This commit is contained in:
Brayden Zhong
2026-07-22 07:44:47 +08:00
committed by GitHub
co-authored by root
parent 1b4cb6b8c1
commit 2f4f2362fb
15 changed files with 63 additions and 237 deletions
+28 -1
View File
@@ -30,6 +30,19 @@ register_kernel(
description="FP8 scaled matmul (sgl_kernel wheel).",
)
)
register_kernel(
KernelSpec(
op="gemm.bmm_fp8",
backend=KernelBackend.FLASHINFER,
target="sglang.srt.layers.quantization.fp8_utils:bmm_fp8",
capabilities=_CUDA,
format_signature=FormatSignature(
supported_dtypes=("float8_e4m3fn", "float8_e5m2"),
description="batched (3D) per-tensor-scale FP8 matmul: D = A_fp8 @ B_fp8 * A_scale * B_scale",
),
description="Batched FP8 matmul (flashinfer cuBLAS backend, torch.compile-safe wrapper).",
)
)
register_kernel(
KernelSpec(
op="gemm.dsv3_fused_a_gemm",
@@ -84,6 +97,20 @@ def fp8_scaled_mm(
)
def bmm_fp8(
A: torch.Tensor,
B: torch.Tensor,
A_scale: torch.Tensor,
B_scale: torch.Tensor,
dtype: torch.dtype,
out: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Batched (3D) per-tensor-scale FP8 matmul, via flashinfer's cuBLAS backend."""
return get_kernel("gemm.bmm_fp8", KernelBackend.FLASHINFER)(
A, B, A_scale, B_scale, dtype, out
)
def dsv3_fused_a_gemm(
mat_a: torch.Tensor,
mat_b: torch.Tensor,
@@ -106,7 +133,7 @@ def dsv3_router_gemm(
return impl(hidden_states, router_weights, out_dtype, output)
__all__ = ["fp8_scaled_mm", "dsv3_fused_a_gemm", "dsv3_router_gemm"]
__all__ = ["fp8_scaled_mm", "bmm_fp8", "dsv3_fused_a_gemm", "dsv3_router_gemm"]
# LoRA SGMV Triton kernels migrated into this group (from lora/triton_ops);
@@ -173,6 +173,36 @@ if _is_cuda:
N = mat_b.shape[-1]
return mat_a.new_empty((M, N), dtype=out_dtype)
from flashinfer import bmm_fp8 as _raw_bmm_fp8_batched
@register_custom_op(op_name="flashinfer_bmm_fp8_batched", mutates_args=["out"])
def _bmm_fp8_batched_op(
A: torch.Tensor,
B: torch.Tensor,
out: torch.Tensor,
A_scale: torch.Tensor,
B_scale: torch.Tensor,
) -> None:
_raw_bmm_fp8_batched(A, B, A_scale, B_scale, out.dtype, out)
def bmm_fp8(
A: torch.Tensor,
B: torch.Tensor,
A_scale: torch.Tensor,
B_scale: torch.Tensor,
dtype: torch.dtype,
out: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Batched (3D) per-tensor-scale FP8 matmul, via flashinfer's cuBLAS backend."""
if out is None:
out = torch.empty(
(A.shape[0], A.shape[1], B.shape[2]),
device=A.device,
dtype=dtype,
)
_bmm_fp8_batched_op(A, B, out, A_scale, B_scale)
return out
use_triton_w8a8_fp8_kernel = get_bool_env_var("USE_TRITON_W8A8_FP8_KERNEL")
@@ -88,30 +88,7 @@ class MlaBmmFusionPlan:
if _is_cuda:
from sgl_kernel import bmm_fp8 as _raw_bmm_fp8
# TODO(yuwei): remove this wrapper after sgl-kernel registers its own fake/meta impl
# Wrap bmm_fp8 as a custom op so torch.compile does not trace into
# torch.cuda.current_blas_handle() (which returns a non-Tensor).
@register_custom_op(mutates_args=["out"])
def _bmm_fp8_op(
A: torch.Tensor,
B: torch.Tensor,
out: torch.Tensor,
A_scale: torch.Tensor,
B_scale: torch.Tensor,
) -> None:
_raw_bmm_fp8(A, B, A_scale, B_scale, out.dtype, out)
def bmm_fp8(A, B, A_scale, B_scale, dtype, out=None):
if out is None:
out = torch.empty(
(A.shape[0], A.shape[1], B.shape[2]),
device=A.device,
dtype=dtype,
)
_bmm_fp8_op(A, B, out, A_scale, B_scale)
return out
from sglang.kernels.ops.gemm import bmm_fp8
if _use_aiter:
@@ -21,7 +21,7 @@ if TYPE_CHECKING:
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA
if _is_cuda:
from sgl_kernel import bmm_fp8
from sglang.kernels.ops.gemm import bmm_fp8
if _is_hip:
from sglang.kernels.ops.attention.rocm_mla_decode_rope import (
+1 -26
View File
@@ -43,32 +43,7 @@ from sglang.srt.utils import add_prefix, is_cuda
from sglang.srt.utils.hf_transformers_utils import get_rope_config
if is_cuda():
from sgl_kernel import bmm_fp8 as _raw_bmm_fp8
from sglang.srt.utils.custom_op import register_custom_op
# TODO(yuwei): remove this wrapper after sgl-kernel registers its own fake/meta impl
# Wrap bmm_fp8 as a custom op so torch.compile does not trace into
# torch.cuda.current_blas_handle() (which returns a non-Tensor).
@register_custom_op(mutates_args=["out"])
def _bmm_fp8_op(
A: torch.Tensor,
B: torch.Tensor,
out: torch.Tensor,
A_scale: torch.Tensor,
B_scale: torch.Tensor,
) -> None:
_raw_bmm_fp8(A, B, A_scale, B_scale, out.dtype, out)
def bmm_fp8(A, B, A_scale, B_scale, dtype, out=None):
if out is None:
out = torch.empty(
(A.shape[0], A.shape[1], B.shape[2]),
device=A.device,
dtype=dtype,
)
_bmm_fp8_op(A, B, out, A_scale, B_scale)
return out
from sglang.kernels.ops.gemm import bmm_fp8
class MiniCPM3MLP(nn.Module):
+2 -1
View File
@@ -81,9 +81,10 @@ _is_cublas_ge_129 = is_nvidia_cublas_version_ge_12_9()
if _is_cuda:
try:
from sgl_kernel import bmm_fp8, merge_state_v2
from sgl_kernel import merge_state_v2
from sglang.jit_kernel.concat_mla import concat_mla_k
from sglang.kernels.ops.gemm import bmm_fp8
from sglang.kernels.ops.quantization.fp8_kernel import per_tensor_quant_mla_fp8
_has_fp8_support = True