CANN ops-math 算子详解:aclnnAmpUpdateScale 动态损失缩放接口的原理、两段式调用与源码剖析
2026/9/20 22:04:49 网站建设 项目流程

CANN ops-math 算子详解:aclnnAmpUpdateScale 动态损失缩放接口的原理、两段式调用与源码剖析

【免费下载链接】ops-math本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-math

aclnnAmpUpdateScale是 CANN ops-math 数学算子库(math/amp_update_scale)提供的 AMP(Automatic Mixed Precision)训练动态 Scale 更新算子,它根据当前 loss scale、growth tracker 计数器以及 Inf/NaN 检测标志,在 NPU 上完成"发现溢出则回退、连续正常则增长"的标量更新逻辑。本文以 aclnnAmpUpdateScale 接口文档 为主体,结合仓库内的算子定义、Tiling 实现、Kernel 代码与单测用例,完整讲解其功能公式、两段式接口签名、参数约束、错误码以及可直接编译运行的调用示例,帮助你快速在 FP16/BF16 混合精度训练框架中接入动态损失缩放能力。

功能与原理:AMP 训练中的动态损失缩放

在 FP16/BF16 混合精度训练中,较小的梯度数值在低精度表示下容易发生下溢(underflow),因此训练框架通常先对 loss 乘以一个缩放因子(loss scale),再执行反向传播与梯度更新。动态损失缩放(Dynamic Loss Scaling)的核心思想是:周期性检测梯度中是否出现 Inf/NaN,据此动态放大或缩小 scale,从而在避免溢出的同时尽可能保持梯度精度。

aclnnAmpUpdateScale即负责其中的"Scale 更新"环节:输入当前的 scale 值、连续未出现 Inf/NaN 的步数计数器、以及本次是否发现 Inf/NaN 的标志,输出更新后的 scale 与计数器。文档给出的计算公式如下:

$$ \text{updated_scale} = \begin{cases} \text{current_scale} \times \text{backoff_factor} & \text{if found_inf} \neq 0 \ \text{current_scale} \times \text{growth_factor} & \text{if growth_tracker + 1 = growth_interval and new_scale is finite} \ \text{current_scale} & \text{otherwise} \end{cases} $$

$$ \text{updated_growth_tracker} = \begin{cases} 0 & \text{if found_inf} \neq 0 \text{ or growth triggered} \ \text{growth_tracker} + 1 & \text{otherwise} \end{cases} $$

公式中各符号含义如下:

符号含义
current_scale当前的 loss scale 值(标量)
found_inf是否检测到 Inf/NaN 的标志,0 表示正常,非 0 表示发现 Inf/NaN
growth_tracker连续未出现 Inf/NaN 的步数计数器
growth_factorscale 增长因子,通常设置为 2.0
backoff_factorscale 回退因子,通常设置为 0.5
growth_interval触发 scale 增长的间隔步数

规则可以概括为四点:

  1. found_inf不为 0 时,scale乘以backoff_factor回退,growth_tracker重置为 0;
  2. found_inf为 0 且growth_tracker + 1等于growth_interval时,scale乘以growth_factor增长;
  3. 如果增长后的新scale溢出(inf/nan),则保持当前scale不变,growth_tracker重置为 0(溢出保护);
  4. 其他情况下,scale保持不变,growth_tracker递增 1。

该算子在 math/amp_update_scale/README.md 中的功能说明与文档完全一致,两者可以互为印证。

产品支持情况

接口文档明确给出了该算子在各产品上的支持矩阵:

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

该矩阵在 README.md 中同样存在。从源码侧看,amp_update_scale_def.cpp 中通过AICore().AddConfig()仅为ascend910bascend910_93ascend950三个平台注册了 AICore 配置,与"Atlas A2/A3 训练推理系列及 Ascend 950 系列支持"的支持范围一致;对应的算子二进制配置见 ascend910b/amp_update_scale_binary.json、ascend910_93 与 ascend950。

函数原型:两段式接口

与其他 aclnn 算子一致,AmpUpdateScale采用两段式接口(详见 两段式接口说明):必须先调用aclnnAmpUpdateScaleGetWorkspaceSize获取计算所需的 workspace 大小以及封装了算子计算流程的执行器,再调用aclnnAmpUpdateScale执行实际计算。

