Prhub

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

原始 PR 作者 b8zhong 合并时间 2026-07-22 07:44 文件变更 15 提交数 7 评论 4 代码增减 +63 / -237

执行摘要

删除 sgl-kernel AOT bmm_fp8,统一使用 flashinfer 实现

sgl-kernel 中的 bmm_fp8.cu 几乎完全复制自 flashinfer 的实现(commit 消息说明),flashinfer 已是硬依赖且支持相同的计算能力(SM89+),故无理由保留独立副本。删除后减少内核维护负担,同时保证功能完全一致。

值得一读,展示了如何安全地移除重复 AOT 实现并统一依赖,同时利用 register_custom_op 保持 torch.compile 兼容性。

讨论亮点

此 PR 无 review 评论,直接由 BBuf 批准合并。但 commit 修复了 capability -> capabilities 拼写(关联 #31292),说明 kernel 注册接口有演进。

实现拆解

  1. 删除 sgl-kernel 的 bmm_fp8 AOT 实现:移除 bmm_fp8.cusgl_kernel_ops.h 中声明、common_extension.cc/common_extension_musa.cc 中注册,以及 gemm.py 中的 _bmm_fp8_internalbmm_fp8 函数。
  2. fp8_utils.py 中新增基于 flashinfer.bmm_fp8 的封装:通过 register_custom_op 注册 _bmm_fp8_batched_op 确保 torch.compile 安全,并暴露 bmm_fp8 函数。
  3. 将所有调用点统一到新入口:修改 minicpm3.pyforward_mla.pysarvam_moe.pyforward_mla_fused_rope_rocm.py,从 sglang.kernels.ops.gemm 导入 bmm_fp8
  4. sglang/kernels/ops/gemm/__init__.py 中注册 bmm_fp8 作为 KernelBackend.FLASHINFER 条目。
  5. 删除 sgl-kernel/tests/test_bmm_fp8.py(flashinfer 已有覆盖)。
文件 模块 状态 重要度
sgl-kernel/python/sgl_kernel/gemm.py 内核层 modified 7.47
python/sglang/srt/layers/quantization/fp8_utils.py 量化层 modified 7.32
python/sglang/srt/models/minicpm3.py 模型层 modified 7.58
sgl-kernel/tests/test_bmm_fp8.py 内核测试 removed 6.64
sgl-kernel/csrc/gemm/bmm_fp8.cu CUDA 内核 removed 5.35
python/sglang/kernels/ops/gemm/__init__.py Kernel 注册 modified 4.91

关键符号

_bmm_fp8_internal bmm_fp8 _bmm_fp8_batched_op _bmm_fp8_op

关键源码片段

python/sglang/srt/layers/quantization/fp8_utils.py core-logic

新增了基于 flashinfer.bmm_fp8 的封装,包括 custom_op 注册和 bmm_fp8 函数,是统一入口的核心文件。

# fp8_utils.py 新增部分:基于 flashinfer 的 bmm_fp8 封装
# 确保 torch.compile 不会跟踪 cuBLAS handlefrom 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:
    """封装 flashinfer.bmm_fp8,通过 custom_op 避免 torch.compile 报错。"""
    _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

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

核心风险在于 flashinfer 版本的 bmm_fp8 行为是否完全一致。作者已验证 bit-identical 输出和匹配延迟。对于非 CUDA 平台(AMD/MThreads),该功能原本就在 if _is_cuda: 条件下,因此无影响。若 flashinfer 未来变更接口,依赖同步更新即可。删除的测试文件减轻了维护负担。

对最终用户透明,功能无变化。对开发者:减少约 240 行代码,消除重复 kernel 维护负担,统一依赖至 flashinfer。对系统:无性能影响。

依赖外部库(flashinfer) 已验证 bit-identical

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论