| 文件 | 最后提交记录 | 最后更新时间 |
|---|---|---|
| 17 天前 | ||
| 9 天前 | ||
| 9 天前 | ||
| 17 天前 | ||
| 17 天前 | ||
| 1 个月前 | ||
| 9 天前 |
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算子。 |