【DGXSpark优化加速】sm120上sglang的fp8的triton内核的矩阵乘性能问题

admin 2026-10-05 04:36:11 网络安全文章 来源:ZONE.CI 全球网 0 阅读模式

文章总结: 本文分析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(...) {
&nbsp; using CTAShapeDefault = Shape<_128, _128, _128>; &nbsp; // ⚠ 问题点 2:只有这一种 tile,M/N/K 方向都是 128
&nbsp; using ClusterShapeDefault = Shape<_1, _1, _1>;
&nbsp; using MainloopScheduleType = cutlass::gemm::collective::KernelScheduleAuto;
&nbsp; using EpilogueScheduleType = cutlass::epilogue::collective::EpilogueScheduleAuto;
&nbsp; using TileSchedulerType = void; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp;// ⚠ 问题点 3:没有 split-K 或 stream-K 调度,
&nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp;// &nbsp; 每个 CTA 要独自沿 K 方向从头读到尾
&nbsp; ...
&nbsp; using GemmDefault = DeviceGemmFp8RowwiseSm120<..., CTAShapeDefault, ...>;
&nbsp; ...
}

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

&nbsp; if (bias) {
&nbsp; &nbsp; if (mp2 <= 16) {
&nbsp; &nbsp; &nbsp; // m in [1, 16]
&nbsp; &nbsp; &nbsp; return launch_sm100_fp8_scaled_mm<BiasGemm16, true>(out, a, b, scales_a, scales_b, bias);
&nbsp; &nbsp; } else if (mp2 <= 64) {
&nbsp; &nbsp; &nbsp; // m in (16, 64]
&nbsp; &nbsp; &nbsp; return launch_sm100_fp8_scaled_mm<BiasGemm64, true>(out, a, b, scales_a, scales_b, bias);
&nbsp; &nbsp; } else if (mp2 <= 256) {
&nbsp; &nbsp; &nbsp; // m in (64, 256]
&nbsp; &nbsp; &nbsp; return launch_sm100_fp8_scaled_mm<BiasGemm256, true>(out, a, b, scales_a, scales_b, bias);
&nbsp; &nbsp; } else {
&nbsp; &nbsp; &nbsp; // m in (256, inf]
&nbsp; &nbsp; &nbsp; return launch_sm100_fp8_scaled_mm<BiasGemmDefault, true>(out, a, b, scales_a, scales_b, bias);
&nbsp; &nbsp; }
&nbsp; } else {
&nbsp; &nbsp; if (mp2 <= 16) {
&nbsp; &nbsp; &nbsp; // m in [1, 16]
&nbsp; &nbsp; &nbsp; return launch_sm100_fp8_scaled_mm<Gemm16, false>(out, a, b, scales_a, scales_b, bias);
&nbsp; &nbsp; } else if (mp2 <= 64) {
&nbsp; &nbsp; &nbsp; // m in (16, 64]
&nbsp; &nbsp; &nbsp; return launch_sm100_fp8_scaled_mm<Gemm64, false>(out, a, b, scales_a, scales_b, bias);
&nbsp; &nbsp; } else if (mp2 <= 256) {
&nbsp; &nbsp; &nbsp; // m in (64, 256]
&nbsp; &nbsp; &nbsp; return launch_sm100_fp8_scaled_mm<Gemm256, false>(out, a, b, scales_a, scales_b, bias);
&nbsp; &nbsp; } else {
&nbsp; &nbsp; &nbsp; return launch_sm100_fp8_scaled_mm<GemmDefault, false>(out, a, b, scales_a, scales_b, bias);
&nbsp; &nbsp; }
&nbsp; }
}

所以这就导致在blackwell芯片上,sglang的fp8推理速度非常慢 修改方案,ai给的:

@triton.jit
def _fp8_decode_gemv_kernel(x_ptr, w_ptr, xs_ptr, ws_ptr, out_ptr, N, K, ...,
&nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr):
&nbsp; &nbsp; offs_n = tl.program_id(0) * BLOCK_N + tl.arange(0, BLOCK_N) &nbsp; # 每个 program 负责 8 列输出
&nbsp; &nbsp; acc0 = acc1 = acc2 = acc3 = 0(fp32) &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp;# 每行激活一个累加器
&nbsp; &nbsp; for k0 in range(0, K, BLOCK_K): &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; # 沿 K 方向每次走 512
&nbsp; &nbsp; &nbsp; &nbsp; w = tl.load(权重块 [8 × 512], eviction_policy="evict_first") &nbsp;# 这块权重只从显存读一次
&nbsp; &nbsp; &nbsp; &nbsp; acc0 += sum(w * x[0, k0:k0+512]) &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; # 用同一块权重
&nbsp; &nbsp; &nbsp; &nbsp; if M > 1: acc1 += sum(w * x[1, ...]) &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; # 依次乘 1~4 行激活
&nbsp; &nbsp; &nbsp; &nbsp; if M > 2: acc2 += ... &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp;# M 是编译期常量,
&nbsp; &nbsp; &nbsp; &nbsp; if M > 3: acc3 += ... &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp;# 用不到的分支编译时直接删掉
&nbsp; &nbsp; out[m, n] = acc_m * x_scale[m] * w_scale[n] &nbsp;→ bf16 &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp; &nbsp;# 最后统一乘 scale

提升:


免责声明:

本文所载程序、技术方法仅面向合法合规的安全研究与教学场景,旨在提升网络安全防护能力,具有明确的技术研究属性。

任何单位或个人未经授权,将本文内容用于攻击、破坏等非法用途的,由此引发的全部法律责任、民事赔偿及连带责任,均由行为人独立承担,本站不承担任何连带责任。

本站内容均为技术交流与知识分享目的发布,若存在版权侵权或其他异议,请通过邮件联系处理,具体联系方式可点击页面上方的联系我。

本文转载自:冲鸭安全 huoji huoji《【DGX Spark优化加速】sm120上sglang的fp8的triton内核的矩阵乘性能问题》

403Forbidden绕过 网络安全文章

403Forbidden绕过

文章总结: 本文分析了403Forbidden错误的工作原理和常见原因,包括IP地址封锁、权限配置错误、代理配置错误等,并分享了绕过403错误的技术,如HTTP
评论:0   参与:  0