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

SGD

产品支持情况

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

上表写的是SGD在各产品形态上的可得性,不是本次交付的架构范围。本仓的Ascend C实现只适配 Ascend 950PR/Ascend 950DT(sgd_def.cpp中仅AddConfig("ascend950"));其余产品形态上的SGD由CANN内置的TBE实现提供,语义一致,但不由本算子承载。

功能说明

  • 算子功能:带动量的随机梯度下降(SGD)优化器更新算子,训练迭代中就地更新一组权重。

  • 计算公式

    dddampeningwdwdweightDecaylrlrlearningRate[0]mmmomentum[0],逐元素计算:

    步骤一 权重衰减(仅wd≠0wd \neq 0时执行,否则grad=gradientgrad = gradient):

    grad=gradient+parameters×wdgrad = gradient + parameters \times wd

    步骤二 动量累积(无条件执行):

    accumt=accum×m+gradaccum_t = accum \times m + grad

    步骤三 阻尼修正(仅d≠0d \neq 0时执行)。statstat逐元素的首步标记,取值1表示该元素处于首步、不施加阻尼:

    accumt=accumt−grad×(1−stat)×daccum_t = accum_t - grad \times (1 - stat) \times d

    步骤四 权重更新(无条件写出):

    parametersout={parameters−(grad×lr+accumt×m×lr),nesterov=trueparameters−accumt×lr,nesterov=falseparameters_{out} = \begin{cases} parameters - (grad \times lr + accum_t \times m \times lr), & nesterov = true \\ parameters - accum_t \times lr, & nesterov = false \end{cases}

    步骤五 动量与标记回写,受m≠0m \neq 0掩码控制:

    accumout,statout={accumt, 0,m≠0保持输入原值(不回写),m=0accum_{out}, stat_{out} = \begin{cases} accum_t,\ 0, & m \neq 0 \\ \text{保持输入原值(不回写)}, & m = 0 \end{cases}

  • 计算精度:中间计算在float32域进行,结果按就近偶数舍入(round-half-to-even)回目标数据类型。learningRatemomentumparameters同数据类型,故float16/bfloat16下这两个标量本身已被量化。

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
parameters 输入 / 输出(原地) 待更新的权重。无条件被改写。维度数(rank)须在[1, 8]内;不支持空tensor。 FLOAT、FLOAT16、BFLOAT16 ND
gradient 输入 梯度。shape与数据类型须与parameters一致。 FLOAT、FLOAT16、BFLOAT16 ND
learning_rate 输入 学习率。标量(元素个数为1),数据类型须与parameters一致。 FLOAT、FLOAT16、BFLOAT16 ND
accum 输入 / 输出(原地) 动量累积量。仅momentum ≠ 0时被改写;momentum = 0时逐位保持原值。shape与数据类型须与parameters一致。 FLOAT、FLOAT16、BFLOAT16 ND
momentum 输入 动量因子。标量(元素个数为1),数据类型须与parameters一致。取值为0(含-0.0)时触发“不回写”语义。 FLOAT、FLOAT16、BFLOAT16 ND
stat 输入 / 输出(原地) 逐元素首步标记,取值1表示该元素处于首步、不施加阻尼。仅momentum ≠ 0时被改写为0;momentum = 0时逐位保持原值。shape与数据类型须与parameters一致。 FLOAT、FLOAT16、BFLOAT16 ND
dampening 属性 动量阻尼系数,默认值0.0。nesterov为true时必须为0。 FLOAT -
weight_decay 属性 权重衰减系数,默认值0.0。必须大于或等于0。 FLOAT -
nesterov 属性 是否启用Nesterov动量,默认值false。 BOOL -
parameters 输出 更新后的权重,与输入parameters为同一块内存。shape、数据类型、数据格式均与输入parameters一致。 FLOAT、FLOAT16、BFLOAT16 ND

约束说明

  • 三路原地回写,但图上仅声明1个输出:算子实际就地更新parametersaccumstat三个张量,而图原型只声明parameters一个输出,accumstat通过覆写其输入内存返回。调用方必须把这三者都视为可写。此形态与CANN内置实现一致。

  • momentum = 0时的回写语义momentum为0(含-0.0)时,accumstat 完全不被写入,逐位保持输入原值(包括NaN的具体位模式、±inf-0.0);parameters不受该掩码影响,任何momentum取值下都照常计算并写出。momentum为极小非零值(如1e-81e-30)时按非零处理,正常回写。

  • ⚠️ 从PyTorch迁移的差异告警:本算子momentum = 0时的“不回写”方向与torch.optim.SGD一致(PyTorch在momentum == 0时整块跳过动量更新)。但 PyTorch的dampening施加在该判断之内,本算子(与CANN内置实现一致)施加在判断之外。因此当momentum = 0dampening > 0stat = 0时,parameters的更新量与PyTorch相差(1−dampening)(1 - dampening)倍。仅当dampening = 0stat = 1时两者一致。从PyTorch迁移的调用方须感知此差异。

  • rank与空tensorparameters的维度数须在[1, 8]内,0维标量被拒绝不支持空tensor —— 任意一轴或多轴为0均判为非法并返回错误码,不存在“空进空出”语义(accum/stat的原地回写在元素数为0时无定义)。

  • 属性取值nesterov = truedampening必须为0;weight_decay必须大于或等于0。违反者返回参数非法错误码。

  • inf/NaN:按IEEE 754语义传播,不做钳制或特判。特别地,accum±infmomentum = 0时,accum×momentumaccum \times momentum产生的NaN会按IEEE语义传播进parameters

  • 确定性:输出逐位可复现。算子为纯逐元素计算、无跨元素累加,多核切分不改变任一元素的计算顺序。

  • 张量连续性:所有输入须为连续张量。本算子不提供aclnn接口,无接口层做转连续/回填,非连续视图由调用方(GE图编译期)负责处理。

调用说明

不提供aclnn单算子接口。 SGD在CANN上游本就是纯图模式算子:CANN 9.1.0的 include/aclnnop/下无aclnn_sgd.h(只有语义不同的aclnn_fused_sgd.h), libopapi.so未导出aclnnSgd*符号,canndev全仓亦无aclnn_sgd定义。 本算子与之对齐,只支持GE图模式下发。

调用方式 样例代码 说明
图模式 test_geir_sgd.cpp 通过GE图方式调用SGD算子。

参考资源

  • 《Ascend C算子开发》:算子开发的概念原理与编程模型。
  • 算子列表:本项目全部算子的分类、调用方式与功能说明。
  • 算子调用快速入门:算子样例的编译与运行步骤。
  • apply_momentum:同族的动量优化器算子,与本算子结构最相近,可对照阅读。
  • fused_sgd:语义不同的另一个算子(多TensorList融合、dampening施加在momentum分支之内、无stat)。名称相近,请勿混用。