news 2026/10/3 2:27:10

CANN ops-nn AddRmsNormCast 算子深度解析:Add 融合 RmsNorm 与 Cast 的归一化算子上手指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
CANN ops-nn AddRmsNormCast 算子深度解析:Add 融合 RmsNorm 与 Cast 的归一化算子上手指南
  • 人工智能
  • 算子库
  • 深度学习
  • CANN
  • Ascend

【免费下载链接】ops-nn

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

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

导读

AddRmsNormCast 是 CANN ops-nn 神经网络算子库中面向大模型(LLM)场景的融合归一化算子:它把Add(逐元素相加)、RmsNorm(均方根归一化)与Cast(类型转换)三个操作融合为一次 Kernel 执行,通过减少数据在内存与计算单元之间的搬入搬出次数来降低算子调用开销。本文以 norm/add_rms_norm_cast/README.md 与配套的 aclnnAddRmsNormCast 接口文档 为主线,结合该算子在仓库中的算子定义、shape 推导、tiling 策略、Kernel 分发与测试代码,完整讲解其产品支持情况、数学原理、输入输出参数、aclnn 两段式调用方法与图模式构图方式。读完本文,你将掌握如何在 Atlas A2/A3、Ascend 950 系列与 Kirin 系列产品上正确构造并调用该算子,并能理解其底层实现机制。

算子定位:为什么要把 Add、RmsNorm、Cast 融合在一起

RmsNorm(Root Mean Square Normalization)是大模型常用的归一化操作。在典型的 Transformer 网络结构中,残差相加(Add)之后往往紧跟 RmsNorm,而部分计算图还会在 RmsNorm 之后衔接类型转换(Cast)以满足后续算子对精度的要求。如果按照原始计算图逐个算子执行,Add、RmsNorm、Cast 各自都需要独立申请内存、启动 Kernel、完成数据搬入搬出,会引入可观的额外开销。

AddRmsNormCast 算子的核心设计目标正是消除这些开销:将 AddRmsNorm 之后的 Cast 算子融合进归一化计算中,减少搬入搬出操作(见 README.md 功能说明)。一次 Kernel 调用即可完成"求和 → 归一化 → 类型转换",同时顺带输出 RmsNorm 中常用的中间量rstd(标准差的倒数)与x(归一化前的数据和),方便上层框架在反向传播中复用,避免重复计算。

从仓库目录结构看,该算子是一个完整的 CANN 算子工程,包含接口文档、调用示例、构图原型、Host 侧定义/推导/tiling、Kernel 侧实现与单元/系统测试:

  • 接口文档:aclnn 单算子调用接口说明;
  • 调用示例:完整的可编译示例;
  • 算子 IR 原型:图模式构图接口;
  • op_host、op_kernel、tests:Host 侧与 Kernel 侧实现及测试。

产品支持情况

根据 README.md,AddRmsNormCast 算子在不同产品上的支持情况如下:

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

这一支持矩阵与算子定义文件 op_host/add_rms_norm_cast_def.cpp 中注册的 AICore 配置一一对应:ascend910b(Atlas A2 系列)、ascend910_93(Atlas A3 系列)、ascend950(Ascend 950 系列)、kirinx90与kirin9030。同时 op_host/config 目录下也分别维护了ascend910b、ascend910_93、ascend950、kirinx90、kirin9030五份算子二进制配置(add_rms_norm_cast_binary.json),与上表的产品范围一致。

需要注意的版本差异:Kirin X90 处理器系列产品与 Kirin 9030 处理器系列产品上,x1、x2、gamma、y2和x的数据类型不支持 BFLOAT16。这一点在算子定义中也有体现——Kirin 系列专属配置为x1、x2、gamma、y2、x只注册了DT_FLOAT16,而其他产品则同时注册DT_FLOAT16与DT_BF16。

功能说明与计算公式

计算流程

AddRmsNormCast 一次执行完成三步计算(此处参数命名遵循 README.md,与 aclnn 接口文档中的命名差异见下文"参数说明"):

  1. 对两个输入做逐元素求和,得到x;
  2. 对x做 RmsNorm 归一化,得到归一化结果y2;
  3. 将y2转换为更高精度的y1。

