CANN ops-nn 中 ApplyCenteredRMSProp 算子的原理、约束与 aclnn 调用实践
2026/9/20 13:37:27 网站建设 项目流程

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 更新;
  • lrrhomomentumepsilon为 0-D 或 1 元素 1-D 标量 Tensor;
  • 逐元素独立计算,无跨元素/跨核依赖,天然适合多核并行切分。

源码中的实现印证

上述公式在 Kernel 类 apply_centered_rms_prop.h 的Compute()中逐条落地,并且有两个值得注意的工程细节:

  1. fp16 路径先提升精度再回写:当输入为half时,CopyIn后先Cast(fp16 -> fp32),全部中间计算在 fp32 精度下进行(Muls/Mul/Add/Sub/Sqrt/Div),写回前再Cast(fp32 -> fp16)(使用RoundMode::CAST_RINT)。
  2. 对负数开方做数值保护:由于 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, FLOATND
mg输入公式中的 mg,一阶梯度指数移动平均(Ref Tensor,原地更新)。shape 与 var/ms/mom/grad 一致。FLOAT16, FLOATND
ms输入公式中的 ms,二阶梯度指数移动平均(Ref Tensor,原地更新)。shape 与 var/mg/mom/grad 一致。FLOAT16, FLOATND
mom输入公式中的 mom,动量项(Ref Tensor,原地更新)。shape 与 var/mg/ms/grad 一致。FLOAT16, FLOATND
lr输入公式中的 lr,学习率。0-D 或 1 元素 1-D Tensor。FLOAT16, FLOATND
rho输入公式中的 rho,指数衰减系数。0-D 或 1 元素 1-D Tensor。FLOAT16, FLOATND
momentum输入公式中的 momentum,动量系数。0-D 或 1 元素 1-D Tensor。FLOAT16, FLOATND
epsilon输入公式中的 epsilon,数值稳定项(> 0)。0-D 或 1 元素 1-D Tensor。FLOAT16, FLOATND
grad输入公式中的 grad,当前步的梯度张量。shape 与 var/mg/ms/mom 一致。FLOAT16, FLOATND
var_out输出更新后的参数,与输入 var 共享存储(inplace 更新)。FLOAT16, FLOATND
mg_out输出更新后的一阶梯度均值,与输入 mg 共享存储(inplace 更新)。FLOAT16, FLOATND
ms_out输出更新后的二阶梯度均值,与输入 ms 共享存储(inplace 更新)。FLOAT16, FLOATND
mom_out输出更新后的动量项,与输入 mom 共享存储(inplace 更新)。FLOAT16, FLOATND

算子定义层的参数声明

参数规格与算子定义文件 apply_centered_rms_prop_def.cpp 中的声明一一对应:

  • 9 个输入按顺序声明为varmgmsmomlrrhomomentumepsilongrad,每个输入均限定DataType({ge::DT_FLOAT16, ge::DT_FLOAT})Format({ge::FORMAT_ND, ge::FORMAT_ND})
  • 4 个输出声明为var_outmg_outms_outmom_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_outvarmg_outmgms_outmsmom_outmom);InferDataType4ApplyCenteredRMSProp则将输出 dtype 直接继承对应输入,即输出类型与输入一致(FLOAT16 或 FLOAT)。

约束说明

使用该算子时须满足以下约束:

  • 仅支持 Ascend 950PR/Ascend 950DT(arch35 / DAV_3510),不适配其他芯片代际;
  • 支持float16float32数据类型,所有输入 Tensor 的 dtype 必须一致;
  • varmgmsmomgrad五者 shape 必须完全一致,且均为连续排布的 ND Tensor;
  • lrrhomomentumepsilon必须为 0-D 或 1 元素 1-D 的标量 Tensor;
  • 调用方需保证epsilon > 0ms - mg^2 + epsilon > 0;当denom == 0rsqrt输出 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()会做三层防御性校验:

  1. dtype 校验:只接受ge::DT_FLOATge::DT_FLOAT16,其余类型直接返回GRAPH_FAILED
  2. 主 Tensor 形状一致性:逐一校验mg/ms/mom/grad(输入索引 1/2/3/8)的numel必须等于varnumel,否则报错并给出具体张量名与元素数;
  3. 标量形状自防御:逐一校验lr/rho/momentum/epsilon(输入索引 4/5/6/7)的numel必须为 1(0-D 视为 1),空标量会被拒绝——因为 Kernel 的LoadScalar会读取element[0],空标量将导致越界访问。

调用说明

调用方式调用样例说明
aclnn 调用test_aclnn_apply_centered_rms_propAscend 950 上通过 aclnn 两段式接口aclnnApplyCenteredRMSPropGetWorkspaceSizeaclnnApplyCenteredRMSProp调用。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 比对),其完整流程如下:

  1. ACL 初始化aclInit(nullptr)aclrtSetDevice(0)aclrtCreateStream(&stream)
  2. 构造 Host 数据:初始化var/mg/ms/mom/grad(示例中ms[i] = 0.20f + 0.01f * i保证恒正),并取lr=1e-2frho=0.9fmomentum=0.9fepsilon=1e-6f;同时在 H2D 拷贝前用ComputeGolden()在 CPU 端按公式算出期望结果;
  3. 分配 Device 缓冲并 H2D:5 个主 Tensor 按N * sizeof(float)分配,4 个标量按sizeof(float)分配,逐一aclrtMemcpy到 Device;
  4. 构建 aclTensor:用aclCreateTensor构造 9 个输入 Tensor;标量用空的std::vector<int64_t>构造 0-D 张量;
  5. 关键一步——构造共享存储的占位输出t_var_o = MakeTensor(d_var, shape, ACL_FLOAT)等四个输出 Tensor必须复用与对应 Ref 输入相同的 Device 指针,否则无法观察 inplace 更新结果;
  6. 两段式调用:先aclnnApplyCenteredRMSPropGetWorkspaceSize(t_var, ..., &ws_size, &executor)查询并分配 workspace,再aclnnApplyCenteredRMSProp(ws_ptr, ws_size, executor, stream)执行,最后aclrtSynchronizeStream同步;
  7. D2H 回读比对:把更新后的var/mg/ms/mom拷回 Host,与 golden 按atol=1e-5 + rtol=1e-4 * |expect|的容差逐元素比对,输出Result (var): 16/16 passed与最终的ALL PASS/FAIL
  8. 资源清理:依次aclDestroyTensoraclrtFree(workspace 与各 Device 缓冲)、aclrtDestroyStreamaclrtResetDeviceaclFinalize

底层 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_propACLNNTYPE 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 > 0ms - mg^2 + epsilon > 0;三是在 Host 侧显式构造与 Ref 输入共享 Device 地址的四个占位输出,以获得正确的 inplace 语义。Kernel 内部对 fp16 路径采用"fp32 精度计算 + 负值钳制"的双重数值保护,进一步保证了更新结果在低精度输入下的稳定性。

【免费下载链接】ops-nn本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-nn

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

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

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

立即咨询