CANN ops-math FusedMulAdd 算子详解:Mul+Add 子图融合的广播计算实现与调用实践

📅 发布时间:2026/9/20 12:49:27
CANN ops-math FusedMulAdd 算子详解:Mul+Add 子图融合的广播计算实现与调用实践
算子库人工智能CANN【免费下载链接】ops-math本项目是CANN提供的数学类基础计算算子库实现网络在NPU上加速计算。项目地址https://gitcode.com/cann/ops-math点击查看免费下载FusedMulAdd 是 CANN ops-math 数学算子库中一类典型的子图融合算子它将神经网络中频繁出现的Mul与Add两个逐元素算子合并为单个算子按 NumPy 广播规则对齐三个输入后一次性完成y x1 * x2 x3的计算从而减少算子下发与访存开销。本文以 math/fused_mul_add/README.md 为核心结合仓库内算子原型、Host 侧定义与 Tiling、Kernel 侧 DAG 描述及端到端示例源码完整解析该算子的产品支持范围、参数与约束、分 dtype 的实现通路、广播调度机制与图模式调用方式帮助读者既会用也能理解其底层为什么这样实现。功能定位为什么要做 Mul Add 融合在 TensorFlow、PyTorch 等框架的图模式网络中y x1 * x2 x3这类乘加结构非常常见例如归一化、残差连接、融合激活前的线性变换等。若按常规实现需要先执行Mul算子产出中间结果再执行Add算子两次算子下发、两次全局内存读写。FusedMulAdd 的思路是把这两个逐元素算子合并为一个算子对x1、x2、x3三个输入按 NumPy 广播规则对齐后在一个 Kernel 内完成乘加计算并直接写出结果。这在 op_graph/fused_mul_add_proto.h 的算子注释中也被明确标注为与 TensorFlow/PyTorch 的 Mul 后接 Add 图融合兼容Compatible with TensorFlow/PyTorch graph fusion of Mul followed by Add。产品支持情况当前仓库中的 FusedMulAdd 算子对应 Host 侧 fused_mul_add_def.cpp 中注册的ascend950配置支持以下产品其余产品暂不支持产品是否支持Ascend 950PR / Ascend 950DT√Atlas A3 训练系列产品 / Atlas A3 推理系列产品√Atlas A2 训练系列产品 / Atlas A2 推理系列产品√Atlas 200I/500 A2 推理产品×Atlas 推理系列产品×Atlas 训练系列产品×从实现细节看该算子的 Kernel 与 Tiling 均位于arch35目录如 fused_mul_add_dag.h、fused_mul_add_tiling_arch35.cpp代码注释中亦直接提及 Ascend 950 实测场景与上表的支持矩阵一致。需要说明的是产品支持矩阵以当前仓库注册与目录结构为准若需在 Atlas A3/A2 之外的平台使用应核实所发布 CANN 版本对应的支持清单。功能与计算公式算子功能将Mul、Add子图融合为单个算子对三个输入按 NumPy 广播规则对齐后逐元素计算乘加。计算公式$$ y x_1 \times x_2 x_3 $$其中计算顺序固定为先乘后加。x1 * x2的结果作为中间值参与x3的加法最终输出y。参数说明算子的四个张量参数三输入一输出如下表所示。注意x2的 shape 需与x1可广播x3的 shape 需与x1*x2的结果可广播参数名输入/输出/属性描述数据类型数据格式x1输入公式中的乘法输入张量 x1。FLOAT16, FLOAT, INT32NDx2输入公式中的乘法输入张量 x2shape 需与 x1 可广播。同 x1NDx3输入公式中的加法输入张量 x3shape 需与 x1*x2 的结果可广播。同 x1NDy输出公式中的输出张量 yshape 为 x1、x2、x3 广播后的统一形状。同 x1ND以上参数定义可在 fused_mul_add_proto.h 中找到一一对应的注册REG_OP(FusedMulAdd)声明x1/x2/x3三个INPUT与y一个OUTPUT类型均为TensorType({DT_FLOAT16, DT_FLOAT, DT_INT32})。Host 侧 fused_mul_add_def.cpp 中进一步约束三个输入与输出均为ParamType(REQUIRED)即全部为必选参数不存在可选输入数据格式为FORMAT_ND并显式设置了UnknownShapeFormat({FORMAT_ND, FORMAT_ND, FORMAT_ND})保证动态 shape 场景下仍以 ND 格式处理DynamicCompileStaticFlag(true)表示动态 shape 场景按静态编译处理、DynamicFormatFlag(false)不启用动态格式、DynamicRankSupportFlag(true)与DynamicShapeSupportFlag(true)开启动态 rank 与动态 shape 支持、PrecisionReduceFlag(true)允许精度收敛优化ExtendCfgInfo(opFile.value, fused_mul_add_apt)将算子 Kernel 入口指向 fused_mul_add_apt.cpp。约束说明dtype 一致性x1、x2、x3、y 必须为同一种数据类型不支持混合数据类型。这一约束在 Tiling 阶段被再次校验CheckDtype()见 fused_mul_add_tiling_arch35.cpp逐一比较四者的DataType任一不一致即返回GRAPH_FAILED并输出包含四种实际 dtype 的错误日志。rank 上限x1、x2、x3 的维度数rank均不能超过 8超过时在 tiling 阶段校验失败。对应实现为CheckShape()中的FUSED_MUL_ADD_MAX_DIM_NUM 8常量与dimNum 8即报错的逻辑。广播能力支持任意 NumPy 广播形态包括标量如[]或[1]、单维 broadcast、跨 rank broadcast同时支持动态 shape 与动态 rank。该能力由DoOpTiling()中调用的Ops::Base::BroadcastBaseTilingOpDag提供其为 ops-math 中统一的广播 tiling 基座。实现方案从图原语到 Kernel 的全链路拆解README 以一张分层表格概括了算子的实现构成结合仓库源码可将其展开为一条完整的调用链层文件仓库根目录相对路径说明计算图原型math/fused_mul_add/op_graph/fused_mul_add_proto.hREG_OP(FusedMulAdd)三输入一输出算子定义math/fused_mul_add/op_host/fused_mul_add_def.cppOpDef::AddConfig(ascend950, ...)InferShapemath/fused_mul_add/op_host/fused_mul_add_infershape.cpp复用Ops::Base::InferShape4Broadcast(ctx, 3)Tilingmath/fused_mul_add/op_host/arch35/fused_mul_add_tiling_arch35.h / .cpp按 dtype 分支调用Ops::Base::BroadcastBaseTilingOpDagDAGmath/fused_mul_add/op_kernel/arch35/fused_mul_add_dag.hfp32/fp16 通路在 fp32 中间精度下用Vec::Mul Vec::Addint32 通路用Vec::Mul Vec::AddStructmath/fused_mul_add/op_kernel/arch35/fused_mul_add_struct.hBRC_TEMP_SCH_MODE_KEY_DECL/SELKernel 入口math/fused_mul_add/op_kernel/fused_mul_add_apt.cppKERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY)BroadcastSchschMode, OpDag计算图原型op_graphfused_mul_add_proto.h 以REG_OP(FusedMulAdd)声明算子的图级描述三个输入与一个输出的 dtype 集合均为{DT_FLOAT16, DT_FLOAT, DT_INT32}。该头文件同时被端到端示例 test_geir_fused_mul_add.cpp 直接 include作为图模式GE-IR构图时的算子定义来源。算子定义与 InferShapeop_hostfused_mul_add_def.cpp 通过OpDef::AddConfig(ascend950, aicoreConfig)将算子绑定到 ascend950 平台的 AICore 配置详见上文参数说明一节。fused_mul_add_infershape.cpp 中InferShape4FusedMulAdd直接复用Ops::Base::InferShape4Broadcast(context, INPUT_NUM_THREE)即三输入广播推断InferDataType4FusedMulAdd将输出y的 dtype 固定为输入x1的 dtypeAlign with op_proto: output dtype is the same as input x1与参数表中同 x1的语义一致。Tiling按 dtype 分支选择 OpDagfused_mul_add_tiling_arch35.cpp 的DoOpTiling()流程为先执行CheckShape()rank ≤ 8 校验再执行CheckDtype()四者 dtype 一致校验依据x1的 dtype 三分支选择对应的 OpDag 模板并调用BroadcastBaseTilingOpDagDT_FLOAT→FusedMulAddFloatOpfloat::OpDagDT_FLOAT16→FusedMulAddFloatOpOps::Base::half::OpDagDT_INT32→FusedMulAddInt32Opint32_t::OpDag通过GET_TPL_TILING_KEY(brcBaseTiling.GetSchMode())记录调度模式schMode作为 Kernel 侧模板实例化的 tilingKey。此外TilingPrepareForFusedMulAdd负责在编译期收集平台信息通过platform_ascendc::PlatformAscendC获取 AIV 核数GetCoreNumAiv与 UB 内存大小GetCoreMemSize(UB, ...)写入FusedMulAddCompileInfo{coreNum, ubSize}供 tiling 决策使用。UT 用例见 test_fused_mul_add_tiling.cpp中给出的compileInfo {64, 245760}即典型的 64 个 AIV core、245760 字节 UB 的平台参数样例。Kernel 入口与广播调度fused_mul_add_apt.cpp 是实际的__global__ __aicore__KernelKERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY)声明仅由 AIV 向量核执行通过if constexpr (std::is_sameDTYPE_X1, int32_t::value)在编译期区分 int32 与浮点通路分别实例化FusedMulAddInt32Opint32_t::OpDag或FusedMulAddFloatOpDTYPE_X1::OpDag统一使用Ops::Base::BroadcastSchschMode, OpDag sch(tiling); sch.Process(x1, x2, x3, y);完成广播调度与执行。schMode模板参数与 Tiling 阶段记录的 tilingKey 对应由 fused_mul_add_struct.h 中BRC_TEMP_SCH_MODE_KEY_DECL/SEL宏体系完成模板特化选择。fp16 / fp32 通路升到 fp32 中间精度显式 Mul Addfp16/fp32 通路的计算图DAG如下In0/In1/In2 -- CopyInBrc -- Cast(-fp32) -- Vec::Mul(x1,x2) -- Vec::Add(x3) -- Cast(-T,RINT) -- CopyOut -- Out0对应 fused_mul_add_dag.h 中FusedMulAddFloatOpT的定义OpInputX1/X2/X3CopyInBrcT将三个输入按广播语义拷贝进 UBIn0/In1/In2为占位符broadcast 时同一份 UB 数据可被多 tile 复用OpCastX1/X2/X3全部输入先Castfloat, T, CAST_NONE_MODE提升到 fp32避免半精度中间误差累积OpMul Vec::Mulfloat(OpCastX1, OpCastX2)、OpAdd Vec::Addfloat(OpMul, OpCastX3)两段式按y x1 * x2 x3顺序计算OpCastRes Vec::CastT, float, CAST_RINT_MODE结果以CAST_RINT_MODEround-to-nearest与 golden.py 中astype的舍入语义对应转回原 dtypeT后CopyOut写出。为什么此处不使用Vec::FusedMulAdd这是本算子实现中最值得注意的工程细节。DAG 头文件中的注释明确给出了原因AscendC::FusedMulAdd底层实现会in-place 写回 src2 buffer实现形式为FusedMulAdd(src2, alpha, src1, count)后跟Copy(dst, src2)。当 src2 是一个由CopyInBrc产生的占位符如OpCastX1且广播调度在大 tensor 场景跨 tile 复用同一 UB buffer 时第一 tile 的写回会污染下一 tile 的输入数据导致 fp32/fp16 大规模广播场景出现精度错误源码注释记录在 Ascend 950 上实测可观察到约 50% / 0% / -50% 量级的错误输出。因此实现刻意退化为普通Vec::Mul Vec::Add保证所有由占位符派生的 cast 结果在 UB 中只读从而兼顾正确性与广播复用效率。int32 通路无硬件 FMA直接 Mul Addint32 通路的 DAG 为In0/In1/In2 -- CopyInBrc -- Vec::Mul(x1,x2) -- Vec::Add(x3) -- CopyOut -- Out0FusedMulAddInt32OpT不做 fp32 提升整数运算无需中间精度直接在Tint32_t上执行Vec::Mul与Vec::Add同样不使用硬件 FMA 指令源码注释 no hardware FMA available, fall back to explicit Mul then Add。这一行为与 golden.py 中 int32 参考实现(x1 * x2 x3).astype(np.int32)完全对齐。测试与验证体系仓库为该算子提供了三层验证参考实现goldentests/assets/golden.py 中的fused_mul_add_golden以 NumPy 实现参考结果——浮点路径先astype(float32)升精度计算再 cast 回原 dtypeint32 路径直接整型乘加与 Kernel 两条通路一一对应可作为数值精度测试的比对基准。Tiling 单测tests/ut/op_host/arch35/test_fused_mul_add_tiling.cpp 使用TilingContextPara构造 fp16/fp32/int32 同 shape 输入shape{16,16}ND 格式以ExecuteTestCaseForEle断言 tiling 返回GRAPH_SUCCESS。注释说明由于BroadcastBaseTiling框架的 schMode 哈希可能随 CANN 版本变化单测刻意跳过 key/data 精确断言精确数值校验交给端到端测试。端到端 GE-IR 示例examples/arch35/test_geir_fused_mul_add.cpp 构建单算子图 设备运行 输出比对的完整流程覆盖 13 个用例是理解算子行为的最佳实践素材同 shapefp32[2,2]x12, x23, x34 → y10、fp16[4]→ 3.5、int32[3]→ 11典型广播fp32[4,1] × [1,4] 标量 → [4,4]→ 7、[1,5] × [5,1]行列互播、3D 互播{2,1,4}×{1,3,4}{4}、全轴广播{4,1,1}×{1,3,1}{1,1,5} → [4,3,5]、跨 rank 广播{3,4,5}×{4,5}{5}动态/混合fp16 2D 互播、int32 2D 互播、fp16 列向量广播{4,1}×{4,6}边界场景空 tensor{0}全空输入、{0,4}携带空前导维的广播。示例对每个 case 逐元素比较输出与期望值浮点使用 absTol如 fp32 取1e-5、fp16 取5e-3全部通过才返回 0可直接作为图模式下算子的正确性验收样板。调用说明图模式GE-IR示例README 给出的调用方式为图模式通过算子 IR 构图调用 FusedMulAdd调用方式样例代码说明图模式math/fused_mul_add/examples/arch35/test_geir_fused_mul_add.cpp通过 算子IR 构图方式调用 FusedMulAdd 算子。其核心调用链可归纳为对应示例中BuildFusedMulAddGraph与RunOneCase以ge::op::FusedMulAdd(fma_tag)创建算子节点并通过set_input_x1/x2/x3接入三个ge::op::Data占位节点为每个输入构造TensorDesc(shape, FORMAT_ND, dtype)以GenConstTensor填充常量数据示例支持 fp32/fp16/int32 三种 dtype 的填充与半精度位级转换graph.SetInputs(inputs).SetOutputs(outputs)后通过Session-AddGraph(graphId, graph)与Session-RunGraph(graphId, inputData, output)完成编译与执行对输出逐元素与期望值比对含 shape 校验与 absTol 容差任一 case 失败即整体返回失败。该示例同样展示了如何在 host 侧完成GEInitialize/GEFinalize生命周期管理{ge.exec.deviceId, 0}, {ge.graphRunMode, 1}可作为在 NPU 上快速验证 FusedMulAdd 行为的最小可运行模板。实践要点小结语义简单约束严格算子只有同 dtype rank≤8 ND 格式三个硬约束广播形态完全交给 NumPy 规则适配面广混合 dtype 会在 Tiling 阶段被拦截并给出明确报错。精度设计是有意为之fp16/fp32 一律提升到 fp32 中间精度计算最后以 RINT 舍入回写既抑制误差累积又规避了Vec::FusedMulAddin-place 写回在广播复用场景下的正确性缺陷——这是一个以少量计算开销换取确定正确性的典型取舍。扩展与验证路径清晰若需要新增 dtype 或平台支持改动点集中在 def 注册、Tiling dtype 分支 与 DAG 定义验证可复用 golden.py 参考实现、tiling 单测 与 GE-IR 端到端示例 三层体系。赞分享算子库人工智能CANN【免费下载链接】ops-math本项目是CANN提供的数学类基础计算算子库实现网络在NPU上加速计算。项目地址https://gitcode.com/cann/ops-math点击查看免费下载相关推荐深入解析 libcurl 的 CURLOPT_ALTSVC_CTRL用位掩码精确控制 Alt-Svc 行为深入解析 libcurl 的 CURLOPT_ALTSVC_CTRL用位掩码精确控制 Alt Svc 行为 导读 本文以 curl 仓库官方文档 CURLO人工智能算子库深度学习CANNAscendCANN ops-transformer 算子实战aclnnMoeDistributeCombineAddRmsNormV2 通信整合与 AddRmsNorm 融合算子详解CANN ops transformer 算子实战aclnnMoeDistributeCombineAddRmsNormV2 通信整合与 AddRmsNor算子库人工智能深度学习AscendCANN ops-math 算子详解Expint 指数积分算子原理、实现与图模式调用实践CANN ops math 算子详解Expint 指数积分算子原理、实现与图模式调用实践 导读 Expint 是 CANN ops math 数学算子库中用于算子库人工智能CANN上一篇Astroship开发环境搭建从零开始的完整配置指南下一篇Pyphoon本地化功能揭秘支持30语言的月球术语翻译表详解创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考