- 人工智能
- 算子库
- 深度学习
- CANN
- Ascend
【免费下载链接】ops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
本技术指南以 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 执行内完成三个步骤:
- Add:
x_i = float(x1_i) + float(x2_i),两个输入相加; - RmsNorm:对求和结果按 Norm 维度做均方根归一化,并乘上可选的缩放因子
gamma; - 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 采用两段式接口(详见 两段式接口说明):
- 第一段
aclnnAddRmsNormCastGetWorkspaceSize:完成入参校验,计算算子执行所需的 workspace 大小,并创建包含完整计算流程的aclOpExecutor执行器; - 第二段
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 | 支持空 Tensor | BFLOAT16、FLOAT16 | ND | 1-8 | √ |
| x2(aclTensor*) | 输入 | Add 计算的第二个输入,对应公式x2 | 支持空 Tensor;shape 与数据类型需与 x1 一致 | FLOAT16、BFLOAT16 | ND | 1-8 | √ |
| gamma(aclTensor*) | 输入 | RmsNorm 缩放因子(权重),对应公式gamma | 支持空 Tensor;数据类型与 x1 一致;shape 需与 x1 后几维(即 Norm 维度)一致 | FLOAT16、BFLOAT16 | ND | 1-8 | √ |
| epsilon(double) | 输入 | 分母附加项,保证数值稳定,对应公式epsilon | 建议值 1e-6 | - | - | - | - |
| y1Out(aclTensor*) | 输出 | 归一化后的输出,对应公式y1Out | 支持空 Tensor;shape、数据格式需与 x1 一致 | FLOAT32 | ND | 1-8 | × |
| y2Out(aclTensor*) | 输出 | 归一化并类型转换后的输出,对应公式y2Out | 支持空 Tensor;shape、数据格式、数据类型均需与 x1 一致 | FLOAT16、BFLOAT16 | ND | 1-8 | × |
| rstdOut(aclTensor*) | 输出 | 归一化标准差的倒数,对应Rms(x)的倒数 | 支持空 Tensor;数据格式与 x1 一致;维度规则见下文 | FLOAT32 | ND | 1-8 | × |
| xOut(aclTensor*) | 输出 | Add 计算结果,对应公式x | 支持空 Tensor;shape、数据格式、数据类型均需与 x1 一致 | FLOAT16、BFLOAT16 | ND | 1-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_NULLPTR | 161001 | 必选输入、输出或必选属性传入空指针 |
| ACLNN_ERR_PARAM_INVALID | 161002 | 输入或输出的数据类型不在支持范围之内 |
| ACLNN_ERR_INNER_TILING_ERROR | 561002 | 输入和输出不符合参数说明内的要求(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 实现。其核心思路:
- 切分视角:把
x1看作numRow × numCol的二维矩阵,其中numCol = gamma 元素总数(Norm 维),numRow = 其余维度元素之积(行维)。avgFactor = 1 / numCol作为均值系数预先算好写入 tiling 数据; - 核数计算:按
numRow与 AI Core 总数计算blockFactor(每个核处理的行数)与useCoreNum(实际使用的核数),写入block_dim; - 模式选择:根据
numCol大小与数据类型在多种模式间切换(源码注释给出 5 种模式:0 Normal、1 SplitD、2 MergeN、3 SingleN、4 MultiN):- 当
numCol超过 UB 容量阈值时进入SplitD模式,把 Norm 维再切分成多段,逐段累加平方和; - 当仅用单核且平台非 310P 时进入SingleN模式,一行一核高效处理;
- FLOAT16 且
numCol恰好对齐时使用Normal模式; - 更多核时还有MultiN模式用于多行/多核协同。
- 当
- 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等参数; - 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内部完成完整计算链:
- CopyIn:
DataCopyCustom把x1、x2对应行数据搬入 UB(Vector 单元),Add求和后Cast到 FLOAT32(BF16 路径则先把两个输入分别 Cast 到 FP32 再相加,最后 Cast 回原类型写出xOut); - 平方累加:
Mul(sqx, xFp32, xFp32)求平方 →Muls(sqx, sqx, avgFactor)乘均值系数 →ReduceSumCustom做归约求和 →Adds加epsilon→Sqrt开方得到Rms(x)→Div(1, Rms)得到 rstd; - 归一化与缩放:
rstd以Brcb(broadcast)广播到整行,x * rstd * gamma完成归一化与缩放; - 双输出:归一化结果 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)最大程度复用中间结果。实际接入时建议遵循以下要点:
- 按两段式接口流程调用:先
aclnnAddRmsNormCastGetWorkspaceSize完成校验并申请 workspace,再aclnnAddRmsNormCast下发执行,结束时同步 Stream 并回收资源; - 严格遵循 shape 规则:
x2、y1Out、y2Out、xOut的 shape 与x1一致;gamma对应x1的后几维(Norm 维);rstdOut的 Norm 维置 1; - 注意平台差异:Ascend 950PR&950DT、Atlas A2/A3 系列支持 FLOAT16 与 BFLOAT16;Kirin 系列仅支持 FLOAT16;Atlas 推理/训练系列与 200I/500 A2 推理产品不支持本算子;
- 关注边界约束:各维大小不超过 INT32 最大值,避免"非空行 + 空 Norm 维"的组合,输入含 Inf/NaN 时结果会原样传递。
如需在更多产品上扩展支持或深入性能调优,可进一步阅读 add_rms_norm_cast_tiling_arch35.cpp(Regbase 高性能调度)与 arch35 目录(单核/多核/拆分归约等专用实现)。
- 人工智能
- 算子库
- 深度学习
- CANN
- Ascend
【免费下载链接】ops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
相关推荐
CANN ops-nn AddRmsNorm 算子深度解析:Add 与 RmsNorm 融合的实现原理与 aclnn 接口实战
CANN ops nn AddRmsNorm 算子深度解析:Add 与 RmsNorm 融合的实现原理与 aclnn 接口实战 AddRmsNorm 是 CAN
人工智能算子库深度学习CANNAscendCANN ops-nn 算子指南:aclnnAddRmsNorm 融合算子(Add + RmsNorm)的接口原理与实战调用
CANN ops nn 算子指南:aclnnAddRmsNorm 融合算子(Add + RmsNorm)的接口原理与实战调用 AddRmsNorm 是 CANN
人工智能算子库深度学习CANNAscendOpenArk:Windows内核级安全分析的开源工具箱
OpenArk:Windows内核级安全分析的开源工具箱 任务管理器里那行占着30% CPU的陌生进程,你大概率不知道它从哪来。OpenArk 是一款面向 Wi
人工智能算子库深度学习CANNAscend
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考