文章总结: 本文分析sglang在Blackwell芯片上fp8推理速度慢的问题,指出fp8scaledmm默认使用128tile大小,导致M较小时效率低下。文章对比了SM100的动态调度实现,并给出了基于Triton的自定义内核修改方案,通过编译期常量M和权重复用提升性能。建议针对小M场景优化tile大小或采用动态调度策略。
综合评分: 75
文章分类: 其他
【DGX Spark优化加速】sm120上sglang的fp8的triton内核的矩阵乘性能问题
huoji
huoji
冲鸭安全
2026年10月2日 17:10
北京
在小说阅读器读本章
去阅读
在公众号小说中沉浸阅读
前言
问题起因是我发现sglang上在blackwell速度和H系列一样了,最搞的是,不量化速度反而快于是深入跑压测研究了一下。让AI跑了一下性能测试,发现原因在于这里
原因
在推理服务端中,有个参数是 M,M 是这一次前向计算里要处理的 token 行数
输出 [M × N] = 输入激活 [M × K] × 权重 [K × N]
- K:输入维度,比如 hidden size 5120;
- N:输出维度,比如 MLP 的 17408;
- M:这次送进来多少个 token,每个 token 是一行。
K 和 N 由模型结构决定,每一层都是固定的。M 取决于这一步调度器塞进来多少 token,每一步都可能不一样
而sglang的fp8实现中:
https://github.com/sgl-project/sglang/blob/main/python/sglang/srt/layers/quantization/fp8_utils.py#L2320
output = fp8_scaled_mm(
qinput,
weight,
x_scale,
weight_scale,
out_dtype=output_dtype,
bias=bias,
)
问题在于,不管 M 多大都直接调 fp8_scaled_mm,而fp8_scaled_mm默认是128
void sm120_fp8_dispatch_bias(...) {
using CTAShapeDefault = Shape<_128, _128, _128>; // ⚠ 问题点 2:只有这一种 tile,M/N/K 方向都是 128
using ClusterShapeDefault = Shape<_1, _1, _1>;
using MainloopScheduleType = cutlass::gemm::collective::KernelScheduleAuto;
using EpilogueScheduleType = cutlass::epilogue::collective::EpilogueScheduleAuto;
using TileSchedulerType = void; // ⚠ 问题点 3:没有 split-K 或 stream-K 调度,
// 每个 CTA 要独自沿 K 方向从头读到尾
...
using GemmDefault = DeviceGemmFp8RowwiseSm120<..., CTAShapeDefault, ...>;
...
}
128 行的 tile 里只有 1–4 行有效,还只能拆出 40–272 个 CTA,读权重的速度上不去.
有意思的是,官方 SM100的写了这种动态的
https://github.com/sgl-project/sglang/blob/41bd213c5f2d7d13b77bd2bd3b370e46022139b3/python/sglang/kernels/aot/csrc/gemm/fp8_gemm_kernel.cu#L791
if (bias) {
if (mp2 <= 16) {
// m in [1, 16]
return launch_sm100_fp8_scaled_mm<BiasGemm16, true>(out, a, b, scales_a, scales_b, bias);
} else if (mp2 <= 64) {
// m in (16, 64]
return launch_sm100_fp8_scaled_mm<BiasGemm64, true>(out, a, b, scales_a, scales_b, bias);
} else if (mp2 <= 256) {
// m in (64, 256]
return launch_sm100_fp8_scaled_mm<BiasGemm256, true>(out, a, b, scales_a, scales_b, bias);
} else {
// m in (256, inf]
return launch_sm100_fp8_scaled_mm<BiasGemmDefault, true>(out, a, b, scales_a, scales_b, bias);
}
} else {
if (mp2 <= 16) {
// m in [1, 16]
return launch_sm100_fp8_scaled_mm<Gemm16, false>(out, a, b, scales_a, scales_b, bias);
} else if (mp2 <= 64) {
// m in (16, 64]
return launch_sm100_fp8_scaled_mm<Gemm64, false>(out, a, b, scales_a, scales_b, bias);
} else if (mp2 <= 256) {
// m in (64, 256]
return launch_sm100_fp8_scaled_mm<Gemm256, false>(out, a, b, scales_a, scales_b, bias);
} else {
return launch_sm100_fp8_scaled_mm<GemmDefault, false>(out, a, b, scales_a, scales_b, bias);
}
}
}
所以这就导致在blackwell芯片上,sglang的fp8推理速度非常慢
修改方案,ai给的:
@triton.jit
def _fp8_decode_gemv_kernel(x_ptr, w_ptr, xs_ptr, ws_ptr, out_ptr, N, K, ...,
M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr):
offs_n = tl.program_id(0) * BLOCK_N + tl.arange(0, BLOCK_N) # 每个 program 负责 8 列输出
acc0 = acc1 = acc2 = acc3 = 0(fp32) # 每行激活一个累加器
for k0 in range(0, K, BLOCK_K): # 沿 K 方向每次走 512
w = tl.load(权重块 [8 × 512], eviction_policy="evict_first") # 这块权重只从显存读一次
acc0 += sum(w * x[0, k0:k0+512]) # 用同一块权重
if M > 1: acc1 += sum(w * x[1, ...]) # 依次乘 1~4 行激活
if M > 2: acc2 += ... # M 是编译期常量,
if M > 3: acc3 += ... # 用不到的分支编译时直接删掉
out[m, n] = acc_m * x_scale[m] * w_scale[n] → bf16 # 最后统一乘 scale
提升:
免责声明:
本文所载程序、技术方法仅面向合法合规的安全研究与教学场景,旨在提升网络安全防护能力,具有明确的技术研究属性。
任何单位或个人未经授权,将本文内容用于攻击、破坏等非法用途的,由此引发的全部法律责任、民事赔偿及连带责任,均由行为人独立承担,本站不承担任何连带责任。
本站内容均为技术交流与知识分享目的发布,若存在版权侵权或其他异议,请通过邮件联系处理,具体联系方式可点击页面上方的联系我。
本文转载自:冲鸭安全 huoji
huoji《【DGX Spark优化加速】sm120上sglang的fp8的triton内核的矩阵乘性能问题》