| 文件 | 最后提交记录 | 最后更新时间 |
|---|---|---|
| 14 天前 | ||
| 14 天前 | ||
| 14 天前 | ||
| 14 天前 | ||
| 14 天前 | ||
| 14 天前 | ||
| 12 天前 |
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)优化器第三阶段的一阶矩更新值,以及行、列和全局归约结果。
-
计算公式:
对二维输入张量
u和m,以及标量输入eps、beta1、clip_threshold、sum_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_moment为true时,输出m为m_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。eps、beta1、clip_threshold和sum_square_u必须为FLOAT32类型的标量或单元素一维张量。global_shape为可选INT64类型输入,必须是一维长度为2的张量[global_n, global_m];未提供时使用输入u的二维形状进行归约计算。
调用说明
| 调用方式 | 调用样例 | 说明 |
|---|---|---|
| 图模式调用 | test_geir_apply_came_part3.cpp | 通过算子IR构图方式调用ApplyCamePart3算子。 |