Prhub

#50185 attn_res kernel latency improvements

原始 PR 作者 gnovack 合并时间 2026-08-06 23:53 文件变更 1 提交数 3 评论 2 代码增减 +84 / -39

执行摘要

优化 attn_res 内核延迟,微基准最高提速约 27%

Kimi K3 的注意力残差内核(attn_res)在 decode 与 prefill 场景都被频繁调用,是线上 TPOT/TTFT 的关键路径之一。PR body 明确列出四项小改进以降低内核延迟:q_cache 向量化加载、B/N 编译期常量、低 batch 路径 NC=3、num_chunks 主循环展开,并给出逐项叠加的微基准数据,说明这是以低风险手段换取可量化延迟收益的性能优化。

值得精读。它展示了 CUDA 内核延迟优化的经典组合拳:编译期参数化 + 向量化访存 + 循环展开 + 正确的 memory 语义,且每步都有可量化的微基准验证,适合作为内核性能优化的参考范式。建议关注 mbarrier memory clobber 的并发正确性解释与模板分派的扩展性。

讨论亮点

review 阶段三位维护者(gau-nernst、zyongye、ZJY0516)直接 APPROVED,无实质代码讨论。值得注意的实现决策是:mbarrier 内联汇编补 memory clobber 属于循环展开带来的并发正确性配套,虽未在讨论中被展开,但它是本次展开优化的必要前提。

实现拆解

实现围绕 csrc/libtorch_stable/kimi_k3/attn_res_kernel.cu 展开,共 4 步:

  1. 模板参数化 B 与 N:将 attn_res_fwd_online_v2_kernel 的模板从 <int H, int NC, ...> 扩展为 <int H, int N, int NC, int B, ...>,把运行时的 N(隐藏维归一化长度)与 B(batch)提升为编译期常量,从而让 num_chunks = (N + N_CHUNK - 1) / N_CHUNK 变成 constexpr,支持后续主循环完全展开;launch_fwd 同步去掉运行时 NB 参数。

  2. q_cache 向量化填充:在 H==7168 特化分支与通用分支中,用 int2/int4 向量读入 rms_wres_w,再经 __nv_bfloat162__bfloat1622float2 一次性计算两组/四组乘积,替代逐元素标量加载,减少指令数与访存次数。

  3. 低 batch 路径 NC=3:在 kimi_k3_attn_res 的通用分支中,将运行时 num_blocksswitch 分派(case 1、case 2、default)到不同的编译期 NSRC,并对每组调用 launch_fwd<7168, NSRC, 3, ...>,使 chunk 数、每 chunk 源数、前缀 chunk 判断全部编译期化。

  4. 主循环展开:对 for (int ci = 0; ci < num_chunks; ci++, gci++)#pragma unroll;同时给 mbarrier_wait/mbarrier_arrive 的内联汇编补 memory clobber,避免编译器对屏障后的访存做错误重排(为循环展开下的正确性兜底)。

测试配套:仅依赖现有 tests/models/kimi_k3/test_attn_res.py,未新增测试文件;PR 提供了微基准与 8×B300 端到端 TPOT/TTFT、gsm8k 精度数据作为验证。

文件 模块 状态 重要度
csrc/libtorch_stable/kimi_k3/attn_res_kernel.cu 内核实现 modified 4.98

关键符号

attn_res_fwd_online_v2_kernel launch_fwd kimi_k3_attn_res mbarrier_wait mbarrier_arrive

关键源码片段

csrc/libtorch_stable/kimi_k3/attn_res_kernel.cu core-logic

唯一变更文件,包含全部四项延迟优化:B/N 模板化、q_cache 向量化、NC=3 低 batch 分派、num_chunks 展开,以及 mbarrier 汇编 memory clobber 配套。

// 内核模板新增 N、B 编译期参数,使 num_chunks 在编译期确定,
// 从而允许主循环完全展开;运行时参数表只保留 T 与各 stride。
template <int H, int N, int NC = N_CHUNK_DEFAULT, int B = 1,
          bool RELEASE_TMEM = false, bool HAS_DELTA = false,
          bool HAS_OUTPUT_NORM = false, bool OUTPUT_NORM_IN_SMEM = false>
__global__ void __launch_bounds__(BLK, 1) attn_res_fwd_online_v2_kernel(
    const bf16_t* __restrict__ block_res, bf16_t* __restrict__ layer_res,
    const bf16_t* __restrict__ delta, const bf16_t* __restrict__ res_w,
    const bf16_t* __restrict__ rms_w, bf16_t* __restrict__ output, int T,
    int block_stride_m, int block_stride_r, float rms_eps,
    const bf16_t* __restrict__ output_norm_weight, float output_norm_eps) {
  constexpr int num_chunks = (N + N_CHUNK - 1) / N_CHUNK; // 编译期常量  // q_cache 填充:int2/int4 向量读入后按 __nv_bfloat162 成对转 float
  // 并累乘,替代逐元素标量加载,减少指令数与访存次数。
  int4 rms_v = *reinterpret_cast<const int4*>(rms_w + h_base);
  int4 res_v = *reinterpret_cast<const int4*>(res_w + h_base);
  auto* rms2 = reinterpret_cast<__nv_bfloat162*>(&rms_v);
  auto* res2 = reinterpret_cast<__nv_bfloat162*>(&res_v);
  #pragma unroll
  for (int k = 0; k < 4; k++) {
    float2 rf = __bfloat1622float2(rms2[k]);
    float2 sf = __bfloat1622float2(res2[k]);
    q_cache[si * VEC + 2 * k] = rf.x * sf.x;
    q_cache[si * VEC + 2 * k + 1] = rf.y * sf.y;
  }  // 主循环全展开;循环展开依赖上面的 constexpr num_chunks。
  #pragma unroll
  for (int ci = 0; ci < num_chunks; ci++, gci++) {
    int ns = ci * N_CHUNK;
    int an = min(N_CHUNK, N - ns);
    // ... 每 chunk 的 mbarrier 同步与计算主体 ...
  }
}

评论区精华

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

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

风险与影响

  1. 正确性风险mbarrier_wait/mbarrier_arrive 的汇编增加 memory clobber 是并发正确性的关键改动,若语义理解有误可能影响 barrier 与访存顺序;不过该改动方向是收紧而非放松编译器约束,风险较低。
  2. 性能回归风险:B/N 模板化与 switch(num_blocks) 分派增加了模板实例化与二进制体积,极端 num_blocks 组合可能落入 default 分支,性能回退到接近基线;不同 GPU 架构(如 sm_100 vs sm_90)上的展开收益与寄存器压力可能不同。
  3. 精度回归风险:gsm8k strict-match 从 0.9636 降至 0.9591,幅度较小,但 PR 未给出详细误差分析,需要关注是否属于数值噪声。
  4. 缺失覆盖:仅 tests/models/kimi_k3/test_attn_res.py 一个既有测试,未覆盖 B>1、长 prefill 等路径的组合矩阵。

影响范围集中于 Kimi K3 模型(H=7168 特化路径与通用路径)在 NVIDIA 平台上的注意力残差内核。微基准显示 1-2048 token 区间延迟下降 9%-27%,长序列(4096/8192)下降约 7%;端到端 8×B300 上 TPOT 普遍下降 0.2-0.32ms,TTFT 在高并发(4/8)下改善明显(256.48→198.94ms、258.03→242.25ms)。对用户与团队而言属于低风险的性能收益,但仅限 kimi_k3 相关模型。

并发内存序改动 精度数值回归 模板膨胀与分派覆盖 单测试覆盖

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论