Prhub

#30247 Optimize LongCat-Flash router GEMM with the HPC-Ops bf16xfp32 kernel

原始 PR 作者 BBuf 合并时间 2026-07-21 20:05 文件变更 4 提交数 14 评论 10 代码增减 +336 / -7

执行摘要

用 HPC-Ops bf16xfp32 内核加速 LongCat-Flash 路由器 GEMM

LongCat-Flash 模型的路由器 GEMM(bf16 激活 × fp32 权重)是 moe 路由部分的计算热点。当前实现使用 fp32 矩阵乘法(将激活转换为 float),不能充分利用 bf16 tensor core。HPC-Ops 库的 gemm_bf16xfp32 内核通过分解权重并融合两个 bf16 GEMM,在保持 fp32 权重精度的同时获得大幅加速。PR body 中展示了详细性能数据,对关键形状加速比达 2-4 倍。

值得精读,尤其是权重拆分缓存的设计、通过 min-M 阈值在性能与兼容性之间平衡的策略,以及外部核集成模式。同时关注后续 #31943 对缓存一致性的修复,以填补本 PR 遗留的安全缺口。

讨论亮点

Review 中主要讨论了三个核心问题:

  • 权重缓存一致性问题:VAthree 指出在 param.data.copy_() 进行在线权重更新时,缓存键中的 _version 不会递增,导致缓存未失效而返回旧权重。BBuf 确认并在后续 PR #31943 中修复,改为在权重更新 API 中主动刷新缓存或直接报错。
  • 统一 GEMM 接口建议:DarkSharpness 建议将多种 GEMM 统一到一个类似 torch.mm 的接口,BBuf 同意并跟踪到 issue #29630。
  • 性能优化建议:gemini-code-assist[bot] 建议缓存 get_device_capability 调用;但代码中已通过 functools.cache 缓存了整个可用性检查,覆盖了此建议。

实现拆解

  1. python/sglang/jit_kernel/dsv4/gemm.py 中新增 HPC-Ops 分支:添加 _hpc_gemm_bf16xfp32_available() 检查 HPC-Ops 是否安装且 GPU 为 Hopper (SM90);_can_use_hpc_gemm_bf16xfp32() 进行形状、连续性等检查;_get_bf16xfp32_weight_split() 将 fp32 权重拆分为两个 bf16 部分并缓存到权重张量上(使用 _sglang_bf16xfp32_weight_cache 属性);_linear_bf16_fp32_hpc() 调用 HPC-Ops 内核,若检查不通过返回 None。修改 linear_bf16_fp32() 入口,新增 hpc_kernel_min_m 参数以支持按形状调度。
  2. python/sglang/srt/models/longcat_flash.py 中定义每个路由器形状的最小 M 阈值字典 _LONGCAT_FLASH_ROUTER_HPC_GEMM_MIN_M,在 LongcatFlashRouter.forward() 中对匹配形状且无 router_bias、权重 dtype 为 fp32 的情况直接调用 linear_bf16_fp32 并传入阈值,否则回退至原始 classifier 前向路径。
  3. 添加两个测试文件:test/registered/gemm/test_linear_bf16_fp32_hpc.py 验证 HPC 路径与 fp32 参考的数值一致性、min_M 下界行为以及权重拆分缓存的复用;test/registered/unit/models/test_longcat_flash_router_hpc_gemm.py 通过 mock 验证 LongCat 路由器向 HPC 内核的正确调度以及未受支持形状的回退。
  4. 配置开关:通过环境变量 SGLANG_OPT_BF16_FP32_GEMM_ALGO=hpc 可全局启用,但 LongCat 模型默认使用内部阈值路径;无 HPC-Ops 时所有路径自动回退至 cublas,对其他用户无影响。
文件 模块 状态 重要度
python/sglang/jit_kernel/dsv4/gemm.py JIT 内核 modified 8.61
python/sglang/srt/models/longcat_flash.py 模型定义 modified 6.98
test/registered/gemm/test_linear_bf16_fp32_hpc.py 单元测试 added 7.57
test/registered/unit/models/test_longcat_flash_router_hpc_gemm.py 单元测试 added 8.07

关键符号

linear_bf16_fp32 _linear_bf16_fp32_hpc _get_bf16xfp32_weight_split _can_use_hpc_gemm_bf16xfp32 _hpc_gemm_bf16xfp32_available _linear_bf16_fp32_cublas LongcatFlashRouter.forward

分析完成后,这里会展示 LLM 生成的相对完整源码片段和详细注释。

评论区精华

权重缓存在线更新时可能变 stale 正确性

VAthree 发现在 `param.data.copy_()` 进行在线权重更新时,缓存键中的 `_version` 不会递增,且 `data_ptr`、形状等均不变,导致缓存返回旧权重。BBuf 确认这是一个静默正确性问题,并在后续 PR #31943 中修复。

结论:已解决:在权重更新 API 中主动刷新权重拆分缓存,或在无法刷新时快速失败。 · 已解决

统一 GEMM 接口以改善代码复用 设计

DarkSharpness 建议将多种 GEMM(cublas、deep_gemm、aiter、HPC-Ops)统一到一个类似 `torch.mm` 的接口,自动根据形状和后端调度。BBuf 同意,但认为应作为独立的后续重构,避免本 PR 范围膨胀。

结论:创建 issue #29630 跟踪统一 GEMM 接口重构。 · follow-up

缓存 get_device_capability 以减少 Python 开销 性能

gemini-code-assist[bot] 建议使用 `functools.lru_cache` 缓存 `get_device_capability` 查询,避免每次前向时重复调用。实际上代码中 `_hpc_gemm_bf16xfp32_available` 已通过 `@functools.cache` 缓存整个函数的返回结果,隐含覆盖了此建议。

结论:代码已覆盖,无需额外修改。 · 已采纳

风险与影响

  • 权重缓存 stale 风险:在线权重更新场景下,缓存不能正确失效,可能导致静默错误。已在 #31943 中通过强制刷新或失败快速修复,但本 PR 合并时尚未包含该修复。
  • HPC-Ops 依赖与平台限制:内核仅运行于 Hopper (SM90) GPU,且需要安装 HPC-Ops 库。无此环境时自动回退,但回退路径必须经过严格测试。
  • min-M 阈值泛化性:阈值基于 H200 基准测试,在 H100 等其他 Hopper 设备上可能不是最优,但仅影响加速比例,不影响正确性。
  • CUDA graph 兼容性:权重拆分缓存的固定地址可能被 CUDA graph 捕获,若图重播时权重发生变化可能导致错误。已在 #31943 中通过持久化缓存和失败快速处理。
  • 用户:LongCat-Flash 模型用户无需修改配置即可获得 prefill 吞吐提升(2-5%),decode 阶段因 token 数不足 min_M 阈值而保持原性能。其他模型用户不受影响。
  • 系统:需要 Hopper GPU 和 HPC-Ops 库(已由 #31390 集成到镜像)。
  • 团队:维护负担低,内核由 HPC-Ops 上游维护;后续内核调优只需更新 pins 版本。
权重缓存 stale 风险 HPC-Ops 外部依赖 仅 Hopper GPU 加速

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论