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