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_factor | scale 增长因子,通常设置为 2.0 |
backoff_factor | scale 回退因子,通常设置为 0.5 |
growth_interval | 触发 scale 增长的间隔步数 |
规则可以概括为四点:
- 当
found_inf不为 0 时,scale乘以backoff_factor回退,growth_tracker重置为 0; - 当
found_inf为 0 且growth_tracker + 1等于growth_interval时,scale乘以growth_factor增长; - 如果增长后的新
scale溢出(inf/nan),则保持当前scale不变,growth_tracker重置为 0(溢出保护); - 其他情况下,
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()仅为ascend910b、ascend910_93、ascend950三个平台注册了 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、BFLOAT16 | ND | 1 | × |
| growthTracker(aclTensor*) | 输入 | 连续未出现Inf/NaN的步数计数器 | shape为标量 [1] | INT32 | ND | 1 | × |
| foundInf(aclTensor*) | 输入 | 是否检测到Inf/NaN的标志 | shape为标量 [1]。0表示正常,非0表示发现Inf/NaN。数据类型需要与currentScale一致 | FLOAT、FLOAT16、BFLOAT16 | ND | 1 | × |
| 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、BFLOAT16 | ND | 1 | × |
| updatedGrowthTracker(aclTensor*) | 输出 | 更新后的growth tracker计数器 | shape为 [1] | INT32 | ND | 1 | × |
| workspaceSize(uint64_t*) | 输出 | 返回需要在Device侧申请的workspace大小 | - | - | - | - | - |
| executor(aclOpExecutor**) | 输出 | 返回op执行器,包含了算子计算流程 | - | - | - | - | - |
参数语义与源码的对应关系
上述参数在算子注册层面与 amp_update_scale_def.cpp 一一对应:current_scale、growth_tracker、found_inf三个输入,updated_scale、updated_growth_tracker两个输出,以及growth_factor(FLOAT 属性)、backoff_factor(FLOAT 属性)、growth_interval(INT 属性)三个必填属性。其中current_scale/found_inf/updated_scale支持ge::DT_FLOAT、ge::DT_FLOAT16、ge::DT_BF16,growth_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 结构),并将growthFactor、backoffFactor、growthInterval写入 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_ERROR | 561002 | 输入currentScale、growthTracker、foundInf的shape不是标量[1]。 |
| ACLNN_ERR_INNER_TILING_ERROR | 561002 | growthInterval超出取值范围[1, 2147483647]。 |
| ACLNN_ERR_PARAM_NULLPTR | 161001 | 传入的currentScale、growthTracker、foundInf、updatedScale、updatedGrowthTracker是空指针。 |
| ACLNN_ERR_PARAM_INVALID | 161002 | currentScale的数据类型不在支持的范围之内。 |
| ACLNN_ERR_PARAM_INVALID | 161002 | foundInf的数据类型与currentScale不一致。 |
| ACLNN_ERR_PARAM_INVALID | 161002 | updatedScale的数据类型与currentScale不一致。 |
| ACLNN_ERR_PARAM_INVALID | 161002 | growthTracker的数据类型不是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_scale与found_inf的数据类型必须一致;updated_scale的数据类型必须与current_scale一致;growth_tracker与updated_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, ¤tScaleDeviceAddr, ACL_FLOAT, ¤tScale); 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.0f、growthTracker = 900、foundInf = 0.0f、growthFactor = 2.0f、backoffFactor = 0.5f、growthInterval = 1000。由于found_inf == 0且growth_tracker + 1 = 901 ≠ 1000,命中公式中的"其他情况"分支:updated_scale保持65536.0f不变,updated_growth_tracker变为901。若将growthTracker改为999,则满足growth_tracker + 1 == growth_interval,输出将变为updated_scale = 131072.0f、updated_growth_tracker = 0;若foundInf非 0,则无论计数器取值如何都会执行回退,输出updated_scale = 32768.0f、updated_growth_tracker = 0。
编译与运行
示例的具体编译和执行过程可参考 编译与运行样例。代码本身遵循标准的 aclnn 两段式接口调用范式:aclInit/aclrtSetDevice/aclrtCreateStream初始化资源,aclCreateTensor构造 ND 格式的标量张量,第一段接口返回workspaceSize与executor,第二段接口在指定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),仅供参考