☰
CANN ops-nn AddRmsNorm 算子深度解析:Add 与 RmsNorm 融合的实现原理与 aclnn 接口实战
2026/10/3 2:14:50 网站建设 项目流程
  • 人工智能
  • 算子库
  • 深度学习
  • CANN
  • Ascend

【免费下载链接】ops-nn

本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。

项目地址:https://gitcode.com/cann/ops-nn
点击查看免费下载

AddRmsNorm 是 CANN ops-nn 算子库中面向大模型(LLM)场景提供的一个融合算子:它将残差连接(Add)与 RMS 归一化(RmsNorm)合并为一次内核调度,从而减少张量在全局内存与片上缓存之间的搬入搬出操作。本文以 norm/add_rms_norm 目录下的 README.md 与 aclnnAddRmsNorm.md 为主体,结合算子定义、tiling 与 kernel 源码,完整讲解其产品支持情况、数学模型、参数语义、两段式 aclnn 调用流程以及 Device 侧执行细节,读完即可在 NPU 上正确配置并调用该算子。

一、算子背景与设计动机

在大模型 Transformer 类网络中,残差连接加归一化是每个 Attention 与 MLP 子层几乎必然出现的计算模式。传统的实现方式是先用独立的 Add 算子完成x = x1 + x2,再把结果整体读入内存交给归一化算子处理,同一份数据需要经历"写出-再读入"两个过程,带来额外的访存开销。

RmsNorm(Root Mean Square Normalization)本身是大模型常用的归一化操作,相比 LayerNorm 去掉了减去均值的部分,仅基于均方根对向量做缩放。AddRmsNorm 算子将 RmsNorm 之前的 Add 算子直接融合进来,把加法结果留在片上完成后续归一化,从而减少搬入搬出(数据搬运)操作。从源码结构看,该算子内部存在 ADD_RMS_NORM、PRE_RMS_NORM、POST_RMS_NORM 等多种计算模式(见 op_host/add_rms_norm_tiling.cpp),分别对应完整输出、仅输出归一化结果等不同场景,灵活支撑前向计算与后续反向(如 InplaceAddRmsNorm)的复用。

二、产品支持情况

根据 README.md 的产品支持矩阵,AddRmsNorm 的硬件适配情况如下:

产品是否支持
Ascend 950PR & 950DT 系列产品√
Atlas A3 系列产品√
Atlas A2 系列产品√
Atlas 200I/500 A2 推理产品×
Atlas 推理系列产品√
Atlas 训练系列产品×
Kirin X90 处理器系列产品√
Kirin 9030 处理器系列产品√

这一支持矩阵与算子注册代码一一对应。在 op_host/add_rms_norm_def.cpp 中可以看到,算子通过AICore().AddConfig()注册了ascend910b、ascend910_93、ascend310p、kirinx90、kirin9030、ascend950、ascend350等多套 AICore 配置;其中ascend310p(Atlas 推理系列)及其同配置的 Kirin 平台对输入输出数据类型做了裁剪,仅支持 FLOAT16 与 FLOAT32(不支持 BFLOAT16)。

注意:README 的表格中"Atlas 训练系列产品(910 系列)"标注为不支持,但源码同时注册了ascend910b与ascend910_93两套配置,二者对应关系以实际发布版本的适配清单为准;本文以仓库源码与文档共同表述的事实为准。

三、数学模型与功能说明

AddRmsNorm 的完整计算分为两步:

  1. 残差相加:将两个输入逐元素相加。

$$ x_i = x1_i + x2_i $$

  1. 均方根归一化:对相加结果按最后一组需要归一的维度计算均方根,并用其倒数缩放,最后乘上缩放因子(权重)。

$$ \operatorname{RmsNorm}(x_i)=\frac{x_i}{\operatorname{Rms}(\mathbf{x})} g_i, \quad \text { where } \operatorname{Rms}(\mathbf{x})=\sqrt{\frac{1}{n} \sum_{i=1}^n x_i^2+\varepsilon} $$

其中 $\varepsilon$ 为epsilon,g为gamma(缩放因子),n为参与归一化的最后一维(或后几维)的元素个数。

