文件最后提交记录最后更新时间
12 天前
30 天前
30 天前
30 天前
30 天前
30 天前
14 天前
6 个月前
4 个月前
README

ApplyAdamW

产品支持情况

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

功能说明

  • 算子功能:实现adamW优化器功能。

  • 计算公式:

    gt={−gt if maxmize=truegt if maxmize=falseg_t=\begin{cases}-g_t & \text{ if } maxmize= true\\ g_t & \text{ if } maxmize=false \end{cases}

    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}

    β1t=β1t−1×β1\beta_{1}^{t}=\beta_{1}^{t-1}\times\beta_{1}

    β2t=β2t−1×β2\beta_{2}^{t}=\beta_{2}^{t-1}\times\beta_{2}

    vt={max⁡(maxGradNorm,vt) if amsgrad=truevt if amsgrad=falsev_t=\begin{cases}\max(maxGradNorm, v_t) & \text{ if } amsgrad = true\\ v_t & \text{ if } amsgrad = false \end{cases}

    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}} \\

    θ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}

<tr>
  <tr>
  <td>epsilon</td>
  <td>计算输入</td>
  <td>防止除数为0。</td>
  <td>FLOAT16、BFLOAT16、FLOAT32</td>
  <td>ND</td>
</tr>

<tr>
  <tr>
  <td>grad</td>
  <td>计算输入</td>
  <td>梯度数据,公式中的g_t。</td>
  <td>FLOAT16、BFLOAT16、FLOAT32</td>
  <td>ND</td>
</tr>
<tr>
  <tr>
  <td>max_grad_norm</td>
  <td>计算输入</td>
  <td>保存v参数的最大值,公式中的v。</td>
  <td>FLOAT16、BFLOAT16、FLOAT32</td>
  <td>ND</td>
</tr>
<tr>
  <td>amsgrad</td>
  <td>属性</td>
  <td>是否使用maxGradNormOptional变量</td>
  <td>BOOL</td>
  <td>ND</td>
</tr>
<tr>
  <td>maximize</td>
  <td>属性</td>
  <td>是否对梯度grad取反,应用梯度上升方向优化权重使损失函数最大化</td>
  <td>BOOL</td>
  <td>ND</td>
</tr>
参数名 输入/输出 描述 数据类型 数据格式
var 计算输入/计算输出 待计算的权重输入同时也是输出,公式中的theta FLOAT16、BFLOAT16、FLOAT32 ND
m 计算输入/计算输出 adamw优化器中m参数,公式中的m FLOAT16、BFLOAT16、FLOAT32 ND
v 计算输入/计算输出 adamw优化器中v参数,公式中的v FLOAT16、BFLOAT16、FLOAT32 ND
beta1_power 计算输入 beta1^(t-1)参数 FLOAT16、BFLOAT16、FLOAT32< ND
beta2_power 计算输入 beta2^(t-1)参数 FLOAT16、BFLOAT16、FLOAT32 ND
lr 计算输入 学习率,公式中的eta FLOAT16、BFLOAT16、FLOAT32 ND
weight_decay 计算输入 权重衰减系数 FLOAT16、BFLOAT16、FLOAT32 ND
beta1 计算输入 beta1参数。 FLOAT16、BFLOAT16、FLOAT32 ND
beta2 计算输入 beta2参数。 FLOAT16、BFLOAT16、FLOAT32 ND

约束说明

  • 输入张量的数据类型应保持一致,数据类型支持FLOAT16、BFLOAT16、FLOAT32。

  • 输入张量beta1Power、beta2Power、lr、weightDecay、beta1、beta2、eps的shape大小应为1。

  • 输入布尔值maximize为true时,maxGradNormOptional参数必选且数据类型和shape应与varRef一致时。

  • 确定性计算:

    • aclnnApplyAdamW默认确定性实现。

调用说明

调用方式 调用样例 说明
aclnn调用 test_aclnn_apply_adam_w 通过aclnnApplyAdamW接口方式调用ApplyAdamW算子。