CANN ops-nn 算子实战:AdamApplyOneWithDecayAssign 的 aclnn 调用与 NPU 内核实现解析
CANN ops-nn 算子实战AdamApplyOneWithDecayAssign 的 aclnn 调用与 NPU 内核实现解析【免费下载链接】ops-nn本项目是CANN提供的神经网络类计算算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-nn本篇技术指南以 CANN 神经网络算子库 ops-nn 中 experimental/optim 目录下的 AdamApplyOneWithDecayAssign 算子文档为主线系统讲解该算子的功能语义、计算公式、全部 14 个输入输出参数、支持的硬件产品与数据类型并结合仓库内 op_host、op_kernel 与 tests 源码深入剖析其 Tiling 策略、昇腾向量指令内核实现与 aclnn 调用样例。读完本文你将掌握如何在 Atlas A2/A3 训练与推理系列产品上通过 aclnn 接口完成带权重衰减的 Adam 优化单步更新并能看懂该算子的 shape 校验、工作区分配与多核数据切分逻辑。一、算子定位带权重衰减的 Adam 单步更新AdamApplyOneWithDecayAssign 是 CANN ops-nn 仓库中位于 experimental/optim 目录下的一类优化器算子。它的功能是对模型中的一个参数例如权重张量完成 Adam 优化算法的单步计算与更新并且把更新后的结果就地写回Assign 语义因此常用于训练循环中每步迭代的参数刷新无需额外的赋值算子。与通用 Adam 算子不同AdamApplyOneWithDecayAssign 把一整套 Adam 更新流程“压平”为一条融合算子一次内核调用同时完成动量一阶矩、二阶矩的更新以及参数修正避免了多次读写全局内存GM的往返从而减少 NPU 上的调度与访存开销。该算子在仓库中的完整实现位于算子定义OpDefop_host/adam_apply_one_with_decay_assign_def.cpp内核实现op_kernel/adam_apply_one_with_decay_assign.cpp 与 op_kernel/adam_apply_one_with_decay_assign.hTiling 计算op_host/adam_apply_one_with_decay_assign_tiling.cppShape 推导op_host/adam_apply_one_with_decay_assign_infershape.cpp该算子的贡献信息记录在 README.md 的贡献说明小节中由个人开发者于 2026 年 4 月适配开源仓。二、产品支持情况原文档明确了该算子的硬件支持范围如下表所示产品是否支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品√Atlas A3 训练系列产品 / Atlas A3 推理系列产品√在源码层面算子定义文件 中通过this-AICore().AddConfig(ascend910b);声明了该算子的 AICore 配置表明内核面向昇腾 910B 系列 AICore 架构编译而 Tiling 侧通过GetCoreNumAiv()获取可用 AIV 核数并在GetCoreMemSize(CoreMemType::UB, ubSize)中读取统一缓冲区UB容量用于后续的多核切分具体可参见 tiling 源码。三、功能说明与计算公式对模型中的单个参数完成 Adam 优化算法的单步计算和更新三个输出对应的计算公式如下$$output_0 input_0^2 \times mul_3_x input_1 \times mul_2_x$$$$output_1 input_2 \times mul_0_x input_0 \times mul_1_x$$$$output_2 input_3 - (\frac{output_1}{\sqrt{output_0} add_2_y} input_3 \times mul_4_x) \times input_4$$结合标准 Adam 更新式可以从公式语义上作如下对应属于对公式结构的解读供理解参考output0是二阶矩估计含 bias correction 相关缩放项先对input0平方再与mul3_x相乘并加上input1 × mul2_xoutput1是一阶矩估计input2 × mul0_x input0 × mul1_xoutput2是最终参数更新值input3减去“学习率 ×一阶矩/sqrt(二阶矩) epsilon 权重衰减项”的乘积其中input4承担学习率缩放、mul4_x承担权重衰减系数、add2_y承担防除零的 epsilon 小量。需要特别说明的是上述映射关系是根据公式形态与参数命名mul*_x、add2_y推断的语义解释仓库文档与源码并未给出各参数的数学定义在工程使用中应以上层框架如 MindSpore 或 PyTorch 的 AdamW 类优化器实际传入的张量为准。内核中的计算顺序与公式完全一致可从 内核头文件 的Compute函数中得到印证。四、参数说明14 个张量参数算子共接收 11 个输入张量与 3 个输出张量全部要求 ND 数据格式。下表完整列出文档中的参数定义参数名输入/输出/属性描述数据类型数据格式input0输入待进行 adam_apply_one_with_decay_assign 计算的入参公式中的 input0BFLOAT16、FLOAT16、FLOATNDinput1输入待进行 adam_apply_one_with_decay_assign 计算的入参公式中的 input1BFLOAT16、FLOAT16、FLOATNDinput2输入待进行 adam_apply_one_with_decay_assign 计算的入参公式中的 input2BFLOAT16、FLOAT16、FLOATNDinput3输入待进行 adam_apply_one_with_decay_assign 计算的入参公式中的 input3BFLOAT16、FLOAT16、FLOATNDinput4输入待进行 adam_apply_one_with_decay_assign 计算的入参公式中的 input4BFLOAT16、FLOAT16、FLOATNDmul0_x输入待进行 adam_apply_one_with_decay_assign 计算的入参公式中的 mul0_xBFLOAT16、FLOAT16、FLOATNDmul1_x输入待进行 adam_apply_one_with_decay_assign 计算的入参公式中的 mul1_xBFLOAT16、FLOAT16、FLOATNDmul2_x输入待进行 adam_apply_one_with_decay_assign 计算的入参公式中的 mul2_xBFLOAT16、FLOAT16、FLOATNDmul3_x输入待进行 adam_apply_one_with_decay_assign 计算的入参公式中的 mul3_xBFLOAT16、FLOAT16、FLOATNDmul4_x输入待进行 adam_apply_one_with_decay_assign 计算的入参公式中的 mul4_xBFLOAT16、FLOAT16、FLOATNDadd2_y输入待进行 adam_apply_one_with_decay_assign 计算的入参公式中的 add2_yBFLOAT16、FLOAT16、FLOATNDoutput0输出待进行 adam_apply_one_with_decay_assign 计算的出参公式中的 output0BFLOAT16、FLOAT16、FLOATNDoutput1输出待进行 adam_apply_one_with_decay_assign 计算的出参公式中的 output1BFLOAT16、FLOAT16、FLOATNDoutput2输出待进行 adam_apply_one_with_decay_assign 计算的出参公式中的 output2BFLOAT16、FLOAT16、FLOATND上述参数约束与算子定义文件中的注册信息完全一致adam_apply_one_with_decay_assign_def.cpp 为每个输入输出均声明了REQUIRED参数类型、{ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT}数据类型、FORMAT_ND格式与AutoContiguous()属性。Tiling 侧的 GetShapeAttrsInfo 也通过supportedDtype集合再次校验了这三种数据类型。约束说明原文档明确该算子“无”额外约束即不限制 shape 的维度数量与各维度大小。但从源码可以推断使用时的实际约束体现在以下两点所有输入 shape 必须一致Shape 推导 InferShape4AdamApplyOneWithDecayAssign 会逐一校验 11 个输入的 shape 是否与 input0 完全相同Tiling 侧 checkShape 也做了同样的校验输出 shape 等于输入 shape三个输出张量的 shape 直接拷贝自 input0即输出与输入逐元素一一对应逐元素运算语义。五、aclnn 调用示例原文档给出的调用方式是 aclnn 调用调用样例为 examples/test_aclnn_adam_apply_one_with_decay_assign.cpp。该样例是理解本算子使用方式的最佳入口其调用流程遵循 CANN aclnn 接口的标准五步范式5.1 环境初始化auto ret aclInit(nullptr); ret aclrtSetDevice(deviceId); ret aclrtCreateStream(stream);即先aclInit初始化 ACL 运行环境再aclrtSetDevice指定设备样例中为 deviceId0最后aclrtCreateStream创建计算流。5.2 构造 14 个 aclTensor样例使用模板函数CreateAclTensor完成“申请显存 → 拷贝数据 → 创建张量”三步auto size GetShapeSize(shape) * sizeof(T); aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST); // 1. 申请 device 内存 aclrtMemcpy(*deviceAddr, size, hostData.data(), size, ACL_MEMCPY_HOST_TO_DEVICE); // 2. 拷入数据 // 3. 依据 shape 计算 strides 后创建张量 *tensor aclCreateTensor(shape.data(), shape.size(), dataType, strides.data(), 0, aclFormat::ACL_FORMAT_ND, shape.data(), shape.size(), *deviceAddr);样例采用的 shape 为{64, 8}两维、共 512 个元素11 个输入与 3 个输出的初始数据均为 0~7 的浮点序列数据类型使用aclDataType::ACL_FLOAT与算子支持的 FLOAT 对齐。5.3 获取 workspace 大小并申请aclnn 接口第一步是查询工作空间uint64_t workspaceSize 0; aclOpExecutor* executor; ret aclnnAdamApplyOneWithDecayAssignGetWorkspaceSize( input0, input1, input2, input3, input4, mul0_x, mul1_x, mul2_x, mul3_x, mul4_x, add2_y, output0, output1, output2, workspaceSize, executor);随后按需申请 workspacevoid* workspaceAddr nullptr; if (workspaceSize 0) { aclrtMalloc(workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); }5.4 执行算子并同步ret aclnnAdamApplyOneWithDecayAssign(workspaceAddr, workspaceSize, executor, stream); ret aclrtSynchronizeStream(stream); // 等待内核执行完成5.5 结果回拷与资源释放样例中的PrintOutResult通过aclrtMemcpy(..., ACL_MEMCPY_DEVICE_TO_HOST)把 device 结果拷回 host 并打印前 10 个元素随后依次aclDestroyTensor销毁张量、aclrtFree释放显存、aclrtDestroyStream/aclrtResetDevice/aclFinalize完成收尾。关于 workspace 的底层来源可以在 Tiling 源码 中看到GetWorkspaceSize固定申请WS_SYS_SIZE 16U * 1024U * 1024U16 MiB的 workspace供内核运行时使用。六、内核实现解析逐元素向量流水线内核部分是本算子性能的关键。入口函数在 adam_apply_one_with_decay_assign.cpp 中定义以schMode作为模板参数编译出调度模式运行时通过REGISTER_TILING_DEFAULT与GET_TILING_DATA_WITH_STRUCT获取 Tiling 数据并实例化NsAdamApplyOneWithDecayAssign::AdamApplyOneWithDecayAssignDTYPE_INPUT0算子类随后执行op.Init(...)与op.Process()。6.1 类结构与数据流类定义位于 adam_apply_one_with_decay_assign.h核心结构包括14 条队列11 条VECIN输入队列input0~input4、mul0_x~mul4_x、add2_y与 3 条VECOUT输出队列BUFFER_NUM 1单缓冲4 块VECCALC临时缓冲区tmp0~tmp3用于 BF16/FP16 精度下的中间计算14 个GlobalTensorT分别指向 GM 中的各输入输出。数据流遵循经典的 CopyIn → Compute → CopyOut 三级流水void Process() { int32_t loopCount this-tileNum; this-processDataNum this-tileDataNum; for (int32_t i 0; i loopCount - 1; i) { CopyIn(i); // GM - UBDataCopy 搬入 Compute(i); // UB 上完成 Adam 数学运算 CopyOut(i); // UB - GMDataCopy 搬出 } // 最后一趟以尾块数据量处理 this-processDataNum this-tailDataNum; CopyIn(loopCount - 1); Compute(loopCount - 1); CopyOut(loopCount - 1); }6.2 计算公式在内核中的实现Compute 函数 针对不同数据类型做了分支floatFLOAT路径直接使用向量指令链完成计算例如output0 input0² × mul3_x input1 × mul2_x对应Mul → Mul → Mul → Add的组合output2部分则依次执行Sqrt对 output0 开方、Add加 add2_y、Div一阶矩除以分母、Mul乘 mul4_x 与 input4与Subinput3 减更新量最终写入 output2halfFLOAT16路径开方前先将 UB 中的半精度数据Cast到 float 临时缓冲tmp0计算再将结果Cast回半精度CAST_RINT舍入以提高中间精度bfloat16BFLOAT16路径四个中间量全部提升到 float 精度计算最后一次Cast写回充分保证 BF16 低精度下的数值稳定性。6.3 多核切分与 TilingCalcTilingData 根据 UB 大小、block 大小、数据类型与总元素数计算出如下切分结果写入 AdamApplyOneWithDecayAssignTilingData 结构体Tiling 字段含义finalCoreNum实际启用的核数不超过数据所需 block 数smallCoreDataNum / bigCoreDataNum分到较少/较多数据的核各自处理的元素数tailBlockNum需要多处理一个 block 的核数量tileDataNum单次流水处理的基本 tile 元素数smallTailDataNum / bigTailDataNum两类核最后一趟尾块的元素数finalSmallTileNum / finalBigTileNum两类核的总循环趟数切分的核心思想是“按 block 均摊”先把总元素按blockSize对齐粒度折算成 block 总数平均分给 AIV 核余数由前tailBlockNum个核各多承担一个 block使各核负载尽量均衡。内核 Init 中根据GetBlockIdx()与tailBlockNum的大小关系区分“大核/小核”并据此计算各自的全局缓冲区偏移globalBufferIndex同时按数据类型初始化临时缓冲区BF16 需要 4 个 float 临时块FP16 需要 1 个。七、UT 测试验证仓库为算子提供了完整的单元测试用于验证内核与 Tiling 的正确性内核级测试tests/ut/op_kernel/test_adam_apply_one_with_decay_assign.cpp 基于 gtest 编写直接包含内核头文件并使用AscendC::GmAlloc分配 GM 内存通过填充AdamApplyOneWithDecayAssignTilingData例如tileDataNum 8176、bigCoreDataNum 136等驱动内核执行验证不同数据量下的计算结果Tiling 级测试tests/ut/op_host/test_adam_apply_one_with_decay_assign_tiling.cpp 用于验证 shape 校验、核数计算与 tiling 数据生成的正确性。结合测试用例与样例中的 shape{64, 8}可以看出该算子对 shape 的维度与大小没有固定要求只要所有输入 shape 一致即可具备较强的动态 shape 适配能力。八、总结AdamApplyOneWithDecayAssign 是 ops-nn 仓库中面向 Atlas A2/A3 系列产品的融合型 Adam 更新算子它通过一次内核调用完成一阶矩、二阶矩的更新与带权重衰减的参数修正支持 BFLOAT16、FLOAT16、FLOAT 三种数据类型与 ND 格式。从源码结构看其实现包含了完整的“算子定义 → Shape 推导 → Tiling 计算 → AICore 内核 → aclnn 接口 → UT 测试”全链路可以作为在 CANN 上编写融合优化器算子或学习昇腾向量内核编程的参考样例。若需在真实环境中运行可按照 examples 目录 中的样例代码在安装好 CANN 与 ops-nn 环境的 Atlas 训练/推理设备上编译执行即可。【免费下载链接】ops-nn本项目是CANN提供的神经网络类计算算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-nn创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考