从算子 IR 定义 op_graph/add_rms_norm_proto.h 中可看到更简洁的等价描述:x = x1 + x2;rstd = rsqrt(mean(x^2, reduce_axis, keepdims=True) + epsilon);y = gamma * (x * rstd)。这清晰地揭示了三个输出y、rstd、x各自的语义:

  • y:归一化后的最终结果;
  • rstd:归一化后标准差的倒数,即Rms(x)的倒数,对应反向传播中常用的中间量;
  • x:Add 相加的结果,同样常被反向计算复用。

融合算子一次性产出三个输出,避免反向阶段重新计算 Add 与统计量,是其在 LLM 训练场景下价值所在。

四、参数说明(算子图模式)

下表完整列出 AddRmsNorm 在图模式(算子 IR)下的输入、输出与属性参数(与 README.md 的参数表一致):

参数名输入/输出/属性描述数据类型数据格式
x1输入用于 Add 计算的第一个输入,对应公式中的x1。FLOAT32、FLOAT16、BFLOAT16ND
x2输入用于 Add 计算的第二个输入,对应公式中的x2。FLOAT32、FLOAT16、BFLOAT16ND
gamma输入RmsNorm 的缩放因子(权重),对应公式中的g。shape 需要与x1后几维保持一致,后几维为x1需要 norm 的维度。FLOAT32、FLOAT16、BFLOAT16ND
epsilon可选属性添加到分母中的值,用于数值稳定、防止除 0 错误,需大于等于零,对应公式中的eps。默认值为 1e-6。FLOAT-
y输出最终的归一化输出,对应公式中的RmsNorm(x)。FLOAT32、FLOAT16、BFLOAT16ND
rstd输出归一化后标准差的倒数,对应公式中Rms(x)的倒数。FLOAT32ND
x输出Add 计算的结果,对应公式中的x。FLOAT32、FLOAT16、BFLOAT16ND

平台差异约束(需在配置参数时注意):

  • Atlas 推理系列产品:x1、x2、gamma、y、x的数据类型不支持 BFLOAT16。
  • Kirin X90 处理器系列产品、Kirin 9030 处理器系列产品:x1、x2、gamma、y、x的数据类型同样不支持 BFLOAT16。

上述约束与 add_rms_norm_def.cpp 中ascend310p配置仅注册{ge::DT_FLOAT16, ge::DT_FLOAT}的实现完全吻合。

关于gamma的 shape 约定,可以这样理解:gamma的维度数决定了"被归一化的后几维",其余维度视为 batch/行方向,逐行独立做归一化。例如x1shape 为(2, 3, 4, 8)时,gamma可取(8)或(4, 8),分别表示只对最后一维、或对后两维联合做均方根归一化。

五、aclnn 接口调用实战

5.1 两段式接口原型

与 CANN 其他算子一致,aclnnAddRmsNorm 采用两段式接口设计(仓库内对应说明见 docs/zh/context/two_phase_api.md):必须先调用aclnnAddRmsNormGetWorkspaceSize获取计算所需的 workspace 大小与执行器,再调用aclnnAddRmsNorm真正执行计算。接口声明位于 op_host/op_api/aclnn_add_rms_norm.h。

aclnnStatus aclnnAddRmsNormGetWorkspaceSize( const aclTensor *x1, const aclTensor *x2, const aclTensor *gamma, double epsilon, aclTensor *yOut, aclTensor *rstdOut, aclTensor *xOut, uint64_t *workspaceSize, aclOpExecutor **executor)
aclnnStatus aclnnAddRmsNorm( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream)

5.2 第一段接口参数详解

第一段接口aclnnAddRmsNormGetWorkspaceSize的参数与图模式参数一一对应,但以aclTensor形式承载,并补充了维度与连续性约束。完整参数语义如下:

参数输入/输出描述数据类型数据格式维度(shape)非连续 Tensor
x1输入Add 的第一个输入,对应公式x1。FLOAT32、FLOAT16、BFLOAT16ND1-8√
x2输入Add 的第二个输入;shape 与数据类型需与x1一致。FLOAT32、FLOAT16、BFLOAT16ND1-8√
gamma输入RmsNorm 缩放因子;数据类型需与x1一致。FLOAT32、FLOAT16、BFLOAT16ND1-8√
epsilon输入分母附加值,需大于等于零,建议值为 1e-6。----
yOut输出归一化输出;shape、数据类型与x1一致。FLOAT32、FLOAT16、BFLOAT16ND1-8√
rstdOut输出标准差的倒数;维度数与x1一致,需要 norm 的维度为 1(见下例)。FLOAT32ND1-8√
xOut输出Add 的结果;shape、数据类型与x1一致。FLOAT32、FLOAT16、BFLOAT16ND1-8√
workspaceSize输出返回需要在 Device 侧申请的 workspace 大小。----
executor输出返回 op 执行器,包含算子计算流程。----

