已关闭
【社区任务】ApplyAdamW算子设计文档 #4058
StarLightOn创建于 4月21日关闭于 5月15日
【社区任务】ApplyAdamW算子设计文档 #4058
已关闭
StarLightOn创建于 4月21日关闭于 5月15日
StarLightOn
4月21日

一、 需求背景

1.1 需求来源

通过社区任务完成开源仓算子贡献的需求,补充完善 Ascend C 算子库。

1.2 背景介绍

基于 ApplyAdamW 算子历史 TBE 版本使用 Ascend C 编程语言进行优化并替换。

1.2.1 TBE 算子现状分析

  • 算子信息库与源码路径
    TBE 源码参考文件为 ${INSTALL_PATH}/opp/built-in/op_impl/ai_core/tbe/impl/ops_legacy/dynamic/apply_adam_w.py
    算子原型库:${INSTALL_PATH}/opp/built-in/op_graph/inc/ops_proto_nn.h
  • 支持的数据类型和数据格式
    根据 TBE 源码分析,当前算子支持的数据类型为:float16float32bfloat16
    支持的数据格式为:ND(无需感知具体维度,当作一维向量处理)。
    部分标量输入(如 beta1_power, beta2_power, lr, weight_decay, beta1, beta2, epsilon)的 shape 为 [1]

1.2.2 TBE 算子实现描述

通过对 ApplyAdamW 算子 TBE 版本代码分析,核心实现逻辑与计算公式如下:

  1. 梯度预处理 (gt)
    若属性 maximize 为 True,则 gt=gradgt = -grad;否则 gt=gradgt = grad
  2. 动量更新 (m_out, v_out)
    一阶动量:m_out=m×β1(β11)×gtm\_out = m \times \beta_1 - (\beta_1 - 1) \times gt
    二阶动量:v_out=v×β2(β21)×gt×gtv\_out = v \times \beta_2 - (\beta_2 - 1) \times gt \times gt
  3. 权重衰减处理 (var_t)
    var_t=var×(1lr×weight_decay)var\_t = var \times (1 - lr \times weight\_decay)
  4. 偏差校正因子计算
    β1_power_out=β1_power×β1\beta_1\_power\_out = \beta_1\_power \times \beta_1
    β2_power_out=β2_power×β2\beta_2\_power\_out = \beta_2\_power \times \beta_2
  5. 分母计算 (denom)
    若属性 amsgrad 为 True:
    max_grad_norm_out=max(max_grad_norm,v_out)max\_grad\_norm\_out = \max(max\_grad\_norm, v\_out)
    denom=max_grad_norm_out/(β2_power_out1)+ϵdenom = \sqrt{-max\_grad\_norm\_out / (\beta_2\_power\_out - 1)} + \epsilon
    amsgrad 为 False:
    denom=v_out/(β2_power_out1)+ϵdenom = \sqrt{-v\_out / (\beta_2\_power\_out - 1)} + \epsilon
  6. 权重更新 (var_out)
    var_out=var_t+lrβ1_power_out1×m_outdenomvar\_out = var\_t + \frac{lr}{\beta_1\_power\_out - 1} \times \frac{m\_out}{denom}

1.2.3 TBE 算子实现流程图

image.png

二、 需求分析

2.1 外部组件依赖

不涉及。

2.2 内部适配模块

适配 aclnn 接口调用框架。

2.3 需求模块设计

2.3.1 Ascend C 算子原型

名称 类别 数据类型 format shape
var 输入 float16, float32, bfloat16 ND all
m 输入 float16, float32, bfloat16 ND all
v 输入 float16, float32, bfloat16 ND all
beta1_power 输入 float16, float32, bfloat16 ND [1]
beta2_power 输入 float16, float32, bfloat16 ND [1]
lr 输入 float16, float32, bfloat16 ND [1]
weight_decay 输入 float16, float32, bfloat16 ND [1]
beta1 输入 float16, float32, bfloat16 ND [1]
beta2 输入 float16, float32, bfloat16 ND [1]
epsilon 输入 float16, float32, bfloat16 ND [1]
grad 输入 float16, float32, bfloat16 ND all
max_grad_norm 输入(可选) float16, float32, bfloat16 ND all
var_out 输出 float16, float32, bfloat16 ND all
m_out 输出 float16, float32, bfloat16 ND all
v_out 输出 float16, float32, bfloat16 ND all

