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

【免费下载链接】ops-nn

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

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

本技术指南以 CANN ops-nn 仓库中 aclnnAddRmsNormCast 接口文档 为主体,结合该算子在仓库内的算子定义、Shape 推导、Tiling 与 Kernel 实现,系统讲解 AddRmsNormCast 融合算子的数学原理、两段式 aclnn 接口用法、参数约束、编译运行与源码级实现。读完本文,你将能够独立完成 AddRmsNormCast 算子的 aclnn 接口调用、理解其内部行/列切分调度策略,并掌握在 Ascend 950、Atlas A2/A3 系列产品上使用该算子的完整方法。

一、算子背景:为什么需要 AddRmsNormCast 融合算子

RmsNorm(Root Mean Square Layer Normalization)是大模型(LLM)训练与推理中最常用的归一化操作之一。在实际网络结构中,残差连接(x1 + x2)之后通常紧跟着 RmsNorm 归一化,随后输出又往往需要经过数据类型转换(Cast)以满足后续算子的精度要求。如果按原始图逐个算子执行,数据需要在 Host 侧与 Device 侧之间反复"搬入搬出",带来可观的内存带宽开销。

AddRmsNormCast 算子的核心设计目标,正如 接口文档 与 README 所述:

"将 AddRmsNorm 后的 Cast 算子融合起来,减少搬入搬出操作。"

即在一次 Kernel 执行内完成三个步骤:

  1. Add:x_i = float(x1_i) + float(x2_i),两个输入相加;
  2. RmsNorm:对求和结果按 Norm 维度做均方根归一化,并乘上可选的缩放因子gamma;
  3. Cast:把归一化结果(FLOAT32 中间精度)转换回 FLOAT16 / BFLOAT16 输出y2Out。

同时该算子还会顺带输出中间量:归一化后(未 Cast)的y1Out(FLOAT32)、标准差的倒数rstdOut、以及 Add 求和结果xOut,供反向传播等下游算子复用,避免重复计算。

产品支持情况

依据 接口文档 及 README,该算子支持的产品如下:

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

在 add_rms_norm_cast_def.cpp 的算子注册代码中可以看到对应的AICore().AddConfig("ascend910b")(即 Atlas A2)、"ascend910_93"(Atlas A3)、"ascend950"、"kirinx90"、"kirin9030"配置项,与文档中的产品支持矩阵一一对应。

二、计算公式与算子 IR 定义

2.1 数学公式

接口文档给出的计算公式如下(i表示逐元素下标,n为 Norm 维度的元素总数):

$$ x_i=float(x1_{i})+float(x2_{i}) $$

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

$$ y2Out=cast(y1Out) $$

需要说明的是,文档公式中的float()表示先把 FLOAT16/BFLOAT16 输入提升到 FLOAT32 中间精度参与累加与归一化,避免累加误差,最后 y2Out 再 Cast 回原输入类型。这一提升在 Kernel 实现中通过Cast指令显式完成(详见后文源码分析)。

2.2 图模式 IR 定义

在 add_rms_norm_cast_proto.h 中,算子通过REG_OP宏注册了完整的输入输出与属性:

REG_OP(AddRmsNormCast) .INPUT(x1, TensorType({DT_FLOAT16, DT_BF16})) .INPUT(x2, TensorType({DT_FLOAT16, DT_BF16})) .INPUT(gamma, TensorType({DT_FLOAT16, DT_BF16})) .OUTPUT(y1, TensorType({DT_FLOAT})) .OUTPUT(y2, TensorType({DT_FLOAT16, DT_BF16})) .OUTPUT(rstd, TensorType({DT_FLOAT})) .OUTPUT(x, TensorType({DT_FLOAT16, DT_BF16})) .ATTR(epsilon, Float, 1e-6f) .OP_END_FACTORY_REG(AddRmsNormCast)

该 IR 头文件中的注释还给出了另一种等价的计算描述,便于理解算子内部逻辑顺序:

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)

其中rstd即"标准差(RMS)的倒数",也就是公式中1/Rms(x)的数值化表达。除了图模式(IR 构图)调用方式外,该算子还提供 aclnn 单算子调用方式,README 调用说明 给出了两条路径的索引:aclnn 接口见 examples 样例,图模式见 算子 IR。

三、函数原型与两段式接口调用范式