rstdOut 的 shape 推导规则:维度数与x1保持一致,其中不需要 norm 的维度与x1对应维度相同,需要 norm 的维度(与gamma维度数相同的后几维)全部置为 1。官方示例:

  • x1shape(2, 3, 4, 8)、gammashape(8)→rstdOutshape(2, 3, 4, 1);
  • x1shape(2, 3, 4, 8)、gammashape(4, 8)→rstdOutshape(2, 3, 1, 1)。

这一规则与 add_rms_norm_infershape.cpp 中的 InferShape 实现一致:rstd的每个维度,当rmsIdx < xDimNum - gammaDimNum时继承x1对应维度,否则置 1。

特殊语义:rstdOut与xOut均可传入nullptr。当rstdOut传入nullptr时该输出无效;xOut传入nullptr时同理。yOut、x1、x2、gamma支持空 Tensor(shape 中有 0 维度)。

5.3 返回码与校验

第一段接口完成入参校验,出现以下场景时报错:

返回码错误码描述
ACLNN_ERR_PARAM_NULLPTR161001传入的x1、x2、gamma、yOut为空指针;或当rstdOut传入的预置值不为nullptr时,xOut传入的预置值为nullptr
ACLNN_ERR_PARAM_INVALID161002输入或输出的数据类型不在支持范围之内;或输入输出参数不满足参数说明中的约束

aclnn 返回码的完整含义可参考仓库文档 docs/zh/context/aclnn_return_code.md。

5.4 第二段接口参数详解

参数输入/输出描述
workspace输入在 Device 侧申请的 workspace 内存地址。
workspaceSize输入在 Device 侧申请的 workspace 大小,由第一段接口获取。
executor输入op 执行器,包含了算子计算流程。
stream输入指定执行任务的 Stream。

5.5 完整调用示例

仓库提供了可直接编译运行的示例程序 examples/test_aclnn_add_rms_norm.cpp,UT 测试 tests/ut/op_host/op_api/test_aclnn_add_rms_norm.cpp 也复用了同样的调用骨架。下面给出带注释的完整流程(以 FLOAT32、shape(2, 16)为例):

