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

ApplyCamePart3

产品支持情况

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

功能说明

  • 算子功能:计算CAME(Confidence-guided Adaptive Memory Efficient)优化器第三阶段的一阶矩更新值,以及行、列和全局归约结果。

  • 计算公式

    对二维输入张量 um,以及标量输入 epsbeta1clip_thresholdsum_square_u,设全局行数为 global_n、全局列数为 global_m,先计算:

    s=max⁡(1,sum_square_uglobal_n×global_m×clip_threshold)s = \max\left(1,\frac{sum\_square\_u}{global\_n \times global\_m \times clip\_threshold}\right)

    实现遵循A2的条件分支语义:仅当缩放值大于1时使用该值,其余情况使用1。

    m_updatei,j=(1−beta1)×ui,js+beta1×mi,jm\_update_{i,j} = (1 - beta1) \times \frac{u_{i,j}}{s} + beta1 \times m_{i,j}

    use_first_momenttrue时,输出mm_update;为false时,输出m保留输入m

    再令:

    xi,j=(ui,js−m_updatei,j)2+epsx_{i,j}=\left(\frac{u_{i,j}}{s}-m\_update_{i,j}\right)^2+eps

    分别进行行、列和全局归约:

    sum_u_ri=∑jxi,jsum_u_cj=∑ixi,jsum_u_rc=∑i∑jxi,j\begin{aligned} sum\_u\_r_i &= \sum_j x_{i,j} \\ sum\_u\_c_j &= \sum_i x_{i,j} \\ sum\_u\_rc &= \sum_i \sum_j x_{i,j} \end{aligned}

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
u 输入 二维输入张量,公式中的u FLOAT32 ND
m 输入 一阶矩输入张量,形状必须与u相同,公式中的m BFLOAT16、FLOAT16、FLOAT32 ND
eps 输入 数值稳定项,标量或单元素一维张量。 FLOAT32 ND
beta1 输入 一阶矩衰减系数,标量或单元素一维张量。 FLOAT32 ND
clip_threshold 输入 裁剪阈值,标量或单元素一维张量。 FLOAT32 ND
sum_square_u 输入 u平方和或对应的全局统计值,标量或单元素一维张量。 FLOAT32 ND
global_shape 可选输入 全局二维形状,包含全局行数和列数。 INT64 ND
use_first_moment 属性 是否输出更新后的一阶矩,默认为false Bool -
m 输出 一阶矩输出张量,形状和数据类型与输入m相同。 BFLOAT16、FLOAT16、FLOAT32 ND
sum_u_r 输出 按列归约后的行结果,形状为[N] FLOAT32 ND
sum_u_c 输出 按行归约后的列结果,形状为[M] FLOAT32 ND
sum_u_rc 输出 全局归约结果,形状为[1] FLOAT32 ND

约束说明

  • u和输入m必须为形状相同的二维ND张量,两个维度都必须大于0,且每个维度不超过INT32_MAX
  • epsbeta1clip_thresholdsum_square_u必须为FLOAT32类型的标量或单元素一维张量。
  • global_shape为可选INT64类型输入,必须是一维长度为2的张量[global_n, global_m];未提供时使用输入u的二维形状进行归约计算。

调用说明

调用方式 调用样例 说明
图模式调用 test_geir_apply_came_part3.cpp 通过算子IR构图方式调用ApplyCamePart3算子。