计算公式

求和:

$$ x_i=x1_{i}+x2_{i} $$

RmsNorm 归一化:

$$ y_2=\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+eps} $$

类型转换:

$$ y_1=float(y_2) $$

其中:

  • x1、x2:需要归一化的原始数据输入;
  • g:数据缩放因子(即gamma),逐元素乘到归一化结果上;
  • eps(即epsilon):添加到分母根号内的常数,用于防止除零并保证数值稳定;
  • y2:归一化后、未做类型转换的输出(FLOAT16 / BFLOAT16);
  • y1:归一化后经过类型转换的输出(FLOAT32);
  • 中间量Rms(x)的倒数即输出rstd,x1+x2的和即输出x。

该计算流程与算子 IR 原型 op_graph/add_rms_norm_cast_proto.h 中注释描述的实现完全一致:

x = float(x1) + float(x2) rstd = np.rsqrt(np.mean(np.power(x,2), reduce_axis, keepdims=True) + epsilon) y1 = gamma * (x * rstd) y2 = cast(y1)

此外,tests/assets/golden.py 中的 golden 函数复用add_rms_norm算子的 golden 实现,并通过_post_action='cast'追加类型转换,以纯 NumPy 复现"Add → RmsNorm → Cast"的完整计算链,是验证算子数值正确性的参考实现。

参数说明

下表完整列出算子各参数的定义(数据格式均为 ND):

参数名输入/输出/属性描述数据类型数据格式
x1输入需要归一化的原始数据输入,公式中的输入x1。FLOAT16、BFLOAT16ND
x2输入需要归一化的原始数据输入,公式中的输入x2。FLOAT16、BFLOAT16ND
gamma可选输入数据缩放因子,公式中的输入g。shape 需要与x1后几维保持一致,后几维为x1需要 norm 的维度。FLOAT16、BFLOAT16ND
epsilon可选属性添加到分母中的值,以确保数值稳定,用于防止除 0 错误,对应公式中的eps。默认值为 1e-6。FLOAT32-
y1输出归一化后经过类型转换的输出数据,公式中的输出y1。FLOAT32ND
y2输出归一化后的输出数据,公式中的输出y2。FLOAT16、BFLOAT16ND
rstd输出x 的标准差,公式中的输出Rms(x)。FLOAT32ND
x输出归一化的数据和,公式中的输出x。FLOAT16、BFLOAT16ND

参数背后的实现约束

  • gamma 与 norm 维度:gamma 的 shape 必须与x1的后几维保持一致,其中后几维就是执行归一化(reduce)的维度。例如x1shape 为(2,3,4,8)时,若gammashape 为(8),则只有最后一维参与归一化;若gammashape 为(4,8),则后两维参与归一化。
  • rstd 的 shape 推导:rstd 需要与x1数据格式一致,维度数与x1相同;其中不需要 norm 的维度(x1的维度数减去gamma的维度数后的前几维)与x1对应维度保持一致,需要 norm 的维度(与gamma维度数相同的后几维)均为 1。举例:若x1shape 为(2,3,4,8)、gammashape 为(8),则rstdshape 为(2,3,4,1);若gammashape 为(4,8),则rstdshape 为(2,3,1,1)。
  • y1 / y2 / x 的 shape:与x1保持一致;x2的 shape 与数据类型也需要与x1保持一致。
  • 数据类型联动:y1固定为 FLOAT32,y2与x的数据类型跟随x1,rstd固定为 FLOAT32。在 op_host/add_rms_norm_cast_infershape.cpp 的InferDataType4AddRmsNormCast中可以看到这一规则的实现:y1与rstd直接置为DT_FLOAT,y2与x取x1的输入类型。

关于接口文档中输出命名的说明