#include <iostream> #include <vector> #include "acl/acl.h" #include "aclnnop/aclnn_add_rms_norm.h" #define CHECK_RET(cond, return_expr) \ do { \ if (!(cond)) { \ return_expr; \ } \ } while (0) #define LOG_PRINT(message, ...) \ do { \ printf(message, ##__VA_ARGS__); \ } while (0) int64_t GetShapeSize(const std::vector<int64_t>& shape) { int64_t shape_size = 1; for (auto i : shape) { shape_size *= i; } return shape_size; } int Init(int32_t deviceId, aclrtStream* stream) { // 固定写法,资源初始化 auto ret = aclInit(nullptr); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclInit failed. ERROR: %d\n", ret); return ret); ret = aclrtSetDevice(deviceId); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSetDevice failed. ERROR: %d\n", ret); aclFinalize(); return ret); ret = aclrtCreateStream(stream); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtCreateStream failed. ERROR: %d\n", ret); aclrtResetDevice(deviceId); aclFinalize(); return ret); return 0; } template <typename T> int CreateAclTensor( const std::vector<T>& hostData, const std::vector<int64_t>& shape, void** deviceAddr, aclDataType dataType, aclTensor** tensor) { auto size = GetShapeSize(shape) * sizeof(T); // 申请 device 侧内存 auto ret = aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtMalloc failed. ERROR: %d\n", ret); return ret); // 将 host 侧数据拷贝到 device 侧内存 ret = aclrtMemcpy(*deviceAddr, size, hostData.data(), size, ACL_MEMCPY_HOST_TO_DEVICE); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtMemcpy failed. ERROR: %d\n", ret); return ret); // 计算连续 tensor 的 strides std::vector<int64_t> strides(shape.size(), 1); for (int64_t i = shape.size() - 2; i >= 0; i--) { strides[i] = shape[i + 1] * strides[i + 1]; } // 调用 aclCreateTensor 接口创建 aclTensor *tensor = aclCreateTensor( shape.data(), shape.size(), dataType, strides.data(), 0, aclFormat::ACL_FORMAT_ND, shape.data(), shape.size(), *deviceAddr); return 0; } int main() { // 1.(固定写法)device/stream 初始化,参考 acl API 手册 int32_t deviceId = 0; aclrtStream stream; auto ret = Init(deviceId, &stream); CHECK_RET(ret == 0, LOG_PRINT("Init acl failed. ERROR: %d\n", ret); return ret); // 2. 构造输入与输出 std::vector<int64_t> xShape = {2, 16}; std::vector<int64_t> gammaShape = {16}; std::vector<int64_t> yShape = {2, 16}; std::vector<int64_t> rstdShape = {2, 1}; // 注意:需要 norm 的维度为 1 void* x1DeviceAddr = nullptr; void* x2DeviceAddr = nullptr; void* gammaDeviceAddr = nullptr; void* yDeviceAddr = nullptr; void* rstdDeviceAddr = nullptr; void* xDeviceAddr = nullptr; aclTensor* x1 = nullptr; aclTensor* x2 = nullptr; aclTensor* gamma = nullptr; aclTensor* y = nullptr; aclTensor* rstd = nullptr; aclTensor* x = nullptr; std::vector<float> x1HostData = {0, 1, 2, 3, 4, 5, 6, 7, 0, 1, 2, 3, 4, 5, 6, 7, 0, 1, 2, 3, 4, 5, 6, 7, 0, 1, 2, 3, 4, 5, 6, 7}; std::vector<float> x2HostData = {0, 1, 2, 3, 4, 5, 6, 7, 0, 1, 2, 3, 4, 5, 6, 7, 0, 1, 2, 3, 4, 5, 6, 7, 0, 1, 2, 3, 4, 5, 6, 7}; std::vector<float> gammaHostData = {0, 1, 2, 3, 4, 5, 6, 7, 0, 1, 2, 3, 4, 5, 6, 7}; std::vector<float> yHostData = {0, 1, 2, 3, 4, 5, 6, 7, 0, 1, 2, 3, 4, 5, 6, 7, 0, 1, 2, 3, 4, 5, 6, 7, 0, 1, 2, 3, 4, 5, 6, 7}; std::vector<float> rstdHostData = {1, 2}; std::vector<float> xHostData = {0, 1, 2, 3, 4, 5, 6, 7, 0, 1, 2, 3, 4, 5, 6, 7, 0, 1, 2, 3, 4, 5, 6, 7, 0, 1, 2, 3, 4, 5, 6, 7}; float epsilon = 1e-6; ret = CreateAclTensor(x1HostData, xShape, &x1DeviceAddr, aclDataType::ACL_FLOAT, &x1); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(x2HostData, xShape, &x2DeviceAddr, aclDataType::ACL_FLOAT, &x2); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(gammaHostData, gammaShape, &gammaDeviceAddr, aclDataType::ACL_FLOAT, &gamma); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(yHostData, yShape, &yDeviceAddr, aclDataType::ACL_FLOAT, &y); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(rstdHostData, rstdShape, &rstdDeviceAddr, aclDataType::ACL_FLOAT, &rstd); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(xHostData, xShape, &xDeviceAddr, aclDataType::ACL_FLOAT, &x); CHECK_RET(ret == ACL_SUCCESS, return ret); // 3. 调用 CANN 算子库 API(两段式) uint64_t workspaceSize = 0; aclOpExecutor* executor; // 第一段接口:计算 workspace 大小并获取执行器 ret = aclnnAddRmsNormGetWorkspaceSize(x1, x2, gamma, epsilon, y, rstd, x, &workspaceSize, &executor); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnAddRmsNormGetWorkspaceSize failed. ERROR: %d\n", ret); return ret); // 根据第一段接口计算出的 workspaceSize 申请 device 内存 void* workspaceAddr = nullptr; if (workspaceSize > 0) { ret = aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("allocate workspace failed. ERROR: %d\n", ret); return ret); } // 第二段接口:执行计算 ret = aclnnAddRmsNorm(workspaceAddr, workspaceSize, executor, stream); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnAddRmsNorm failed. ERROR: %d\n", ret); return ret); // 4.(固定写法)同步等待任务执行结束 ret = aclrtSynchronizeStream(stream); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSynchronizeStream failed. ERROR: %d\n", ret); return ret); // 5. 获取输出值:将 device 侧结果拷贝至 host 侧 auto size = GetShapeSize(yShape); std::vector<float> resultData(size, 0); ret = aclrtMemcpy( resultData.data(), resultData.size() * sizeof(resultData[0]), yDeviceAddr, size * sizeof(float), ACL_MEMCPY_DEVICE_TO_HOST); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("copy result from device to host failed. ERROR: %d\n", ret); return ret); for (int64_t i = 0; i < size; i++) { LOG_PRINT("y result[%ld] is: %f\n", i, resultData[i]); } // 6. 释放 aclTensor aclDestroyTensor(x1); aclDestroyTensor(x2); aclDestroyTensor(gamma); aclDestroyTensor(y); aclDestroyTensor(rstd); aclDestroyTensor(x); // 7. 释放 device 资源 aclrtFree(x1DeviceAddr); aclrtFree(x2DeviceAddr); aclrtFree(xDeviceAddr); aclrtFree(gammaDeviceAddr); aclrtFree(yDeviceAddr); aclrtFree(rstdDeviceAddr); if (workspaceSize > 0) { aclrtFree(workspaceAddr); } aclrtDestroyStream(stream); aclrtResetDevice(deviceId); aclFinalize(); return 0; }

