资讯详情

CANN ops-nn 算子解读:AdamApplyOneWithDecay 权重衰减 Adam 优化算子的单步实现与 aclnn 调用实战

📅 2026/9/20 22:12:29 | 华诺云谱 👁 阅读
CANN ops-nn 算子解读:AdamApplyOneWithDecay 权重衰减 Adam 优化算子的单步实现与 aclnn 调用实战
人工智能算子库深度学习CANNAscend【免费下载链接】ops-nn本项目是CANN提供的神经网络类计算算子库实现网络在NPU上加速计算。项目地址https://gitcode.com/cann/ops-nn点击查看免费下载AdamApplyOneWithDecay 是 CANN ops-nn 神经网络算子库experimental/optim 目录中面向 NPU 训练场景的优化器类算子它把带权重衰减的 Adam 单步更新一阶矩、二阶矩与权重的就地更新融合为一次 AICore 内核执行。本文以算子目录下的 README.md 为骨架结合算子定义、Tiling、内核源码与单元测试系统讲解该算子的计算公式、全部入参出参语义、aclnn 单算子调用流程与底层实现原理帮助开发者快速完成接入、调试与二次开发。一、算子概述与产品支持情况AdamApplyOneWithDecay 的算子功能是对模型中的一个参数如权重完成 Adam 优化算法的单步计算与更新。它属于“单参数更新型”优化算子——一次调用只处理一个参数的完整 Adam 更新过程且一阶矩、二阶矩和权重三个张量作为独立输入、独立输出计算过程在算子内部一次性完成。根据 README.md 的产品支持矩阵该算子当前的产品支持情况如下产品是否支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品√在源码层面产品支持范围与算子定义文件中的 AICore 配置严格对应adam_apply_one_with_decay_def.cpp 中通过this-AICore().AddConfig(ascend910b)注册了内核运行平台ascend910b即对应 Atlas A2 系列产品的 AI Core 架构。二、计算公式与 Adam 语义映射2.1 官方计算公式README 以三个公式精确定义了算子的计算行为假设逐元素计算input0input4、mul0_xmul4_x、add2_y为同 Shape 张量output0 input0 × input0 × mul3_x input1 × mul2_x output1 input2 × mul0_x input0 × mul1_x output2 input3 - (output1 / (sqrt(output0) add2_y) input3 × mul4_x) × input4注意第三个公式中出现了mul4_x而 README 的参数表格只列出了mul0_xmul3_x。实际上算子的完整输入为11 个input0input4共 5 个 mul0_xmul4_x共 5 个 add2_y共 1 个mul4_x是参与权重衰减乘法的系数其存在可通过 adam_apply_one_with_decay_def.cpp 中显式注册的Input(mul4_x)确认。下文参数表已将其补充完整。2.2 与标准 Adam 优化器的对应关系README 未显式声明各入参的物理语义但对照标准 Adam 优化算法含解耦权重衰减的 AdamW 形式可以推断出如下合理对应关系属于由公式推导出的语义映射非文档明文公式片段推断语义对应 Adam 量input0当前梯度g参与一阶矩、二阶矩更新input1上一时刻二阶矩v_{t-1}动量累积input2上一时刻一阶矩m_{t-1}动量累积input3当前权重w_{t-1}待更新参数input4学习率lr步长mul0_x/mul1_x一阶矩系数beta1/1 - beta1一阶矩衰减mul2_x/mul3_x二阶矩系数beta2/1 - beta2二阶矩衰减mul4_x权重衰减系数weight_decay解耦权重衰减add2_y数值稳定项eps分母保护项对照关系一目了然output1 input2 × mul0_x input0 × mul1_x即一阶矩更新m_t beta1·m_{t-1} (1 - beta1)·goutput0 input0² × mul3_x input1 × mul2_x即二阶矩更新v_t (1 - beta2)·g² beta2·v_{t-1}output2 input3 - (m_t / (sqrt(v_t) eps) wd·w_{t-1}) × lr即带权重衰减的权重更新w_t w_{t-1} - lr·(m_t/(sqrt(v_t)eps) wd·w_{t-1})。由于公式中一阶矩、二阶矩均未除以偏差校正项1 - beta^t该算子实现的是无偏差校正的 Adam-with-decay 单步变体融合了权重衰减项正好对应其名称中的 WithDecay。三、参数说明以下参数表完整覆盖算子的 11 个输入与 3 个输出在 README 表格基础上补入了公式使用、算子定义中存在的mul4_x。所有输入输出均为REQUIRED 必选参数数据类型与数据格式均一致参数名输入/输出描述数据类型数据格式input0输入待进行 adam_apply_one_with_decay 计算的入参公式中的 input0BFLOAT16、FLOAT16、FLOATNDinput1输入待进行 adam_apply_one_with_decay 计算的入参公式中的 input1BFLOAT16、FLOAT16、FLOATNDinput2输入待进行 adam_apply_one_with_decay 计算的入参公式中的 input2BFLOAT16、FLOAT16、FLOATNDinput3输入待进行 adam_apply_one_with_decay 计算的入参公式中的 input3BFLOAT16、FLOAT16、FLOATNDinput4输入待进行 adam_apply_one_with_decay 计算的入参公式中的 input4BFLOAT16、FLOAT16、FLOATNDmul0_x输入待进行 adam_apply_one_with_decay 计算的入参公式中的 mul0_xBFLOAT16、FLOAT16、FLOATNDmul1_x输入待进行 adam_apply_one_with_decay 计算的入参公式中的 mul1_xBFLOAT16、FLOAT16、FLOATNDmul2_x输入待进行 adam_apply_one_with_decay 计算的入参公式中的 mul2_xBFLOAT16、FLOAT16、FLOATNDmul3_x输入待进行 adam_apply_one_with_decay 计算的入参公式中的 mul3_xBFLOAT16、FLOAT16、FLOATNDmul4_x输入待进行 adam_apply_one_with_decay 计算的入参公式中的 mul4_xBFLOAT16、FLOAT16、FLOATNDadd2_y输入待进行 adam_apply_one_with_decay 计算的入参公式中的 add2_yBFLOAT16、FLOAT16、FLOATNDoutput0输出待进行 adam_apply_one_with_decay 计算的出参公式中的 output0BFLOAT16、FLOAT16、FLOATNDoutput1输出待进行 adam_apply_one_with_decay 计算的出参公式中的 output1BFLOAT16、FLOAT16、FLOATNDoutput2输出待进行 adam_apply_one_with_decay 计算的出参公式中的 output2BFLOAT16、FLOAT16、FLOATND上述约束在 adam_apply_one_with_decay_def.cpp 中逐一登记11 个输入、3 个输出全部ParamType(REQUIRED)数据类型集合为{ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT}数据格式为ge::FORMAT_ND并设置了UnknownShapeFormat与AutoContiguous()动态 Shape 下自动保证存储连续。四、约束说明README 中约束说明一节标注为无算子层面无额外约束。但结合源码可以进一步确认如下隐含的使用前提属于实现层面的约束Shape 一致性adam_apply_one_with_decay_infershape.cpp 中CheckInputShapeEqual会逐一校验 11 个输入的 Shape 必须完全相等否则报错shape ... must be equal for AdamApplyOneWithDecay!Tiling 阶段 adam_apply_one_with_decay_tiling.cpp 的checkShape同样对输入与输出的维度和各维尺寸做了等价校验。输出 Shape3 个输出张量的 Shape 与input0完全一致逐元素运算无广播。数据类型仅支持 BFLOAT16、FLOAT16、FLOAT 三种类型Tiling 侧通过supportedDtype集合再次校验非法类型直接返回失败。数据格式仅支持 ND 格式非 NCHW/NHWC 等重排格式。五、调用说明aclnn 单算子调用实战5.1 调用方式概览README 给出的唯一官方调用方式为aclnn 模式调用即通过 CANN 的 aclnnAscendCL Neural Network单算子接口执行对应示例为 test_aclnn_adam_apply_one_with_decay.cpp编译需链接 CANN 的aclnn与acl运行时。5.2 调用流程拆解该示例完整演示了标准 aclnn 单算子调用五步法步骤 1初始化运行环境。调用aclInit初始化 ACLaclrtSetDevice指定设备示例使用deviceId 0aclrtCreateStream创建任务流auto ret aclInit(nullptr); CHECK_RET(ret ACL_SUCCESS, ...); ret aclrtSetDevice(deviceId); CHECK_RET(ret ACL_SUCCESS, ...); ret aclrtCreateStream(stream);步骤 2构造输入/输出张量。示例定义了shape {64, 8}通过辅助函数CreateAclTensor完成Host 数据 →aclrtMalloc设备内存 →aclrtMemcpy拷入 →aclCreateTensor创建aclTensor的全过程。张量使用aclDataType::ACL_FLOAT、aclFormat::ACL_FORMAT_ND并按 Shape 自动推导连续 strides。11 个输入与 3 个输出均需逐个创建示例中 Host 数据以 8 个元素初始化仅用于演示接口调用流程实际业务使用时应按 Shape 元素总数准备完整数据。步骤 3查询 workspace 大小并获取执行器。这是 aclnn 接口的固定两段式结构先调用aclnnAdamApplyOneWithDecayGetWorkspaceSize获取执行器与所需 workspace 大小再按需aclrtMalloc分配 workspace本算子固定申请 16 MB与 Tiling 中WS_SYS_SIZE 16U * 1024U * 1024U一致uint64_t workspaceSize 0; aclOpExecutor* executor; ret aclnnAdamApplyOneWithDecayGetWorkspaceSize( input0, input1, input2, input3, input4, mul0_x, mul1_x, mul2_x, mul3_x, mul4_x, add2_y, output0, output1, output2, workspaceSize, executor); ... if (workspaceSize 0) { ret aclrtMalloc(workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); }步骤 4执行算子并同步。调用aclnnAdamApplyOneWithDecay下发计算随后aclrtSynchronizeStream等待完成ret aclnnAdamApplyOneWithDecay(workspaceAddr, workspaceSize, executor, stream); CHECK_RET(ret ACL_SUCCESS, ...); ret aclrtSynchronizeStream(stream);步骤 5取回结果并释放资源。示例用aclrtMemcpyACL_MEMCPY_DEVICE_TO_HOST将三个输出拷回 Host并打印每个输出的前 10 个元素最后依次aclDestroyTensor销毁张量、aclrtFree释放设备内存含 workspace、aclrtDestroyStream、aclrtResetDevice、aclFinalize完成资源回收。从 op_kernel/adam_apply_one_with_decay.cpp 的内核入口签名可以看出内核实际接收 16 个 GM 地址参数11 输入 3 输出 workspace tiling与 aclnn 接口的 14 个张量参数11 入 3 出一一对应workspace 与 tiling 数据由运行时框架注入开发者无需直接接触。六、源码级实现原理6.1 算子定义Def层adam_apply_one_with_decay_def.cpp 以OpDef派生类完成算子注册OP_ADD(AdamApplyOneWithDecay)集中声明了 11 个输入、3 个输出以及各自的数据类型、格式、动态 Shape 策略和 AICore 平台配置ascend910b。这是算子能被 GE图引擎识别、校验并参与构图的基础。6.2 形状推导InferShape层adam_apply_one_with_decay_infershape.cpp 通过IMPL_OP_INFERSHAPE注册推导逻辑InferShape4AdamApplyOneWithDecay首先校验所有输入 Shape 相等然后将 3 个输出的 Shape 直接赋值为input0的 Shape。这意味着算子天然要求 m、v、w、g 等所有参与量 Shape 完全一致符合优化器逐元素更新的语义。6.3 Tiling 切分策略adam_apply_one_with_decay_tiling.cpp 是算子性能的核心其要点如下平台信息获取通过GetPlatformInfo读取 AIV 核数GetCoreNumAiv与 UB 内存大小GetCoreMemSize两者均为 0 时报错。UB 数据槽位预算由于算子一次处理 14 个张量11 输入 3 输出且 BF16 路径还需额外 4 个 float 临时缓冲、FP16 路径需 1 个 float 临时缓冲Tiling 按数据类型设置了不同的 UB 数据槽位常数UBDataNumberFloat 14、UBDataNumberFp16 18、UBDataNumberBFp16 22据此推导每次搬运的tileDataNum。多核负载均衡把总数据量按 AIV 核数均分余量部分由前tailBlockNum个核承担bigCoreDataNum其余核承担smallCoreDataNum最终通过context-SetBlockDim(finalCoreNum)设置实际核数切分结果大/小核数据量、tile 数、尾块数等 8 个字段写入 adam_apply_one_with_decay_tiling_data.h 定义的AdamApplyOneWithDecayTilingData结构体。Workspace 申请固定申请 16 MBWS_SYS_SIZE。TilingKey 选择模板参数schMode支持ELEMENTWISE_TPL_SCH_MODE_0 / _1两种调度模式见 adam_apply_one_with_decay_tiling_key.h当前 Tiling 函数实际下发模式 0GET_TPL_TILING_KEY(ELEMENTWISE_TPL_SCH_MODE_0)。6.4 AI Core 内核计算内核入口 op_kernel/adam_apply_one_with_decay.cpp 读取 tiling 数据后实例化 adam_apply_one_with_decay.h 中的NsAdamApplyOneWithDecay::AdamApplyOneWithDecayT模板类按经典CopyIn → Compute → CopyOut三段流水Process处理每个 tile末 tile 使用tailDataNum处理尾块。Compute 阶段按数据类型走两条实现路径BF16 路径由于 BF16 精度有限先Cast到 4 个 float 临时缓冲tmp0tmp3依次完成output1 input2·mul0_x input0·mul1_x、output0 input0²·mul3_x input1·mul2_x、output2 input3 - (output1/(sqrt(output0)add2_y) input3·mul4_x)·input4最后统一以CAST_RINT舍入模式Cast回 BF16 写出。FP32 / FP16 路径FP32 直接在原张量上完成 Mul/Add/Sqrt/Div/Sub 运算FP16 路径则先Cast到 1 个 float 临时缓冲完成Sqrt再CAST_RINT转回 FP16避免 FP16 求平方根精度损失。内核使用AscendC::TQue队列与AscendC::TPipe流水BUFFER_NUM 为 1单缓冲通过DataCopy完成 GM ↔ Local 搬运。此外内核支持动态 ShapeDTYPE_INPUT0模板参数 运行期 tiling 数据并声明AscendC::GlobalTensorT按核偏移globalBufferIndex定位各核数据段。七、测试与验证仓库为算子提供了 host 侧与 kernel 侧两层单元测试均可作为验证与二次开发的参考基准Tiling 单元测试test_adam_apply_one_with_decay_tiling.cpp 使用TilingContextFaker构造{1, 2, 8, 16}四维输入11 输入 3 输出FLOAT/ND模拟UB_SIZE 196608、CORE_NUM 48等平台编译信息通过OpImplRegistry取到注册的 Tiling 函数并断言其返回GRAPH_SUCCESS。Kernel 单元测试test_adam_apply_one_with_decay.cpp 基于__CCE_KT_TEST__ICPU 仿真模式分配 11 输入 3 输出 workspace tiling 共 16 块 GM 内存手工填充AdamApplyOneWithDecayTilingData如smallCoreDataNum 128、bigCoreDataNum 136、tileDataNum 8176等以AIV_MODE运行内核并完成资源释放验证内核在给定 tiling 下的可执行性与正确性框架。测试工程通过 CMakeLists.txt 在ENABLE_TEST或BENCHMARK开启时纳入构建读者可按仓库根目录的构建指引在使能测试的配置下运行这两组用例。八、贡献说明README 记录了算子的开源贡献信息贡献者贡献方贡献算子贡献时间贡献内容CyndiZ个人开发者AdamApplyOneWithDecay2026/04/27AdamApplyOneWithDecay 算子适配开源仓该信息同时表明本算子属于 ops-nn 仓库中由社区贡献、已完成适配合入的优化器算子其目录结构examples/op_host/op_kernel/tests遵循仓库统一的算子工程组织规范可对照 CONTRIBUTING.md 了解算子移植与合入的整体流程。赞分享人工智能算子库深度学习CANNAscend【免费下载链接】ops-nn本项目是CANN提供的神经网络类计算算子库实现网络在NPU上加速计算。项目地址https://gitcode.com/cann/ops-nn点击查看免费下载相关推荐CANN ops-nn WeightQuantBatchMatmulExperiment 算子 aclnn 单算子调用样例深度解析CANN ops nn WeightQuantBatchMatmulExperiment 算子 aclnn 单算子调用样例深度解析 本指南以 experimen人工智能算子库深度学习CANNAscendCANN ops-nn 中 MatmulFp32 算子 aclnn 单算子调用样例全解析CANN ops nn 中 MatmulFp32 算子 aclnn 单算子调用样例全解析 导读 本文基于 CANN ops nn 开源仓库中 experimen人工智能算子库深度学习CANNAscendCANN ops-nn 中 GaussianNllLossGrad 算子的梯度计算原理与 ACLNN 调用实战CANN ops nn 中 GaussianNllLossGrad 算子的梯度计算原理与 ACLNN 调用实战 GaussianNllLossGrad 是 CA人工智能算子库深度学习CANNAscend创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
📝

华诺云谱内容团队

资深建站顾问 · 行业研究员

10年+企业数字化服务经验,专注智能建站、SEO优化与品牌营销,持续输出建站技巧、行业洞察与营销干货,已帮助5000+企业实现数字化增长。

你可能需要的服务

订阅华诺云谱资讯周报

每周一封,精选建站技巧、SEO与营销干货,直达邮箱。已有 8,000+ 企业主订阅,助你少走弯路。