ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

CANN ops-nn 算子融合实践:SoftmaxGradExtV2FusionPass 图融合原理与约束详解

CANN ops-nn 算子融合实践:SoftmaxGradExtV2FusionPass 图融合原理与约束详解 CANN ops-nn 算子融合实践SoftmaxGradExtV2FusionPass 图融合原理与约束详解【免费下载链接】ops-nn本项目是CANN提供的神经网络类计算算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-nn导读本文围绕 CANN ops-nn 算子库中activation/softmax_grad_ext模块的图融合 Pass——SoftmaxGradExtV2FusionPass 展开讲解如何将 Softmax 反向梯度计算过程中由 Mul、ReduceSum、Sub 组成的小算子子图在计算图编译阶段融合为单一的 SoftmaxGradExt 算子从而减少算子调度开销、提升 NPU 上的执行效率。读完本文你将掌握该融合模式的四种图结构变体、触发融合的完整约束条件输入连接关系、shape/format/数据类型、平台与动态 shape 限制以及融合 Pass 在源码中的实现路径Patterns → MeetRequirements → Replacement与配套的算子定义、Tiling 策略和测试验证方法。融合模式SoftmaxGradExtV2FusionPass 的目标是把符合特定图结构的 Mul、ReduceSum、Sub 小算子融合为一个 SoftmaxGradExt 算子。其数学本质对应 Softmax 反向传播中的梯度计算y x2 * x1 * (grad - ReduceSum(grad * x1, axes))即先用grad与x1逐元素相乘Mul对乘积沿指定轴求和ReduceSum用grad减去求和结果Sub再将结果依次与x1、缩放因子x2相乘得到最终输出。融合前该子图在计算图中表现为多个独立算子节点融合后由单个 SoftmaxGradExt 算子整体承接避免了中间张量的反复搬运与多次 Kernel 启动。与第一版 SoftmaxGradExtFusionPassv1单一模式Mul_1 Mul(x2, x1)不同V2 版本在源码中注册了4 个图结构变体softmax_grad_ext_fusion_pass.cpp用于覆盖实际模型图里因算子输入顺序差异x1/sub谁在前、mul1/input2谁在前而出现的不同子图排布提高模式匹配的覆盖率。四个场景示意如下场景一场景二场景三场景四从源码实现看四种变体共享相同的骨架mul Mul(grad, x1) // input0 * input1 sum ReduceSum(mul, axes-1) // 对最后一维求和 sub Sub(grad, sum) // input0 - sum区别仅在于其后两个 Mul 的输入顺序MakePatternSoftmaxGradExtV2变体mul1 构造mulGrad 构造variant 0Mul(input1, sub)Mul(mul1, input2)variant 1Mul(sub, input1)Mul(mul1, input2)variant 2Mul(input1, sub)Mul(input2, mul1)variant 3Mul(sub, input1)Mul(input2, mul1)需要注意的是V2 的四种变体与 v1 的图结构Mul_1 Mul(x2, x1)即mul1 Mul(input2, input1)互不匹配单元测试中专门验证了「v1 的 Pass 无法匹配 v2 的图」与「v2 的 Pass 无法匹配 v1 的图」两个场景test_softmax_grad_ext_fusion_pass.cpp二者作为两个独立的融合 Pass 并存。使用约束融合 Pass 只在满足以下约束时才会执行否则保持原图不做改动返回GRAPH_NOT_CHANGED。这些约束既写入了文档也在融合 Pass 的MeetRequirements与替换逻辑中得到代码级落实。输入连接关系约束Mul_1 节点的输入与 Sub 节点的第一个输入共用 input0即grad。Mul_1 节点的输入与 Mul_2 节点的输入共用 input1即x1。ReduceSum 的 axis 参数必须为 -1 或最后一维。从源码看pattern 中 ReduceSum 的axes通过内部常量节点CreateConst({-1})固定为最后一维BuildPatternReduceSumkeep_dims与noop_with_empty_axes均置为true。在替换阶段Pass 会从匹配到的原 ReduceSum 节点上读取真实的axes常量输入与keep_dims属性解析后传递给新建的 SoftmaxGradExt 节点GetAxisFromReduceSum因此最终算子的axes属性与原始 ReduceSum 保持一致。axes常量支持 INT64 与 INT32 两种数据类型且必须是单个标量。数据格式与 shape 约束input0、input1 的 shape 必须是 1D6D。input0 和 input1 的数据格式为 ND。input2 可以是 scalarshape 为(1,)也可以是与 input0 相同 shape 的张量。对应到算子定义softmax_grad_ext_def.cpp中将三个输入与输出y的格式全部限定为FORMAT_NDsoftmax_grad_ext_def.cppTiling 侧在GetDimsAndCheckShapeValid中校验输入维数不超过 6MAX_DIMS、拒绝空 tensor并检查 input2 为 scalar 时与 input0/input1/output 逐维相等input2 为张量时各维与 input0 保持一致softmax_grad_ext_tiling_base_arch35.cpp。融合 Pass 侧则通过CheckInputsShapeValid检查子图所有输入均不含未知维度值为 -1 的 dim一旦发现未知 shape 直接跳过融合softmax_grad_ext_fusion_pass.cpp。动态 shape 限制不支持动态 shape 场景。该限制由融合 Pass 的MeetRequirements强制执行任何输入含未知维度时返回false融合不生效计算图保持原样。单元测试v1_unknown_shape_not_changed与v2_unknown_shape_not_changed均以dims {-1, 32, 128}的输入验证了此行为test_softmax_grad_ext_fusion_pass.cpp。数据类型约束input0、input1、input2 的数据类型需要保持一致。数据类型支持 FLOAT16、FLOAT32、BFLOAT16。该约束在算子原型softmax_grad_ext_proto.h与算子定义中均以{DT_FLOAT16, DT_FLOAT, DT_BF16}声明Tiling 阶段GetAndCheckDtypes会校验三个输入与输出的 dtype 完全一致且仅接受上述三种类型softmax_grad_ext_tiling_base_arch35.cpp。算子二进制配置文件按数据类型拆分为三个二进制SoftmaxGradExt_float16、SoftmaxGradExt_float32、SoftmaxGradExt_bfloat16softmax_grad_ext_binary.json所有输入输出均为FormatAgnostic的 ND 格式、shape [-2]动态 rank。融合 Pass 的实现机制SoftmaxGradExtV2FusionPass 继承自 GE 图引擎的PatternFusionPass基类softmax_grad_ext_fusion_pass.h通过三个虚函数完成「模式匹配 → 条件校验 → 图替换」的标准融合流程Patterns()构造子图模式。V2 版本循环生成 4 个变体 patternkPatternV2VariantCount 4每个 pattern 通过es::EsGraphBuilder以 ES 图构建 DSL 描述Mul → ReduceSum → Sub → Mul → Mul的连接关系并CaptureTensor捕获内部的 ReduceSum 输出供替换阶段读取 axes 信息softmax_grad_ext_fusion_pass.cpp。MeetRequirements()融合前置校验包含两道关卡IsTargetPlatform()读取平台信息仅当short_soc_version Ascend950时放行softmax_grad_ext_fusion_pass.cpp其他平台直接跳过融合CheckInputsShapeValid()检查子图所有输入 shape 均为已知无 -1 维度。 单元测试v1_unsupported_platform_not_changed将平台切换为Ascend910_93后验证融合被拒绝test_softmax_grad_ext_fusion_pass.cpp。Replacement()执行图替换。SoftmaxGradExtReplacementCommon依次完成从捕获的 ReduceSum 节点读取axes与keep_dims→ 收集子图边界输入至少 3 个→ 以输入 0/1/2 分别映射为grad/x1/x2构造 SoftmaxGradExt 节点 → 对新图执行 shape 推导softmax_grad_ext_fusion_pass.cpp。替换后的输入顺序固定为SoftmaxGradExt(grad, x1, x2)与算子原型的输入定义一致。注册阶段两个 Pass 均通过REG_FUSION_PASS(...).Stage(CustomPassStage::kAfterInferShape)注册即在计算图完成 shape 推导之后执行融合softmax_grad_ext_fusion_pass.cpp。此外实现中还针对图匹配的健壮性做了防御性处理单元测试覆盖了多种异常场景Mul 输入顺序调换不应误匹配v1_mul_input_order_swapped_not_changed、控制边可能引发环时跳过融合v1_control_edge_cycle_not_changed、axes 常量被子图外节点共享时保证外节点不被误删v1_shared_axes_const_success、同一图中存在多个重叠子图时迭代匹配不崩溃v1_two_subgraphs_no_crash、输出同时存在数据消费者与控制消费者、以及 axes 来自依赖子图输出的外部算子时不得产生环v1_axes_from_op_depending_on_mulgrad_cycle这些用例均位于 test_softmax_grad_ext_fusion_pass.cpp。融合产物SoftmaxGradExt 算子落地融合后的 SoftmaxGradExt 算子由仓库中完整的「原型 → 算子定义 → Infershape → Tiling → Kernel → 配置 → 测试」链路支撑算子原型softmax_grad_ext_proto.h三个输入grad、x1、x2输出y数据类型均限定{DT_FLOAT16, DT_FLOAT, DT_BF16}属性axesInt默认 1与keep_dimsBool默认 true。原型注释还指出该算子在 Ascend 950 上格式必须为 ND。算子定义softmax_grad_ext_def.cpp声明输入输出均为 ND 格式AICore 配置同时注册了ascend950与ascend350两个型号开启动态编译静态化、动态 rank 支持、动态 shape 支持与精度收敛标志。Infershapesoftmax_grad_ext_infershape.cpp输出y的 shape 与 dtype 直接继承自输入x1。Tiling针对不同数据规模提供三种切分策略——SoftmaxGradExtARSmallR小 R 场景TILINGKEY 500、SoftmaxGradExtAR常规TILINGKEY 1000与SoftmaxGradExtARRecompute重计算策略TILINGKEY 2000实现位于 op_host/arch35 目录Kernel 入口按 TILING_KEY 分发softmax_grad_ext_apt.cpp核心计算模板位于 op_kernel/arch35。ST 测试ttk_kernel_softmax_grad_ext_st.csv覆盖 float32/float16/bfloat16 三种数据类型、2D4D shape、axes取最后一维如 1、2、3、keep_dimsTrue、以及 input2 为 scalar(1,)与同 shape 张量两种形态验证了文档所述约束在实际 Kernel 执行中的一致性。支持的型号融合 Pass 与算子本身的平台适配保持一致。融合 Pass 通过平台信息校验仅作用于Ascend950short_soc_version Ascend950因此该融合在以下型号上生效Ascend 950PR / Ascend 950DT与此同时算子本身不经融合直接单独调用时在 AICore 配置上同时支持ascend950与ascend350两个系列见 softmax_grad_ext_def.cpp并在对应目录下各有一份算子二进制配置ascend950/softmax_grad_ext_binary.json 与 ascend350/softmax_grad_ext_binary.json。验证与观察方式单元测试融合 Pass 的行为可通过 op_graph 单元测试 验证。测试会在融合前后分别将图导出为 ONNXDumpToFile(Graph::DumpFormat::kOnnx, ...)并断言融合后的图仅包含 3 个 Data 输入、1 个 SoftmaxGradExt 节点与 1 个 NetOutput节点总数 5且图中不再残留任何 Mul/Sub/ReduceSum 节点同时核对axes、keep_dims属性与输入映射关系input0grad、input1x1、input2x2。整库测试入口softmax_grad_ext 的 UT/ST 均挂接在 tests/CMakeLists.txt 与算子目录的 CMakeLists.txt 下可随 ops-nn 仓库的既有构建与测试流程运行。【免费下载链接】ops-nn本项目是CANN提供的神经网络类计算算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-nn创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表