| 文件 | 最后提交记录 | 最后更新时间 |
|---|---|---|
| 23 天前 | ||
| 9 天前 | ||
| 23 天前 | ||
| 23 天前 | ||
| 11 天前 | ||
| 23 天前 | ||
| 23 天前 | ||
| 18 天前 |
ApplyCenteredRMSProp
产品支持情况
| 产品 | 是否支持 |
|---|---|
| Ascend 950PR/Ascend 950DT | √ |
| Atlas A3 训练系列产品/Atlas A3 推理系列产品 | √ |
| Atlas A2 训练系列产品/Atlas A2 推理系列产品 | √ |
| Atlas 200I/500 A2 推理产品 | × |
| Atlas 推理系列产品 | √ |
| Atlas 训练系列产品 | √ |
功能说明
-
算子功能:实现带"中心化"修正的RMSProp优化器更新步骤,用于深度学习训练阶段的模型参数更新。在标准RMSProp基础上,额外维护一阶梯度的指数移动平均(mg),以
ms - mg²作为方差估计替代原始二阶矩ms,提供更稳定的自适应学习率。 -
计算公式:
Step 1 — 一阶矩更新:
mgt=ρ⋅mgt−1+(1−ρ)⋅gtmg_t = \rho \cdot mg_{t-1} + (1 - \rho) \cdot g_t
Step 2 — 二阶矩更新:
mst=ρ⋅mst−1+(1−ρ)⋅gt2ms_t = \rho \cdot ms_{t-1} + (1 - \rho) \cdot g_t^2
Step 3 — 方差估计:
variancet=mst−mgt2\text{variance}_t = ms_t - mg_t^2
Step 4 — 分母计算(epsilon加在sqrt内部):
denomt=variancet+ϵ\text{denom}_t = \sqrt{\text{variance}_t + \epsilon}
Step 5 — 动量更新(momentum > 0时):
momt=μ⋅momt−1+lr⋅gtdenomtmom_t = \mu \cdot mom_{t-1} + lr \cdot \frac{g_t}{\text{denom}_t}
vart=vart−1−momtvar_t = var_{t-1} - mom_t
Step 6 — 直接更新(momentum == 0时):
vart=vart−1−lr⋅gtdenomtvar_t = var_{t-1} - lr \cdot \frac{g_t}{\text{denom}_t}
其中:
var:待更新的模型参数mg:梯度一阶矩的指数移动平均ms:梯度二阶矩的指数移动平均mom:动量缓冲g_t:当前步梯度lr:学习率ρ(rho):衰减系数μ(momentum):动量系数ε(epsilon):数值稳定性常数
注意:epsilon加在sqrt内部(
√(x + ε)),与TensorFlowtf.raw_ops.ApplyCenteredRMSProp语义一致。与PyTorch的√x + ε存在细微数值差异。
参数说明
| 参数名 | 输入/输出 | 描述 | 数据类型 | 数据格式 |
|---|---|---|---|---|
| var | 输入 | 待更新的模型参数。支持空Tensor。shape与mg/ms/mom/grad一致。 | FLOAT、FLOAT16、BFLOAT16 | ND |
| mg | 输入 | 梯度一阶矩的指数移动平均。支持空Tensor。shape与var一致。 | FLOAT、FLOAT16、BFLOAT16 | ND |
| ms | 输入 | 梯度二阶矩的指数移动平均。支持空Tensor。shape与var一致。 | FLOAT、FLOAT16、BFLOAT16 | ND |
| mom | 输入 | 动量缓冲。支持空Tensor。shape与var一致。 | FLOAT、FLOAT16、BFLOAT16 | ND |
| lr | 输入 | 学习率。必须为0-D标量Tensor。 | FLOAT、FLOAT16、BFLOAT16 | ND |
| rho | 输入 | 衰减系数。必须为0-D标量Tensor。 | FLOAT、FLOAT16、BFLOAT16 | ND |
| momentum | 输入 | 动量系数。必须为0-D标量Tensor。 | FLOAT、FLOAT16、BFLOAT16 | ND |
| epsilon | 输入 | 数值稳定性常数。必须为0-D标量Tensor。 | FLOAT、FLOAT16、BFLOAT16 | ND |
| grad | 输入 | 当前步梯度。支持空Tensor。shape与var一致。 | FLOAT、FLOAT16、BFLOAT16 | ND |
| use_locking | 属性 | 可选属性,是否对更新加锁。默认值为false。 | BOOL | - |
| var | 输出 | 更新后的模型参数,与输入var共享存储(inplace)。mg/ms/mom亦在输入上原地更新,与canndev IR一致,不作为独立输出声明。 | FLOAT、FLOAT16、BFLOAT16 | ND |
约束说明
- var、mg、ms、mom、grad五个ND Tensor的shape必须完全一致。
- lr、rho、momentum、epsilon必须为0-D标量Tensor。
- 所有输入和输出的数据类型必须相同(FLOAT、FLOAT16或BFLOAT16)。
- 算子仅声明1个输出var,与输入var共享存储;mg/ms/mom在输入上原地更新。
- 不支持稀疏梯度。
- FLOAT16/BFLOAT16输入时,内部走FP32混合精度计算路径,最终Cast回原始dtype。
- fp16上溢边界:grad值接近fp16最大值65504时,内部FP32计算路径可防止中间溢出,但最终Cast回fp16时若结果超出范围则饱和截断。
调用说明
| 调用方式 | 调用样例 | 说明 |
|---|---|---|
| 图模式调用 | test_geir_apply_centered_rms_prop | 通过算子IR构图方式调用ApplyCenteredRMSProp算子。 |