ApplyRMSProp
产品支持情况
| 产品 | 是否支持 |
|---|---|
| Ascend 950PR/Ascend 950DT | √ |
| Atlas A3 训练系列产品/Atlas A3 推理系列产品 | √ |
| Atlas A2 训练系列产品/Atlas A2 推理系列产品 | √ |
| Atlas 200I/500 A2 推理产品 | × |
| Atlas 推理系列产品 | √ |
| Atlas 训练系列产品 | √ |
功能说明
-
算子功能:执行RMSProp优化器(非centered版本)的单步参数更新。根据当前梯度、梯度平方移动平均和动量累积,更新参数张量
var、ms、mom,全部采用 inplace 语义。对标TensorFlowtf.raw_ops.ApplyRMSProp接口的计算语义。 -
计算公式:
mst=ρ⋅mst−1+(1−ρ)⋅grad2momt=momentum⋅momt−1+lr⋅gradmst+ϵvart=vart−1−momt\begin{aligned} ms_{t} &= \rho \cdot ms_{t-1} + (1 - \rho) \cdot grad^2 \\ mom_{t} &= momentum \cdot mom_{t-1} + lr \cdot \frac{grad}{\sqrt{ms_{t} + \epsilon}} \\ var_{t} &= var_{t-1} - mom_{t} \end{aligned}
其中
lr为学习率,rho为衰减系数(取值范围 [0, 1)),momentum为动量系数,epsilon为数值稳定常数(必须 > 0)。
参数说明
| 参数名 | 输入/输出/属性 | 描述 | 数据类型 | 数据格式 |
|---|---|---|---|---|
| var | 输入 / 输出(inplace) | 待更新的权重参数,对应公式中的var。Kernel内inplace更新,GE IR单输出视图与输入var共享Device内存。 | FLOAT、FLOAT16、BFLOAT16 | ND |
| ms | 输入(inplace更新) | 梯度平方移动平均,对应公式中的ms。shape/dtype必须与var一致;Kernel内显式写回输入GM地址。 | FLOAT、FLOAT16、BFLOAT16 | ND |
| mom | 输入(inplace更新) | 动量累积,对应公式中的mom。shape/dtype必须与var一致;Kernel内显式写回输入GM地址。 | FLOAT、FLOAT16、BFLOAT16 | ND |
| lr | 输入 | 学习率,对应公式中的lr。shape={1} 的1元素scalar Tensor,dtype必须与var一致。 | FLOAT、FLOAT16、BFLOAT16 | ND |
| rho | 输入 | 衰减系数,对应公式中的rho。shape={1} 的1元素scalar Tensor,取值范围 [0, 1)。 | FLOAT、FLOAT16、BFLOAT16 | ND |
| momentum | 输入 | 动量系数,对应公式中的momentum。shape={1} 的1元素scalar Tensor。 | FLOAT、FLOAT16、BFLOAT16 | ND |
| epsilon | 输入 | 数值稳定常数,对应公式中的epsilon。shape={1} 的1元素scalar Tensor,必须 > 0。 | FLOAT、FLOAT16、BFLOAT16 | ND |
| grad | 输入 | 当前梯度Tensor,对应公式中的grad。shape/dtype必须与var一致。 | FLOAT、FLOAT16、BFLOAT16 | ND |
| use_locking | 属性 | 是否在更新时加锁。默认false。当前实现不强制互斥锁,仅作语义占位。 | BOOL | - |
| var (output) | 输出 | 更新后的var Tensor,与输入var共享Device内存(inplace)。 | FLOAT、FLOAT16、BFLOAT16 | ND |
约束说明
无
调用说明
| 调用方式 | 调用样例 | 说明 |
|---|---|---|
| 图模式 | test_geir_apply_rms_prop | 通过 算子IR构图方式调用ApplyRMSProp算子。 |