属性 (Attributes):

  • amsgrad (bool, default=False)
  • maximize (bool, default=False)

2.3.2 Ascend C 算子相关约束

  1. 暂不考虑输入 Tensor 之间的非标量广播(即 var, m, v, grad, max_grad_norm 的 shape 必须一致)。
  2. 标量输入(如 lr, beta1, epsilon 等)在 Host 侧解析时其 shape 必须为 1,在 Kernel 侧计算时作为标量或借助标量广播指令参与向量计算。
  3. 如果属性 amsgrad 为 True,输入 max_grad_norm 必须有效。

三、 需求详细设计

3.1 使能方式

上层框架 涉及的框架勾选
TF训练/推理
Pytorch训练/推理
ATC推理
Aclnn直调
OPAT调优
SGAT子图切分

3.2 需求总体设计

3.2.1 Host侧设计

3.2.1.1 分核策略
采取均分满核策略。在 Host 侧由于所有计算都是 element-wise,可以将输入视为 1D 数组。
根据 var 的元素总数(inputLength)和硬件可用核心数(coreNum)进行切分。对于不能整除的部分,将剩余的数据块(tailBlockNum)优先分配给前面的计算核心(大核),剩余核心为小核,确保每个核心的计算量差异不超过一个数据块,最大化硬件利用率。

3.2.1.2 数据分块和内存优化策略
该算子输入输出较多,需要充分利用 UB(Unified Buffer)空间。

  • 空间复用var, m, v, grad 为主要输入,计算中涉及大量临时变量,可通过 Ascend C 的 TBuf 资源池和对 LocalTensor 的重用来减少总空间占用。针对标量输入(lr, beta1 等),可以直接解析到标量寄存器或开辟极小的 UB 空间存放,节约主流水线内存。
  • 双缓冲流水线:计算 UB 空间可容纳的最大单次处理元素量 tileDataNum,默认开启 DOUBLE_BUFFER 隐藏 MTE2(搬入)和 MTE3(搬出)的内存访问延迟。

3.2.1.3 tilingKey规划策略
为了避免在 Kernel 侧的每个循环内部引入 bool 类型的分支跳转(影响流水线排布),利用 TilingKey 在 Host 侧进行静态分支分发:
根据 amsgrad (0/1) 和 maximize (0/1) 属性,可组合出 4 种 TilingKey(如:NORMAL, MAXIMIZE, AMSGRAD, MAXIMIZE_AMSGRAD)。Kernel 侧通过模板参数或 if constexpr 实例化出独立且最优的指令流分支。

3.2.2 Kernel侧设计

3.2.2.1 kernel侧实现描述
分为 InitProcess 两步:

  • Init:获取 Host 下发的 Tiling 信息,计算出 coreDataNum 及偏移。初始化 TPipe,申请相应的 TQueTBuf(包括辅助标量张量和计算使用的临时掩码/运算中转空间)。
  • Process:以流水线方式遍历分块(Tile)。
    1. CopyIn (MTE2):从 GM 并发异步读取 var, m, v, grad(及可选的 max_grad_norm)到 UB 中。对于标量如 lr 仅需在最外层或 Host 侧读取一次。
    2. Compute (V):依照 1.2.2 节的数学公式,大量调用 Mul, Adds, Sub, Div, Sqrt, Max 等高性能向量计算 API 计算 m_out, v_out, var_out。在 bfloat16 场景下,由于部分 API 精度限制,需先 cast 到 float32 计算再 cast 回 bfloat16
    3. CopyOut (MTE3):将更新后的 var_out, m_out, v_out 异步写回 GM。

3.2.2.2 Ascend C 实现流程图
image.png