该示例的编译与运行方式遵循 CANN 通用样例流程,可参考仓库文档 docs/zh/context/compile_and_run_sample.md。工程侧还存在对应的 ST 用例(tests/st/aclnnAddRmsNorm/executor_aclnnAddRmsNorm.py)与 golden 脚本(tests/assets/golden.py),可对照验证结果。

六、数据类型与输出组合约束

6.1 各产品的输入输出组合

aclnn 接口下,x1、x2、gamma、yOut、rstdOut、xOut的支持组合因产品而异:

Atlas A2 系列产品、Atlas A3 系列产品:

x1x2gammayOutrstdOutxOut
FLOAT32FLOAT32shape 与x1后几维一致;FLOAT32FLOAT32必选必选,FLOAT32
FLOAT16FLOAT16shape 与x1后几维一致;FLOAT16FLOAT16必选必选,FLOAT16
BFLOAT16BFLOAT16shape 与x1后几维一致;BFLOAT16BFLOAT16必选必选,BFLOAT16
FLOAT16FLOAT16shape 为[1, x1 的最后一维];FLOAT16FLOAT16空指针空指针,FLOAT16
BFLOAT16BFLOAT16shape 为[1, x1 的最后一维];BFLOAT16BFLOAT16空指针空指针,BFLOAT16
FLOAT16FLOAT16shape 为[1, x1 的最后一维];FLOAT16FLOAT16空指针必选,FLOAT16
BFLOAT16BFLOAT16shape 为[1, x1 的最后一维];BFLOAT16BFLOAT16空指针必选,BFLOAT16

Ascend 950PR & 950DT 系列产品:

x1x2gammayOutrstdOutxOut
FLOAT32FLOAT32shape 与x1后几维一致;FLOAT32FLOAT32必选必选,FLOAT32
FLOAT16FLOAT16shape 与x1后几维一致;FLOAT16FLOAT16必选必选,FLOAT16
BFLOAT16BFLOAT16shape 与x1后几维一致;BFLOAT16BFLOAT16必选必选,BFLOAT16

Atlas 推理系列产品:

x1x2gammayOutrstdOutxOut
FLOAT32FLOAT32shape 与x1后几维一致;FLOAT32FLOAT32必选必选,FLOAT32
FLOAT16FLOAT16shape 与x1后几维一致;FLOAT16FLOAT16必选必选,FLOAT16
FLOAT16FLOAT16shape 为[1, x1 的最后一维];FLOAT16FLOAT16空指针空指针,FLOAT16
FLOAT16FLOAT16shape 为[1, x1 的最后一维];FLOAT16FLOAT16空指针必选,FLOAT16