README 的参数表将 FLOAT32 输出命名为y1、FLOAT16/BFLOAT16 输出命名为y2;而 aclnnAddRmsNormCast 接口文档 中 FLOAT32 输出命名为y1Out(归一化输出)、FLOAT16/BFLOAT16 输出命名为y2Out(cast 输出)。两者对"归一化结果"与"类型转换结果"的命名顺序相反,但类型映射一致(FLOAT32 一个、FLOAT16/BFLOAT16 一个),实际调用时以你使用的接口文档为准、按数据类型对应即可。

约束说明

  • 输出不支持非连续 Tensor:y1Out、y2Out、rstdOut、xOut均要求连续。
  • 维度边界:x1、x2、gamma、y1、y2、rstd、x的 shape 中每一维大小都不大于 INT32 最大值 2147483647;各张量维度数限制为 1~8(tiling 校验代码 op_host/add_rms_norm_cast_tiling.cpp 中的MAX_DIM_NUM=8与之对应)。
  • 空 Tensor 边界:不支持"非 Norm 维度元素总数大于 0 且 Norm 维度元素总数为 0"的空 Tensor 场景。
  • 特殊数值传递:输入为 Inf 时输出为 Inf,输入为 NaN 时输出为 NaN。
  • 确定性计算:aclnnAddRmsNormCast默认确定性实现,同一输入多次运行结果可复现。
  • epsilon 取值:建议值为 1e-6;tiling 中还会校验 epsilon 不小于 0。
  • Kirin 系列 dtype 限制:Kirin X90 / Kirin 9030 上x1、x2、gamma、y2、x不支持 BFLOAT16。

调用说明

AddRmsNormCast 支持两种调用方式:通过aclnnAddRmsNormCast接口(单算子调用,对应 示例代码)与通过算子 IR 构图(图模式,对应 add_rms_norm_cast_proto.h)。

调用方式样例代码说明
aclnn接口test_aclnn_add_rms_norm_cast通过 aclnnAddRmsNormCast 接口方式调用 AddRmsNormCast 算子。
图模式-通过 算子IR 构图方式调用 AddRmsNormCast 算子。

aclnn 两段式接口

每个算子分为两段式接口:必须先调用aclnnAddRmsNormCastGetWorkspaceSize获取计算所需 workspace 大小以及包含了算子计算流程的执行器(executor),再调用aclnnAddRmsNormCast执行计算。

aclnnStatus aclnnAddRmsNormCastGetWorkspaceSize( const aclTensor *x1, const aclTensor *x2, const aclTensor *gamma, double epsilon, const aclTensor *y1Out, const aclTensor *y2Out, const aclTensor *rstdOut, const aclTensor *xOut, uint64_t *workspaceSize, aclOpExecutor **executor)
aclnnStatus aclnnAddRmsNormCast( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream)
aclnnAddRmsNormCastGetWorkspaceSize 参数说明
参数名输入/输出描述使用说明数据类型数据格式维度(shape)非连续Tensor
x1(aclTensor*)输入表示用于 Add 计算的第一个输入。对应公式中的x1。支持空 Tensor。BFLOAT16、FLOAT16ND1-8√
x2(aclTensor*)输入表示用于 Add 计算的第二个输入。对应公式中的x2。支持空 Tensor;shape 和数据类型需要与x1的 shape 和数据类型保持一致。FLOAT16、BFLOAT16ND1-8√
gamma(aclTensor*)输入表示 RmsNorm 的缩放因子(权重)。对应公式中的gamma。支持空 Tensor;数据类型与x1保持一致;shape 需要与x1后几维保持一致。FLOAT16、BFLOAT16ND1-8√
epsilon(double)输入表示添加到分母中的值,以确保数值稳定。对应公式中的epsilon。建议值为 1e-6。----
y1Out(aclTensor*)输出表示归一化后的输出数据。支持空 Tensor;shape、数据格式与入参x1保持一致。FLOAT32ND1-8×
y2Out(aclTensor*)输出表示归一化后经过类型转换的输出数据。支持空 Tensor;shape、数据格式、数据类型均与入参x1保持一致。FLOAT16、BFLOAT16ND1-8×
rstdOut(aclTensor*)输出表示归一化后的标准差的倒数。支持空 Tensor;与入参x1数据格式一致;不需要 norm 的维度与x1对应维度一致,需要 norm 的维度均为 1。FLOAT32ND1-8×
xOut(aclTensor*)输出表示 Add 计算的结果。支持空 Tensor;shape、数据格式、数据类型均与入参x1保持一致。FLOAT16、BFLOAT16ND1-8×
workspaceSize(uint64_t*)输出返回需要在 Device 侧申请的 workspace 大小。-----
executor(aclOpExecutor**)输出返回 op 执行器,包含了算子计算流程。-----
返回码

