CANN ops-nn 算子解读:AdamApplyOneWithDecay 权重衰减 Adam 优化算子的单步实现与 aclnn 调用实战
2026/9/20 22:12:29 网站建设 项目流程
  • 人工智能
  • 算子库
  • 深度学习
  • CANN
  • Ascend

【免费下载链接】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 以三个公式精确定义了算子的计算行为(假设逐元素计算,input0input4mul0_xmul4_xadd2_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)·g
  • output0 = 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、FLOATND
input1输入待进行 adam_apply_one_with_decay 计算的入参,公式中的 input1BFLOAT16、FLOAT16、FLOATND
input2输入待进行 adam_apply_one_with_decay 计算的入参,公式中的 input2BFLOAT16、FLOAT16、FLOATND
input3输入待进行 adam_apply_one_with_decay 计算的入参,公式中的 input3BFLOAT16、FLOAT16、FLOATND
input4输入待进行 adam_apply_one_with_decay 计算的入参,公式中的 input4BFLOAT16、FLOAT16、FLOATND
mul0_x输入待进行 adam_apply_one_with_decay 计算的入参,公式中的 mul0_xBFLOAT16、FLOAT16、FLOATND
mul1_x输入待进行 adam_apply_one_with_decay 计算的入参,公式中的 mul1_xBFLOAT16、FLOAT16、FLOATND
mul2_x输入待进行 adam_apply_one_with_decay 计算的入参,公式中的 mul2_xBFLOAT16、FLOAT16、FLOATND
mul3_x输入待进行 adam_apply_one_with_decay 计算的入参,公式中的 mul3_xBFLOAT16、FLOAT16、FLOATND
mul4_x输入待进行 adam_apply_one_with_decay 计算的入参,公式中的 mul4_xBFLOAT16、FLOAT16、FLOATND
add2_y输入待进行 adam_apply_one_with_decay 计算的入参,公式中的 add2_yBFLOAT16、FLOAT16、FLOATND
output0输出待进行 adam_apply_one_with_decay 计算的出参,公式中的 output0BFLOAT16、FLOAT16、FLOATND
output1输出待进行 adam_apply_one_with_decay 计算的出参,公式中的 output1BFLOAT16、FLOAT16、FLOATND
output2输出待进行 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,并设置了UnknownShapeFormatAutoContiguous()(动态 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同样对输入与输出的维度和各维尺寸做了等价校验。
  • 输出 Shape:3 个输出张量的 Shape 与input0完全一致(逐元素运算,无广播)。
  • 数据类型:仅支持 BFLOAT16、FLOAT16、FLOAT 三种类型,Tiling 侧通过supportedDtype集合再次校验,非法类型直接返回失败。
  • 数据格式:仅支持 ND 格式(非 NCHW/NHWC 等重排格式)。

五、调用说明:aclnn 单算子调用实战

5.1 调用方式概览

README 给出的唯一官方调用方式为aclnn 模式调用,即通过 CANN 的 aclnn(AscendCL Neural Network)单算子接口执行,对应示例为 test_aclnn_adam_apply_one_with_decay.cpp,编译需链接 CANN 的aclnnacl运行时。

5.2 调用流程拆解

该示例完整演示了标准 aclnn 单算子调用五步法:

步骤 1:初始化运行环境。调用aclInit初始化 ACL,aclrtSetDevice指定设备(示例使用deviceId = 0),aclrtCreateStream创建任务流:

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_FLOATaclFormat::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)、aclrtDestroyStreamaclrtResetDeviceaclFinalize完成资源回收。

从 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 = 14UBDataNumberFp16 = 18UBDataNumberBFp16 = 22,据此推导每次搬运的tileDataNum
  • 多核负载均衡:把总数据量按 AIV 核数均分,余量部分由前tailBlockNum个核承担(bigCoreDataNum),其余核承担smallCoreDataNum,最终通过context->SetBlockDim(finalCoreNum)设置实际核数;切分结果(大/小核数据量、tile 数、尾块数等 8 个字段)写入 adam_apply_one_with_decay_tiling_data.h 定义的AdamApplyOneWithDecayTilingData结构体。
  • Workspace 申请:固定申请 16 MB(WS_SYS_SIZE)。
  • TilingKey 选择:模板参数schMode支持ELEMENTWISE_TPL_SCH_MODE_0 / _1两种调度模式(见 adam_apply_one_with_decay_tiling_key.h),当前 Tiling 函数实际下发模式 0(GET_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::AdamApplyOneWithDecay<T>模板类,按经典CopyIn → Compute → CopyOut三段流水(Process)处理每个 tile,末 tile 使用tailDataNum处理尾块。

Compute 阶段按数据类型走两条实现路径:

  • BF16 路径:由于 BF16 精度有限,先Cast到 4 个 float 临时缓冲(tmp0tmp3),依次完成output1 = input2·mul0_x + input0·mul1_xoutput0 = input0²·mul3_x + input1·mul2_xoutput2 = 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 搬运。此外内核支持动态 Shape(DTYPE_INPUT0模板参数 + 运行期 tiling 数据),并声明AscendC::GlobalTensor<T>按核偏移(globalBufferIndex)定位各核数据段。

七、测试与验证

仓库为算子提供了 host 侧与 kernel 侧两层单元测试,均可作为验证与二次开发的参考基准:

  • Tiling 单元测试:test_adam_apply_one_with_decay_tiling.cpp 使用TilingContextFaker构造{1, 2, 8, 16}四维输入(11 输入 + 3 输出,FLOAT/ND),模拟UB_SIZE = 196608CORE_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 = 128bigCoreDataNum = 136tileDataNum = 8176等),以AIV_MODE运行内核并完成资源释放,验证内核在给定 tiling 下的可执行性与正确性框架。

测试工程通过 CMakeLists.txt 在ENABLE_TESTBENCHMARK开启时纳入构建,读者可按仓库根目录的构建指引在使能测试的配置下运行这两组用例。

八、贡献说明

README 记录了算子的开源贡献信息:

贡献者贡献方贡献算子贡献时间贡献内容
CyndiZ个人开发者AdamApplyOneWithDecay2026/04/27AdamApplyOneWithDecay 算子适配开源仓

该信息同时表明本算子属于 ops-nn 仓库中由社区贡献、已完成适配合入的优化器算子,其目录结构(examples/op_host/op_kernel/tests)遵循仓库统一的算子工程组织规范,可对照 CONTRIBUTING.md 了解算子移植与合入的整体流程。

  • 人工智能
  • 算子库
  • 深度学习
  • CANN
  • Ascend

【免费下载链接】ops-nn

本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。

项目地址:https://gitcode.com/cann/ops-nn
点击查看免费下载

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

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

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

立即咨询