6.2 边界值行为

  • 当输入是 Inf 时,输出为 Inf。
  • 当输入是 NaN 时,输出为 NaN。

6.3 确定性

aclnnAddRmsNorm 默认采用确定性实现,多次运行同一输入可得到一致的输出结果。

七、Device 侧执行原理(源码级解析)

7.1 从算子注册到图模式构图

除了 aclnn 接口调用方式,AddRmsNorm 还支持图模式调用:通过算子 IR op_graph/add_rms_norm_proto.h 中REG_OP(AddRmsNorm)注册的AddRmsNorm算子节点构图,图中可直接设置epsilon属性。对应地,框架侧还提供了 ONNX 插件 framework/npu_add_rms_norm_onnx_plugin.cpp,用于将 ONNX 模型中的 Add + RmsNorm 子图识别并融合为 AddRmsNorm 算子节点。

7.2 Shape 推导

op_host/add_rms_norm_infershape.cpp 实现了InferShape4AddRmsNorm与InferDataType4AddRmsNorm:

  • 输出y与x的 shape 直接继承x1的 shape;
  • 输出rstd的维度数与x1一致,前xDimNum - gammaDimNum维继承x1对应维度,后gammaDimNum维全部置 1;
  • 输出y与x的数据类型继承x1,输出rstd固定为DT_FLOAT(FLOAT32)。

同一实现同时注册给了AddRmsNorm与InplaceAddRmsNorm(原地版本,用于支持反向传播场景下的复用),说明该算子与同目录 norm 家族的其他融合算子共享了部分归一化基础设施。

7.3 Tiling 调度策略

tiling 在 Host 侧完成,负责把全局 shape 拆分为适配各 AI Core 的切块并写入 tiling data。op_host/add_rms_norm_tiling.h 定义了多套 tiling 数据结构,每一套对应一种内核策略:

  • AddRMSNormTilingData:通用策略,字段包括num_row、num_col、block_factor(每核行数)、row_factor(每次循环处理行数)、ub_factor(单次进 UB 的列数)、epsilon、avg_factor以及 fp16/fp32 分离的乘法循环参数等;
  • AddRMSNormRegbaseRFullLoadTilingData(key=1000):R(行维度)全部载入;
  • AddRMSNormRegbaseTilingData(key=2000):基础 regbase 策略;
  • AddRMSNormRegbaseSplitARTilingData(key=3000):按 R 分核,每核遍历全部归一化行;
  • AddRMSNormRegbaseTransTilingData(key=4000):物理转置后沿归一化维度向量化;
  • AddRMSNormRegbaseReduceEmptyTilingData(key=5000):R 为空时仅切分并写出rstd。

从 add_rms_norm_tiling.cpp 的常量可以推断,tiling 会根据数据类型(FP16/FP32/BF16 分别对应 dtypeKey 1/2/3)、归一化列数(num_col)、UB 容量与 16/8 元素对齐粒度(BLOCK_ALIGN_NUM=16、FLOAT_BLOCK_ALIGN_NUM=8)等条件,在多种模式(MODE_NORMAL、MODE_SPLIT_D、MODE_MERGE_N、MODE_SINGLE_N、MODE_MULTI_N)中选择合适的切分方案,并通过getPerformanceFlag在 910B 上针对 2/3 维、低行数、FP16/BF16 的典型 LLM shape 开启性能模式。

7.4 Kernel 执行流程

Device 侧内核模板 op_kernel/add_rms_norm.h 中,KernelAddRmsNorm<T, MODE>是核心实现:

  1. Init 阶段:根据 tiling data 读取numRow、numCol、blockFactor、rowFactor、ubFactor、epsilon等参数;按GetBlockIdx()计算当前核负责的行区间(blockIdx < GetBlockNum()-1时处理blockFactor行,最后一个核处理剩余尾行),并按行偏移建立x1Gm、x2Gm、gammaGm、yGm、rstdGm、xGm的全局内存视图;
  2. UB 缓冲区:为输入队列、gamma 队列、输出 y 队列、rstd 队列申请片上缓冲区,FP16/BF16 输入额外申请 FP32 转换缓冲(xFp32Buf)与平方和缓冲(sqxBuf),供归一化统计计算使用;
  3. Process 阶段:先CopyInGamma载入权重,再按rowFactor分块循环调用SubProcess,每个子块内逐行完成CopyIn → Compute → CopyOutY,在 ADD_RMS_NORM 模式下最后统一CopyOutRstd写出标准差倒数;
  4. 多模式支持:通过MODE == ADD_RMS_NORM_MODE / PRE_RMS_NORM_MODE等编译期分支决定是否写出rstd与x,与 tiling 阶段根据输出 shape 是否为 0 计算的norm_key(0 表示完整模式、100 表示 PRE、1000 表示 POST)相配合。

