Pull Request已成功合入, 合并人@ascend-robot
(感谢 SorryNaCN 的贡献)变更摘要
该 PR 为 DVM 后端(torch_npu._inductor.dvm)新增了矩阵乘法(matmul)的 template fusion 能力。核心思路是:通过环境变量 INDUCTOR_DVM_ENABLE_MATMUL_FUSION=1 按需启用,为 aten.mm、aten.bmm、aten.addmm、aten.baddbmm 注册 DVM template lowering,将 matmul 及其合法的 pointwise epilogue 纵向融合为单个 DVM kernel;同时补齐了 baddbmm 的转置标注、图构建参数透传和 codegen 支持,并新增了 DvmTemplateBuffer 来保存 template FX 图与逻辑 placeholder 到实际 IR 输入的绑定关系,确保 wrapper 调用时能恢复真实参数。
主要改动
-
新增
DvmTemplateBuffer与 matmul template lowering 注册(template.py):新增DvmTemplateBuffer类继承ir.TemplateBuffer,保存 traced graph 及 input_bindings;通过_register_dvm_mm_template_lowerings为mm/bmm/addmm/baddbmm注册 lowering,在 lowering 阶段调用mm_rule校验合法性,K=1 场景下沉为逐元素mul(或mul+add组合),不满足条件时回退到原有 fallback 路径。 -
调度器融合策略扩展(
mlir_fusion.py):在NpuDvmScheduling中修改can_fuse_vertical和can_fuse_horizontal,当模板节点为DvmTemplateBuffer时仅允许 pointwise epilogue 纵向融合、禁止横向融合和 template-to-template 融合;新增codegen_template方法,复用NpuMetaScheduling的 traced-graph 构图能力生成 MLIR kernel,并在调用前通过input_bindings恢复真实输入参数。 -
新增
enable_matmul_fusion配置开关(config.py与mlir_fusion.py):在config.py中增加enable_matmul_fusion配置项,通过环境变量INDUCTOR_DVM_ENABLE_MATMUL_FUSION=1控制,默认关闭;在DvmMlirFusionPatch初始化时检测该开关,仅开启时才调用patch_dvm_matmul_template_fusion()注册 lowering 和修改npu_ir.subtract_graph。 -
baddbmm全链路支持与mm_rule规则调整(op_emitter.py、graph_build.py、fx_pass.py):mm_rule扩展支持aten.baddbmm.default,移除了对小输出维度(SMALL_OUTPUT_MAX=256)的硬性限制;addmmDVM op 注册同时覆盖addmm和baddbmm,重构addmmcodegen 逻辑为先计算 matmul 再分别应用alpha/beta缩放后相加;DvmCodegenInterpreter和annotate_mm_transpose_flags同步扩展baddbmm的转置标注与参数透传。 -
新增 template fusion 测试覆盖(
test_dvm_mlir_fusion.py):新增test_matmul_uses_dvm_fusion、test_k1_matmul_lowers_to_mul、test_k1_addmm_lowers_to_pointwise、test_matmul_fusion_output_with_multiple_users、test_matmul_does_not_fuse_view_only_epilogue、test_matmul_fuses_view_with_pointwise_epilogue、test_bmm_with_view_input_uses_dvm_fusion、test_bmm_same_buffer_views_keep_distinct_input_meta共 8 个测试用例,覆盖 mm/bmm/addmm/baddbmm 四种算子的 fusion 行为、K=1 下沉、多用户输出、view-only epilogue 不融合等边界场景。


