文件最后提交记录最后更新时间
26 天前
10 天前
26 天前
26 天前
26 天前
26 天前
2 个月前
26 天前
README

ApplyAdaMax

产品支持情况

产品 是否支持
Ascend 950PR/Ascend 950DT
Atlas A3 训练系列产品/Atlas A3 推理系列产品
Atlas A2 训练系列产品/Atlas A2 推理系列产品
Atlas 200I/500 A2 推理产品 ×
Atlas 推理系列产品
Atlas 训练系列产品

功能说明

  • 算子功能:执行AdaMax优化器的单步参数更新。AdaMax是Adam优化器的变体,使用无穷范数 L∞L_{\infty} 代替二阶矩估计,对权重var、一阶矩m、无穷范数v进行原地更新(inplace语义)。对标TensorFlow tf.raw_ops.ApplyAdaMax接口和PyTorch torch.optim.Adamax的计算语义。

  • 计算公式

    给定时间步 tt 的梯度 gtg_t,衰减系数 β1,β2\beta_1, \beta_2,学习率 lrlr,数值稳定常数 ϵ\epsilon,以及外部传入的偏差校正因子 β1t\beta_1^t

    mt=β1⋅mt−1+(1−β1)⋅gtvt=max⁡(β2⋅vt−1, ∣gt∣)vart=vart−1−lr1−β1t⋅mtvt+ϵ\begin{aligned} m_{t} &= \beta_1 \cdot m_{t-1} + (1 - \beta_1) \cdot g_t \\ v_{t} &= \max(\beta_2 \cdot v_{t-1},\ |g_t|) \\ var_{t} &= var_{t-1} - \frac{lr}{1 - \beta_1^t} \cdot \frac{m_t}{v_t + \epsilon} \end{aligned}

    算子原型对齐canndev REG_OP(ApplyAdaMax):9输入 + 1输出 + 1属性。m / v不显式作为图输出端口,通过输入GM地址inplace写回完成更新(GE Variable inplace语义)。

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
var 输入 待更新的权重参数张量,对应公式中的var。与图输出端口var共享GM地址(inplace更新)。1-8维ND格式。 FLOAT16、FLOAT ND
m 输入 一阶矩张量,对应公式中的m。shape/dtype必须与var一致;由Kernel写回输入GM地址完成inplace更新,无显式输出端口。 FLOAT16、FLOAT ND
v 输入 无穷范数估计张量($L_{\infty}$),对应公式中的v。shape/dtype必须与var一致;由Kernel写回输入GM地址完成inplace更新,无显式输出端口。 FLOAT16、FLOAT ND
beta1_power 输入 偏差校正因子 $\beta_1^t$,由外部传入。shape必须为 [1] 的scalar Tensor。须保证 $1 - \beta_1^t \neq 0$。 FLOAT16、FLOAT ND
lr 输入 学习率,对应公式中的lr。shape必须为 [1] 的scalar Tensor。 FLOAT16、FLOAT ND
beta1 输入 一阶矩衰减系数 $\beta_1$,取值范围 [0, 1)。shape必须为 [1] 的scalar Tensor。 FLOAT16、FLOAT ND
beta2 输入 无穷范数衰减系数 $\beta_2$,取值范围 [0, 1)。shape必须为 [1] 的scalar Tensor。 FLOAT16、FLOAT ND
epsilon 输入 数值稳定常数 $\epsilon$,必须大于0。shape必须为 [1] 的scalar Tensor。 FLOAT16、FLOAT ND
grad 输入 当前梯度 $g_t$,对应公式中的grad。shape/dtype必须与var一致。 FLOAT16、FLOAT ND
use_locking 属性 语义占位的bool属性,默认false。当前实现不强制互斥锁。 BOOL -
var 输出 更新后的权重张量,与输入var共享Device内存(inplace)。shape/dtype与输入var完全相同。 FLOAT16、FLOAT ND

约束说明

调用说明

调用方式 调用样例 说明
图模式 test_geir_apply_adamax 通过算子IR构图方式调用ApplyAdaMax算子。