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:
@@ -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:
|
||||
|
||||
+1
-1
@@ -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 (
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user