3.2.2.3 Ascend C 实现流程图与 TBE 流程图存在的差异点和原因

  1. 静态分支剥离:TBE 在计算图中直接通过 if-else 组合算子图;Ascend C 流程中,maximizeamsgrad 的判断被前置到 Host 侧,通过 TilingKey 控制 Kernel 的模板实例化。原因:在 Ascend C 的 V 单元中做标量级别的 if 分支会打断指令流水线,降低并行执行效率。
  2. 标量与 Tensor 运算处理:TBE 中直接使用 tbe.broadcast 将标量撑成与 Tensor 同等大小后计算;Ascend C 中,出于节省 UB 内存考虑,大量操作会直接使用标量-向量混合的 API(如 Adds, Muls 传入标量操作数),避免冗余的 broadcast 内存消耗。
  3. 流水线编排:Ascend C 显式展现了 MTE2 -> V -> MTE3 的 WaitFlag / SetFlag 数据流同步和 Ping-Pong 双缓冲处理机制,这是底层算子获取极致性能的关键,而 TBE 是由 TVM Auto-schedule 隐式生成的。

3.3 支持硬件

  • Atlas A2 训练系列产品
  • Atlas 800I A2 推理产品

3.4 算子约束限制

  • 输入的 var, m, v, grad 以及 max_grad_norm (当 amsgrad 为 True 时) 数据类型需保持一致。
  • 标量输入 shape 必须为 [1] 或为标量实体。
  • 由于不考虑广播,所有相关 Tensor 输入的 shape(降维至 1D 元素总量后)需严格相同。

四、 特性交叉分析

不涉及

五、 可维可测分析

5.1 精度标准/性能标准

验收标准 描述 来源
精度标准 精度不低于历史 TBE 版本,满足 ACL 算子精度要求 历史 TBE 对标
性能标准 吞吐量与执行耗时不劣于历史 TBE 版本 历史 TBE 对标

5.2 兼容性分析

本需求为新增 Ascend C 算子替换原有 TBE 算子。
对外接口类型、参数列表均对齐原有 TBE 实现规范,向上层提供一致的 aclnn 调用表现,不存在兼容性断裂问题。

likedislike
当前Pull Request已关闭, 关闭人@fulltower
SStarLightOn
4月21日 创建了 pull request,commit c4fc6e2d
CANN-robot
CANN-robot成员
4月21日 评论:

Thanks for your pull-request.
The full list of commands accepted by me can be found at here
You can get sig-info at here


PR Approval Progress

⚠️ This PR does not yet meet the following requirements:lgtm (requires ≥ 2 person(s) per module)、approve (requires ≥ 1 person(s) per module)

Module Approval Details

module lgtm status approve status
experimental/quant ❌ (0/2)(You can also ask: liujie12345678, Chen_HaoWen, 胡碧霞, zhang-wu, wangyongguang) ❌ (0/1)(You can also ask: 王子韬, Chen_HaoWen, zhang-wu, 胡碧霞, 唐玮玮)

💡 Tip:

  • Committer can comment /approve or /lgtm
  • Commenting /approve implies both code review (lgtm) and intent to merge (approve)

CLA Signature Pass

StarLightOn, thanks for your pull request. All authors of the commits have signed the CLA. 👍

likedislike
CANN-robotCANN-robot成员
4月21日 添加了label:cann-cla/yes
CANN-robotCANN-robot成员
4月21日 将sxb154714,chenfeng61,crystalhu,yangyang016,fanqirui,zhou-qilong,zhajianqing123,chenqi317,liubo75,tangweiwei2,wangzitao_leo,liujie12345678,Chen_HaoWen,wangyongguang,zhang-wu设为评审人
CANN-robotCANN-robot成员
4月21日 将sxb154714,chenfeng61,crystalhu,zhou-qilong,zhajianqing123,chenqi317,liubo75,tangweiwei2,wangzitao_leo,liujie12345678,Chen_HaoWen,wangyongguang,zhang-wu设为审查人
SStarLightOn
4月21日 修改了pull request 的描述
SStarLightOn
4月21日 修改了pull request 的描述
SStarLightOn
4月21日 修改了pull request 的描述
SStarLightOn
4月21日 取消了草稿状态
忧莫晓
4月21日 评论:

评审通过

likedislike
fulltower成员
5月15日 评论:

您好,当前pr创建的时间比较久了,我们为了方便管理,先关闭此pr,如果您有需要,可以再自己打开

likedislike
Ffulltower成员
5月15日 关闭了 pull request