文件最后提交记录最后更新时间
23 天前
9 天前
23 天前
23 天前
11 天前
23 天前
23 天前
18 天前
README

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 + ε)),与TensorFlow tf.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算子。