Prhub

#29117 [CPU] optimize GDN prefill performance

原始 PR 作者 mingfeima 合并时间 2026-06-25 09:04 文件变更 3 提交数 19 评论 1 代码增减 +1340 / -910

执行摘要

优化 CPU GDN prefill 性能

Optimize the CPU GDN (Gated Delta Network) kernels in fla.cpp to better match the fused Triton pipeline and improve prefill/decode performance on AVX512/AMX CPUs. Achieve ~1.5x speedup against current implementation.

值得精读,特别是 fla.cpp 中的循环融合策略和 AVX512 优化技巧,对 CPU 推理场景很有参考价值。

讨论亮点

该 PR 没有 review 讨论。

实现拆解

  1. 算法融合:在 chunk_gated_delta_rule_kernel_impl 中将 intra-chunk 的 KKT 求解、下三角求解和 w/u 重计算融合为一体,避免物化矩阵 A;将 inter-chunk 的状态更新和输出计算融合,避免对中间变量 h 和 v_new 的显式存储。
  2. AVX512 核心:编写专用的 AVX512 内核完成 decay mask 生成、tril solve 以及 beta/g 掩码操作,减少循环开销。
  3. VNNI2 打包与状态布局对齐:新增 pack_vnni2 函数完成 FP32 到 BF16 的 VNNI2 格式转换并融合 element-wise 缩放;将状态布局从 KV 改为 VK 以匹配 Triton 实现,允许使用 amx-bf16 矩阵乘更新状态。
  4. 测试调整:在 test/registered/cpu/test_mamba.py 中更新参考函数和测试用例,纳入状态布局转置,确保精度测试通过。
文件 模块 状态 重要度
sgl-kernel/csrc/cpu/mamba/fla.cpp CPU 内核 modified 7.48
sgl-kernel/csrc/cpu/vec.h 工具函数 modified 6.34
test/registered/cpu/test_mamba.py 测试 modified 5.83

关键符号

pack_vnni2 l2norm_kernel::apply transpose_16x16_16bit chunk_gated_delta_rule_kernel_impl

关键源码片段

sgl-kernel/csrc/cpu/mamba/fla.cpp core-logic

核心优化实现,大幅重写 GDN 内核

// 将矩阵从 [K/2, 2, N] FP32 转换为 [K/2, N, 2] BF16,
// 同时将 src 乘以 exp(g_last) 以融合 element-wise 缩放。
template <typename scalar_t, int K, int N>
void pack_vnni2(scalar_t* __restrict__ dst, float* __restrict__ src,
                const float g_last, int ld_src, int ld_dst) {
  static_assert(K % 32 == 0);
  static_assert(N % 32 == 0);
  const float scale = std::exp(g_last);
#if defined(CPU_CAPABILITY_AVX512)
  constexpr int KB = K / 2;
  constexpr int NB = N / 32;
  __m512i s0, s1, d0, d1;
  __m512 vd = _mm512_set1_ps(scale);
  const auto trans = [&](auto i) {
    constexpr int kb = i / NB;
    constexpr int nb = i % NB;
    constexpr int k0 = kb * 2 + 0;
    constexpr int k1 = kb * 2 + 1;
    // 加载 k0 行和 k1 行的各一半(32 个元素每份)
    __m512 v00 = _mm512_loadu_ps(src + k0 * ld_src + nb * 32);
    __m512 v01 = _mm512_loadu_ps(src + k0 * ld_src + nb * 32 + 16);
    __m512 v10 = _mm512_loadu_ps(src + k1 * ld_src + nb * 32);
    __m512 v11 = _mm512_loadu_ps(src + k1 * ld_src + nb * 32 + 16);
    // FP32 -> BF16 并交错打包
    s0 = (__m512i)_mm512_cvtne2ps_pbh(v01, v00);
    s1 = (__m512i)_mm512_cvtne2ps_pbh(v11, v10);
    // 转置为 [K/2, N/32, 32, 2] 布局
    std::tie(d0, d1) = transpose_2x32_16bit(s0, s1);
    _mm512_storeu_si512(dst + kb * ld_dst * 2 + nb * 32 * 2, d0);
    _mm512_storeu_si512(dst + kb * ld_dst * 2 + nb * 32 * 2 + 32, d1);
    // 在原 src 上更新乘以 exp(g_last) 以用于后续计算
    _mm512_storeu_ps(src + k0 * ld_src + nb * 32, _mm512_mul_ps(v00, vd));
    _mm512_storeu_ps(src + k0 * ld_src + nb * 32 + 16, _mm512_mul_ps(v01, vd));
    _mm512_storeu_ps(src + k1 * ld_src + nb * 32, _mm512_mul_ps(v10, vd));
    _mm512_storeu_ps(src + k1 * ld_src + nb * 32 + 16, _mm512_mul_ps(v11, vd));
  };
  Unroll<KB * NB>{}(trans);
#else
  // 通用路径:逐元素转换与缩放
  for (int k = 0; k < K; k += 2) {
    for (int n = 0; n < N; ++n) {
      const float v0 = src[(k + 0) * ld_src + n];
      const float v1 = src[(k + 1) * ld_src + n];
      dst[(k >> 1) * ld_dst * 2 + n * 2 + 0] = static_cast<scalar_t>(v0);
      dst[(k >> 1) * ld_dst * 2 + n * 2 + 1] = static_cast<scalar_t>(v1);
      src[(k + 0) * ld_src + n] = v0 * scale;
      src[(k + 1) * ld_src + n] = v1 * scale;
    }
  }
#endif
}
sgl-kernel/csrc/cpu/vec.h core-logic

新增 16 位矩阵转置函数支持 VNNI2 packing

// 使用 AVX512 对 16x16 的 16 位矩阵进行原位转置(输入 / 输出均为 __m256i 数组)。
inline void transpose_16x16_16bit(__m256i* v) {
  __m256i v1[16];
  // 第一层:交错 16 位通道
  v1[0] = _mm256_unpacklo_epi16(v[0], v[1]);
  v1[1] = _mm256_unpackhi_epi16(v[0], v[1]);
  v1[2] = _mm256_unpacklo_epi16(v[2], v[3]);
  v1[3] = _mm256_unpackhi_epi16(v[2], v[3]);
  // ... 重复到 v[14], v[15]
  //(完整实现见源文件)
  // 第二层:交错 32 位通道
  v[0] = _mm256_unpacklo_epi32(v1[0], v1[2]);
  // ...
  // 最终通过 _mm256_permute2x128_si256 交换 128 位通道得到转置结果
}

评论区精华

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

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

风险与影响

  1. 状态布局变更:从 KV 改为 VK 可能影响其他依赖状态格式的组件(如与 GPU 共享状态的路径),但已通过测试验证。
  2. AVX512 路径依赖:优化核心主要针对 AVX512/AMX,若在不支持这些指令集的 CPU 上运行,会回退到通用路径,性能提升可能不明显。
  3. 精度与数值稳定性:融合运算和近似计算可能引入微小误差,需持续跟踪。

直接影响 CPU 上使用 GDN 注意力(如 Mamba 类模型)的推理性能,prefill 阶段加速约 1.5 倍,提升用户体验。不影响 GPU 或其他后端。代码仅涉及 sgl-kernel CPU 模块,团队需关注后续维护成本。

状态布局变更带来兼容性风险 非 AVX512 路径性能待验证

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论