文件最后提交记录最后更新时间
27 天前
10 天前
27 天前
27 天前
27 天前
27 天前
2 个月前
1 个月前
README

ApplyProximalGradientDescent

产品支持情况

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

功能说明

  • 算子功能:完成带L1/L2正则的近端梯度下降(Proximal Gradient Descent)一步更新。给定待更新权重var、学习率alpha、L1正则系数l1、L2正则系数l2和梯度delta,按元素计算输出。

  • 对标接口:TensorFlow tf.raw_ops.ResourceApplyProximalGradientDescent

  • 计算公式:

    prox_v=var−alpha×deltaprox\_v = var - alpha \times delta

    varOut=sign⁡(prox_v)1+alpha×l2×max⁡(∣prox_v∣−alpha×l1, 0)varOut = \dfrac{\operatorname{sign}(prox\_v)}{1 + alpha \times l2} \times \max\bigl(|prox\_v| - alpha \times l1,\ 0\bigr)

    其中sign(0) = 0,与TensorFlow语义一致。

    退化情形:

    条件 等价公式
    l1 = 0, l2 = 0 varOut = var - alpha * delta(标准SGD)
    l1 = 0 varOut = (var - alpha * delta) / (1 + alpha * l2)

参数说明

参数名 输入/输出 描述 数据类型 数据格式
var 输入 待更新权重张量,支持1-8维tensor,支持非连续tensor。shape需要与delta的shape完全相同。 FLOAT16、FLOAT32 NC1HWC0、C1HWNCoC0、ND、FRACTAL_Z
alpha 输入 学习率标量张量(0-D或shape=[1],非负)。 FLOAT16、FLOAT32 ND
l1 输入 L1正则系数标量张量(0-D或shape=[1],非负)。 FLOAT16、FLOAT32 ND
l2 输入 L2正则系数标量张量(0-D或shape=[1],非负)。 FLOAT16、FLOAT32 ND
delta 输入 梯度张量,shape与var完全相同,支持非连续tensor。 FLOAT16、FLOAT32 NC1HWC0、C1HWNCoC0、ND、FRACTAL_Z
var(输出) 输出 输出张量,shape/dtype/format与输入var一致。 FLOAT16、FLOAT32 NC1HWC0、C1HWNCoC0、ND、FRACTAL_Z

约束说明

  • 所有输入的数据类型必须一致,仅支持FLOAT16、FLOAT32。
  • var和delta的shape必须完全相同,不支持广播。
  • alpha / l1 / l2必须为0-D或shape=[1] 的标量张量。
  • FP16输入在内部会提升至FP32精度计算,结果再转换回原始精度。
  • 默认确定性实现。

调用说明

调用方式 样例代码 说明
图模式 test_geir_apply_proximal_gradient_descent.cpp 通过图模式方式调用ApplyProximalGradientDescent算子。