aclnnStatus aclnnAmpUpdateScaleGetWorkspaceSize( const aclTensor* currentScale, const aclTensor* growthTracker, const aclTensor* foundInf, double growthFactor, double backoffFactor, int64_t growthInterval, const aclTensor* updatedScale, const aclTensor* updatedGrowthTracker, uint64_t* workspaceSize, aclOpExecutor** executor)
aclnnStatus aclnnAmpUpdateScale( void* workspace, uint64_t workspaceSize, aclOpExecutor* executor, aclrtStream stream)

两段接口的职责划分:第一段接口完成入参校验与算子编译/执行器构建,第二段接口在指定 Stream 上提交计算任务。调用时需包含头文件aclnnop/aclnn_amp_update_scale.h

aclnnAmpUpdateScaleGetWorkspaceSize 参数说明

第一段接口的参数如下:

参数名输入/输出描述使用说明数据类型数据格式维度(shape)非连续Tensor
currentScale(aclTensor*)输入当前的loss scale值shape为标量 [1]FLOAT、FLOAT16、BFLOAT16ND1×
growthTracker(aclTensor*)输入连续未出现Inf/NaN的步数计数器shape为标量 [1]INT32ND1×
foundInf(aclTensor*)输入是否检测到Inf/NaN的标志shape为标量 [1]。0表示正常,非0表示发现Inf/NaN。数据类型需要与currentScale一致FLOAT、FLOAT16、BFLOAT16ND1×
growthFactor(float)输入scale增长因子当连续growth_interval步未检测到Inf/NaN时,scale将乘以该因子。通常设置为2.0----
backoffFactor(float)输入scale回退因子当检测到Inf/NaN时,scale将乘以该因子。通常设置为0.5----
growthInterval(int64_t)输入触发scale增长的间隔步数即连续多少步未检测到Inf/NaN后将增大scale。取值范围 >= 1----
updatedScale(aclTensor*)输出更新后的loss scale值shape为 [1]。数据类型需要与currentScale的数据类型一致FLOAT、FLOAT16、BFLOAT16ND1×
updatedGrowthTracker(aclTensor*)输出更新后的growth tracker计数器shape为 [1]INT32ND1×
workspaceSize(uint64_t*)输出返回需要在Device侧申请的workspace大小-----
executor(aclOpExecutor**)输出返回op执行器,包含了算子计算流程-----

参数语义与源码的对应关系

上述参数在算子注册层面与 amp_update_scale_def.cpp 一一对应:current_scalegrowth_trackerfound_inf三个输入,updated_scaleupdated_growth_tracker两个输出,以及growth_factor(FLOAT 属性)、backoff_factor(FLOAT 属性)、growth_interval(INT 属性)三个必填属性。其中current_scale/found_inf/updated_scale支持ge::DT_FLOATge::DT_FLOAT16ge::DT_BF16growth_tracker/updated_growth_tracker固定为ge::DT_INT32,格式统一为 ND,与接口文档参数表完全一致。

Tiling 阶段(amp_update_scale_tiling.cpp)会做两类关键校验:

  • growthInterval取值范围校验为[1, INT32_MAX](即 [1, 2147483647]),非法值返回GRAPH_FAILED
  • 三个输入 Tensor 的存储 shape 均需为标量[1]EnsureNotScalar兼容维度数为 0 的标量表示),shapeSize 必须为 1。

Tiling 同时根据current_scale的数据类型设置 tiling key:FLOAT 对应0、FLOAT16 对应1、BF16 对应2(见 amp_update_scale_tiling.cpp 与 amp_update_scale_tiling.h 中定义的 TilingData 结构),并将growthFactorbackoffFactorgrowthInterval写入 TilingData,最后SetBlockDim(1)—— 由于所有数据均为标量,算子以单核方式执行。

Kernel 侧(amp_update_scale.h)的ComputeScaleUpdate直接实现了文档公式:

