已关闭
[Requirement|需求建议]: aclnn_fused_add_rmsnorm算子AscendC实现 #3397
wuxs68创建于  6月17日关闭于  7月17日
wuxs68
6月17日 创建

Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.

Backgroud(背景信息)

MiniCPM 模型中存在 residual scaled add 与 RMSNorm 连续计算场景,当前该计算通常由多个基础算子组合完成:

xOut = x1 * scale + x2
rstd = rsqrt(mean(xOut * xOut) + epsilon)
yOut = xOut * rstd * gamma

该计算模式在 MiniCPM 多层 Transformer Block 中重复出现。若拆分为多个基础算子执行,会产生额外的 Host 调度开销和中间 Tensor 读写开销,影响推理性能。因此希望新增一个面向 MiniCPM 场景的融合算子 MinicpmScaledAddRmsNorm,将 ScaledAdd 与 RMSNorm 融合为一个算子执行。

Origin(信息来源)

MiniCPM 模型中存在如下 residual 计算逻辑:
residual = residual + hidden_states * (scale_depth / sqrt(num_hidden_layers))
在 MiniCPM 典型配置中:

hidden_size = 2304
num_hidden_layers = 40
scale_depth = 1.4
epsilon = 1e-5

因此模型侧缩放系数为:
scale = 1.4 / sqrt(40)
CANN ops-nn 仓库已有相似算子 AddRmsNorm,可作为实现参考:
https://gitcode.com/cann/ops-nn/blob/master/norm/add_rms_norm/docs/aclnnAddRmsNorm.md
本算子与官方 AddRmsNorm 的主要差异为:
1、xOut 为必选输出,用于返回 ScaledAdd 后的 residual 结果。
2、前置加法不是普通 Add,而是 ScaledAdd:
xOut = x1 * scale + x2

Benefit / Necessity (价值/作用)

1、减少算子调度次数:将 Mul + Add + RMSNorm 等连续计算融合为一个算子,降低多次 kernel launch 和 Host 调度开销。

2、减少中间 Tensor 读写:ScaledAdd 的中间结果可在融合算子内部直接参与 RMSNorm 计算,减少显存读写和中间结果搬运。

3、适配 MiniCPM 模型推理场景:MiniCPM 中该计算模式在多层结构中重复出现,融合后可提升模型推理性能。

4、补充 CANN ops-nn 算子能力:当前官方已有 AddRmsNorm,但缺少 MiniCPM 所需的 ScaledAdd + RMSNorm 变体,因此新增该算子具有实际模型适配价值。

Design(设计方案)

本算子计划新增 MinicpmScaledAddRmsNorm 融合算子,用于 MiniCPM 模型中 ScaledAdd 与 RMSNorm 连续计算场景。整体计算公式为:

xOut = x1 * scale + x2
rstd = rsqrt(mean(xOut * xOut) + epsilon)
yOut = xOut * rstd * gamma

其中 xOut 为必选输出,用于保留 ScaledAdd 后的 residual 结果;yOut 为 RMSNorm 输出;rstdOut 为 RMSNorm 中间倒标准差。算子输入为 x1、x2、gamma,属性为 epsilon、scale,输出为 yOut、rstdOut、xOut,计划支持 FLOAT32、FLOAT16、BFLOAT16 数据类型和 ND 数据格式。

Host 侧设计
1、Host 侧主要负责参数校验、shape 推导和 tiling 生成。需要校验 x1 与 x2 shape 一致,gamma 与 x1 的尾部归一化维度匹配,yOut 和 xOut 与输入 shape/dtype 一致,rstdOut shape 根据归一化维度推导。
2、Shape 和 tiling 设计上,将输入抽象为二维结构:

num_row = x1 元素总数 / gamma 元素数
num_col = gamma 元素数

其中 num_row 表示 RMSNorm 行数,num_col 表示每行归一化长度。Tiling 分核遵循“优先满核”原则,数据量较小时单核处理,数据量较大时满核并行;若核间不能均分,则将余数分配到前几个 core。单 core 内切分遵循“充分使用 UB 空间”原则,结合 UB 大小、double buffer、临时 Tensor 数量等因素确定 row_factor、ub_factor 等参数。
TilingKey 根据 kernel 是否需要走不同分支决定是否规划。如果 dtype、硬件架构、归一化维度大小或 split_d 场景会影响 kernel 分支,则使用 tiling key 区分;若 kernel 侧逻辑通用,则不额外引入 tiling key。对于不支持 BFLOAT16 -> FLOAT32 Cast 的硬件,Host 侧提前返回不支持错误,避免 kernel 进入非法路径。

Kernel 侧设计
1、Kernel 侧采用 Init 和 Process 两阶段设计,Process 内部包括 CopyIn、Compute、CopyOut。CopyIn 将 x1、x2、gamma 从 GM 搬入 UB,Compute 在 UB 中完成 ScaledAdd、平方求和、rstd 和 RMSNorm 计算,CopyOut 将 xOut、rstdOut、yOut 写回 GM。
2、计算流程为先执行:
x = x1 * scale + x2
并将结果写回 xOut;随后计算:

rstd = 1 / sqrt(sum(x * x) / num_col + epsilon)
y = x * rstd * gamma

