Prhub

#46353 [CPU][Perf] Accelerate unquantized MoE for AArch64

原始 PR 作者 fadara01 合并时间 2026-06-24 22:14 文件变更 15 提交数 2 评论 5 代码增减 +672 / -32

执行摘要

AArch64 MoE 加速,启用 NEON BFMMLA 微 kernel

AArch64 平台缺乏未量化 MoE 的硬件加速,导致 gpt-oss 和 gemma4 等模型的推理性能受限。PR body 提供的基准显示原有实现吞吐量低,本 PR 利用 AdvSIMD BFMMLA 指令实现融合 MoE kernel,以提升性能。

推荐架构师和性能工程师阅读。该 PR 展示了如何在 vLLM 的 CPU kernel 框架中高效添加新的 ISA 分支,其 pack 策略和宏调度设计值得在未来的 kernel 优化中复用。

讨论亮点

reviewer bigPYJ1151 指出在 w2 GEMM 中,当输入已被 packing 时,w2_input_tile_size 可能不再需要,因为 GEMM kernel 只看到 packed buffer。作者 fadara01 同意该观点,并引入了 w2_input_buffer_size 条件分支。但最终代码仍保留了 w2_input_tile_size 作为未使用变量,可能是为未来解 packed 路径预留。该讨论体现了对缓存预留的精细化控制。

实现拆解

实现分为 5 步:

  1. 新增 NEON 微 GEMM kernel:创建 csrc/cpu/micro_gemm/cpu_micro_gemm_neon.hpp,实现 8x8 BFMMLA 微 GEMM,支持输入矩阵 packing,并提供初始化/存储累加器的辅助函数。
  2. 集成到 fused MoE 调度:在 csrc/cpu/cpu_fused_moe.cpp 中添加 NEON_DISPATCH 宏,扩展 CPU_ISA_DISPATCH_IMPL 以包含 NEON 分支;修改 fused_moe_impl 根据 pack_a 标志调整 tile 大小,避免为 packed 输入预留过多缓存。
  3. 激活函数适配:优化 swigluoai_and_mul 等激活函数,为 AArch64 提供专用的偶奇分离加载(load_even_odd)以避免 gather 操作;用通用 tanh() 替代 x86 特定的 Sleef_tanhf16_u10,确保跨架构正确性。
  4. Python 层 ISA 发现:修改 vllm/model_executor/layers/fused_moe/cpu_fused_moe.py,在 check_grouped_gemm 中检测 ARM 架构并返回 "neon" ISA,使得 grouped GEMM 路径能为 AArch64 选择 NEON 微 kernel。
  5. 基准与测试更新:更新 benchmarks/kernels/cpu/benchmark_cpu_fused_moe.pytests/kernels/moe/test_cpu_fused_moe.py,支持 NEON ISA 选择;在 csrc/cpu/utils.hpp 中添加 NEON ISA 枚举,在 csrc/cpu/cpu_types_arm.hpp 中新增 load_even_odd 和取负操作。
文件 模块 状态 重要度
csrc/cpu/micro_gemm/cpu_micro_gemm_neon.hpp 微 GEMM added 8.35
csrc/cpu/cpu_fused_moe.cpp 融合 MoE modified 7.21
vllm/model_executor/layers/fused_moe/cpu_fused_moe.py MoE 调度 modified 6.52
benchmarks/kernels/cpu/benchmark_cpu_fused_moe.py 基准脚本 modified 6.31
csrc/cpu/cpu_types_arm.hpp ARM 类型 modified 6.17
csrc/cpu/utils.hpp CPU 枚举 modified 6.14

关键符号

MicroGemm gemm_micro_bfmmla_8x8_packed_a init_acc_rowpair store_acc_rowpair swigluoai_and_mul fused_moe_impl check_grouped_gemm load_even_odd operator-(FP32Vec16) get_isa

关键源码片段

csrc/cpu/micro_gemm/cpu_micro_gemm_neon.hpp core-logic

新增 NEON 微 GEMM kernel,实现基于 BFMMLA 指令的 8x8 微 GEMM,包含 pack 和 tile 支持,是本次加速的核心。

