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

ApplyAddSign

产品支持情况

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

功能说明

  • 算子功能:执行AddSign优化器的单步参数更新。根据当前梯度、动量、学习率和符号衰减系数,计算参数更新量并原地更新权重参数(inplace 语义)。对标TensorFlow中tf.raw_ops.ApplyAddSign 接口的计算语义。

  • 计算公式

    mt=β⋅mt−1+(1−β)⋅gradupdate=(α+sign_decay⋅sign(grad)⋅sign(mt))⋅gradvart=vart−1−lr⋅update\begin{aligned} m_{t} &= \beta \cdot m_{t-1} + (1 - \beta) \cdot grad \\ update &= (\alpha + sign\_decay \cdot \text{sign}(grad) \cdot \text{sign}(m_{t})) \cdot grad \\ var_{t} &= var_{t-1} - lr \cdot update \end{aligned}

    其中beta为动量衰减系数(取值范围 [0, 1]),alpha为更新缩放系数,sign_decay为符号衰减系数,lr为学习率,sign(\cdot)为符号函数(NaN保留NaN)。

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
var 输入 / 输出 (inplace) 待更新的权重参数,对应公式中的var。Kernel内inplace 更新,GE IR单输出视图与输入var共享Device内存。 FLOAT16、FLOAT、BFLOAT16 ND
m 输入 (inplace 更新) 梯度一阶矩动量,对应公式中的m。shape/dtype 必须与var一致;Kernel内显式写回输入GM地址。 FLOAT16、FLOAT、BFLOAT16 ND
lr 输入 学习率,对应公式中的lr。shape={1} 的 1 元素scalar Tensor,dtype必须与var一致。 FLOAT16、FLOAT、BFLOAT16 ND
alpha 输入 更新缩放系数,对应公式中的alpha。shape={1} 的1元素 scalar Tensor。 FLOAT16、FLOAT、BFLOAT16 ND
sign_decay 输入 符号衰减系数,对应公式中的sign_decay。shape={1} 的 1 元素scalar Tensor。 FLOAT16、FLOAT、BFLOAT16 ND
beta 输入 动量衰减系数,对应公式中的beta。shape={1} 的 1 元素scalar Tensor,取值范围 [0, 1]。 FLOAT16、FLOAT、BFLOAT16 ND
grad 输入 当前梯度Tensor,对应公式中的grad。shape/dtype 必须与var一致。 FLOAT16、FLOAT、BFLOAT16 ND
use_locking 属性 是否在更新时加锁。默认false。当前实现不强制互斥锁,仅作语义占位。 BOOL -
var (output) 输出 更新后的var Tensor,与输入var共享Device内存(inplace)。 FLOAT16、FLOAT、BFLOAT16 ND

约束说明

调用说明

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