CANN ops-nn 算子融合规则解析QuantBatchMatmulV3TransposeFusionPass 转置融合原理与实践【免费下载链接】ops-nn本项目是CANN提供的神经网络类计算算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-nn导读QuantBatchMatmulV3TransposeFusionPass是 CANN ops-nn 仓库中针对量化矩阵乘算子QuantBatchMatmulV3设计的一类图编译器融合规则其作用是把输入支路上显式的Transpose/TransposeD以及部分等价于转置的Reshape操作合入矩阵乘计算本身通过调整图连接并取反算子的transpose_x1/transpose_x2属性来消除一次独立的数据搬移。本文以 QuantBatchMatmulV3TransposeFusionPass.md 为主线结合仓库内融合规则的源码实现quant_batch_matmul_v3_transpose_fusion_pass.cpp与单元测试用例test_quant_batch_matmul_v3_transpose_fusion_pass.csv系统讲解全部融合模式、使用约束、平台与版本限制以及源码级判定流程帮助读者理解该规则何时触发、如何工作、在哪些产品上生效。一、融合规则定位与核心思想QuantBatchMatmulV3完成量化的矩阵乘计算其最小支持输入维度为 2 维最大为 6 维详见 QuantBatchMatmulV3 算子说明。在实际量化推理/训练图中x1或x2支路经常先接一个显式转置如权重为(K, N)而矩阵乘要求(N, K)随后再做矩阵乘。显式转置意味着一次独立的数据重排会带来额外的内存读写开销。该融合规则正是针对这一场景的图优化将QuantBatchMatmulV3的x1或x2输入支路中的显式转置合入矩阵乘计算同时也可以处理两条支路。可识别的转置节点类型为Transpose和TransposeD。融合完成后不再需要单独执行转置算子融合x1支路时将QuantBatchMatmulV3节点的transpose_x1属性取反融合x2支路时将该节点的transpose_x2属性取反。“取反”是指将false改为true或将true改为false。这一操作在源码中由SetTransposeAttr实现读取当前属性值后写回其逻辑非见 quant_batch_matmul_v3_transpose_fusion_pass.cpp。需要说明的是本规则是一条图编译器自定义融合 Pass从源码结构看它继承自ge::fusion::FusionBasePass通过REG_FUSION_PASS宏注册到图编译器的自定义 Pass 阶段源码中注册阶段为kCompatibleInherited或kAfterInferShape取决于编译器版本见 quant_batch_matmul_v3_transpose_fusion_pass.h 与同名.cpp的注册部分。二、融合模式总览2.1 “旁路”与图示约定本文用“旁路”表示以下连接调整将转换节点的输入直接连接到该转换节点的后继节点使当前支路不再经过该转换节点。旁路后转换节点仍有其他数据使用者时保留否则删除。后继节点可以是QuantBatchMatmulV3也可以是保留的Bitcast。图示约定如下输入对应关系分别展示x1、x2及两路scale。x1_scale表示输入pertoken_scaleIR 索引 5x2_scale表示输入scaleIR 索引 2这两个名称是图示别名不是新增算子参数也不是算子属性。这一索引约定与源码中的常量一致X1_INDEX 0、X2_INDEX 1、X2_SCALE_INDEX 2、X1_SCALE_INDEX 5。符号定义t1、t2分别表示融合前QuantBatchMatmulV3的transpose_x1、transpose_x2属性值!表示逻辑取反。源张量融合前后的同名源张量表示同一个张量。省略内容图中省略bias、offset、指定转置维度顺序的输入及指定重塑后形状的输入只展示与融合相关的数据连接。原文档共给出 7 幅模式图位于 docs/zh/figures文件名为quant_batch_matmul_v3_transpose_fusion_pass_1.png至_7.png分别对应下文 2.22.8 的七种场景。2.2 仅融合 x1 侧转置当x1支路和x1_scalepertoken_scale支路中各存在一个转置节点时分别旁路这两个转置节点将它们的输入直接连接到QuantBatchMatmulV3的x1和pertoken_scale输入端口将QuantBatchMatmulV3节点的transpose_x1属性取反即!t1x2和scale输入的连接以及该节点的transpose_x2属性t2保持不变。对应源码行为GetScaleTransNode只在数据侧转置节点被识别到nodeTransX1 ! nullptr时才去检查对应scale支路的转换节点二者都满足时才会被加入待旁路/待删除集合。2.3 仅融合 x2 侧转置当x2支路和x2_scalescale支路中各存在一个转置节点时处理逻辑与 x1 侧对称分别旁路两个转置节点将它们的输入直接连接到QuantBatchMatmulV3的x2和scale输入端口将QuantBatchMatmulV3节点的transpose_x2属性取反!t2x1和pertoken_scale输入的连接以及该节点的transpose_x1属性t1保持不变。2.4 同时融合两侧转置当x1、x2、x1_scale、x2_scale四条支路中各存在一个转置节点且x1、x2两侧均满足使用约束时分别旁路四个转置节点四个转置节点的输入分别连接到QuantBatchMatmulV3的x1、x2、pertoken_scale和scale输入端口QuantBatchMatmulV3节点的transpose_x1和transpose_x2属性均取反!t1、!t2。需要注意仅匹配图中结构不足以触发融合还需满足数据类型、维度和平台能力等约束详见下文“使用约束”。2.5 普通 scale 的 Reshape 融合某些场景下scale支路上的转置在数据上等价于一次Reshape例如某维为 1交换最后两维不改变元素顺序。当x1、x2支路中各存在一个转置节点两路scale支路中各存在一个等价于所需转置的Reshape时融合会分别旁路两个转置节点和两个Reshape四条支路连接到各自对应的QuantBatchMatmulV3输入端口同时transpose_x1和transpose_x2属性均取反。静态普通scale的识别要求见下文“scale 的 Reshape 要求”。图中展示的是两侧均可融合的情形实际规则也可仅处理满足约束的一侧。对应源码IsReshapeTrans中对于普通 tensor只有当其最后两维存在一维为 1时最后两维的转置才可以等价为Reshape并被识别。2.6 MX scale 的 Reshape 融合MXMicroscaling量化场景中FLOAT8_E8M0类型的scale使用三维 shape如(1, N, 2)其Reshape将(1, N, 2)变为(N, 1, 2)。以x2侧为例x2输入前为Transpose或TransposeDx2_scale为FLOAT8_E8M0其Reshape将(1, N, 2)变为(N, 1, 2)。融合时旁路x2支路的转置节点和x2_scale支路的Reshape将二者的输入分别连接到QuantBatchMatmulV3的x2和scale输入端口并将transpose_x2属性取反x1及x1_scale保持原连接。x1侧满足相应条件时采用相同的处理逻辑。对应源码IsReshapeTrans中针对DT_FLOAT8_E8M0的 MX 分支要求Reshape输入、输出均为 3 维前两维互换且其中至少一维为 1shapeInput.GetDim(0) shapeOut.GetDim(1)且shapeInput.GetDim(1) shapeOut.GetDim(0)且shapeInput.GetDim(0) 1 || shapeInput.GetDim(1) 1。源码注释进一步说明MX 量化场景中FLOAT8_E8M0scale 使用(m, ceil(k/64), 2)或(ceil(k/64), n, 2)等 3 维 shape当前两维任意一维为 1 时交换这两维可以用Reshape表示。图中 MX 示例的末维 2 来自该示例的量化布局并非条件 B 额外要求的固定值见原文档“scale 的 Reshape 要求”一节的说明。2.7 保留 Bitcast 的转置融合Bitcast按目标数据类型重新解释数据的位表示不改动内存内容。当四路输入各自按“转置节点 →Bitcast→QuantBatchMatmulV3”连接时融合会将四个转置节点的输入分别直接连接到后面的Bitcast保留四个Bitcast及其到QuantBatchMatmulV3的连接并将transpose_x1和transpose_x2属性取反。融合后Bitcast输出的数据类型保持不变。补充约束各路输入不要求同时存在Bitcast。转置节点可以直接连接矩阵乘也可以通过一个Bitcast连接矩阵乘不支持越过连续多个Bitcast进行融合。对于scale支路Bitcast前的节点也可以是满足要求的Reshape。平台限制和动态Reshape的适用边界见下文“使用约束”。对应源码GetTransposeCandidateNode在isBitcastPattern为真且输入节点是Bitcast时会再向上取一层输入作为转置候选RelinkNode在 Bitcast 场景下把源节点重连到Bitcast的输入端口dstInputPort 0并先更新Bitcast输入 desc使其匹配类转置节点的输入再保留Bitcast输出侧 dtype 后更新QuantBatchMatmulV3的输入 desc。2.8 动态 scale 的 Reshape 融合当x1或x2存在未知维度shape中出现-1等未知维时按动态场景处理。以x2侧为例旁路x2支路的转置节点和x2_scale支路的Reshape将二者的输入分别连接到QuantBatchMatmulV3的x2和scale输入端口将transpose_x2属性取反x1和x1_scale保持原连接。动态场景不执行静态Reshape的维度条件检查但输入图仍须保证该Reshape与所需的scale转置等价。对应源码IsReshapeTrans中if (isDynamic) return true;直接放行IsDynamicNode通过检查x1/x2的 shape 中是否存在ge::UNKNOWN_DIM、ge::UNKNOWN_DIM_NUM或小于 0 的维度来判定动态场景。2.9 关于 scale 支路处理的通用说明各图中的scale转换节点是所示场景的具体结构不要求每一路scale都存在转换节点只有相应数据输入的Transpose或TransposeD被融合时才处理该路scalepertoken_scale未连接时不处理该输入源码中通过GetInputIndexByName(pertoken_scale)判断其是否存在不存在时直接返回nullptr避免对 IR 索引 5 直接访问导致报错被旁路的节点仍有其他数据使用者时保留否则从图中删除。三、使用约束3.1 输入与转置要求x1和x2的原始shape维数均不得小于 2至少一路数据输入存在Transpose或TransposeD否则不触发融合。数据输入前仅有Reshape不能触发融合源码CheckTranspose只对数据侧识别真实转置节点CheckNodeShape要求两路输入的OriginShape维数均不小于 2。x1、x2的数据类型均支持INT8、INT4、FLOAT8_E4M3FN、FLOAT8_E5M2、HIFLOAT8、FLOAT4_E2M1输出数据类型支持INT8、FLOAT16、BF16、INT32、FLOAT32。实际组合还需满足目标产品的算子约束。源码中CheckNodeDtype维护了legalInputDtypes与legalOutDtypes两个静态列表逐一校验x1、x2与输出 desc 的数据类型。支持动态shapex1或x2的shape存在未知维度时按动态场景处理。数据输入的转置必须只交换最后两维其他维度的顺序保持不变。例如形状为(B, M, K)的输入转置为(B, K, M)。交换批次维度等其他转置方式不属于本规则的适用范围。对应的scale转换也须保持量化系数与数据的对应关系。pertoken_scale是可选输入未连接时不处理该路scale。3.2 Bitcast 场景的限制部分平台的矩阵乘指令支持 INT8 与 INT4 混合输入。在这类平台上源码中以平台 intrinsic 映射中是否包含Intrinsic_mmad的s8s8/s8s4能力作为判定依据见GetPlatformSupport只要四路输入中的一路符合下表结构本规则就保留整个QuantBatchMatmulV3节点及其所有输入支路不执行本次转置融合。即使当前节点没有采用 INT8 与 INT4 混合输入上述限制也适用。输入支路触发上述限制的连接结构x1、x2或两路scale中的任一路Transpose或TransposeD→Bitcast→QuantBatchMatmulV3两路scale中的任一路满足融合要求的Reshape→Bitcast→QuantBatchMatmulV3仅有Bitcast、其前方没有上述转换节点时不会因本条限制跳过融合。对应源码IsBitcastPattern遍历四个输入端口若某输入是Bitcast且其前方存在转置/等价 Reshape 则判定为 Bitcast 模式CheckFusionPreconditions中在平台支持s8s4混合输入时supportMmadS8S4为真直接返回GRAPH_NOT_CHANGED。动态Reshape与Bitcast组合使用时还存在识别限制如果只有scale支路包含“动态Reshape→Bitcast→ 矩阵乘”且其他支路都没有可识别的上述带Bitcast结构该动态Reshape不保证被融合。因此原文档的动态Reshape示例仅展示未经过Bitcast的连接。3.3 scale 的 Reshape 要求scale输入前的Reshape需满足下表条件。静态场景还要求输入图提供Reshape输入、输出的形状和数据类型信息且对应QuantBatchMatmulV3输入至少为 2 维。场景识别条件条件 A静态通用条件对应QuantBatchMatmulV3输入shape的最后两维至少一维为 1此判断不限制scale的数据类型条件 B静态 MX 条件数据类型为FLOAT8_E8M0Reshape输入、输出均为 3 维前两维互换且其中至少一维为 1例如(1, N, 2)变为(N, 1, 2)动态shape不要求满足上述静态形状条件但Reshape仍须与所需的scale转置等价说明静态场景满足条件 A 或条件 B 中的任意一项即可。条件 B 指FLOAT8_E8M0三维scale的前两维交换条件不满足条件 B 时仍可通过条件 A 识别。动态场景按上表“动态shape”行处理。满足表中的形状条件后仍须保证Reshape与所需的scale转置等价不能改变量化系数与数据的对应关系。本规则不单独删除用于计算动态形状的Shape、Gather、Pack等节点。3.4 维度与平台限制QuantBatchMatmulV3TransposeFusionPass是普通转置融合规则QuantBatchMatmulV3TransposeLimitFusionPass是针对特定维度场景的补充规则源码中 Limit Pass 继承自普通 Pass二者Run时以limitPass布尔参数区分见 quant_batch_matmul_v3_transpose_fusion_pass.h。两者分别注册均将符合条件的显式转置合入矩阵乘计算通过调整输入连接并取反QuantBatchMatmulV3的相应转置属性实现融合区别在于适用的产品和维度条件。Limit 规则用于在普通规则被关闭时仍为符合其条件的支路提供转置融合处理。令outer为融合前QuantBatchMatmulV3对应数据输入原始shape的倒数第二维inner为最后一维。两个规则均分别检查x1、x2可只融合满足条件的一侧。本文所列产品的处理条件如下实际融合还须满足前文的输入、转置节点及scale等使用约束。产品普通规则Limit 补充规则Ascend 950PR/Ascend 950DTx1、x2均不受本节 65535 维度阈值限制满足其他使用约束时执行转置融合跳过本补充规则转置融合由普通规则处理Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品x1要求outer≤ 65535x2要求outer≤ 65535 或格式为FRACTAL_NZ分别检查x1、x2仅处理满足 0 outer≤ 65535 且inner 65535 的支路Limit 规则的维度条件同样适用于FRACTAL_NZ格式不因格式而放宽。参与判断的outer或inner为未知值-1时该支路不满足 Limit 规则的维度条件。对应源码GetIntrinsicsLimit中普通规则为x1: outer 65535 || supportL12btBf16x2: outer 65535 || supportL12btBf16 || 格式为 FRACTAL_NZLimit 规则为outer 0 outer 65535 inner 65535。其中supportL12btBf16来自平台 intrinsic 映射中Intrinsic_data_move_l12bt是否支持bf16的查询结果。源码中INNER_SHAPE_LIMIT 65535即为该阈值。3.5 版本要求编译和运行时的图编译器版本号均须不小于 90100000。编译版本不满足要求时不启用这两个转置规则运行版本不满足要求时两个规则均保持输入图不变。对应源码IsTargetVersion()通过aclsysGetVersionNum(ge-compiler, version)获取编译器版本并与TAEGET_VERSION 90100000比较两个 Pass 的Run入口在版本不满足时直接返回GRAPH_NOT_CHANGED。同时REG_FUSION_PASS注册被#if GE_COMPILER_VERSION_NUM 90100000宏条件保护。四、支持的型号产品的数据类型与格式支持范围参见 QuantBatchMatmulV3 算子说明。本规则按平台指令能力判断融合条件算子支持某产品并不表示该产品支持本规则的全部融合模式。按原文档声明本规则支持的型号如下Atlas A2 训练系列产品/Atlas A2 推理系列产品Atlas A3 训练系列产品/Atlas A3 推理系列产品Ascend 950PR/Ascend 950DT以原文档中的型号注释标记为准上述三组型号与 README 中“产品支持情况”表格所列QuantBatchMatmulV3算子本身的支持面并不完全等价融合规则的生效范围以“按平台指令能力判断”为准。五、源码实现剖析融合判定主流程融合 Pass 的核心实现在 quant_batch_matmul_v3_transpose_fusion_pass.cpp其工作流可概括为以下阶段版本检查IsTargetVersion编译器版本低于 90100000 时直接返回不执行任何修改。节点收集遍历graph-GetDirectNode()仅缓存所有QuantBatchMatmulV3节点避免融合删除Transpose节点后继续访问失效的 GNode。前置条件检查CheckFusionPreconditionsCheckQuantBatchMatmulV3节点类型必须为QuantBatchMatmulV3GetPlatformSupport查询平台Intrinsic_data_move_l12bt是否支持 bf16与Intrinsic_mmad是否支持 s8s4 混合输入能力IsBitcastPattern检测是否存在带Bitcast的连接结构若平台支持s8s4且命中 Bitcast 模式则直接跳过对应 3.2 节限制CheckNodeDtype/CheckNodeShape校验输入输出数据类型与维度数IsDynamicNode判定是否为动态 shape 场景CheckTransposex1/x2中至少一路存在真实Transpose/TransposeD节点。维度判定GetIntrinsicsLimit按普通规则或 Limit 规则分别计算x1、x2两侧是否允许删除转置对应 3.4 节产品差异。候选节点识别GetTransposeNode/GetScaleTransNode找出数据侧与 scale 侧待旁路的转换节点scale侧允许Reshape等价转置普通/ MX / 动态三种识别分支对应 3.3 节。执行融合CommitFusionSetTransposeAttrs对命中支路取反transpose_x1/transpose_x2RelinkFusionNodes删除类转置节点到后继节点的输出边再把其源节点直接连到后继普通场景重连到QuantBatchMatmulV3Bitcast 场景重连到Bitcast并同步更新输入 descReportTransposeFusion通过GraphFuseInspectorUtils::ReportFuse上报融合结果RemoveFusionNodes删除已无输出使用的类转置节点RemoveNode会先清理输入边再删除节点。值得注意的实现细节源码注释指出“torchair框架会为可选输入创建占位节点GEIR 框架不会”因此对可选输入pertoken_scale的访问通过GetInputIndexByName解析实际索引避免直接按 IR 索引 5 访问时报错这也是文档中“pertoken_scale未连接时不处理该输入”约束的工程化落地。六、测试用例验证仓库为该规则提供了参数化单元测试test_quant_batch_matmul_v3_transpose_fusion_pass.cpp 从 CSV 文件 test_quant_batch_matmul_v3_transpose_fusion_pass.csv 读取用例构造包含Transpose/TransposeD/Bitcast/Reshape的图运行 Pass 后断言返回状态、图中剩余TransposeTransposeD数量、剩余Reshape数量以及融合后transpose_x1/transpose_x2属性值。从 CSV 用例可以看出规则的覆盖范围与预期行为case_name → 关键配置 → 期望结果基础融合quant_bmm_v3_with_x1_transpose、quant_bmm_v3_with_x2_transpose、quant_bmm_v3_with_both_transpose动态 shape-1 -1均期望SUCCESS转置节点被全部消除expected_transpose_count 0对应属性被置为 1。三维输入quant_bmm_v3_with_x1_transpose_3d、quant_bmm_v3_with_both_transpose_3d验证(B, M, K)形状下仅交换最后两维的转置可融合。Reshape 等价转置quant_bmm_v3_with_reshape_transposeReshape (1,128)、quant_bmm_v3_with_reshape_known_shapeReshape (1,16)scale来自reshape_scale_fromscale、quant_bmm_v3_reshape_last_dim_1x2形状1 16验证条件 A。MX 三维 scalequant_bmm_v3_with_reshape_transpose_mx、quant_bmm_v3_with_reshape_mx_scale_3dscale形状1 120 2FLOAT8_E8M0Reshape 到1 120 2、quant_bmm_v3_mx_reshape_dim0_eq_1验证条件 B 的前两维交换识别quant_bmm_v3_with_x2_transpose_fp4验证FLOAT4_E2M1数据类型。Bitcast 场景quant_bmm_v3_with_both_transpose_bitcast四路带 Bitcast 期望融合、quant_bmm_v3_bitcast_scale_reshapex1转置 Bitcast scale Reshape、quant_bmm_v3_with_bitcast_reshape_scale验证保留 Bitcast 的融合模式。Limit 规则quant_bmm_v3_limit_inner_gt65535x2形状70000 128outer70000 不满足普通规则但满足 Limit 规则inner65535分支期望融合quant_bmm_v3_limit_outer_gt_65535x2形状128 70000outer128 ≤ 65535、inner70000 65535期望SUCCESSquant_bmm_v3_limit_no_transpose_no_change无转置时不改变。不触发场景quant_bmm_v3_no_transpose_no_change无转置、quant_bmm_v3_x1_1d_shapex1 为 1 维、quant_bmm_v3_with_x1_dtype_invalid/quant_bmm_v3_with_x2_dtype_invalid/quant_bmm_v3_with_out_dtype_invalid非法数据类型、quant_bmm_v3_x1_dtype_bf16输入不支持 BF16、quant_bmm_v3_out_dtype_hifloat8输出不支持 HIFLOAT8均期望GRAPH_NOT_CHANGED且节点保持不变。NZ 格式quant_bmm_v3_x2_nz_format验证x2为FRACTAL_NZ格式时可执行普通规则融合。这些用例与文档的使用约束一一对应为“何时触发、何时保持原图”提供了可执行的验证基准。此外op_host侧 tiling 实现quant_batch_matmul_v3_tiling.cpp中存在相关提示信息说明未启用该融合 Pass 时 tiling 阶段也会给出引导性日志。七、实践要点总结触发前提x1/x2至少一侧存在真实Transpose/TransposeD节点且转置只交换最后两维shape维数不小于 2输入输出数据类型落在规则允许集合内编译器版本不小于 90100000。scale 联动数据侧转置被融合时对应scale支路若存在等价Reshape普通条件 A / MX 条件 B / 动态场景也会一并融合pertoken_scale未连接时不处理。Bitcast 特例在支持 INT8/INT4 混合输入的平台上带Bitcast的转置结构不执行融合普通平台则保留Bitcast完成位级重解释后再融合。产品差异Ascend 950 系列由普通规则全量处理Atlas A2/A3 系列由普通规则65535 阈值与 Limit 补充规则0 outer ≤ 65535 且 inner 65535配合覆盖。结果判定融合成功意味着图中不再残留被旁路的转换节点除非其仍有其他数据使用者QuantBatchMatmulV3的transpose_x1/transpose_x2被取反等价于把转置语义下沉到矩阵乘指令内部。如需进一步了解算子的输入输出定义与调用方式可参考 QuantBatchMatmulV3 算子说明 及其 docs 目录下的 aclnnQuantMatmulV3.md、aclnnQuantMatmulV4.md 等接口文档。【免费下载链接】ops-nn本项目是CANN提供的神经网络类计算算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-nn创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考