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

ApplyProximalAdagrad

产品支持情况

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

功能说明

  • 算子功能:ApplyProximalAdagrad是结合Adagrad自适应学习率与FOBOS(Forward-Backward Splitting)Proximal近端算法的优化器算子,功能对标tf.raw_ops.ApplyProximalAdagrad。基于梯度平方累加器自适应调整学习率,并通过软阈值(L1正则化)与缩放(L2正则化)对模型参数进行原地更新。

  • 计算公式:

    accumt=accumt−1+gradt2ηt=lraccumtproxt=vart−1−ηt⋅gradtvart=sign(proxt)1+ηt⋅l2⋅max⁡ ⁣(∣proxt∣−ηt⋅l1, 0)\begin{aligned} \text{accum}_t &= \text{accum}_{t-1} + \text{grad}_t^2 \\ \eta_t &= \frac{\text{lr}}{\sqrt{\text{accum}_t}} \\ \text{prox}_t &= \text{var}_{t-1} - \eta_t \cdot \text{grad}_t \\ \text{var}_t &= \frac{\text{sign}(\text{prox}_t)}{1 + \eta_t \cdot \text{l2}} \cdot \max\!\left(|\text{prox}_t| - \eta_t \cdot \text{l1},\ 0\right) \end{aligned}

    当L1 = 0时简化为:

    vart=proxt1+ηt⋅l2\text{var}_t = \frac{\text{prox}_t}{1 + \eta_t \cdot \text{l2}}

  • 说明:

    • var(参数)与accum(梯度平方累加器)均为Ref Tensor,算子执行后原地更新
    • lrl1l2为0-D标量Tensor,分别要求lr > 0l1 ≥ 0l2 ≥ 0
    • 逐元素独立计算,天然确定性,无跨元素/跨核依赖。
    • float16 / bfloat16输入在算子内部Cast为float32完成中间计算,再Cast回float16 / bfloat16输出,与PyTorch / TensorFlow实现约定一致。
    • bfloat16路径上lr / l1 / l2标量读取与accum tail padding 1.0构造均通过LocalTensor借道 + Vector Cast实现,规避bisheng编译器后端在scalar路径上对bf16类型转换的限制("not support bf16 type cast")。

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
var 输入 公式中的var,待更新的模型参数(Ref Tensor,原地更新)。shape与accum/grad一致。 FLOAT32, FLOAT16, BFLOAT16 ND
accum 输入 公式中的accum,梯度平方累加器(Ref Tensor,原地更新)。shape与var/grad一致,要求各元素非负。 FLOAT32, FLOAT16, BFLOAT16 ND
lr 输入 公式中的lr,学习率。0-D或1元素1-D Tensor,要求lr > 0。 FLOAT32, FLOAT16, BFLOAT16 ND
l1 输入 公式中的l1,L1正则化强度。0-D或1元素1-D Tensor,要求l1 ≥ 0。 FLOAT32, FLOAT16, BFLOAT16 ND
l2 输入 公式中的l2,L2正则化强度。0-D或1元素1-D Tensor,要求l2 ≥ 0。 FLOAT32, FLOAT16, BFLOAT16 ND
grad 输入 公式中的grad,当前步的梯度张量。shape与var/accum一致。 FLOAT32, FLOAT16, BFLOAT16 ND
var (output) 输出 更新后的参数,与输入var共享存储(inplace更新)。算子仅暴露此1个输出端口。 FLOAT32, FLOAT16, BFLOAT16 ND

约束说明

  • 支持float32 / float16 / bfloat16三种数据类型;所有tensor(var / accum / lr / l1 / l2 / grad)数据类型必须严格一致。
  • varaccumgrad三者shape必须完全一致,且均为连续排布的ND Tensor。
  • lrl1l2必须为0-D或1元素1-D的标量Tensor。
  • 调用方需保证accum ≥ 0lr > 0l1 ≥ 0l2 ≥ 0;算子内部不做运行时值域校验。
  • accum + grad^2 == 0rsqrt输出Inf/NaN,行为与PyTorch / TensorFlow原生实现一致,需由上游调用方规避。
  • varaccum均为Ref Tensor:var通过显式输出端口inplace写回,accum直接通过input ref端口原地更新,调用方需将其视为可被算子修改的存储。

调用说明

调用方式 样例代码 说明
图模式 test_geir_apply_proximal_adagrad.cpp 通过GE IR图模式构建并运行ApplyProximalAdagrad算子。示例按var → accum → lr → l1 → l2 → grad顺序串接6个Data/标量输入,输出端口仅声明varaccum通过input ref端口原地更新。可执行参数fp16 / bf16切换数据类型。