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

LambNextRight

产品支持情况

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

功能说明

  • 算子功能:BERT LAMB优化器图融合算子:计算二阶矩next_v及其偏差校正后的分母sqrt(v_unbiased)+epsilon。

  • 计算公式:

    y1=input_mul2×mul2_x+input_square2×mul3_x(next_v)y1 = input\_mul2 \times mul2\_x + input\_square^2 \times mul3\_x\quad(next\_v)

    y2=y1×truediv1_recip+add2_yy2 = \sqrt{y1 \times truediv1\_recip} + add2\_y

参数说明

参数名 输入/输出 描述 数据类型 数据格式
input_square 输入 支持空Tensor。公式中的input_square(梯度g),主张量,shape需与input_mul2满足broadcast关系。 FLOAT16、FLOAT ND
input_mul2 输入 支持空Tensor。公式中的input_mul2(二阶矩v),主张量,shape需与input_square满足broadcast关系,其broadcast结果决定各输出的shape。 FLOAT16、FLOAT ND
mul2_x 输入 不支持空Tensor。公式中的mul2_x(beta2),标量。 FLOAT16、FLOAT ND
mul3_x 输入 不支持空Tensor。公式中的mul3_x(1-beta2),标量。 FLOAT16、FLOAT ND
truediv1_recip 输入 不支持空Tensor。公式中的truediv1_recip(偏差校正分母的倒数),标量。 FLOAT16、FLOAT ND
add2_y 输入 不支持空Tensor。公式中的add2_y(epsilon),标量。 FLOAT16、FLOAT ND
y1 输出 支持空Tensor。公式中的y1(next_v),shape取input_square与input_mul2的broadcast结果。 FLOAT16、FLOAT ND
y2 输出 支持空Tensor。公式中的y2(偏差校正分母),shape取input_square与input_mul2的broadcast结果。 FLOAT16、FLOAT ND

约束说明

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

调用说明

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