代码审查
审查总结
我逐一审查了所有 7 个变更文件:
| 文件 | 审查结果 |
|---|---|
test/_inductor/test_dvm_mlir_fusion.py |
发现 1 个 P1 问题(环境变量泄漏) |
torch_npu/_inductor/dvm/config.py |
无问题(仅新增 enable_matmul_fusion 配置项) |
torch_npu/_inductor/dvm/fx_pass.py |
无问题(annotate_mm_transpose_flags 扩展支持 baddbmm) |
torch_npu/_inductor/dvm/graph_build.py |
无问题(重构 ktype 逻辑并扩展 addmm/baddbmm 的转置参数透传) |
torch_npu/_inductor/dvm/mlir_fusion.py |
无问题(新增 codegen_template 方法,集成模板融合调度) |
torch_npu/_inductor/dvm/op_emitter.py |
无问题(移除 check_output、扩展 mm_rule 到 baddbmm、调整 addmm 运算顺序) |
torch_npu/_inductor/dvm/template.py |
发现 2 个 P3 问题(硬编码索引、跳过 super().__init__()) |
发现统计:
- P1: 1 个
- P2: 0 个
- P3: 2 个
整体风险判断: 中等风险。核心功能逻辑(模板 lowering、融合规则、codegen)设计合理,边界检查充分。主要风险在于测试辅助方法 _run_and_get_code_with_dvm 缺少异常安全保护,若测试中途失败会污染后续测试环境变量,可能导致级联测试失败。两个 P3 问题属于代码可维护性改进建议,不影响当前功能正确性。
| 类型 | 数量 |
|---|---|
| 🔴 阻塞 | 1 |
| 🟡 建议 | 0 |
⛔ 需要修改


Pull Request 已合并或已关闭。
If you want to solve this problem, you can click here to do it in the FAQs.




【合入来源】
关联图模式 Issue:https://gitcode.com/Ascend/pytorch/issues/1978
移植来源:https://gitcode.com/Ascend/pytorch/merge_requests/40027
【修改方案】
在
torch_npu._inductor.dvm.config中增加enable_matmul_fusion。该开关默认关闭,仅在设置INDUCTOR_DVM_ENABLE_MATMUL_FUSION=1时注册 DVM matmul template lowering,未启用时保持原有 DVM/Inductor 路径不变。为
aten.mm、aten.bmm、aten.addmm与aten.baddbmm注册 DVM template lowering,并新增DvmTemplateBuffer保存 matmul template FX 图及逻辑 placeholder 到实际 IR 输入的绑定关系。生成 wrapper 调用前恢复真实参数;不满足 DVM shape/type 规则的场景继续走原有 fallback。保留
addmm的原始算子形态以进入 template lowering,补齐baddbmm的转置标注、图构建参数透传和 DVM codegen;保持addmm/baddbmm的alpha、beta语义。K=1 的mm、bmm下沉为逐元素mul,K=1 的addmm下沉为mul与add组合。在
NpuDvmScheduling的 template codegen 中复用NpuMetaScheduling的 traced-graph 构图与回退能力。DVM matmul template 仅支持合法 pointwise epilogue 的纵向融合;prologue、horizontal fusion、reduction、template-to-template、group/numel 不一致及不支持的 broadcast 场景均不融合。移除
_is_view_only_graph对纯view、reshape、_unsafe_viewepilogue 的额外拒绝逻辑,统一由既有 pointwise、shape、依赖和广播合法性检查决定是否进入 DVM template 融合路径。补充
mm、bmm、addmm、baddbmm的 template fusion 回归覆盖,并覆盖 K=1、multi-user 输出、view/view+pointwise epilogue、view 输入及同一 buffer 多 view 场景。【资料变更】
不涉及。
【接口变更】
不涉及客户可见接口变更。
【功能验证】
python3 -m py_compile,通过。git diff --check,通过。性能收益
该优化将 matmul 与下游 pointwise/view 融合为单个 DVM mix kernel,减少中间张量读写和 kernel launch。以下为 v2.9.0 同功能实测(端到端统计已排除首次编译与 warm-up):
GLM-4-9B Chat 对应吞吐约提升 9.34%。GPT-OSS-20B 的 device kernel duration 为 -0.62%,当前不作为性能收益结论。
详细测试配置与完整数据见:https://gitcode.com/Ascend/pytorch/merge_requests/40027
【CheckList】