aclnnStatus返回状态码。第一段接口完成入参校验,出现以下场景时报错:

返回码错误码描述
ACLNN_ERR_PARAM_NULLPTR161001如果传入参数是必选输入、输出或者必选属性,且是空指针。
ACLNN_ERR_PARAM_INVALID161002输入或输出的数据类型不在支持的范围之内。
ACLNN_ERR_INNER_TILING_ERROR561002输入和输出不符合上述参数说明内的要求。
aclnnAddRmsNormCast 参数说明
参数名输入/输出描述
workspace输入在 Device 侧申请的 workspace 内存地址。
workspaceSize输入在 Device 侧申请的 workspace 大小,由第一段接口aclnnAddRmsNormCastGetWorkspaceSize获取。
executor输入op 执行器,包含了算子计算流程。
stream输入指定执行任务的 Stream。

完整调用示例

以下代码取自 examples/test_aclnn_add_rms_norm_cast.cpp,展示了从环境初始化、Tensor 构造、两段式接口调用到结果回拷与资源释放的完整流程:

#include <iostream> #include <vector> #include "acl/acl.h" #include "aclnnop/aclnn_add_rms_norm_cast.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; } // 固定写法:acl 资源初始化 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; } // 申请 device 内存、拷贝数据并创建 ND 格式 aclTensor 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); 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); 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]; } *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 初始化 int32_t deviceId = 0; // 根据自己的实际 device 填写 aclrtStream stream; auto ret = Init(deviceId, &stream); CHECK_RET(ret == 0, LOG_PRINT("Init acl failed. ERROR: %d\n", ret); return ret); // 2. 构造输入与输出:示例 shape 为 x1/x2/y1/y2/x = (2,16),gamma = (16),rstd = (2,1) 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}; void* x1DeviceAddr = nullptr; void* x2DeviceAddr = nullptr; void* gammaDeviceAddr = nullptr; void* y1DeviceAddr = nullptr; void* y2DeviceAddr = nullptr; void* rstdDeviceAddr = nullptr; void* xDeviceAddr = nullptr; aclTensor* x1 = nullptr; aclTensor* x2 = nullptr; aclTensor* gamma = nullptr; aclTensor* y1 = nullptr; aclTensor* y2 = nullptr; aclTensor* rstd = nullptr; aclTensor* x = nullptr; std::vector<short> x1HostData = {0x0000, 0x3C00, 0x4000, 0x4200, 0x4400, 0x4500, 0x4600, 0x4700, 0x0000, 0x3C00, 0x4000, 0x4200, 0x4400, 0x4500, 0x4600, 0x4700, 0x0000, 0x3C00, 0x4000, 0x4200, 0x4400, 0x4500, 0x4600, 0x4700, 0x0000, 0x3C00, 0x4000, 0x4200, 0x4400, 0x4500, 0x4600, 0x4700}; std::vector<short> x2HostData = x1HostData; // 实际示例中与 x1HostData 相同 std::vector<short> gammaHostData = {0x0000, 0x3C00, 0x4000, 0x4200, 0x4400, 0x4500, 0x4600, 0x4700, 0x0000, 0x3C00, 0x4000, 0x4200, 0x4400, 0x4500, 0x4600, 0x4700}; std::vector<float> y1HostData = {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<short> y2HostData = x1HostData; // 实际示例中与 x1HostData 相同 std::vector<float> rstdHostData = {1, 2}; std::vector<short> xHostData = x1HostData; // 实际示例中与 x1HostData 相同 float epsilon = 1e-6; // 创建各 aclTensor(ACL_FLOAT16 对应 FLOAT16 输入,ACL_FLOAT 对应 FLOAT32 输出) ret = CreateAclTensor(x1HostData, xShape, &x1DeviceAddr, aclDataType::ACL_FLOAT16, &x1); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(x2HostData, xShape, &x2DeviceAddr, aclDataType::ACL_FLOAT16, &x2); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(gammaHostData, gammaShape, &gammaDeviceAddr, aclDataType::ACL_FLOAT16, &gamma); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(y1HostData, yShape, &y1DeviceAddr, aclDataType::ACL_FLOAT, &y1); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(y2HostData, yShape, &y2DeviceAddr, aclDataType::ACL_FLOAT16, &y2); 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_FLOAT16, &x); CHECK_RET(ret == ACL_SUCCESS, return ret); // 3. 调用第一段接口,获取 workspace 大小与执行器 uint64_t workspaceSize = 0; aclOpExecutor* executor; ret = aclnnAddRmsNormCastGetWorkspaceSize(x1, x2, gamma, epsilon, y1, y2, rstd, x, &workspaceSize, &executor); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnAddRmsNormCastGetWorkspaceSize 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 = aclnnAddRmsNormCast(workspaceAddr, workspaceSize, executor, stream); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnAddRmsNormCast 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]), y1DeviceAddr, size * sizeof(resultData[0]), 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("y1 result[%ld] is: %f\n", i, resultData[i]); } // 6. 释放 aclTensor aclDestroyTensor(x1); aclDestroyTensor(x2); aclDestroyTensor(gamma); aclDestroyTensor(y1); aclDestroyTensor(y2); aclDestroyTensor(rstd); aclDestroyTensor(x); // 7. 释放 device 资源 aclrtFree(x1DeviceAddr); aclrtFree(x2DeviceAddr); aclrtFree(xDeviceAddr); aclrtFree(gammaDeviceAddr); aclrtFree(y2DeviceAddr); aclrtFree(y1DeviceAddr); aclrtFree(rstdDeviceAddr); if (workspaceSize > 0) { aclrtFree(workspaceAddr); } aclrtDestroyStream(stream); aclrtResetDevice(deviceId); aclFinalize(); return 0; }