kernel 目录下还包含多种架构专用实现:arch35/下按 1000/2000/3000/4000/5000 key 拆分的add_rms_norm_regbase*.h系列,以及面向 950 平台的 APT(AoT/预编译)实现 add_rms_norm_apt.cpp,配合 op_host/config 下各产品的 binary 配置文件(add_rms_norm_binary.json,覆盖 ascend910b、ascend910_93、ascend350、ascend950、ascend310p、kirinx90、kirin9030)完成算子的二进制形态分发。

八、测试与验证

仓库为 AddRmsNorm 提供了完整的测试覆盖,可用于验证本文所述行为:

  • UT(Host 侧):tests/ut/op_host/test_AddRmsNorm_infershape.cpp 验证 shape 推导(包括rstd置 1 规则);tests/ut/op_host/test_add_rms_norm_tiling.cpp 验证 tiling 参数计算;tests/ut/op_host/op_api/test_aclnn_add_rms_norm.cpp 验证两段式接口调用;
  • UT(Kernel 侧):tests/ut/op_kernel/test_add_rms_norm.cpp 与 test_add_rms_norm_regbase.cpp 验证 Device 内核计算结果;
  • ST(系统测试):tests/st/aclnnAddRmsNorm/executor_aclnnAddRmsNorm.py 配合 atk_aclnnAddRmsNorm.json 执行端到端用例;arch35 目录下还有ttk_kernel_add_rms_norm_st.csv与ttk_kernel_add_rms_norm_perf.csv记录 35 架构的内核 ST 与性能测试项。

九、使用建议与注意事项

  1. 优先使用融合算子:凡是在大模型中出现的x1 + x2 → RmsNorm模式,直接用 AddRmsNorm 替代 Add + RmsNorm 两个算子,可减少一次全局内存往返;
  2. 正确设置 gamma 与 rstd 的 shape:gamma的维度数决定归一化范围;rstdOut需要 norm 的维度必须为 1,否则会触发参数校验错误(161002);
  3. 留意产品差异:Atlas 推理系列与 Kirin 平台不支持 BFLOAT16;Atlas 200I/500 A2 推理产品与 Atlas 训练系列(910 系列)在 README 中标注不支持;
  4. 合理利用可选输出:不需要rstd或x时传nullptr,让 tiling 走 PRE/POST 裁剪模式,减少写回开销;
  5. epsilon 取值:需大于等于零,默认 1e-6;建议保持默认或使用与训练一致的数值,避免数值稳定性问题;
  6. 确定性:该算子默认确定性实现,适合对结果可复现性有要求的训练场景。

十、小结

AddRmsNorm 通过"Add + RmsNorm"的算子融合,将大模型最频繁的残差归一化模式压缩为单次内核执行,并额外产出rstd与x两个中间量供反向复用。本文从产品支持矩阵、数学模型、图模式参数、两段式 aclnn 接口、Device 侧 tiling 与 kernel 执行链路四个层面做了完整拆解,并给出了可直接运行的调用示例与测试入口。读者可以进一步阅读 README.md、aclnnAddRmsNorm.md,或直接基于 examples/test_aclnn_add_rms_norm.cpp 展开实验。

  • 人工智能
  • 算子库
  • 深度学习
  • CANN
  • Ascend

【免费下载链接】ops-nn

本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。

项目地址:https://gitcode.com/cann/ops-nn
点击查看免费下载

相关推荐

上一篇:大麦自动抢票工具 ticket-purchase 实战:Selenium + Appium 双端,毫秒级下单
下一篇:如何永久保存微信聊天记录?WeChatMsg完整指南帮你轻松实现数据自由

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询