CANN ops-math ReduceStdWithMean 算子深度解析:从 torch.std_mean 到 BatchNorm 统计量的 Two-Pass 归约实现

发布时间:2026/9/19 17:48:41
CANN ops-math ReduceStdWithMean 算子深度解析:从 torch.std_mean 到 BatchNorm 统计量的 Two-Pass 归约实现 算子库人工智能CANN【免费下载链接】ops-math本项目是CANN提供的数学类基础计算算子库实现网络在NPU上加速计算。项目地址https://gitcode.com/cann/ops-math点击查看免费下载本文围绕 CANN ops-math 仓库中 experimental/math/reduce_std_with_mean/README.md 展开系统讲解 ReduceStdWithMean 算子的功能语义、L0 kernel 计算公式、参数约束、aclnn 两段式调用方式以及构建、UT、ST 测试的完整流程。读者读完将掌握该算子在 NPU 上的 Two-Pass 归约实现原理能够独立完成自定义算子包构建、接口调用与示例代码运行。一、产品支持情况与适用平台ReduceStdWithMean 是 CANN ops-math 仓库experimental目录下的实验性数学算子当前支持的产品范围如下产品是否支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品√该产品的数据类型支持 FLOAT、FLOAT16、BFLOAT16共 3 类并且输入、输出数据类型必须一致这是非 RegBase 平台的约束Ascend 950 等 RegBase 平台允许输入输出类型不一致见 aclnn_std_mean_correction.cpp 中CheckDtypeValid的实现。算子配置仅注册了ascend910b见 reduce_std_with_mean_def.cpp 中的AICore().AddConfig(ascend910b)。二、功能说明一个 L0 算子两个 L2 场景算子功能沿给定维度dim归约计算标准差和均值。仓库提供了两个 L2 API 覆盖不同场景aclnnStdMeanCorrection对应torch.std_mean计算标准差 均值支持 Bessel 修正aclnnBatchNormStats对应 BatchNorm 统计量计算均值 标准差倒数1/sqrt(vareps)。两个 L2 API 通过一个统一的 L0 kernelReduceStdWithMean完成核心计算invert属性控制最终输出是std还是1/std。2.1 L2 API 内部管线以aclnnStdMeanCorrection为例L2 API 内部的算子调度管线如下self ──→ ReduceMean ──→ mean ──→ Broadcast ──→ mean(expanded) │ │ └──→ meanOut └──→ ReduceStdWithMean(self, mean) ──→ stdOutReduceMean和ReduceStdWithMean是 L0 kernel。ReduceStdWithMean接收预计算的mean已 broadcast 到self同 shape计算diff self - mean进而得出方差和标准差。预计算 mean 避免了在 kernel 内重复计算均值——这是 Two-Pass 算法的设计意图而非使用限制。2.2 L0 kernel 计算公式ReduceStdWithMeanL0 kernel 的计算公式$$ \begin{aligned} \text{diff} \text{self} - \text{mean} \ \text{sum_sqr} \sum (\text{diff})^2 \ \text{var} \frac{\text{sum_sqr}}{\max(0,\ N - \text{correction})} \ \text{output} \begin{cases} \sqrt{\text{var}}, \text{invert} \text{false} \[6pt] \dfrac{1}{\sqrt{\text{var} \text{eps}}}, \text{invert} \text{true} \end{cases} \end{aligned} $$其中N为归约维度元素总数。上述公式在 reduce_std_with_mean_kernel.h 的Process()中逐行对应实现先逐 tile 累加sum_sqr再计算denom reduceLength - correction负值 clamp 为 0最后依据invert_输出sqrt(var)或1/sqrt(var eps)。三、参数说明3.1 L0 kernel / L2 API 参数总览参数名输入/输出/属性描述数据类型数据格式self (input)输入输入张量。对应公式中 self。支持 1-8 维。aclnnStdMeanCorrection 使用参数名 selfaclnnBatchNormStats 使用参数名 input。FLOAT、FLOAT16、BFLOAT16NDmean输入L0 kernel 的输入L2 API 内部由 ReduceMean 自动计算并传入调用方无需关心。为预计算的均值张量已 broadcast 到与 self 同 shape。与 self 保持一致NDdim属性归约维度。支持单维或多维归约。L2 API 通过 Transpose 自动处理非连续归约维度。INT64-correction属性Bessel 修正值。0 表示总体方差除以 N 1 表示样本方差除以 N - correction。当 correction 超过归约维度元素数时输出结果 clamp 为 0。INT64-keepdim属性是否保留归约维度。true 时输出与输入同维数归约维度长度为 1false 时压缩归约维度。仅 aclnnStdMeanCorrection 使用。BOOL-invert属性输出控制。false 输出标准差stdtrue 输出标准差倒数1/std。仅 L0 API 使用aclnnStdMeanCorrection 固定输出 stdaclnnBatchNormStats 固定输出 1/std。BOOL-eps属性数值稳定性常数加在方差上再开根号避免除零。aclnnStdMeanCorrection 使用 FLOATaclnnBatchNormStats 使用 DOUBLE。FLOAT / DOUBLE-stdOut (invstdOut)输出标准差结果。aclnnStdMeanCorrection 输出标准差std参数名为 stdOutaclnnBatchNormStats 输出标准差倒数invstd参数名为 invstdOut。shape 取决于 keepdim 设置aclnnStdMeanCorrection或固定不保留维度aclnnBatchNormStats。与 self 保持一致NDmeanOut输出均值结果。aclnnStdMeanCorrection 输出归约后的均值aclnnBatchNormStats 输出沿 batch 维归约的均值。shape 取决于 keepdim / 归约维度设置。与 self 保持一致ND3.2 从算子定义看属性次序L0 算子ReduceStdWithMean的属性声明顺序见 reduce_std_with_mean_def.cppdimAttrType OPTIONALListIntunbiasedAttrType OPTIONALBoolkeepdimAttrType OPTIONALBoolinvertAttrType OPTIONALBoolepsAttrType OPTIONALFloatcorrectionAttrType OPTIONALIntHost 端 tiling 在 reduce_std_with_mean_tiling.cpp 的GetShapeAttrsInfo中按索引读取GetInt(5)读 correction、GetBool(3)读 invert、GetFloat(4)读 eps缺省时分别默认为 0、false、0.0f。dim 为空时 tiling 默认归约最后一维reduceLen selfShape.GetDim(rank - 1)。3.3 两个 L2 API 的差异对照差异点aclnnStdMeanCorrectionaclnnBatchNormStats对应语义torch.std_mean(x, dim, correction, keepdim)BatchNorm 统计量归约维度用户通过 dim 指定可多维固定除 channel第 1 维外的所有维度correction用户传入默认场景为 1固定为 0keepdim用户控制固定 falseinvert固定 false输出 std固定 true输出 1/stdepsFLOAT 类型默认 0.0fDOUBLE 类型由调用方传入输出 dtype与 self 一致mean / invstd 固定 FLOAT值得注意的是aclnnStdMeanCorrection的 L2 实现中invert与eps被声明为编译期常量static const bool invert false; static const float eps 0.0f;见 aclnn_std_mean_correction.cpp两个输出均为原始 dtype无需额外 Cast。四、约束说明self(input) 的维度数rank必须在 1 到 8 之间。对于 L2 APImean由内部ReduceMean自动计算和 broadcast调用方无需提供对于 L0 APImean必须与self同 shape、同 dtype由调用方预计算并传入。支持多维归约归约维度通过dim参数指定。非连续归约维度L2 API 通过 Transpose 自动将非连续归约维度移到末尾后再调度 kernel调用方无需手动 transpose。correction 必须为非负整数。当 correction 归约维度元素数时方差为 0输出 0std或 eps 保护下的 1/sqrt(eps)invstd不会除零崩溃。FLOAT16 / BFLOAT16 中间计算全部升精度到 FLOAT32 执行Cast → Sub → Mul → ReduceSum规避半精度中间累加精度损失。仅支持 ND 格式。确定性说明Pre-computed Mean Two-Pass 统一算法路径invert 参数仅影响最终输出选择sqrt 或 1/sqrt核心计算路径完全相同。默认确定性实现相同输入恒产生相同输出。4.1 边界情况的源码级处理在 L2 层面对边界 shape 做了显式处理见 aclnn_std_mean_correction.cpp空 tensorself-IsEmpty()时对 stdOut、meanOut 填充 NANDealEmptymeanOut 不覆盖。归约元素数 correction当shapeProd 1且shapeProd correction时 stdOut 填充 NAN当correction 1且shapeProd correction时 stdOut 填充 INFINITYmeanOut 通过ViewCopy正常输出。BatchNormStats 空输入mean 填 0、invstd 填quiet_NaN见 aclnn_batch_norm_stats.cpp 的ProcessEmptyTensorWithValue。4.2 L2 层的参数校验aclnnStdMeanCorrection的CheckParams依次执行四类校验空指针检查self、stdOut、meanOut 不能为空dtype 校验输入输出必须在 FLOAT/FLOAT16/BF16 支持列表内非 RegBase 平台要求三者一致dim 校验CheckDimValiddim 可为空等价全归约允许负数currentDim selfDimNum归一化不允许越界与重复shape 校验CheckShapestdOut、meanOut 必须严格匹配由 dim keepdim 推导出的归约后 shape。aclnnBatchNormStats的校验略有不同mean、invstd 必须是一维 tensor 且长度等于 channel 数input-GetViewShape()[1]channel 数不能为 0格式不允许私有格式仅支持 ND、NCL、NCHW、NCDHW。五、调用说明本算子通过 aclnn 两段式接口单算子调用。调用方式调用样例说明aclnn 调用StdMeanCorrectiontest_aclnn_std_mean_correction通过 aclnn 两段式接口aclnnStdMeanCorrectionGetWorkspaceSizeaclnnStdMeanCorrection声明见 aclnn_std_mean_correction_experimental.h由自定义算子包custom_math导出计算标准差 均值。对应 PyTorch 的torch.std_mean(x, dim, correctioncorrection, keepdimkeepdim)。aclnn 调用BatchNormStatstest_aclnn_batch_norm_status通过 aclnn 两段式接口aclnnBatchNormStatsGetWorkspaceSizeaclnnBatchNormStats声明见 aclnn_batch_norm_stats_experimental.h计算 BatchNorm 统计量均值 标准差倒数。图模式调用-暂不支持。本算子仅提供 aclnn 接口。5.1 两段式接口的调用模板以 test_aclnn_std_mean_correction.cpp 为例完整调用流程如下// 1. 初始化aclInit aclrtSetDevice aclrtCreateStream // 2. 构造输入输出 aclTensoraclCreateTensor格式 ACL_FORMAT_ND std::vectorint64_t selfShape {2, 3, 4}; std::vectorint64_t stdOutShape {2, 4}; // dim1, keepdimfalse std::vectorint64_t meanOutShape {2, 4}; // 3. 第一段接口计算 workspace 大小 uint64_t workspaceSize 0; aclOpExecutor* executor; ret aclnnStdMeanCorrectionGetWorkspaceSize(self, dim, correction, keepdim, stdOut, meanOut, workspaceSize, executor); // 4. 按 workspaceSize 申请 device 内存 void* workspaceAddr nullptr; if (workspaceSize 0) { ret aclrtMalloc(workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); } // 5. 第二段接口执行计算 ret aclnnStdMeanCorrection(workspaceAddr, workspaceSize, executor, stream); // 6. aclrtSynchronizeStream 同步拷贝结果回 host 并打印 // 7. 释放 aclTensor、workspace、streamaclFinalize示例中self为{2,3,4}的 FLOAT 张量dim {1}correction 1keepdim false输出 shape 为{2,4}——与 InferShape 按 keepdim 折叠/保留归约维度的规则一致。5.2 L2 API 内部实现调用链aclnnStdMeanCorrectionGetWorkspaceSize的内部执行序列非 RegBase 平台源码见 aclnn_std_mean_correction.cppContiguousReFormat(FORMAT_ND)统一输入为连续 ND 张量l0op::ReduceMean(self, dimArray, keepdim)在原始 dtype 上计算均值kernel 内部升 fp32避免引入 fp32 ReduceMean kernel 的二进制依赖边界分支处理NAN/INF 填充见上文GetExpandMeankeepdimfalse 时先UnsqueezeNd还原归约维度再l0op::Expandbroadcast 到 self 同 shape若归约维度不连续非最后若干维构造 perm 将非归约维前置、归约维移到最后对 self 与 mean 分别执行TransposeContiguous并同步改写dimForKernel调度l0op::ReduceStdWithMean(selfForKernel, meanForKernel, dimForKernel, correction, keepdim, invert, eps)keepdimtrue 且发生过 Transpose 时用invPerm对输出做逆 Transpose 还原维度位置两个输出经ViewCopy写入 stdOut、meanOut兼容非连续输出 tensor。aclnnBatchNormStats的管线与之类似但略有差异见 aclnn_batch_norm_stats.cpp输入先Cast到 FLOAT归约维度固定为除 channel 外的所有轴mean 用ReduceMean计算后经UnsqueezeNdBroadcastTo还原L0 kernel 以correction0, keepdimfalse, inverttrue, eps调用直接得到 invstd。在 RegBaseAscend 950平台上BatchNormStats 走统一路径aclnnBatchNormStatsImplUnify改用ReduceVarAdd(eps)Pow(-0.5)实现1/sqrt(vareps)。六、编译运行6.1 构建自定义算子包# 在仓库根目录source CANN/set_env.sh 后 source /usr/local/Ascend/cann-8.5.1/set_env.sh bash build.sh --pkg --experimental --socascend910b --opsreduce_std_with_mean -j16 bash build_out/cann-ops-math-custom_linux-*.run --install-path/usr/local/Ascend/cann-8.5.1 # 同步到 opp 目录必须 V/usr/local/Ascend/cann-8.5.1/vendors/custom_math O/usr/local/Ascend/cann-8.5.1/opp/vendors/custom_math cp $V/op_api/lib/libcust_opapi.so $O/op_api/lib/ patchelf --add-needed libopapi_math.so $O/op_api/lib/libcust_opapi.so cp $V/op_impl/ai_core/tbe/op_tiling/liboptiling.so $O/op_impl/ai_core/tbe/op_tiling/ rm -rf $O/op_impl/ai_core/tbe/kernel/ascend910b/reduce_std_with_mean cp -a $V/op_impl/ai_core/tbe/kernel/ascend910b/reduce_std_with_mean $O/op_impl/ai_core/tbe/kernel/ascend910b/ # 加载环境 source /usr/local/Ascend/cann-8.5.1/vendors/custom_math/bin/set_env.bash注意--opsreduce_std_with_mean将算子限定在单算子包构建范围--experimental与--socascend910b与算子定义中的 AICore 配置一致同步libcust_opapi.so、liboptiling.so与 kernel 二进制到 opp 目录是 aclnn 接口可被外部程序找到的前提。6.2 UT 测试# 在仓库根目录执行 bash build.sh -u --ophost --opapi --opkernel --opsreduce_std_with_mean --experimental -j16UT 覆盖三层op_host 侧的 test_reduce_std_with_mean_infershape.cpp 与 test_reduce_std_with_mean_tiling.cppshape 推导与 tiling 参数校验、op_kernel 侧的 test_reduce_std_with_mean.cpp含 reduce_std_with_mean_ut_tiling_data.h 构造的 tiling 数据、op_api 侧的 test_aclnn_std_mean_correction.cpp 与 test_aclnn_batch_norm_stats.cpp。6.3 ST 测试ATKsource /usr/local/Ascend/cann-8.5.1/set_env.sh source /usr/local/Ascend/cann-8.5.1/vendors/custom_math/bin/set_env.bash # StdMeanCorrection ATK atk aclnn --task accuracy --devices 0 \ experimental/math/reduce_std_with_mean/tests/st/aclnnStdMeanCorrection/atk_aclnnStdMeanCorrection.json # BatchNormStats ATK atk aclnn --task accuracy --devices 0 \ -p experimental/math/reduce_std_with_mean/tests/st/aclnnBatchNormStats/executor_aclnnBatchNormStats.py \ experimental/math/reduce_std_with_mean/tests/st/aclnnBatchNormStats/atk_aclnnBatchNormStats.jsonATK 用例的配置 JSON 与 executor 脚本位于 tests/st 目录其中atk_aclnnBatchNormStats.json定义了精度测试用例与比对阈值executor_aclnnBatchNormStats.py负责数据生成与结果校验。6.4 示例代码编译运行# 在仓库根目录执行需先完成算子包构建和安装 bash build.sh --run_example reduce_std_with_mean eager cust --experimental --vendor_namecustom --socascend910b七、L0 kernel 实现原理7.1 入口与模板调度kernel 入口 reduce_std_with_mean.cpp 通过REGISTER_TILING_DEFAULT注册 tiling 数据结构GET_TILING_DATA_WITH_STRUCT反序列化随后实例化NsReduceStdWithMean::ReduceStdWithMeanDTYPE_SELF, schMode。schMode由 TilingKey 区分三种数据类型见 reduce_std_with_mean_tiling_key.hREDUCE_STD_SCH_FP16 0REDUCE_STD_SCH_FP32 1REDUCE_STD_SCH_BF16 2Host 端 tiling 在 reduce_std_with_mean_tiling.cpp 中按输入 dtype 选择对应 TilingKey 并调用SetTilingKey。7.2 Two-Pass 核心计算路径kernel 类 reduce_std_with_mean_kernel.h 中的Process()实现统一 Two-Pass 算法StdMeanCorrection 与 BatchNormStats 共用仅 invert 不同外层遍历当前核负责的非归约切片mcoreStartM_到coreStartM_ coreM_内层按ubLength将归约维度切分为多个 tile每个 tileDataCopyPad从 GM 搬运 self 与 mean 到 UBFLOAT 路径直接Sub(tBuf, sLocal, mLocal)→Mul(wBuf, tBuf, tBuf)→ReduceSum(dBuf, wBuf, tBuf)FP16/BF16 路径先Cast到 FLOAT32再执行同样的 Sub/Mul/ReduceSum每 tile 的ReduceSum结果累加到sum_sqr计算denom reduceLength - correction负值置 0var sum_sqr / denominvert为 false 输出sqrt(var)为 true 输出1/sqrt(var eps)stdTmp 为 0 时输出 0 保护WriteOutput将结果写回FLOAT 直接SetValueFP16/BF16 用CAST_ROUND舍入回半精度。Buffer 布局self 与 mean 双缓冲BUFFER_NUM 2另含 fp32 的 reduceBuf、scratchBuf、workBuf以及非 FLOAT 类型的 castBuf。归约维度在 L2 层已通过 Transpose 保证连续因此 tiling 无需维护 srcStride见 reduce_std_with_mean_tiling_data.h 中的注释说明。7.3 多核切分与 UB 预算Tiling 策略reduce_std_with_mean_tiling.cpp多核切分沿非归约维blockFactor CeilAlign(CeilDiv(nonReduceNum, coreNum), ubBlockSize)每个核处理至多blockFactor个非归约切片usedCoreNum CeilDiv(nonReduceNum, blockFactor)作为实际 block dimUB tile 沿归约维per-element UB 占用为2 * BUFFER_NUM * typeSizeself mean 输入队列2 * sizeof(float)scratch work 非 FLOAT 时的sizeof(float)cast 缓冲再预留UB_RESERVED 8 * 1024字节ubLength由可用 UB 与 perElemUB 相除并对齐得到且不超过 reduceLen空输入totalNum 或 reduceLen 为 0 时 block dim 置 1仅设置 TilingKey 后直接返回workspaceL0 kernel 不使用额外 workspaceSetWorkspace将 workspace 大小置 0。7.4 InferShape 与图模式说明L0 算子的 shape 推导注册在 reduce_std_with_mean_infershape.cpp通过IMPL_OP_INFERSHAPE(ReduceStdWithMean)挂载dim 为空时视为全归约axes 所有维依据 keepdim 分别调用框架的ReduceDimsWithKeepDims/ReduceDimsWithoutKeepDims推导输出 shape。虽然 L0 层注册了 InferShape但本算子的对外发布形态仅限 aclnn 单算子接口图模式Graph调用暂不支持。八、总结ReduceStdWithMean 以「预计算 mean 的 Two-Pass 算法」为统一内核通过invert与eps属性弹性输出 std 或1/sqrt(vareps)从而同时服务torch.std_mean与 BatchNorm 统计量两类场景。其工程实现值得借鉴的要点包括半精度中间计算统一升 FP32 规避精度损失、L2 层 Transpose 将非连续归约维搬移到末位以保证 kernel 访存连续性、按 dtype 模板化的 TilingKey 调度以及 correction 边界下的 NAN/INF/0 输出保护。开发者可按上文流程构建custom_math算子包并运行示例、UT 与 ATK 用例进一步验证数值与性能表现。赞分享算子库人工智能CANN【免费下载链接】ops-math本项目是CANN提供的数学类基础计算算子库实现网络在NPU上加速计算。项目地址https://gitcode.com/cann/ops-math点击查看免费下载相关推荐CANN ops-math 算子 aclnnGcd 逐元素最大公约数计算实现深度解析CANN ops math 算子 aclnnGcd 逐元素最大公约数计算实现深度解析 aclnnGcd 是 CANN ops math 仓库在 experime算子库人工智能CANNCANN ops-math Triu 算子从计算公式、参数约束到 NPU 分块实现CANN ops math Triu 算子从计算公式、参数约束到 NPU 分块实现 本文围绕 CANN ops math 仓库中 conversion/tri算子库人工智能CANNCANN ops-math ChunkCat 算子深度解析从 aclnnChunkCat 接口到 NPU 内核实现CANN ops math ChunkCat 算子深度解析从 aclnnChunkCat 接口到 NPU 内核实现 本篇技术指南围绕 CANN ops mat算子库人工智能CANN创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考