文件最后提交记录最后更新时间
29 天前
13 天前
29 天前
14 天前
29 天前
19 天前
1 个月前
1 个月前
README

SparseApplyRMSProp

产品支持情况

产品 是否支持
Ascend 950PR/Ascend 950DT
Atlas A3 训练系列产品/Atlas A3 推理系列产品 ×
Atlas A2 训练系列产品/Atlas A2 推理系列产品 ×
Atlas 200I/500 A2 推理产品 ×
Atlas 推理系列产品 ×
Atlas 训练系列产品 ×

功能说明

  • 算子功能:对指定稀疏索引的行执行RMSProp优化算法更新。根据indices向量,从var/ms/mom中gather出对应行,应用RMSProp更新公式后将更新后的行scatter写回,未命中行保持不变。

  • 计算公式

indices中每个索引i

msnew[i]=ρ⋅msold[i]+(1−ρ)⋅grad[i]2momnew[i]=momentum⋅momold[i]+lr⋅grad[i]/msnew[i]+ϵvarnew[i]=varold[i]−momnew[i]\begin{aligned} \text{ms}_{\text{new}}[i] &= \rho \cdot \text{ms}_{\text{old}}[i] + (1 - \rho) \cdot \text{grad}[i]^2 \\ \text{mom}_{\text{new}}[i] &= \text{momentum} \cdot \text{mom}_{\text{old}}[i] + \text{lr} \cdot \text{grad}[i] / \sqrt{\text{ms}_{\text{new}}[i] + \epsilon} \\ \text{var}_{\text{new}}[i] &= \text{var}_{\text{old}}[i] - \text{mom}_{\text{new}}[i] \end{aligned}

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
var 输入 待优化参数矩阵,shape [N, D1, D2, ...]。 FLOAT16、BFLOAT16、FLOAT32 ND
ms 输入 梯度平方的指数移动平均,shape同var。 FLOAT16、BFLOAT16、FLOAT32 ND
mom 输入 动量累计,shape同var。 FLOAT16、BFLOAT16、FLOAT32 ND
lr 输入 学习率标量,shape [1]。 FLOAT16、BFLOAT16、FLOAT32 ND
rho 输入 衰减率标量(典型值0.9),shape [1]。 FLOAT16、BFLOAT16、FLOAT32 ND
momentum 输入 动量系数标量(典型值0.0~0.9),shape [1]。 FLOAT16、BFLOAT16、FLOAT32 ND
epsilon 输入 平滑项标量(典型值1e-8),shape [1]。 FLOAT16、BFLOAT16、FLOAT32 ND
grad 输入 稀疏梯度,shape [M, D1, D2, ...],M = len(indices)。 FLOAT16、BFLOAT16、FLOAT32 ND
indices 输入 稀疏索引,shape [M],值应在[0, N)范围内且应唯一。 INT32、INT64 ND
var 输出 更新后的参数矩阵,shape同输入var。 FLOAT16、BFLOAT16、FLOAT32 ND
ms 输出 更新后的梯度平方 EMA,shape同输入ms。 FLOAT16、BFLOAT16、FLOAT32 ND
mom 输出 更新后的动量累计,shape同输入mom。 FLOAT16、BFLOAT16、FLOAT32 ND
use_locking 属性 是否使用互斥锁保护更新,默认false。 BOOL -

约束说明

  • 所有数据张量(var, ms, mom, lr, rho, momentum, epsilon, grad)必须使用相同的dtype(float16/bfloat16/float32)。
  • indices使用int32或int64。
  • var.shape == ms.shape == mom.shape。
  • grad.shape[0] == indices.shape[0]。
  • var.shape[1:] == grad.shape[1:],即var和grad的尾维度必须匹配。
  • var的rank >= 1。
  • indices的rank == 1。
  • indices中的值应在[0, var.shape[0])范围内且应唯一。若indices含重复值,结果不可预测(与TensorFlow行为一致)。
  • 空indices(len(indices) == 0)时跳过计算,var/ms/mom保持不变。
  • fp16/bf16内部以fp32计算后cast回原类型。

调用说明

调用方式 调用样例 说明
图模式 test_geir_sparse_apply_rms_prop 通过算子IR构图方式调用SparseApplyRMSProp算子。