示例中的关键点:

  • 输出命名对应:示例中y1对应接口文档的y1Out(FLOAT32),y2对应y2Out(FLOAT16/BFLOAT16);
  • rstd 的 shape 构造:x1为(2,16)、gamma为(16)时,rstd为(2,1)——前 1 维与x1相同、norm 维压缩为 1,与 shape 推导规则一致;
  • workspace 处理:只有当workspaceSize > 0时才需要申请内存,执行完毕后同样需要释放;
  • 完整的编译与运行流程可参考 CANN 文档中"编译与运行样例"的相关指引。

图模式构图

在图模式下,通过算子 IR 在计算图中插入 AddRmsNormCast 节点完成调用。算子原型定义在 op_graph/add_rms_norm_cast_proto.h:

  • 输入:x1(FLOAT16/BFLOAT16,ND)、x2(FLOAT16/BFLOAT16,ND)、gamma(FLOAT16/BFLOAT16,ND);
  • 输出:y1(FLOAT32)、y2(FLOAT16/BFLOAT16)、rstd(FLOAT32)、x(FLOAT16/BFLOAT16),格式均为 ND;
  • 属性:epsilon,Float 类型,默认值1e-6f。

即使用REG_OP(AddRmsNormCast)注册的图算子节点,框架侧可通过ge::Operator按上述输入输出与属性构造节点加入图。

源码实现纵深:Host 侧与 Kernel 侧的关键机制

算子定义(op_host/add_rms_norm_cast_def.cpp)

