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

LambNextMVWithDecay

产品支持情况

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

功能说明

  • 算子功能:BERT LAMB优化器图融合算子(含权重衰减):与LambNextMV相同,但y4额外加上权重衰减项input_mul4*mul4_x。

  • 计算公式:

    y3=input_mul2×mul2_x+input_mul3×mul3_sub1(next_v)y3 = input\_mul2 \times mul2\_x + input\_mul3 \times mul3\_sub1\quad(next\_v)

    y2=input_mul0×mul0_x+input_mul1×mul1_sub(next_m)y2 = input\_mul0 \times mul0\_x + input\_mul1 \times mul1\_sub\quad(next\_m)

    y1=input_mul4×mul4_x+y2/input_realdiv0y3/input_realdiv1+add2_yy1 = input\_mul4 \times mul4\_x + \frac{y2/input\_realdiv0}{\sqrt{y3/input\_realdiv1 + add2\_y}}

    y4=input_mul4×mul4_x+y2/input_realdiv0y3/input_realdiv1+add2_yy4 = input\_mul4 \times mul4\_x + \frac{y2/input\_realdiv0}{\sqrt{y3/input\_realdiv1} + add2\_y}

参数说明

参数名 输入/输出 描述 数据类型 数据格式
input_mul3 输入 支持空Tensor。公式中的input_mul3(g^2),主张量,shape需与input_mul0满足broadcast关系。 FLOAT16、FLOAT ND
input_mul2 输入 支持空Tensor。公式中的input_mul2(二阶矩v),主张量。 FLOAT16、FLOAT ND
input_realdiv1 输入 支持空Tensor。公式中的input_realdiv1(1-beta2^t),主张量。 FLOAT16、FLOAT ND
input_mul1 输入 支持空Tensor。公式中的input_mul1(梯度g),主张量。 FLOAT16、FLOAT ND
input_mul0 输入 支持空Tensor。公式中的input_mul0(一阶矩m),主张量,shape需与input_mul3满足broadcast关系,其broadcast结果决定各输出的shape。 FLOAT16、FLOAT ND
input_realdiv0 输入 支持空Tensor。公式中的input_realdiv0(1-beta1^t),主张量。 FLOAT16、FLOAT ND
input_mul4 输入 支持空Tensor。公式中的input_mul4(参数param),主张量。 FLOAT16、FLOAT ND
mul0_x 输入 不支持空Tensor。公式中的mul0_x(beta1),标量。 FLOAT16、FLOAT ND
mul1_sub 输入 不支持空Tensor。公式中的mul1_sub(1-beta1),标量。 FLOAT16、FLOAT ND
mul2_x 输入 不支持空Tensor。公式中的mul2_x(beta2),标量。 FLOAT16、FLOAT ND
mul3_sub1 输入 不支持空Tensor。公式中的mul3_sub1(1-beta2),标量。 FLOAT16、FLOAT ND
mul4_x 输入 不支持空Tensor。公式中的mul4_x(权重衰减系数),标量。 FLOAT16、FLOAT ND
add2_y 输入 不支持空Tensor。公式中的add2_y(epsilon),标量。 FLOAT16、FLOAT ND
y1 输出 支持空Tensor。公式中的y1(update),shape取input_mul3与input_mul0的broadcast结果。 FLOAT16、FLOAT ND
y2 输出 支持空Tensor。公式中的y2(next_m),shape取input_mul3与input_mul0的broadcast结果。 FLOAT16、FLOAT ND
y3 输出 支持空Tensor。公式中的y3(next_v),shape取input_mul3与input_mul0的broadcast结果。 FLOAT16、FLOAT ND
y4 输出 支持空Tensor。公式中的y4,shape取input_mul3与input_mul0的broadcast结果。 FLOAT16、FLOAT ND

约束说明

  • 所有输入的数据类型必须一致,同为FLOAT16或同为FLOAT。
  • input_mul0/input_mul1/input_mul2/input_mul3/input_mul4 为主张量,其shape需保持一致(或可相互广播到同一shape);各输出y1/y2/y3/y4的shape均取该广播结果(实现以input_mul3与input_mul0的broadcast结果为准)。

调用说明

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