CANN ops-math 算子详解:aclnnAmpUpdateScale 动态损失缩放接口的原理、两段式调用与源码剖析
CANN ops-math 算子详解aclnnAmpUpdateScale 动态损失缩放接口的原理、两段式调用与源码剖析【免费下载链接】ops-math本项目是CANN提供的数学类基础计算算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-mathaclnnAmpUpdateScale是 CANN ops-math 数学算子库math/amp_update_scale提供的 AMPAutomatic 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/NaNgrowth_tracker连续未出现 Inf/NaN 的步数计数器growth_factorscale 增长因子通常设置为 2.0backoff_factorscale 回退因子通常设置为 0.5growth_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)非连续TensorcurrentScaleaclTensor*输入当前的loss scale值shape为标量 [1]FLOAT、FLOAT16、BFLOAT16ND1×growthTrackeraclTensor*输入连续未出现Inf/NaN的步数计数器shape为标量 [1]INT32ND1×foundInfaclTensor*输入是否检测到Inf/NaN的标志shape为标量 [1]。0表示正常非0表示发现Inf/NaN。数据类型需要与currentScale一致FLOAT、FLOAT16、BFLOAT16ND1×growthFactorfloat输入scale增长因子当连续growth_interval步未检测到Inf/NaN时scale将乘以该因子。通常设置为2.0----backoffFactorfloat输入scale回退因子当检测到Inf/NaN时scale将乘以该因子。通常设置为0.5----growthIntervalint64_t输入触发scale增长的间隔步数即连续多少步未检测到Inf/NaN后将增大scale。取值范围 1----updatedScaleaclTensor*输出更新后的loss scale值shape为 [1]。数据类型需要与currentScale的数据类型一致FLOAT、FLOAT16、BFLOAT16ND1×updatedGrowthTrackeraclTensor*输出更新后的growth tracker计数器shape为 [1]INT32ND1×workspaceSizeuint64_t*输出返回需要在Device侧申请的workspace大小-----executoraclOpExecutor**输出返回op执行器包含了算子计算流程-----参数语义与源码的对应关系上述参数在算子注册层面与 amp_update_scale_def.cpp 一一对应current_scale、growth_tracker、found_inf三个输入updated_scale、updated_growth_tracker两个输出以及growth_factorFLOAT 属性、backoff_factorFLOAT 属性、growth_intervalINT 属性三个必填属性。其中current_scale/found_inf/updated_scale支持ge::DT_FLOAT、ge::DT_FLOAT16、ge::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 keyFLOAT 对应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后的指数位若全部为 10xFF则判定为 Inf/NaN即非有限数。此外针对 FLOAT16 输入会在读入后转 float 计算、BF16 输入使用Cast指令完成精度转换后计算计算完成再转回原类型写回保证中间计算精度见LoadInputData/StoreOutputData。Kernel 入口 amp_update_scale.cpp 根据 TilingKey 实例化AmpUpdateScalefloat/AmpUpdateScalehalf/AmpUpdateScalebfloat16_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_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::vectorint64_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::vectorT hostData, const std::vectorint64_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::vectorint64_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::vectorint64_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::vectorfloat currentScaleHost {65536.0f}; ret CreateAclTensor(currentScaleHost, scalarShape, currentScaleDeviceAddr, ACL_FLOAT, currentScale); CHECK_RET(ret ACL_SUCCESS, return ret); // 创建growthTracker std::vectorint32_t growthTrackerHost {900}; ret CreateAclTensor(growthTrackerHost, scalarShape, growthTrackerDeviceAddr, ACL_INT32, growthTracker); CHECK_RET(ret ACL_SUCCESS, return ret); // 创建foundInf std::vectorfloat foundInfHost {0.0f}; ret CreateAclTensor(foundInfHost, scalarShape, foundInfDeviceAddr, ACL_FLOAT, foundInf); CHECK_RET(ret ACL_SUCCESS, return ret); // 创建输出updatedScale std::vectorfloat updatedScaleHost {0.0f}; ret CreateAclTensor(updatedScaleHost, scalarShape, updatedScaleDeviceAddr, ACL_FLOAT, updatedScale); CHECK_RET(ret ACL_SUCCESS, return ret); // 创建输出updatedGrowthTracker std::vectorint32_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覆盖 FLOATtiling key 0、FLOAT16tiling key 1、BF16tiling 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),仅供参考