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

ApplyAdamWV2

产品支持情况

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

功能说明

  • 算子功能:

    实现adamW优化器功能。

  • 计算公式:

    if(maximize):gt=−gtif(maximize) : g_{t} = - g_{t}

    mt=β1mt−1+(1−β1)gtm_{t}=\beta_{1} m_{t-1}+\left(1-\beta_{1}\right) g_{t}

    vt=β2vt−1+(1−β2)gt2v_{t}=\beta_{2} v_{t-1}+\left(1-\beta_{2}\right) g_{t}^{2}

    m^t=mt1−β1t\hat{m}_{t}=\frac{m_{t}}{1-\beta_{1}^{t}}

    v^t=vt1−β2t\hat{v}_{t}=\frac{v_{t}}{1-\beta_{2}^{t}}

    if(amsgrad):maxGradNorm=max(maxGradNorm,v^t)if(amsgrad) : maxGradNorm = max(maxGradNorm,\hat{v}_{t})

    θt+1=θt−ηv^t+ϵm^t−η⋅λ⋅θt−1\theta_{t+1}=\theta_{t}-\frac{\eta}{\sqrt{\hat{v}_{t}}+\epsilon} \hat{m}_{t}-\eta \cdot \lambda \cdot \theta_{t-1}

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
varRef 输入/输出 待计算的权重输入同时也是输出,公式中的输入/输出θ FLOAT16、BFLOAT16、FLOAT ND
mRef 输入/输出 待更新参数对应的一阶动量,公式中的输入/输出m。 FLOAT16、BFLOAT16、FLOAT ND
vRef 输入/输出 待更新参数对应的二阶动量,公式中的输入/输出v。 FLOAT16、BFLOAT16、FLOAT ND
maxGradNormOptionalRef 输入/输出 保存v参数的最大值,公式中的maxGradNorm。 FLOAT16、BFLOAT16、FLOAT ND
grad 输入 梯度数据,公式中的输入g。 FLOAT16、BFLOAT16、FLOAT ND
step 输入 迭代次数,公式中的t。 INT64、FLOAT ND
lr 属性
  • 学习率。
  • 取值范围是(0,1),默认为0.1。计算公式中的η。
FLOAT -
beta1 属性
  • beta1参数。
  • 取值范围是(0,1),默认为0.1。计算公式中的β1。
FLOAT -
beta2 属性
  • beta2参数。
  • 取值范围是(0,1),默认为0.1。计算公式中的β2。
FLOAT -
weightDecay 属性
  • 权重衰减系数。
  • 取值范围是(0,1),默认为0.1。计算公式中的λ。
FLOAT -
eps 属性
  • 防除0参数。
  • 默认为1e-8。计算公式中的ϵ。
FLOAT -
amsgrad 属性
  • 是否使用算法的AMSGrad变量,默认为false。
BOOL -
maximize 属性
  • 是否最大化参数,默认为false。
BOOL -

约束说明

调用说明

调用方式 调用样例 说明
aclnn调用 test_aclnn_apply_adam_w_v2 通过aclnnApplyAdamWV2接口方式调用ApplyAdamWV2算子。