与 CANN 其他 aclnn 算子一致,AddRmsNormCast 采用两段式接口(详见 两段式接口说明):

  1. 第一段aclnnAddRmsNormCastGetWorkspaceSize:完成入参校验,计算算子执行所需的 workspace 大小,并创建包含完整计算流程的aclOpExecutor执行器;
  2. 第二段aclnnAddRmsNormCast:真正把算子下发到 Device 侧执行。

函数原型如下:

// 第一段:获取 workspace 大小与执行器 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)

四、第一段接口 GetWorkspaceSize 参数详解

4.1 参数表

第一段接口共有 8 个 Tensor/标量入参和 2 个输出参数,完整说明如下(对应 接口文档参数表):

参数名输入/输出描述使用说明数据类型数据格式维度(shape)非连续Tensor
x1(aclTensor*)输入Add 计算的第一个输入,对应公式x1支持空 TensorBFLOAT16、FLOAT16ND1-8√
x2(aclTensor*)输入Add 计算的第二个输入,对应公式x2支持空 Tensor;shape 与数据类型需与 x1 一致FLOAT16、BFLOAT16ND1-8√
gamma(aclTensor*)输入RmsNorm 缩放因子(权重),对应公式gamma支持空 Tensor;数据类型与 x1 一致;shape 需与 x1 后几维(即 Norm 维度)一致FLOAT16、BFLOAT16ND1-8√
epsilon(double)输入分母附加项,保证数值稳定,对应公式epsilon建议值 1e-6----
y1Out(aclTensor*)输出归一化后的输出,对应公式y1Out支持空 Tensor;shape、数据格式需与 x1 一致FLOAT32ND1-8×
y2Out(aclTensor*)输出归一化并类型转换后的输出,对应公式y2Out支持空 Tensor;shape、数据格式、数据类型均需与 x1 一致FLOAT16、BFLOAT16ND1-8×
rstdOut(aclTensor*)输出归一化标准差的倒数,对应Rms(x)的倒数支持空 Tensor;数据格式与 x1 一致;维度规则见下文FLOAT32ND1-8×
xOut(aclTensor*)输出Add 计算结果,对应公式x支持空 Tensor;shape、数据格式、数据类型均需与 x1 一致FLOAT16、BFLOAT16ND1-8×
workspaceSize(uint64_t*)输出需要在 Device 侧申请的 workspace 大小-----
executor(aclOpExecutor**)输出算子执行器,包含计算流程-----

说明:非连续 Tensor一列中,输入(x1/x2/gamma)标注 √ 表示支持非连续输入,输出(y1Out/y2Out/rstdOut/xOut)标注 × 表示输出必须为连续 Tensor。这一点在 README 约束说明 中也明确为"输出不支持非连续 Tensor"。

4.2 rstdOut 的 shape 推导规则

rstdOut的维度规则是参数中较容易出错的一项。接口文档给出了明确的推导方法:

  • rstdOut的维度数与x1保持一致;
  • 不需要 Norm 的维度(x1维度数减去gamma维度数后的前几维)与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_cast_infershape.cpp 中有逐行对应的实现:前xDimNum - gammaDimNum维继承x1的维度,后gammaDimNum维置 1。同时 infershape 单测 验证了x1(4,1,8) + gamma(8) → rstd(4,1,1)的推导结果,并覆盖了{4,0}空归约、{-2}未知维度等边界场景。

4.3 epsilon 的默认值与实现

epsilon为 double 类型标量,用于防止除 0 错误,建议值(同时也是默认值)为1e-6。在图模式 IR 中它被定义为可选的 Float 属性epsilon(默认1e-6f),在 add_rms_norm_cast_def.cpp 中通过.Attr("epsilon").AttrType(OPTIONAL).Float(1e-6)注册。Tiling 阶段会校验epsilon >= 0(见 tiling 源码),非法负值将返回错误。

4.4 第一段接口返回值与错误码

第一段接口返回aclnnStatus状态码(完整定义参见 aclnn 返回码说明)。入参校验不通过时会报错,典型错误码如下:

返回码错误码描述
ACLNN_ERR_PARAM_NULLPTR161001必选输入、输出或必选属性传入空指针
ACLNN_ERR_PARAM_INVALID161002输入或输出的数据类型不在支持范围之内
ACLNN_ERR_INNER_TILING_ERROR561002输入和输出不符合参数说明内的要求(shape/dtype/格式等)

五、第二段接口 aclnnAddRmsNormCast 参数详解