__aicore__ inline void ComputeScaleUpdate() { if (foundInf_) { currentScale_ *= backoffFactor_; // 发现 Inf/NaN:回退 growthTracker_ = 0; } else { successful_ = growthTracker_ + 1; if (successful_ == growthInterval_) { newScale_ = currentScale_ * growthFactor_; // 达到间隔:尝试增长 if (IsFinite(newScale_)) { // 溢出保护 currentScale_ = newScale_; } growthTracker_ = 0; } else { growthTracker_ = successful_; // 否则计数器 +1 } } }

其中IsFinite通过 IEEE 754 位级判断实现(amp_update_scale.h):取出 float32 的符号位屏蔽掩码0x7FFFFFFF后的指数位,若全部为 1(0xFF)则判定为 Inf/NaN,即非有限数。此外,针对 FLOAT16 输入会在读入后转 float 计算、BF16 输入使用Cast指令完成精度转换后计算,计算完成再转回原类型写回,保证中间计算精度(见LoadInputData/StoreOutputData)。Kernel 入口 amp_update_scale.cpp 根据 TilingKey 实例化AmpUpdateScale<float>/AmpUpdateScale<half>/AmpUpdateScale<bfloat16_t>三个模板分支。

返回值与错误码

aclnnStatus返回状态码的完整说明可参见 aclnn 返回码。第一段接口完成入参校验,出现以下场景时报错:

返回值错误码描述
ACLNN_ERR_INNER_TILING_ERROR561002输入currentScale、growthTracker、foundInf的shape不是标量[1]。
ACLNN_ERR_INNER_TILING_ERROR561002growthInterval超出取值范围[1, 2147483647]。
ACLNN_ERR_PARAM_NULLPTR161001传入的currentScale、growthTracker、foundInf、updatedScale、updatedGrowthTracker是空指针。
ACLNN_ERR_PARAM_INVALID161002currentScale的数据类型不在支持的范围之内。
ACLNN_ERR_PARAM_INVALID161002foundInf的数据类型与currentScale不一致。
ACLNN_ERR_PARAM_INVALID161002updatedScale的数据类型与currentScale不一致。
ACLNN_ERR_PARAM_INVALID161002growthTracker的数据类型不是INT32。

错误码 561002 所对应的两类校验(shape 非标量、growthInterval 越界)在 amp_update_scale_tiling.cpp 的Init中均有对应的OP_CHECK_IF检查实现,可作为排查问题时的对照依据。

aclnnAmpUpdateScale 参数说明

第二段接口的参数如下:

参数名输入/输出描述
workspace输入在Device侧申请的workspace内存地址。
workspaceSize输入在Device侧申请的workspace大小,由第一段接口aclnnAmpUpdateScaleGetWorkspaceSize获取。
executor输入op执行器,包含了算子计算流程。
stream输入指定执行任务的Stream。

返回值同样为aclnnStatus,参见 aclnn 返回码。

约束说明

使用该接口时需遵守以下约束:

  • 确定性计算aclnnAmpUpdateScale默认确定性实现。
  • 数据类型约束current_scalefound_inf的数据类型必须一致;updated_scale的数据类型必须与current_scale一致;growth_trackerupdated_growth_tracker必须为 INT32。
  • shape 约束:所有输入输出张量均为标量,shape 为[1]
  • growthInterval 约束growthInterval取值范围为[1, 2147483647]
  • Inf/NaN 优先级found_inf不为 0 时,直接执行回退逻辑,忽略growth_tracker状态。
  • 溢出保护:当 scale 增长后的新值溢出(inf/nan)时,保持当前 scale 不变,growth_tracker重置为 0。

调用示例与逐步解析

仓库在 examples/test_aclnn_amp_update_scale.cpp 提供了完整的可直接运行的 aclnn 调用样例,接口文档中亦给出了等价的完整示例代码。其运行流程分为"资源初始化 → 构造输入输出 → 第一段接口取 workspace 与执行器 → 申请 workspace → 第二段接口执行 → 同步 → 拷回结果 → 释放资源"八个步骤,完整代码如下:

#include <iostream> #include <vector> #include "acl/acl.h" #include "aclnnop/aclnn_amp_update_scale.h" #define CHECK_RET(cond, return_expr) \ do { \ if (!(cond)) { \ return_expr; \ } \ } while (0) #define LOG_PRINT(message, ...) \ do { \ printf(message, ##__VA_ARGS__); \ } while (0) int64_t GetShapeSize(const std::vector<int64_t>& shape) { int64_t shapeSize = 1; for (auto i : shape) { shapeSize *= i; } return shapeSize; } int Init(int32_t deviceId, aclrtStream* stream) { // 固定写法,资源初始化 auto ret = aclInit(nullptr); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclInit failed. ERROR: %d\n", ret); return ret); ret = aclrtSetDevice(deviceId); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSetDevice failed. ERROR: %d\n", ret); return ret); ret = aclrtCreateStream(stream); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtCreateStream failed. ERROR: %d\n", ret); return ret); return 0; } template <typename T> int CreateAclTensor(const std::vector<T>& hostData, const std::vector<int64_t>& shape, void** deviceAddr, aclDataType dataType, aclTensor** tensor) { auto size = GetShapeSize(shape) * sizeof(T); // 调用aclrtMalloc申请device侧内存 auto ret = aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtMalloc failed. ERROR: %d\n", ret); return ret); // 调用aclrtMemcpy将host侧数据拷贝到device侧内存上 ret = aclrtMemcpy(*deviceAddr, size, hostData.data(), size, ACL_MEMCPY_HOST_TO_DEVICE); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtMemcpy failed. ERROR: %d\n", ret); return ret); // 计算连续tensor的strides std::vector<int64_t> strides(shape.size(), 1); for (int64_t i = shape.size() - 2; i >= 0; i--) { strides[i] = shape[i + 1] * strides[i + 1]; } // 调用aclCreateTensor接口创建aclTensor *tensor = aclCreateTensor(shape.data(), shape.size(), dataType, strides.data(), 0, aclFormat::ACL_FORMAT_ND, shape.data(), shape.size(), *deviceAddr); return 0; } int main() { // 1.(固定写法)device/stream初始化,参考acl API手册 // 根据自己的实际device填写deviceId int32_t deviceId = 0; aclrtStream stream; auto ret = Init(deviceId, &stream); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("Init acl failed. ERROR: %d\n", ret); return ret); // 2. 构造输入与输出,需要根据API的接口自定义构造 std::vector<int64_t> scalarShape = {1}; void* currentScaleDeviceAddr = nullptr; void* growthTrackerDeviceAddr = nullptr; void* foundInfDeviceAddr = nullptr; void* updatedScaleDeviceAddr = nullptr; void* updatedGrowthTrackerDeviceAddr = nullptr; aclTensor* currentScale = nullptr; aclTensor* growthTracker = nullptr; aclTensor* foundInf = nullptr; aclTensor* updatedScale = nullptr; aclTensor* updatedGrowthTracker = nullptr; // 创建currentScale std::vector<float> currentScaleHost = {65536.0f}; ret = CreateAclTensor(currentScaleHost, scalarShape, &currentScaleDeviceAddr, ACL_FLOAT, &currentScale); CHECK_RET(ret == ACL_SUCCESS, return ret); // 创建growthTracker std::vector<int32_t> growthTrackerHost = {900}; ret = CreateAclTensor(growthTrackerHost, scalarShape, &growthTrackerDeviceAddr, ACL_INT32, &growthTracker); CHECK_RET(ret == ACL_SUCCESS, return ret); // 创建foundInf std::vector<float> foundInfHost = {0.0f}; ret = CreateAclTensor(foundInfHost, scalarShape, &foundInfDeviceAddr, ACL_FLOAT, &foundInf); CHECK_RET(ret == ACL_SUCCESS, return ret); // 创建输出updatedScale std::vector<float> updatedScaleHost = {0.0f}; ret = CreateAclTensor(updatedScaleHost, scalarShape, &updatedScaleDeviceAddr, ACL_FLOAT, &updatedScale); CHECK_RET(ret == ACL_SUCCESS, return ret); // 创建输出updatedGrowthTracker std::vector<int32_t> updatedGrowthTrackerHost = {0}; ret = CreateAclTensor(updatedGrowthTrackerHost, scalarShape, &updatedGrowthTrackerDeviceAddr, ACL_INT32, &updatedGrowthTracker); CHECK_RET(ret == ACL_SUCCESS, return ret); // 3. 调用第一段接口,获取workspace大小和执行器 float growthFactor = 2.0f; float backoffFactor = 0.5f; int64_t growthInterval = 1000; uint64_t workspaceSize = 0; aclOpExecutor* executor = nullptr; ret = aclnnAmpUpdateScaleGetWorkspaceSize(currentScale, growthTracker, foundInf, growthFactor, backoffFactor, growthInterval, updatedScale, updatedGrowthTracker, &workspaceSize, &executor); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnAmpUpdateScaleGetWorkspaceSize failed. ERROR: %d\n", ret); return ret); // 4. 根据workspaceSize申请workspace内存 void* workspace = nullptr; if (workspaceSize > 0) { ret = aclrtMalloc(&workspace, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("allocate workspace failed. ERROR: %d\n", ret); return ret); } // 5. 调用第二段接口,执行计算 ret = aclnnAmpUpdateScale(workspace, workspaceSize, executor, stream); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnAmpUpdateScale failed. ERROR: %d\n", ret); return ret); // 6. 同步等待任务执行结束 ret = aclrtSynchronizeStream(stream); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSynchronizeStream failed. ERROR: %d\n", ret); return ret); // 7. 将输出数据从device拷贝到host,并打印结果 float updatedScaleVal = 0.0f; int32_t updatedGrowthTrackerVal = 0; ret = aclrtMemcpy(&updatedScaleVal, sizeof(float), updatedScaleDeviceAddr, sizeof(float), ACL_MEMCPY_DEVICE_TO_HOST); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("copy updatedScale failed. ERROR: %d\n", ret); return ret); ret = aclrtMemcpy(&updatedGrowthTrackerVal, sizeof(int32_t), updatedGrowthTrackerDeviceAddr, sizeof(int32_t), ACL_MEMCPY_DEVICE_TO_HOST); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("copy updatedGrowthTracker failed. ERROR: %d\n", ret); return ret); LOG_PRINT("aclnnAmpUpdateScale result: updatedScale = %f, updatedGrowthTracker = %d\n", updatedScaleVal, updatedGrowthTrackerVal); // 8.(固定写法)释放资源 aclDestroyTensor(currentScale); aclDestroyTensor(growthTracker); aclDestroyTensor(foundInf); aclDestroyTensor(updatedScale); aclDestroyTensor(updatedGrowthTracker); aclrtFree(currentScaleDeviceAddr); aclrtFree(growthTrackerDeviceAddr); aclrtFree(foundInfDeviceAddr); aclrtFree(updatedScaleDeviceAddr); aclrtFree(updatedGrowthTrackerDeviceAddr); if (workspaceSize > 0) { aclrtFree(workspace); } aclrtDestroyStream(stream); auto aclRet = aclrtResetDevice(deviceId); CHECK_RET(aclRet == ACL_SUCCESS, LOG_PRINT("reset device failed. ERROR: %d\n", aclRet); return aclRet); aclRet = aclFinalize(); CHECK_RET(aclRet == ACL_SUCCESS, LOG_PRINT("finalize acl failed. ERROR: %d\n", aclRet); return aclRet); return 0; }

示例数据推演

示例中currentScale = 65536.0fgrowthTracker = 900foundInf = 0.0fgrowthFactor = 2.0fbackoffFactor = 0.5fgrowthInterval = 1000。由于found_inf == 0growth_tracker + 1 = 901 ≠ 1000,命中公式中的"其他情况"分支:updated_scale保持65536.0f不变,updated_growth_tracker变为901。若将growthTracker改为999,则满足growth_tracker + 1 == growth_interval,输出将变为updated_scale = 131072.0fupdated_growth_tracker = 0;若foundInf非 0,则无论计数器取值如何都会执行回退,输出updated_scale = 32768.0fupdated_growth_tracker = 0

编译与运行

示例的具体编译和执行过程可参考 编译与运行样例。代码本身遵循标准的 aclnn 两段式接口调用范式:aclInit/aclrtSetDevice/aclrtCreateStream初始化资源,aclCreateTensor构造 ND 格式的标量张量,第一段接口返回workspaceSizeexecutor,第二段接口在指定stream上提交执行,最后通过aclrtMemcpy将输出拷回 host 并打印,随后依次销毁 tensor、释放 device 内存与 stream。

单测覆盖

仓库为算子提供了 Tiling 层单测(tests/ut/op_host/test_amp_update_scale_tiling.cpp),覆盖 FLOAT(tiling key 0)、FLOAT16(tiling key 1)、BF16(tiling key 2)三种数据类型,分别以growth_interval为 5、3、10 构造标量[1]输入,断言 Tiling 返回GRAPH_SUCCESS、预期的 tiling key 以及workspace大小为 0(单标量计算无需额外 workspace)。此外 Kernel 侧单测见 tests/ut/op_kernel/test_amp_update_scale.cpp,可用于在算子 UT 框架下验证回退、增长与溢出保护等分支行为。

小结

aclnnAmpUpdateScale是一个输入输出均为标量的轻量算子,核心价值在于把 AMP 训练中"周期性增长、溢出回退、计数器维护"这套状态更新逻辑固化到 NPU 侧,避免了在训练框架中逐 step 用 host 侧逻辑拼接多个原子算子。理解本文的公式规则、两段式接口参数、错误码语义与源码实现(算子定义 → Tiling → Kernel),即可在自有混合精度训练流程中正确接入并排查问题。若需进一步了解两段式接口的通用规范、aclnn 返回码含义或示例工程编译方式,可分别查阅 两段式接口说明、aclnn 返回码 与 编译与运行样例。

【免费下载链接】ops-math本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-math

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询