/*
 * 以下两个函数是 NEON 微 GEMM kernel 的核心辅助函数。
 * init_acc_rowpair 从 C 矩阵加载 1-2 行,并将它们交错压缩到 4 个 float32x4_t
 * 寄存器对 (acc01, acc23, acc45, acc67) 中,每个寄存器承载两行的部分元素。
 * 这种布局便于后续 BFMMLA 指令高效计算。
 * store_acc_rowpair 执行逆操作,将累加器拆分并写回 C 矩阵。
 */// 加载 C 矩阵并初始化累加器对
FORCE_INLINE void init_acc_rowpair(float32x4_t& acc01, float32x4_t& acc23,
                                   float32x4_t& acc45, float32x4_t& acc67,
                                   const float* __restrict__ c_ptr,
                                   const int64_t ldc, const int32_t m_rows,
                                   const bool accum_c) {
  if (!accum_c || m_rows == 0) {
    // 清空为 0
    acc01 = vdupq_n_f32(0.0f);
    acc23 = vdupq_n_f32(0.0f);
    acc45 = vdupq_n_f32(0.0f);
    acc67 = vdupq_n_f32(0.0f);
    return;
  }  // 加载行 0 和行 1 ( 或填充 0)
  const float32x4_t row0_0123 = vld1q_f32(c_ptr);
  const float32x4_t row0_4567 = vld1q_f32(c_ptr + 4);
  const float32x4_t row1_0123 =
      (m_rows == 2) ? vld1q_f32(c_ptr + ldc) : vdupq_n_f32(0.0f);
  const float32x4_t row1_4567 =
      (m_rows == 2) ? vld1q_f32(c_ptr + ldc + 4) : vdupq_n_f32(0.0f);  // 交错 : acc01 = [row0_0, row0_1, row1_0, row1_1]
  acc01 = zip1_f32x4(row0_0123, row1_0123);
  acc23 = zip2_f32x4(row0_0123, row1_0123);
  acc45 = zip1_f32x4(row0_4567, row1_4567);
  acc67 = zip2_f32x4(row0_4567, row1_4567);
}// 将累加器对写回 C 矩阵
FORCE_INLINE void store_acc_rowpair(const float32x4_t acc01,
                                    const float32x4_t acc23,
                                    const float32x4_t acc45,
                                    const float32x4_t acc67,
                                    float* __restrict__ c_ptr,
                                    const int64_t ldc, const int32_t m_rows) {
  if (m_rows == 0) return;  // 反交错 : 将 acc01/acc23 解压为行 0 的 0-3 和 4-7 元素
  vst1q_f32(c_ptr, zip1_f32x4(acc01, acc23));
  vst1q_f32(c_ptr + 4, zip1_f32x4(acc45, acc67));  if (m_rows == 2) {
    // 写回行 1
    vst1q_f32(c_ptr + ldc, zip2_f32x4(acc01, acc23));
    vst1q_f32(c_ptr + ldc + 4, zip2_f32x4(acc45, acc67));
  }
}

评论区精华

w2_input_tile_size 在 pack 后的必要性 设计

bigPYJ1151 指出 w2_input_tile_size 在输入被 packed 后不应计入缓存,因为 GEMM 只看到 packed buffer。

结论:作者同意并添加了条件分支 w2_input_buffer_size,但保留了原始 w2_input_tile_size 计算,可能因兼容性考虑。 · 已解决

风险与影响

1) 条件编译 ARM_BF16_SUPPORT__aarch64__ 确保 NEON 分支仅对 AArch64 生效,不会影响 x86/RISC-V 编译。
2) pack 逻辑可能引入内存越界,但 w2_input_buffer_size 基于 round_up 且与 pack_a 联动,风险可控。
3) sleef.h 的移除和 tanh() 替换已在 gelu_tanh_and_mul 中测试,DiffusionGemma 模型正确性已验证。
4) NEON 微 kernel 的 tile 大小(8x8)与其他 ISA(AMX、VEC)不同,需确保所有融合路径的一致性。

直接影响:AArch64 CPU 上 MoE 模型推理吞吐量提升 2x 以上,特别利好 gpt-oss 和 gemma4。间接影响:新增 NEON ISA 分支增加代码维护成本,但通过清晰的宏和枚举设计隔离。对现有 x86 路径无影响。测试和基准覆盖了 NEON 路径,CI 中包含 AArch64 构建。

架构特定分支 新增 NEON kernel 性能关键路径变更 x86 sleef 依赖移除

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论