第二段接口参数全部为执行环境要素:

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

第二段接口同样返回aclnnStatus状态码。调用成功后,算子计算结果会写入用户预先创建的y1Out/y2Out/rstdOut/xOut四个输出 Tensor。

六、约束说明

依据 接口文档约束说明,使用时需注意:

维度边界:参数x1、x2、gamma、y1Out、y2Out、rstdOut、xOut的 shape 中每一维大小均不能超过 INT32 最大值 2147483647。此外从 tiling 实现 可见:x1/x2/gamma 维度数需在 [1, 8] 之间;x1 与 y1Out/y2Out/xOut/x2 的维度数必须一致;x1 的维度数不能小于 gamma 的维度数;x1 与 gamma 的对应维度需要相等。

边界值场景:

  • 当前不支持"非 Norm 维度元素总数大于 0 且 Norm 维度元素总数为 0"的空 Tensor 场景(即 x1 元素总数为 0 时空 Tensor 不被允许,除非 rstd 也为空,tiling 中有对应校验);
  • 当输入是 Inf 时,输出为 Inf;
  • 当输入是 NaN 时,输出为 NaN。

确定性计算:aclnnAddRmsNormCast为默认确定性实现,多次运行结果可复现。

平台差异:README 特别说明:Kirin X90、Kirin 9030 处理器上 x1、x2、gamma、y2 和 x 的数据类型不支持 BFLOAT16(仅支持 FLOAT16),该限制在 算子定义 的 Kirin 专用配置中通过DataType({ge::DT_FLOAT16})落实。

七、完整调用示例(可编译运行)

以下示例代码取自 examples/test_aclnn_add_rms_norm_cast.cpp,与 接口文档调用示例 内容一致。编译与执行的完整流程请参考 编译与运行样例。

#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; } 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); // 调用aclrtMalloc申请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); // 调用aclrtMemcpy将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. 构造输入与输出,需要根据API的接口自定义构造 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 = {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> 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 = {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<float> rstdHostData = {1, 2}; std::vector<short> xHostData = {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}; float epsilon = 1e-6; // 创建x1 aclTensor ret = CreateAclTensor(x1HostData, xShape, &x1DeviceAddr, aclDataType::ACL_FLOAT16, &x1); CHECK_RET(ret == ACL_SUCCESS, return ret); // 创建x2 aclTensor ret = CreateAclTensor(x2HostData, xShape, &x2DeviceAddr, aclDataType::ACL_FLOAT16, &x2); CHECK_RET(ret == ACL_SUCCESS, return ret); // 创建gamma aclTensor ret = CreateAclTensor(gammaHostData, gammaShape, &gammaDeviceAddr, aclDataType::ACL_FLOAT16, &gamma); CHECK_RET(ret == ACL_SUCCESS, return ret); // 创建y1 aclTensor ret = CreateAclTensor(y1HostData, yShape, &y1DeviceAddr, aclDataType::ACL_FLOAT, &y1); CHECK_RET(ret == ACL_SUCCESS, return ret); // 创建y2 aclTensor ret = CreateAclTensor(y2HostData, yShape, &y2DeviceAddr, aclDataType::ACL_FLOAT16, &y2); CHECK_RET(ret == ACL_SUCCESS, return ret); // 创建rstd aclTensor ret = CreateAclTensor(rstdHostData, rstdShape, &rstdDeviceAddr, aclDataType::ACL_FLOAT, &rstd); CHECK_RET(ret == ACL_SUCCESS, return ret); // 创建x aclTensor ret = CreateAclTensor(xHostData, xShape, &xDeviceAddr, aclDataType::ACL_FLOAT16, &x); CHECK_RET(ret == ACL_SUCCESS, return ret); // 3. 调用CANN算子库API uint64_t workspaceSize = 0; aclOpExecutor* executor; // 调用aclnnAddRmsNormCast第一段接口 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); } // 调用aclnnAddRmsNormCast第二段接口 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]); } std::vector<uint16_t> resultData1(size, 0); ret = aclrtMemcpy( resultData1.data(), resultData1.size() * sizeof(resultData1[0]), y2DeviceAddr, size * sizeof(resultData1[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("y2 result[%ld] is: 0x%04x\n", i, resultData1[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; }

