已关闭
[Requirement|需求建议]: 新增SGD优化器算子的Ascend 950(arch35/regbase)实现 #4554
raoliang_sac创建于  8月4日关闭于  8月4日
raoliang_sac成员
8月4日 创建

Backgroud(背景信息)

需要在 ops-nn 仓补齐 SGD(带动量的随机梯度下降)优化器算子在 Ascend 950PR/Ascend 950DT(arch35 / DAV_3510 / regbase)上的 Ascend C 实现。

现状:

  • 本仓 optim 目录下不存在 SGD 算子。已有的 optim/fused_sgd 是语义不同的另一个算子(融合优化器),不能替代。
  • 910B/910C 等产品形态上的 SGD 由 CANN 内置的 TBE 实现承载,算子原型登记在 canndev ops/built-in/op_proto/inc/nn_training_ops.hREG_OP(SGD)
  • Ascend 950 缺失该算子的 Ascend C 实现,训练场景下使用 SGD 优化器的图无法在 950 上落到本仓算子。

需要实现的算子语义(与 910B/910C 基线逐字对齐):

grad     = (wd != 0) ? gradient + parameters * wd : gradient
accum_t  = accum * m + grad                       // 无条件计算
accum_t -= (d != 0) ? grad * (1 - stat) * d : 0
parameters -= nesterov ? (grad * lr + accum_t * m * lr) : (accum_t * lr)
if (m != 0) { accum = accum_t; stat = 0; }        // 回写掩码

其中 d = dampeningwd = weight_decaylr = learning_rate[0]m = momentum[0]

接口形态:六输入 parameters / gradient / learning_rate / accum / momentum / stat,三属性 dampening / weight_decay / nesterov,图上只声明一个输出 parameters —— accumstat 靠覆写输入 GM 原地回写(与 A2 的 TBE 实现 reuse=('accum','parameters','stat') 形态一致)。

Origin(信息来源)

CANN 长尾算子补齐任务,开发分支 SGD-810

Benefit / Necessity (价值/作用)

  • 补齐 Ascend 950 上的训练优化器能力:SGD(含 momentum / Nesterov / weight decay / dampening)是训练场景最基础的优化器之一,950 上缺失会导致相关训练图无法在本仓获得 Ascend C 实现。
  • 与 A2 基线语义一致:dtype 支持面(FLOAT / FLOAT16 / BFLOAT16)、三个属性的默认值、非法组合的拒绝口径(nesterov == truedampening 必须为 0、weight_decay >= 0)全部对齐 910B/910C,同一张图在不同产品形态上行为一致。
  • 性能收益:与 A100 对标 200 例实测,NPU 侧 200/200 更快。

Design(设计方案)

1. 目录与交付形态

新增 optim/sgd/,共 22 个文件、+2719 行,纯新增,不修改任何存量文件(唯一非新增改动是在 docs/zh/op_list.md 追加一行算子登记)。

交付形态为 GE 图模式,不提供 aclnn 接口sgd_def.cpp 保持 ACLNNTYPE aclnn_exclude)。依据:CANN 9.1.0 的 include/aclnnop/aclnn_sgd.hlibopapi.so 未导出任何 aclnnSgd* 符号、canndev 全仓无 aclnn_sgd 定义 —— SGD 在上游本就是纯图模式算子,本算子与之对齐。

2. Kernel(op_kernel/arch35/

基于 ATVOSS DAG + ElementwiseSch

  • sgd_dag.h 用算子级 DAG 描述上述公式;sgd.cpp 在运行期按 momentum == 0 选择两套 DAG。
  • momentum == 0 掩码不作为 TilingKey 维度momentum 是 Device 侧 [1] 张量,Host Tiling 阶段拿不到其数值,只能做运行期分支;两套 DAG 同时存在于同一个 binary 内,binary 数量不变。
  • TilingKey 四维模板参数 schMode / useNesterov / hasWeightDecay / hasDampening。非法组合 nesterov==1 && dampening!=0ASCENDC_TPL_SEL 的两组 ARGS_SEL 在编译期剪除,不生成对应 binary
  • 组合数:业务模板 2×2×2 − 2 = 6(K0~K5)→ TPL_SEL 展开 12 → binary 12 × 3 dtype = 36

3. Host(op_host/

  • sgd_def.cpp:dtype FLOAT / FLOAT16 / BFLOAT16,format 仅 NDDynamicShapeSupportFlag(true) / DynamicRankSupportFlag(true) / PrecisionReduceFlag(false)AddConfig("ascend950"),不触碰任何 A2 配置。format 只做 ND 是跟随本仓 arch35 全族 optim 算子(apply_momentum / apply_ftrl / apply_adam_w_v2 / apply_adamax / apply_centered_rms_prop)的既定约定,这些算子的 def.cpp 一律只声明 ge::FORMAT_ND
  • sgd_infershape.cpp:属性校验(nesterov==truedampening 必须为 0、weight_decay >= 0);rank 限 [1,8](rank-0 拒绝);空 tensor(任一轴为 0)拒绝为 null_input-1(UNKNOWN_DIM)支持、-2(UNKNOWN_RANK)透传。
  • arch35/sgd_tiling.cpp:基于 ElewiseBaseTiling

4. 与 PyTorch 的已知语义分歧(有意保留,对齐 A2)

PyTorch 把 dampening 放在 momentum 块内部,本算子(同 A2)放在外部,故 m==0 && d>0 && stat==0 时结果差 (1-d) 倍。该分歧已在 sgd_proto.hREADME.md 显式记录,golden 参考实现按本算子契约(而非 torch.optim.SGD)编写。

5. 验证方案

测试腿 规模 结果
功能泛化(TTK kernel 模式) 800 例 800/800 PASS,6 个 tilingKey × 3 dtype 零缺失
三方精度(NPU / CPU-fp64 golden / A100-GPU) 50 例 50/50 PASS,ratio_mare 恒为 1.000
Host UT(infershape 11 + arch35 tiling 20) 31 例 31/31 PASS
GE 图模式冒烟(accum / stat 两条回写分支) 2 分支 各 60/60 元素逐一核对
性能对标 A100 200 例 200/200 NPU 更快

合计 1081 项、0 失败。

关联 PR:https://gitcode.com/cann/ops-nn/pull/8155

likedislike
Rraoliang_sac成员
8月4日 添加了label:requirement
CANN-robotCANN-robot成员
8月4日 关闭了 issue
CANN-robotCANN-robot成员
8月4日 添加了label:resolved