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

ApplyCamePart4

产品支持情况

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

功能说明

  • 算子功能:

    CAME优化器第4段(参数更新段)。输入待更新参数param_in、一阶动量m(形状(N,M))与置信因子 r_in(N)、c_in(M),按CAME更新规则回写param_out/r_out/c_out。sum_r(全N行r_in的归约和)与 global_shape(全局N,M)为可选输入:分布式场景由前序算子/全局形状传入;单机场景缺省,kernel 内部完成归约并取本地n/m。

  • 计算公式(n = len(r_in),m = len(c_in);N/M取global_shape(给定)否则n/m):

    sum_r=∑ir_ini(sum_r缺省时kernel内归约)sum\_r=\sum_{i} r\_in_i \quad (\text{sum\_r缺省时kernel内归约})

    r_out=β3⋅r_in+(1−β3)M⋅sum_u_rr\_out=\beta_3 \cdot r\_in+\frac{(1-\beta_3)}{M} \cdot sum\_u\_r

    c_out=β3⋅c_in+(1−β3)N⋅sum_u_cc\_out=\beta_3 \cdot c\_in+\frac{(1-\beta_3)}{N} \cdot sum\_u\_c

    denom=β3⋅sum_rN+(1−β3)⋅sum_u_rcM⋅Ndenom=\beta_3 \cdot \frac{sum\_r}{N}+(1-\beta_3)\cdot\frac{sum\_u\_rc}{M \cdot N}

    param_out=(1−lr⋅weight_decay)⋅param_in−lr⋅mr_out⊗c_out/denomparam\_out=(1-lr\cdot weight\_decay)\cdot param\_in-\frac{lr \cdot m}{\sqrt{r\_out \otimes c\_out / denom}}

    其中 r_out⊗c_outr\_out \otimes c\_out 为(N,1)×(1,M)外积。fp16/bf16路径:输入cast到fp32计算, 输出按RNE(CAST_RINT)round回低精度;param更新以round后的r_out/c_out为输入。

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
param_in 输入 待更新参数,形状(N,M)。 FLOAT、FLOAT16、BFLOAT16 ND
m 输入 一阶动量,形状(N,M),与param_in一致。 FLOAT、FLOAT16、BFLOAT16 ND
r_in 输入 行置信因子,形状(N,)。 FLOAT、FLOAT16、BFLOAT16 ND
c_in 输入 列置信因子,形状(M,)。 FLOAT、FLOAT16、BFLOAT16 ND
weight_decay 输入 权重衰减系数,标量(1,)。 FLOAT ND
lr 输入 学习率,标量(1,)。 FLOAT ND
beta3 输入 置信因子衰减系数,标量(1,)。 FLOAT ND
sum_u_r 输入 r方向更新量归约,形状(N,)。 FLOAT ND
sum_u_c 输入 c方向更新量归约,形状(M,)。 FLOAT ND
sum_u_rc 输入 全局更新量归约,标量(1,)。 FLOAT ND
sum_r 输入(可选) 全N行r_in的归约和,标量(1,);缺省时kernel内归约。 FLOAT ND
global_shape 输入(可选) 全局(N,M),形状(2,);缺省取本地n/m。 INT64 ND
param_out 输出 更新后参数,形状(N,M)。 FLOAT、FLOAT16、BFLOAT16 ND
r_out 输出 更新后行置信因子,形状(N,)。 FLOAT、FLOAT16、BFLOAT16 ND
c_out 输出 更新后列置信因子,形状(M,)。 FLOAT、FLOAT16、BFLOAT16 ND

约束说明

  • param_in/m必须为2D,r_in/c_in必须为1D;r_in长度 = param_in第0维,c_in长度 = param_in第1维(tiling校验,不一致返回GRAPH_FAILED)。
  • 无属性(attr)。
  • sqrt域:s = r_out⊙c_out/denom逐元素需 ≥ 0,denom ≠ 0;否则按IEEE语义产生inf/nan传播(与torch公式语义一致,非算子错误)。
  • 空tensor(N=0或M=0)不做守卫:不崩溃,但非空维输出不写入(未定义),行为对齐A2。
  • 支持关系:Atlas A2系列(ascend910b/ascend910_93)由canndev仓内置ascendc实现支持(ascendc_config.json compute_units=[ascend910b,ascend910_93],kernel binary构建配置binary_json_cfg.ini已收录;TBE旧式aic-ops-info.ini未收录,以动态编译形式提供);Ascend 950PR/950DT由本仓vendor包支持(arch35)。两实现互不影响,vendor包只许安装到Ascend 950环境——其proto/tiling注册无soc维度,误入A2环境会遮蔽canndev内置实现并因tiling数据结构不匹配导致算子不可用。

调用说明

调用方式 调用样例 说明
GE图模式调用 test_geir_apply_came_part4 通过算子IR构图方式调用ApplyCamePart4算子。