示例中的关键点:

  • 示例数据以FLOAT16 十六进制位模式(如0x3C00对应 1.0)构造输入,便于逐位核对输出;
  • rstdShape = {2, 1}正是按上文规则:x1(2,16)减gamma(16)的维度后,前 1 维保持 2、后 1 维置 1;
  • 输出的y1为 FLOAT32、y2为 FLOAT16,分别用float与uint16_t容器回读打印,可以直观看到 y2 是 y1 的 Cast 结果。

八、源码级实现原理:Shape 推导、Tiling 切分与 Kernel 分发

8.1 Shape 与 DataType 推导

add_rms_norm_cast_infershape.cpp 注册了InferShape与InferDataType两个实现:

  • InferShape:y1、y2、x三个输出的 shape 直接继承x1;rstd按"前xDimNum - gammaDimNum维继承 x1、后gammaDimNum维置 1"的规则构造;当x1或gamma为未知维度(unknown rank)时,rstd也置为 unknown;
  • InferDataType:y1、rstd固定为DT_FLOAT(FLOAT32),y2、x继承x1的数据类型。

对应的单测见 test_AddRmsNormCast_infershape.cpp,覆盖了(4,1,8)常规场景、{-2}未知维度场景以及{4,0}空归约场景。

8.2 Tiling:行/列切分与多种计算模式

Tiling 阶段(add_rms_norm_cast_tiling.cpp)负责把输入张量切分成多个 tile,并决定使用多少个 AI Core、选用哪种 Kernel 实现。其核心思路:

  1. 切分视角:把x1看作numRow × numCol的二维矩阵,其中numCol = gamma 元素总数(Norm 维),numRow = 其余维度元素之积(行维)。avgFactor = 1 / numCol作为均值系数预先算好写入 tiling 数据;
  2. 核数计算:按numRow与 AI Core 总数计算blockFactor(每个核处理的行数)与useCoreNum(实际使用的核数),写入block_dim;
  3. 模式选择:根据numCol大小与数据类型在多种模式间切换(源码注释给出 5 种模式:0 Normal、1 SplitD、2 MergeN、3 SingleN、4 MultiN):
    • 当numCol超过 UB 容量阈值时进入SplitD模式,把 Norm 维再切分成多段,逐段累加平方和;
    • 当仅用单核且平台非 310P 时进入SingleN模式,一行一核高效处理;
    • FLOAT16 且numCol恰好对齐时使用Normal模式;
    • 更多核时还有MultiN模式用于多行/多核协同。
  4. Tiling Key:最终以dtypeKey * 10 + modeKey生成 tiling key(如 10=FP16 Normal、11=FP16 SplitD、13=FP16 SingleN、14=FP16 MultiN、30=BF16 Normal、31=BF16 SplitD、33=BF16 SingleN),供 Kernel 侧分发使用。tiling 数据中还写入了epsilon、avg_factor等参数;
  5. Workspace:常规路径申请usrSize(256B) + 16MB 系统 workspace。

Tiling 单测见 test_add_rms_norm_cast_tiling.cpp。此外,Ascend 950 系列(DAV_3510 架构)走独立的 Regbase 路径(add_rms_norm_cast_tiling_arch35.cpp 与 regbase 系列头文件),支持 128B 对齐的高性能调度,并注册了AddRmsNormCast_100/101/102/103/199多组 tiling 数据类。

8.3 Kernel:按 Tiling Key 分发执行的融合计算

Kernel 侧 add_rms_norm_cast.cpp 是一个典型的"按 tiling key 分发"入口:

extern "C" __global__ __aicore__ void add_rms_norm_cast(GM_ADDR x1, GM_ADDR x2, GM_ADDR gamma, GM_ADDR y1, GM_ADDR y2, GM_ADDR rstd, GM_ADDR x, GM_ADDR workspace, GM_ADDR tiling) { TPipe pipe; GET_TILING_DATA(tilingData, tiling); if (TILING_KEY_IS(10)) { GENERAL_OP_IMPL(KernelAddRmsNormCast, half); // FP16 Normal } else if (TILING_KEY_IS(30)) { GENERAL_OP_IMPL(KernelAddRmsNormCast, bfloat16_t); // BF16 Normal } else if (TILING_KEY_IS(11)) { GENERAL_OP_IMPL(KernelAddRmsNormCastSplitD, half); // FP16 SplitD } else if (TILING_KEY_IS(31)) { GENERAL_OP_IMPL(KernelAddRmsNormCastSplitD, bfloat16_t); } else if (TILING_KEY_IS(13)) { GENERAL_OP_IMPL(KernelAddRmsNormCastSingleN, half); // FP16 SingleN } else if (TILING_KEY_IS(33)) { GENERAL_OP_IMPL(KernelAddRmsNormCastSingleN, bfloat16_t); } else if (TILING_KEY_IS(14)) { GENERAL_OP_IMPL(KernelAddRmsNormCastMultiN, half); // FP16 MultiN } }

在默认 Normal 模式的 add_rms_norm_cast.h 中,KernelAddRmsNormCast::Init按 tiling 数据计算每个核负责的行区间与起始偏移,Process按rowFactor分批处理,SubProcess内部完成完整计算链:

  1. CopyIn:DataCopyCustom把x1、x2对应行数据搬入 UB(Vector 单元),Add求和后Cast到 FLOAT32(BF16 路径则先把两个输入分别 Cast 到 FP32 再相加,最后 Cast 回原类型写出xOut);
  2. 平方累加:Mul(sqx, xFp32, xFp32)求平方 →Muls(sqx, sqx, avgFactor)乘均值系数 →ReduceSumCustom做归约求和 →Adds加epsilon→Sqrt开方得到Rms(x)→Div(1, Rms)得到 rstd;
  3. 归一化与缩放:rstd以Brcb(broadcast)广播到整行,x * rstd * gamma完成归一化与缩放;
  4. 双输出:归一化结果 Cast 到 FLOAT16/BFLOAT16 写出y2(CopyOutY),同时把 FLOAT32 结果写出y1,rstd 也单独写出rstd。

整条流水使用TPipe双缓冲队列(BUFFER_NUM深度)并通过PipeBarrier、SetFlag/WaitFlag同步 Vector 与搬运(MTE)单元,实现数据搬运与计算的流水重叠——这正是该算子"减少搬入搬出"性能收益在 Kernel 层的具体体现。

8.4 测试与验证资产

仓库为该算子提供了完整的验证体系:

  • Host 侧单测:test_AddRmsNormCast_infershape.cpp(Shape/DataType 推导)、test_add_rms_norm_cast_tiling.cpp(Tiling 参数);
  • Kernel 侧单测:test_add_rms_norm_cast.cpp 与 test_add_rms_norm_cast_regbase.cpp;
  • ST 用例:arch35 平台的 CSV 用例 ttk_kernel_add_rms_norm_cast_st.csv;
  • Golden 脚本:golden.py,用于生成期望结果,可直接对照验证算子输出;
  • 平台配置:各产品的 binary 配置位于 op_host/config(ascend910b / ascend910_93 / ascend950 / kirin9030 / kirinx90)。

九、总结与使用建议

AddRmsNormCast 是 CANN ops-nn 中面向大模型残差连接场景的典型融合算子,把 Add、RmsNorm、Cast 三步合并为一次 Kernel 执行,并通过四个输出(y1Out/y2Out/rstdOut/xOut)最大程度复用中间结果。实际接入时建议遵循以下要点:

  1. 按两段式接口流程调用:先aclnnAddRmsNormCastGetWorkspaceSize完成校验并申请 workspace,再aclnnAddRmsNormCast下发执行,结束时同步 Stream 并回收资源;
  2. 严格遵循 shape 规则:x2、y1Out、y2Out、xOut的 shape 与x1一致;gamma对应x1的后几维(Norm 维);rstdOut的 Norm 维置 1;
  3. 注意平台差异:Ascend 950PR&950DT、Atlas A2/A3 系列支持 FLOAT16 与 BFLOAT16;Kirin 系列仅支持 FLOAT16;Atlas 推理/训练系列与 200I/500 A2 推理产品不支持本算子;
  4. 关注边界约束:各维大小不超过 INT32 最大值,避免"非空行 + 空 Norm 维"的组合,输入含 Inf/NaN 时结果会原样传递。

如需在更多产品上扩展支持或深入性能调优,可进一步阅读 add_rms_norm_cast_tiling_arch35.cpp(Regbase 高性能调度)与 arch35 目录(单核/多核/拆分归约等专用实现)。

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

【免费下载链接】ops-nn

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

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

相关推荐

上一篇:Gutenberg ConfirmDialog 组件完全指南:基于 Modal 的受控与非受控确认对话框
下一篇:G-Helper深度解析:华硕笔记本硬件控制与性能优化技术方案

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

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

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

立即咨询