最终写回 rstdOut 和 yOut。其中 xOut 必选输出是为了保留 MiniCPM 后续 residual 路径所需的 ScaledAdd 结果。
3、对于 FLOAT16 和 BFLOAT16 输入,Kernel 侧将中间计算转为 FLOAT32,用于平方、累加、sqrt 和归一化,降低数值误差;最终 xOut 和 yOut 按输入 dtype 写回,rstdOut 使用 FLOAT32。由于支持 Ascend C 的硬件中 AscendC::Sqrt 支持 FLOAT32 输入,因此可直接使用 AscendC::Sqrt 完成 sqrt,再通过倒数得到 rsqrt 语义。
4、该设计将 ScaledAdd、Reduce、Sqrt/Rsqrt 和 RMSNorm 融合到一次 kernel 中执行,相比拆分为多个基础算子,可以减少 kernel launch、Host 调度开销和中间 Tensor 的 GM 读写,更适合 MiniCPM 多层重复出现的 residual scaled add + RMSNorm 场景。

当前验证情况
当前本地已完成以下验证:
测试项 结果

op_host UT: 8 tests passed
op_kernel UT: 8 tests passed
op_api UT: 接口生成与编译通过

Python 单算子正确性,单算子性能测试,数据类型为 bf16,hidden size = 2304,单算子性能测试 相比拆分实现约 4.7x 到 5.0x 加速

Shape Baseline FusedAddRmsNorm 加速比
(1, 8, 2304) 0.2311 ms 0.0478 ms 4.83x
(1, 128, 2304) 0.2533 ms 0.0502 ms 5.05x
(1, 512, 2304) 0.2312 ms 0.0491 ms 4.71x

性能测试
单算子性能测试,CANN 单算子执行收益
测试配置:C++ ACLNN 直接调用,warmup=20,repeat=200,统计 NPU event 的 device_avg_ms,shape 为 MiniCPM hidden=2304。

数据类型 输入 Shape Add + RmsNorm 耗时 FusedAddRmsNorm 耗时 加速比
FP16 1 x 1 x 2304 0.031522 ms 0.022131 ms 1.42x
FP16 1 x 8 x 2304 0.030227 ms 0.022941 ms 1.32x
FP16 1 x 128 x 2304 0.032682 ms 0.024936 ms 1.31x
FP16 1 x 512 x 2304 0.035144 ms 0.029465 ms 1.19x
BF16 1 x 1 x 2304 0.029037 ms 0.020586 ms 1.41x
BF16 1 x 8 x 2304 0.029598 ms 0.022068 ms 1.34x
BF16 1 x 128 x 2304 0.030867 ms 0.022121 ms 1.40x
BF16 1 x 512 x 2304 0.034977 ms 0.031735 ms 1.10x

说明:相对于“单独 Add + 单独 RmsNorm”两算子串联,我们的 FusedAddRmsNorm 在这些 MiniCPM 典型 shape 下都有加速,主要收益来自减少一次 kernel 调度/launch 和中间结果读写。小 token/短序列场景收益更明显,约 1.3x~1.4x;长序列 (1,512,2304) 下计算/访存占比提高,融合收益下降到 1.10x~1.19x。这个测试是保守对比,因为单独 Add 没有 scale。如果严格模拟 MiniCPM 的 x = x1 * scale + x2,未融合链路应是 Mul/Muls + Add + RmsNorm,会比这里的 Add + RmsNorm 多一个算子,理论上更能体现融合算子的价值

MiniCPM 模型替换性能测试:

指标 Baseline Fused
平均耗时 84.9105 ms 78.3203 ms
加速比 - 1.0841x

说明:模型级加速受 MatMul、Attention、MLP、RoPE 等其他算子影响,单算子融合收益在完整模型中会被整体推理链路稀释。

likedislike
CANN-robotCANN-robot成员
6月17日 将 m0_74080324 设为负责人
此处折叠了6条事件消息 查看更多
Wwuxs68
6月17日 issue状态由 已确认 改变为 技术评审中
wuxs68
6月17日 评论:

/assign

likedislike
CANN-robotCANN-robot成员
6月17日 将 wuxs68 设为负责人,移除负责人 m0_74080324
Wwuxs68
6月17日 关联了pull request:Rename operator to FusedAddRmsNorm
Wwuxs68
6月17日 关联了pull request:Rename operator to FusedAddRmsNorm
Wwuxs68
6月18日 issue状态由 技术评审中 改变为 进行中
wuxs68
6月18日 评论:

该Issue关联的PR:#6540,请尽快评审

likedislike
Wwuxs68
6月18日 修改了issue 的描述
Wwuxs68
6月25日 关联了pull request:Add experimental fused add rms norm operator
Cchenqi317成员
6月26日 将 sxb154714 设为负责人
Cchenqi317成员
6月26日 移除了负责人 sxb154714
chenqi317成员
7月3日 评论:

建议补充下性能结果

likedislike
Wwuxs68
7月4日 修改了issue 的描述
Cchenqi317成员
7月17日 issue状态由 进行中 改变为 已完成
Cchenqi317成员
7月17日 关闭了 issue
CANN-robotCANN-robot成员
7月17日 添加了label:Accepted