在 add_rms_norm_cast_def.cpp 中:

  • 所有输入输出都声明为REQUIRED(必选),数据类型与格式与文档一致;epsilon声明为OPTIONAL属性,默认1e-6;
  • ascend950配置开启动态编译静态化(DynamicCompileStaticFlag)、动态 rank 与动态 shape 支持(DynamicRankSupportFlag、DynamicShapeSupportFlag),并指定扩展编译配置add_rms_norm_cast_apt(对应 kernel 侧 add_rms_norm_cast_apt.cpp);
  • Kirin 系列(kirinx90、kirin9030)使用专属配置:只支持 FLOAT16、开启动态格式/动态 rank/动态 shape、关闭支持检查(NeedCheckSupportFlag(false))并开启精度降低(PrecisionReduceFlag(true))。

shape 推导(op_host/add_rms_norm_cast_infershape.cpp)

add_rms_norm_cast_infershape.cpp 实现了 shape 与 dtype 推导:

  • y1、y2、x的 shape 直接继承x1;
  • rstd的 shape 按"前xDimNum - gammaDimNum维与x1相同、norm 维为 1"规则构造(x1的维数小于gamma的维数时报错);
  • 数据类型:y1、rstd固定 FLOAT,y2、x跟随x1。

tiling 策略(op_host/add_rms_norm_cast_tiling.cpp)

tiling 是算子性能的关键。从 add_rms_norm_cast_tiling.cpp 可以看到其核心思路:

  • 数据视图:把输入按矩阵方式看待,numRow为x1前x1DimNum - gammaDimNum维的元素乘积(即归一化"行"数),numCol为gamma的元素个数(即归一化"列"数),avgFactor = 1/numCol作为均值因子参与计算;
  • tiling key 编码:tilingKey = dtypeKey * 10 + modeKey,其中dtypeKey(1=FLOAT16、2=FLOAT32、3=BFLOAT16)与modeKey(0=Normal、1=SplitD、2=MergeN、3=SingleN、4=MultiN)组合出多种实现变体。例如 10/30 对应 Normal 模式、11/31 对应 SplitD、13/33 对应 SingleN、14 对应 MultiN;
  • 模式选择:numCol超过 UB 容量(UB_FACTOR_B16=8704等阈值)时切到 SplitD 模式对归一化维做切分;blockFactor 为 1 且非特定 SoC 时走 SingleN 模式;
  • workspace 规划:默认申请约 16MB 系统 workspace 加 256B 用户 workspace;
  • 校验:调用前对x1/x2/y1/y2/x的 shape 一致性、gamma与x1后几维的一致性、rstd与x1前几维的一致性、维度数范围(1~8)以及epsilon >= 0做完整检查。

此外,add_rms_norm_cast_tiling.h 中注册了多套 tiling 数据类(AddRMSNormCastTilingData与AddRmsNormCastRegbaseTilingData),其中AddRmsNormCast_100~_103、_199等 tiling key 对应 arch35 平台的 RegBase 实现,说明该算子在 Atlas A3(arch35)等平台上有独立的 regbase tiling 路径,其 kernel 实现在 op_kernel/arch35 目录下(含add_rms_norm_cast_regbase.h、add_rms_norm_cast_regbase_high_performance.h、add_rms_norm_cast_regbase_single_n.h、add_rms_norm_cast_regbase_spilt_reduce.h等)。

Kernel 分发(op_kernel/add_rms_norm_cast.cpp)

add_rms_norm_cast.cpp 是 kernel 入口,按 tiling key 将计算分发给不同实现类:

  • TILING_KEY_IS(10/30)→KernelAddRmsNormCast(Normal,half / bfloat16_t);
  • TILING_KEY_IS(11/31)→KernelAddRmsNormCastSplitD(归一化维切分);
  • TILING_KEY_IS(13/33)→KernelAddRmsNormCastSingleN;
  • TILING_KEY_IS(14)→KernelAddRmsNormCastMultiN(BF16 多 N 场景)。

各实现类(add_rms_norm_cast_single_n.h、add_rms_norm_cast_multi_n.h、add_rms_norm_cast_split_d.h)都通过统一的Init(x1, x2, gamma, y1, y2, rstd, x, workspace, &tilingData)+Process()接口执行,先 Add 求和、再按 tiling 数据计算rstd与归一化结果、最后完成 cast 输出。

