文件最后提交记录最后更新时间
17 天前
1 个月前
9 天前
9 天前
17 天前
9 天前
1 个月前
9 天前
README

LambApplyOptimizerAssign

产品支持情况

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

功能说明

  • 算子功能:对模型中的一个参数完成LAMB优化器的Adam矩更新与偏差校正后的update计算(含可选权重衰减),并原地更新一阶矩inputm与二阶矩inputv。

  • 计算公式:

    next_v=inputv×mul2_x+grad2×mul3_xnext\_v = inputv \times mul2\_x + grad^2 \times mul3\_x

    next_m=inputm×mul0_x+grad×mul1_xnext\_m = inputm \times mul0\_x + grad \times mul1\_x

    output0=next_m/(1−mul0_xsteps)next_v/(1−mul2_xsteps)+add2_y+input3×weight_decay_rate×do_use_weightoutput0 = \frac{next\_m / (1 - mul0\_x^{steps})}{\sqrt{next\_v / (1 - mul2\_x^{steps})} + add2\_y} + input3 \times weight\_decay\_rate \times do\_use\_weight

参数说明

参数名 输入/输出 描述 数据类型 数据格式
grad 输入 支持空Tensor。公式中的grad(梯度)。允许小于inputv/inputm并向上广播,但其shape必须能broadcast进inputv、inputm的shape。 FLOAT16、FLOAT ND
inputv 输入 支持空Tensor。公式中的inputv(二阶矩)。inputv为**原地(in-place)更新**输出,其shape必须等于所有输入广播后的完整输出shape:即须与inputm同shape,且grad、input3能broadcast进inputv。 FLOAT16、FLOAT ND
inputm 输入 支持空Tensor。公式中的inputm(一阶矩)。inputm为**原地(in-place)更新**输出,其shape必须等于所有输入广播后的完整输出shape:即须与inputv同shape,且grad、input3能broadcast进inputm。 FLOAT16、FLOAT ND
input3 输入 支持空Tensor。公式中的input3(参与权重衰减的参数),主张量。 FLOAT16、FLOAT ND
mul0_x 输入 不支持空Tensor。公式中的mul0_x(beta1),标量。 FLOAT16、FLOAT ND
mul1_x 输入 不支持空Tensor。公式中的mul1_x(1-beta1),标量。 FLOAT16、FLOAT ND
mul2_x 输入 不支持空Tensor。公式中的mul2_x(beta2),标量。 FLOAT16、FLOAT ND
mul3_x 输入 不支持空Tensor。公式中的mul3_x(1-beta2),标量。 FLOAT16、FLOAT ND
add2_y 输入 不支持空Tensor。公式中的add2_y(epsilon),标量。 FLOAT16、FLOAT ND
steps 输入 不支持空Tensor。公式中的steps(步数),标量。 FLOAT16、FLOAT ND
do_use_weight 输入 不支持空Tensor。公式中的do_use_weight(是否使用权重衰减),标量。 FLOAT16、FLOAT ND
weight_decay_rate 输入 不支持空Tensor。公式中的weight_decay_rate(权重衰减率),标量。 FLOAT16、FLOAT ND
output0 输出 支持空Tensor。公式中的output0(update),shape取grad与inputv的broadcast结果。 FLOAT16、FLOAT ND
inputv 输出 支持空Tensor。更新后的inputv(二阶矩,原地更新),shape取grad与inputv的broadcast结果。 FLOAT16、FLOAT ND
inputm 输出 支持空Tensor。更新后的inputm(一阶矩,原地更新),shape取grad与inputm的broadcast结果。 FLOAT16、FLOAT ND

约束说明

  • 所有输入的数据类型必须一致,同为FLOAT16或同为FLOAT。

调用说明

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