CANN ops-nn 中 ApplyCenteredRMSProp 算子的原理、约束与 aclnn 调用实践
【免费下载链接】ops-nn本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-nn
导读
ApplyCenteredRMSProp 是 CANN ops-nn 神经网络算子库中为 Ascend 950 系列(arch35 / DAV_3510)实现的一类"中心化" RMSProp 优化器算子,功能对标tf.raw_ops.ApplyCenteredRMSProp:在经典 RMSProp 基础上额外维护一阶梯度指数移动平均mg,并以ms - mg^2作为方差估计参与归一化,从而获得更稳定的更新步长。本文以算子模块 README 为骨架,结合仓库中的算子定义、Tiling、Kernel 与示例源码,完整讲解其数学原理、参数规格、约束条件以及基于 aclnn 两段式接口的调用与验证方案,读者阅读后可以直接在 Ascend 950 环境上完成该算子的编译、调用与结果校验。
产品支持情况
| 产品 | 是否支持 |
|---|---|
| Ascend 950PR/Ascend 950DT | √ |
| Atlas A3 训练系列产品/Atlas A3 推理系列产品 | × |
| Atlas A2 训练系列产品/Atlas A2 推理系列产品 | × |
| Atlas 200I/500 A2 推理产品 | × |
| Atlas 推理系列产品 | × |
| Atlas 训练系列产品 | × |
从源码角度可以交叉印证这一限制:在算子定义文件 apply_centered_rms_prop_def.cpp 中,OpAICoreConfig只通过this->AICore().AddConfig("ascend950", aiCoreConfig)注册了ascend950一个计算单元;模块级 CMakeLists.txt 同样通过set(SUPPORT_COMPUTE_UNIT "ascend950")声明仅支持 ascend950。因此该算子目前是 Ascend 950 专属实现,不适用于其他芯片代际。
功能说明
算子定位
ApplyCenteredRMSProp 是带"中心化"修正的 RMSProp 优化器算子,功能对标tf.raw_ops.ApplyCenteredRMSProp。与普通 RMSProp 相比,它额外维护一阶梯度指数移动平均mg,并以ms - mg^2作为方差估计参与归一化,从而获得更稳定的步长。var/mg/ms/mom均为 Ref Tensor,算子执行后原地更新。
计算公式
$$ \begin{aligned} mg_t &= \rho \cdot mg_{t-1} + (1 - \rho) \cdot \text{grad}t \ ms_t &= \rho \cdot ms{t-1} + (1 - \rho) \cdot \text{grad}t^2 \ denom_t &= \sqrt{ms_t - mg_t^2 + \epsilon} \ mom_t &= \text{momentum} \cdot mom{t-1} + \text{lr} \cdot \frac{\text{grad}t}{denom_t} \ var_t &= var{t-1} - mom_t \end{aligned} $$
关键说明
var/mg/ms/mom为 Ref Tensor,与对应的*_out输出共享存储以实现 inplace 更新;lr、rho、momentum、epsilon为 0-D 或 1 元素 1-D 标量 Tensor;- 逐元素独立计算,无跨元素/跨核依赖,天然适合多核并行切分。
源码中的实现印证
上述公式在 Kernel 类 apply_centered_rms_prop.h 的Compute()中逐条落地,并且有两个值得注意的工程细节:
- fp16 路径先提升精度再回写:当输入为
half时,CopyIn后先Cast(fp16 -> fp32),全部中间计算在 fp32 精度下进行(Muls/Mul/Add/Sub/Sqrt/Div),写回前再Cast(fp32 -> fp16)(使用RoundMode::CAST_RINT)。 - 对负数开方做数值保护:由于 fp16 截断可能造成
ms_new < mg_new^2,导致Sqrt(NaN)并向下游传播,Kernel 在计算ms_new - mg_new^2后先执行Maxs(tmp2, tmp2, 0.0f)将内部钳制到 ≥ 0,再加epsilon后开方,与示例中 CPU golden 的数值安全防护保持一致(见 test_aclnn_apply_centered_rms_prop.cpp 的denom计算)。
参数说明
下表完整列出算子的 9 个输入与 4 个输出。var/mg/ms/mom四个输入为 Ref Tensor,与其对应的*_out输出共享 Device 存储。
| 参数名 | 输入/输出/属性 | 描述 | 数据类型 | 数据格式 |
|---|---|---|---|---|
| var | 输入 | 公式中的 var,待更新的模型参数(Ref Tensor,原地更新)。shape 与 mg/ms/mom/grad 一致。 | FLOAT16, FLOAT | ND |
| mg | 输入 | 公式中的 mg,一阶梯度指数移动平均(Ref Tensor,原地更新)。shape 与 var/ms/mom/grad 一致。 | FLOAT16, FLOAT | ND |
| ms | 输入 | 公式中的 ms,二阶梯度指数移动平均(Ref Tensor,原地更新)。shape 与 var/mg/mom/grad 一致。 | FLOAT16, FLOAT | ND |
| mom | 输入 | 公式中的 mom,动量项(Ref Tensor,原地更新)。shape 与 var/mg/ms/grad 一致。 | FLOAT16, FLOAT | ND |
| lr | 输入 | 公式中的 lr,学习率。0-D 或 1 元素 1-D Tensor。 | FLOAT16, FLOAT | ND |
| rho | 输入 | 公式中的 rho,指数衰减系数。0-D 或 1 元素 1-D Tensor。 | FLOAT16, FLOAT | ND |
| momentum | 输入 | 公式中的 momentum,动量系数。0-D 或 1 元素 1-D Tensor。 | FLOAT16, FLOAT | ND |
| epsilon | 输入 | 公式中的 epsilon,数值稳定项(> 0)。0-D 或 1 元素 1-D Tensor。 | FLOAT16, FLOAT | ND |
| grad | 输入 | 公式中的 grad,当前步的梯度张量。shape 与 var/mg/ms/mom 一致。 | FLOAT16, FLOAT | ND |
| var_out | 输出 | 更新后的参数,与输入 var 共享存储(inplace 更新)。 | FLOAT16, FLOAT | ND |
| mg_out | 输出 | 更新后的一阶梯度均值,与输入 mg 共享存储(inplace 更新)。 | FLOAT16, FLOAT | ND |
| ms_out | 输出 | 更新后的二阶梯度均值,与输入 ms 共享存储(inplace 更新)。 | FLOAT16, FLOAT | ND |
| mom_out | 输出 | 更新后的动量项,与输入 mom 共享存储(inplace 更新)。 | FLOAT16, FLOAT | ND |
算子定义层的参数声明
参数规格与算子定义文件 apply_centered_rms_prop_def.cpp 中的声明一一对应:
- 9 个输入按顺序声明为
var、mg、ms、mom、lr、rho、momentum、epsilon、grad,每个输入均限定DataType({ge::DT_FLOAT16, ge::DT_FLOAT})与Format({ge::FORMAT_ND, ge::FORMAT_ND}); - 4 个输出声明为
var_out、mg_out、ms_out、mom_out; var/mg/ms/mom/grad等主 Tensor 均调用了.AutoContiguous(),要求连续排布;lr/rho/momentum/epsilon四个标量输入未设置AutoContiguous,其形状合法性由 Host 侧 Tiling 校验。
形状与数据类型推导
在 apply_centered_rms_prop_infershape.cpp 中,InferShape4ApplyCenteredRMSProp将各输出 shape 直接拷贝自对应 Ref 输入(var_out←var,mg_out←mg,ms_out←ms,mom_out←mom);InferDataType4ApplyCenteredRMSProp则将输出 dtype 直接继承对应输入,即输出类型与输入一致(FLOAT16 或 FLOAT)。
约束说明
使用该算子时须满足以下约束:
- 仅支持 Ascend 950PR/Ascend 950DT(arch35 / DAV_3510),不适配其他芯片代际;
- 支持
float16与float32数据类型,所有输入 Tensor 的 dtype 必须一致; var、mg、ms、mom、grad五者 shape 必须完全一致,且均为连续排布的 ND Tensor;lr、rho、momentum、epsilon必须为 0-D 或 1 元素 1-D 的标量 Tensor;- 调用方需保证
epsilon > 0、ms - mg^2 + epsilon > 0;当denom == 0时rsqrt输出 Inf/NaN,行为与 PyTorch / TensorFlow 原生实现一致,需由上游调用方规避; var/mg/ms/mom为 Ref Tensor,Host aclnn 侧必须显式构造四个占位输出 Tensor(var_out/mg_out/ms_out/mom_out),并与各自输入共享 Device 地址以保证 inplace 语义。
Host Tiling 侧的强校验
上述约束并非仅停留在文档层面,Host 侧 Tiling 函数 apply_centered_rms_prop_tiling.cpp 的GetShapeInfo()会做三层防御性校验:
- dtype 校验:只接受
ge::DT_FLOAT与ge::DT_FLOAT16,其余类型直接返回GRAPH_FAILED; - 主 Tensor 形状一致性:逐一校验
mg/ms/mom/grad(输入索引 1/2/3/8)的numel必须等于var的numel,否则报错并给出具体张量名与元素数; - 标量形状自防御:逐一校验
lr/rho/momentum/epsilon(输入索引 4/5/6/7)的numel必须为 1(0-D 视为 1),空标量会被拒绝——因为 Kernel 的LoadScalar会读取element[0],空标量将导致越界访问。
调用说明
| 调用方式 | 调用样例 | 说明 |
|---|---|---|
| aclnn 调用 | test_aclnn_apply_centered_rms_prop | Ascend 950 上通过 aclnn 两段式接口aclnnApplyCenteredRMSPropGetWorkspaceSize→aclnnApplyCenteredRMSProp调用。var_out/mg_out/ms_out/mom_out需与对应 Ref 输入共享 Device 地址以保证 inplace 更新。 |
两段式 aclnn 调用流程详解
仓库在 examples/arch35/test_aclnn_apply_centered_rms_prop.cpp 中提供了最小可运行的 fp32 演示(shape=[16],与 CPU golden 比对),其完整流程如下:
- ACL 初始化:
aclInit(nullptr)→aclrtSetDevice(0)→aclrtCreateStream(&stream); - 构造 Host 数据:初始化
var/mg/ms/mom/grad(示例中ms[i] = 0.20f + 0.01f * i保证恒正),并取lr=1e-2f、rho=0.9f、momentum=0.9f、epsilon=1e-6f;同时在 H2D 拷贝前用ComputeGolden()在 CPU 端按公式算出期望结果; - 分配 Device 缓冲并 H2D:5 个主 Tensor 按
N * sizeof(float)分配,4 个标量按sizeof(float)分配,逐一aclrtMemcpy到 Device; - 构建 aclTensor:用
aclCreateTensor构造 9 个输入 Tensor;标量用空的std::vector<int64_t>构造 0-D 张量; - 关键一步——构造共享存储的占位输出:
t_var_o = MakeTensor(d_var, shape, ACL_FLOAT)等四个输出 Tensor必须复用与对应 Ref 输入相同的 Device 指针,否则无法观察 inplace 更新结果; - 两段式调用:先
aclnnApplyCenteredRMSPropGetWorkspaceSize(t_var, ..., &ws_size, &executor)查询并分配 workspace,再aclnnApplyCenteredRMSProp(ws_ptr, ws_size, executor, stream)执行,最后aclrtSynchronizeStream同步; - D2H 回读比对:把更新后的
var/mg/ms/mom拷回 Host,与 golden 按atol=1e-5 + rtol=1e-4 * |expect|的容差逐元素比对,输出Result (var): 16/16 passed与最终的ALL PASS/FAIL; - 资源清理:依次
aclDestroyTensor、aclrtFree(workspace 与各 Device 缓冲)、aclrtDestroyStream、aclrtResetDevice、aclFinalize。
底层 Kernel 与 Tiling 执行路径
当 aclnn 接口下发后,算子依次经过 Host Tiling 与 Device Kernel 两个阶段:
- Tiling 阶段:
ApplyCenteredRMSPropTilingFunc首先读取 AIV 核数与 UB 容量;随后按blockFactor = CeilAlign(CeilDiv(totalElements, coreNum), ubBlockSize)把总元素数平均切分到各 AI 向量核(块大小向上对齐到 32B DMA 粒度,避免相邻核互相踩踏);再按 dtype 选择 UB 预算(fp32 路径 21 个 fp32-unit/元素,fp16 路径 17 个),得出ubFactor,并以TILE_ELEM_NUM_TARGET = 2048作为单次 UB 迭代目标元素数上限。TilingData 结构体定义于 apply_centered_rms_prop_tiling_data.h,包含totalElements/blockFactor/ubFactor三个字段。工作空间大小为 0(tiling 中WS_SYS_SIZE = 0)。 - Kernel 阶段:Kernel 入口 apply_centered_rms_prop.cpp 先解析 TilingData,再调用
NsApplyCenteredRMSProp::ApplyCenteredRMSProp<D_T_VAR>的Init()与Process()。Init()按blockFactor * GetBlockIdx()计算本核偏移与blockLen_,并把四个标量通过LoadScalar读取为 float(oneMinusRho_ = 1.0f - rho在 Init 时预计算);Process()以ubFactor_为步长循环执行CopyInTile → Compute → CopyOutTile。 - dtype 模板分派:TilingKey 由输入 dtype 唯一决定(apply_centered_rms_prop_tiling_key.h 中
D_T_VAR = C_DT_FLOAT16对应 TilingKey 2,D_T_VAR = C_DT_FLOAT对应 TilingKey 1),通过ASCENDC_TPL_SEL_PARAM(context, dTypeVar)在 Tiling 侧选择模板参数,Kernel 内用if constexpr (std::is_same_v<T, half>)编译期分流两条路径。 - 尾块填充技巧:
CopyInTile对尾块做 pad 感知的DataCopyPad右填充——var/mg/mom/grad以 0 填充(0 参与 Add/Mul 不改变累加结果),而ms以 1 填充(保证sqrt(1 - 0 + eps) > 0,避免除 0),从而保证非对齐长度的末尾元素计算安全;CopyOutTile只回写有效元素数。
构建与运行方式
- 算子本体通过模块级 CMakeLists.txt 纳入构建:
add_modules_sources指定OPTYPE apply_centered_rms_prop、ACLNNTYPE aclnn(生成aclnnApplyCenteredRMSProp系列接口)、COMPUTE_UNIT ascend950,且DISABLE_IN_OPP TRUE(该算子不作为 OPP 内置算子合入,属于独立交付形态)。 - 示例程序按文件头注释通过
examples/run.sh构建与运行(仓库当前该目录下仅含示例源码,运行脚本属于该模块示例配套内容);也可以参照其代码骨架,将shape、dtype 与超参数替换为实际训练场景取值,并保持"四个占位输出与 Ref 输入共享 Device 地址"这一核心写法不变。
总结
ApplyCenteredRMSProp 是 CANN ops-nn 为 Ascend 950 系列新增的中心化 RMSProp 优化器算子,具备完整的算子定义、Host 侧形状/类型推导与 Tiling、Device 侧 fp16/fp32 双路径 Kernel 以及可直接运行验证的 aclnn 两段式调用示例。使用时的三个要点:一是严格保持var/mg/ms/mom/grad五者 shape 与 dtype 一致、四个超参为 0-D 或单元素 1-D;二是由调用方保证epsilon > 0且ms - mg^2 + epsilon > 0;三是在 Host 侧显式构造与 Ref 输入共享 Device 地址的四个占位输出,以获得正确的 inplace 语义。Kernel 内部对 fp16 路径采用"fp32 精度计算 + 负值钳制"的双重数值保护,进一步保证了更新结果在低精度输入下的稳定性。
【免费下载链接】ops-nn本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-nn
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考