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

LarsV2Update

产品支持情况

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

功能说明

  • 算子功能:LARS-V2 优化器梯度更新,根据权重的 L2 范数自适应调整梯度的信任系数,用于大 batch 训练场景。

  • 计算公式:

    w_norm=w_square_sumw\_norm = \sqrt{w\_square\_sum}

    g_norm=g_square_sumg\_norm = \sqrt{g\_square\_sum}

    coeff=hyperpara×w_normweight_decay×w_norm+g_norm+epsiloncoeff = \frac{hyperpara \times w\_norm}{weight\_decay \times w\_norm + g\_norm + epsilon}

    • 若 use_clip = True:

      coeff=max⁡(0, min⁡(coefflearning_rate, 1))coeff = \max\left(0,\ \min\left(\frac{coeff}{learning\_rate},\ 1\right)\right)

    grad_weight=w×weight_decay+ggrad\_weight = w \times weight\_decay + g

    g_new=grad_weight×coeffg\_new = grad\_weight \times coeff

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
w 输入 权重张量。 FLOAT ND
g 输入 梯度张量,与w同型同形。 FLOAT ND
w_square_sum 输入 权重平方和(由SquareSumAll预计算),标量。 FLOAT ND
g_square_sum 输入 梯度平方和,标量。 FLOAT ND
weight_decay 输入 权重衰减系数,标量。 FLOAT ND
learning_rate 输入 学习率,标量。仅在use_clip=True时参与计算。 FLOAT ND
hyperpara 属性 LARS信任系数(eta),默认0.001。 Float -
epsilon 属性 防除零小常数,默认0.00001。 Float -
use_clip 属性 是否将coeff裁剪到[0,1],默认False。 Bool -
g_new 输出 更新后的梯度,与w同型同形。 FLOAT ND

约束说明

  • 输入w和g必须具有相同的形状和数据类型。
  • 输出g_new与w同型同形同dtype。
  • w_square_sum、g_square_sum、weight_decay、learning_rate恒为FLOAT类型标量。
  • 支持维度1~8维。
  • 支持动态shape和动态rank。

调用说明

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