二进制配置与测试

每个产品目录下的 add_rms_norm_cast_binary.json(ascend910b、ascend910_93、ascend950、kirinx90、kirin9030各一份)以bin_filename(如AddRmsNormCast_fp16、AddRmsNormCast_bf16)绑定不同 dtype 的二进制,声明输入输出 shape 为动态(-2)、epsilon默认值0.000001,供编译框架生成对应的算子二进制。

测试方面:

  • Host 侧单测:tests/ut/op_host/test_AddRmsNormCast_infershape.cpp 覆盖 shape 推导(含 rstd 各维度规则),tests/ut/op_host/test_add_rms_norm_cast_tiling.cpp 覆盖 tiling 参数计算;
  • Kernel 侧单测:tests/ut/op_kernel/test_add_rms_norm_cast.cpp 与 tests/ut/op_kernel/test_add_rms_norm_cast_regbase.cpp 覆盖常规与 arch35 regbase 实现;
  • 系统测试:arch35 平台 ST 用例在 tests/st/arch35/ttk_kernel_add_rms_norm_cast_st.csv;
  • 数值 golden:tests/assets/golden.py 复用 add_rms_norm 的 golden 并通过_post_action='cast'生成参考结果。

总结

AddRmsNormCast 是 CANN ops-nn 中一个典型的"融合 + 多输出"归一化算子:它将 Add、RmsNorm、Cast 三合一,减少数据搬入搬出;除归一化结果外还输出rstd与求和结果x,为反向传播提供可复用中间量。使用时需重点把握三点:一是gamma决定归一化维度、rstd的 shape 随之确定;二是通过 aclnn 两段式接口(GetWorkspaceSize+ 执行)完成单算子调用,输出不支持非连续 Tensor;三是按产品选择支持的 dtype(Kirin 系列不支持 BFLOAT16)。从源码看,其 tiling 层针对不同数据规模选择了 Normal / SplitD / SingleN / MultiN 多种实现模式,并在 arch35 平台提供独立的 RegBase 高性能路径,体现了面向大模型归一化场景的性能优化思路。若需在计算图中使用,可直接基于 add_rms_norm_cast_proto.h 的算子 IR 构图。

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

【免费下载链接】ops-nn

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

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

相关推荐

上一篇:Hyperframes Cinematic 字幕模式实战:一个引擎、十种 DNA 视觉语言的纯嵌入字幕管线
下一篇:Remix UI Accordion 组件全解:从 `remix/ui/accordion` 样式组件到 primitives 无头原语

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/10/3 2:23:21

AI领跑与云智融合:大模型落地与工程实践指南

1. 从“会用AI”到“规模化用AI”&#xff0c;2024的拐点在哪2024年过了一半的时候&#xff0c;我已经明显感觉到一个变化&#xff1a;大家早就不聊“AI能不能做”&#xff0c;而是聊“AI怎么在业务里稳定地跑起来”。年初那种“你好我好大家好”的通识科普阶段过去了&#xff…

作者头像 李华
网站建设 2026/10/3 2:22:41

agno Agent 输入输出完整指南:9 个实战示例讲透结构化输入输出

agno Agent 输入输出完整指南&#xff1a;9 个实战示例讲透结构化输入输出 【免费下载链接】agno Build, run, and manage agent platforms. 项目地址: https://gitcode.com/GitHub_Trending/ag/agno 本文以 agno 的 9 个官方示例为线索&#xff0c;一次讲清 agno Agent…

作者头像 李华
网站建设 2026/10/3 2:20:20

linux-command 命令详解:volname 读取 ISO-9660 设备卷名称

文档教程 【免费下载链接】linux-command Linux命令大全搜索工具&#xff0c;内容包含Linux命令手册、详解、学习、搜集。https://git.io/linux 项目地址&#xff1a; https://gitcode.com/GitHub_Trending/linux/linux-command 点击查看 免费下载 本篇技术指南以 command/vol…

作者头像 李华