已开启
DVM 适配图模式(torch.compile)流程总结 #3
hbhu_bin创建于  5月27日
hbhu_bin成员
5月27日 创建

DVM 适配图模式(torch.compile)流程总结

基于 torch_npu v2.7.1 / Ascend NPU + DVM 后端的源码分析

1. 角色与目录结构

DVM 适配 torch.compile 涉及两个并列模块,加上一个上层 IR 框架 anir:

torch_npu/_inductor/
├── dvm/                          DVM 后端实现
│   ├── mlir_fusion.py            DvmMlirFusionPatch — codegen 入口替换
│   ├── op_emitter.py             DVM_OP_REGISTRY + emitter 函数
│   ├── graph_build.py            DvmCodegenInterpreter — fx→DVM IR 翻译
│   ├── graph_fusion.py           DvmOpSupport — 图分组判定
│   ├── decomp.py                 DVM 自定 decomposition (sigmoid/gelu/tanh)
│   └── fx_pass.py                fx graph 改写 pass
│
├── mfusion/                      MFusion 上层图分组框架
│   ├── graph_fusion.py           MFusionPatch — 注册 post_grad_custom_post_pass
│   ├── decomp.py                 mfusion decomposition tweaks
│   └── subgraph_registry.py
│
└── ascend_npu_ir/                anir (Ascend NPU IR) — MLIR 框架
    └── ascend_npu_ir/
        ├── config.py             GENERATE_LIST 默认值定义
        └── npu/
            ├── npu_lowering.py   读 GENERATE_LIST,决定走 lowering 还是 fallback
            └── codegen/mlir.py   NpuMlirKernel / NpuMlirScheduling
模块 职责
MFusion 上层框架:把 fx graph 切成子图(DVM 子图 + ExternKernel 节点)
anir 中层 IR:MLIR-based 表达,控制哪些 op 走 generate vs fallback
DVM 后端实现:把分配到的子图 emit 成 DVM kernel 源码

2. 完整编译流程(按时间线)

Python: torch.compile(model, dynamic=...)
   │
   ▼
[1] Dynamo trace fx graph
   │  节点形如 aten.X.Y(...)
   ▼
[2] AOT Autograd + inductor decomposition
   │  - core_aten_decompositions 把 batch_norm/sigmoid/... 拆成基础 op
   │  - dvm/decomp.py: patch_decomp() 注册 DVM 自家的 decomp (gelu/tanh)
   │  - mfusion/decomp.py: patch_mfusion_decomp() 注册 mfusion tweaks
   ▼
[3] post_grad_custom_post_pass = mfusion_graph_fusion / dvm_graph_fusion
   │  ⚠️ 默认不启用!
   │  - MFusionPatch.enable(): 需 env TORCHINDUCTOR_ENABLE_MFUSION=1
   │  - DvmGraphFusionPatch.enable(): 需 `with DvmGraphFusionPatch():`
   │
   │  若启用 DvmGraphFusionPatch (走 dvm_graph_fusion):
   │  ├─ 扫 fx 节点
   │  ├─ 用 DvmOpSupport.is_node_supported 判定每个节点能否进 DVM 子图
   │  │     dvm/graph_fusion.py:125:
   │  │         if node.target in GRAPH_FUSION_SUPPORT_OP:
   │  │             _, rule = DVM_OP_REGISTRY.get(node.target) ← 读 DVM_OP_REGISTRY [点①]
   │  │             return rule(node)
   │  │         return False
   │  ├─ 用 union-find 把 supported 节点聚成连通子图
   │  └─ 把每个子图替换成一个 "fused" call_module 节点
   │
   │  若启用 MFusionPatch (走 mfusion_graph_fusion):
   │  ├─ fx graph → torch-mlir
   │  ├─ AKG mfusion fuse_and_optimize(mlir_str)  (需 AKG 包)
   │  └─ torch-mlir → fx graph
   │
   │  DCNv2 默认两个都不启用,本步骤实际跳过,直接进入 [4]
   ▼
[4] inductor 标准 lowering 阶段
   │
   │  每个 fx 节点要走 fallback_node_due_to_unsupported_type 判定:
   │     被 mlir_fusion.py:_patch_lowering_type_checks 替换为:
   │         if node.target in DVM_OP_REGISTRY:                 ← 读 DVM_OP_REGISTRY [点②]
   │             return not rule(node)
   │         return not common_rule(node)
   │
   │  返回 True → 走 FallbackKernel (生成 torch.ops.aten.X.default 调用)
   │  返回 False → 调用 register_lowering 注册的 lowering 函数生成 inductor IR
   │
   │  anir npu_lowering.py 还会读 GENERATE_LIST [读取点]:
   │     gen_set = set(config.GENERATE_LIST)
   │     gen_set 之外的 op → fallback
   ▼
[5] inductor scheduler 排程 + DVM 子图 codegen
   │
   │  对每个 mfusion 子图:
   │     调 NpuMlirKernel.codegen_kernel = _codegen_dvm_kernel    ← DvmMlirFusionPatch
   │        ↓
   │     DvmCodegenInterpreter.run() 遍历子图节点
   │        在 call_function 时:
   │           if target not in DVM_OP_REGISTRY:               ← 读 DVM_OP_REGISTRY [点③]
   │              raise NotImplementedError
   │           func, _ = DVM_OP_REGISTRY.get(target)
   │           return func(*args)    # 生成 "k.add(x,y)" 字符串
   │
   │  对每个 FallbackKernel:
   │     生成 torch.ops.aten.X.default(...) Python 调用
   ▼
[6] 生成 output_code.py
   │  - dvm_fused_*_build 函数 (DVM kernel 源码)
   │  - call(args) 函数 (wrapper,含 DVM kernel launch + fallback 调用)
   ▼
[7] 运行时:
   │  call(args) 被反复调用
   │  ├─ DVM kernel: aclrtLaunchKernelWithHostArgs(...) 启动
   │  └─ fallback: ATen dispatcher → torch_npu → aclnnXxx

3. DVM_OP_REGISTRY

位置torch_npu/_inductor/dvm/op_emitter.py:11

结构{aten_op: (emitter_func, rule_func)} 字典

注册方式:用 @register_dvm_op(...) 装饰器:

@register_dvm_op(aten.add.Tensor, aten.add.Scalar)
def add(x, y):
    return f"k.add({x}, {y})"

@register_dvm_op(aten.mm.default, aten.bmm.default, rule=mm_rule)
def matmul(x, y, trans_a, trans_b):
    return f"k.matmul({x}, {y}, {trans_a}, {trans_b})"

装饰器参数(如 aten.add.Tensor)是 registry 的 key;函数名只是 Python 标识符,与 ATen op 名称无关(注意 op_emitter.pydef select(...) 其实注册的是 aten.where,是命名陷阱)。

rule 函数:在 op_emitter.py 中定义:

rule 作用
common_rule 默认规则:args 必须都是 fp16/bf16/fp32(DVM_SUPPORT_FLOAT_TYPE);node 输出 dtype 必须在 DVM_SUPPORT_TYPE
mm_rule mm/addmm/bmm:仅 fp16/bf16;inner axis ≤ 65280;输出至少一维 > 256
where_rule where:args[1:] 是 fp,输出是 DVM_SUPPORT_TYPE
full_rule / cast_rule 各自针对性放宽

作用点(最多 3 处,但默认只生效 2 处)

# 阶段 位置 用法 默认生效?
dvm_graph_fusion 图分组 dvm/graph_fusion.py:125 (DvmOpSupport.is_node_supported) dvm_graph_fusion 流程中判定节点能否进 DVM 子图,需先在 GRAPH_FUSION_SUPPORT_OP 里,再查 registry 取 rule 检查 ❌ 仅 DvmGraphFusionPatch 显式启用时(with DvmGraphFusionPatch():
inductor lowering 前判定 mlir_fusion.py:295 (_fallback_node_due_to_unsupported_type) inductor 调 fallback_node_due_to_unsupported_type 时被替换为这段:查 registry 取 rule,决定是否走 fallback DvmMlirFusionPatch 模块加载时自动启用
DVM codegen dvm/graph_build.py:138 (DvmCodegenInterpreter.call_function) 对 DVM 子图的每个节点查 emitter 函数,生成 k.add(...) 这样的 DVM IR 字符串 ✅ DvmMlirFusion codegen 路径一定走

重要更正:点 ① 实际上 默认不启用MFusionPatch.enable()TORCHINDUCTOR_ENABLE_MFUSION=1DvmGraphFusionPatch.enable() 只在显式 with context 才启用。DCNv2 默认配置下只有 DvmMlirFusionPatch 启用,所以 DVM_OP_REGISTRY 实际只在点 ② 和 ③ 起作用


4. GENERATE_LIST

位置torch_npu/_inductor/dvm/mlir_fusion.py:40(赋值),原始默认值在 ascend_npu_ir/config.py:246

结构:list of aten.X op

当前 DCNv2 上的内容(部分):

anir_config.GENERATE_LIST = [
    aten.mul, aten.add, aten.sub, aten.div,
    aten.clamp_min, aten.clamp_max,
    aten.maximum, aten.minimum,
    aten.abs, aten.reciprocal, aten.log, aten.exp, aten.pow,
    aten.sqrt, aten.rsqrt, aten.neg,
    aten.lt, aten.le, aten.gt, aten.ge, aten.eq, aten.ne,
    aten.where, aten.expand,
    aten.var_mean, aten.sum, aten.mean,
    aten.full, aten.relu, aten.scalar_tensor,
    aten.unsqueeze, aten.squeeze,
    # aten.reshape,        ← 注释掉(启用会触发 dyn_shape 稳定性 bug)
    # aten.clone,
    prims.convert_element_type,
    torch.ops.npu.npu_dtype_cast,
    ...
    triton_kernel_wrapper_mutation,
]

作用点(唯一)ascend_npu_ir/ascend_npu_ir/npu/npu_lowering.py:46

gen_set = set()
for fn in config.GENERATE_LIST:
    gen_set.add(fn)
    if isinstance(fn, torch._ops.OpOverloadPacket):
        for overload in fn.overloads():
            gen_set.add(getattr(fn, overload))

# gen_set 之外的 op 全部走 fallback

含义:anir 层"哪些 op 走 generate 路径,哪些 fallback"的白名单。


5. 两个 list 的差异速查

DVM_OP_REGISTRY GENERATE_LIST
位置 dvm/op_emitter.py dvm/mlir_fusion.py:40(设置)
anir/npu_lowering.py:46(读取)
作用层 DVM 后端 anir 中层 IR 框架
数据结构 dict[op, (emitter, rule)] list[op]
注册方式 @register_dvm_op(...) 装饰器 直接列表赋值
作用点数 3 处 1 处
粒度 每个 op 有 rule(带条件) 仅 op 类型,无条件
影响 ① fallback 判定 ② 图分组 ③ codegen emit anir lowering 阶段 generate vs fallback 决策

两者必须配套:op 在 GENERATE_LIST 里告诉 anir "给我生成代码",在 DVM_OP_REGISTRY 里告诉 DVM "怎么生成"。任何一边缺了,节点最终都会 fallback。


6. 一个 op 如何最终决定命运(决策表)

每个 fx 节点要过四道关:

关 1: 节点 dtype 是否在 DVM_SUPPORT_TYPE / DVM_SUPPORT_FLOAT_TYPE?
关 2: 节点 target 是否在 GENERATE_LIST?
关 3: 节点 target 是否在 DVM_OP_REGISTRY?rule 是否通过?
关 4: 节点 target 是否在 GRAPH_FUSION_SUPPORT_OP?

实际命运分类:

命运 触发条件 典型例子
进 DVM 融合块(kernel 内 emit) dtype OK + GENERATE_LIST 有 + DVM_OP_REGISTRY 有 emitter + GRAPH_FUSION_SUPPORT_OP 有 + rule 通过 fp32 的 add/mul/sub/sqrt/relu
进 DVM 但借道改写 dtype OK + GENERATE_LIST 有原 target,但 NPU lowering 把 target 改写成 DVM 能 emit 的形式 aten.unsqueeze (fp32) → 改写为 aten.reshape
被 decomposition 提前消灭 inductor decomposition 在 fx graph 形成之前就拆成基础 op aten._native_batch_norm_legit_no_training → sub/mul/sqrt/add
走 reinterpret_tensor(短路) 节点是纯 view,下游 ExternKernel 能吃 strided input 启用 reshape 后的 matmul 边界 reshape
fallback (ExternKernel) dtype 不通过 / 不在 GENERATE_LIST / 没 emitter / rule 拒绝 i64 unsqueeze, fp32 mm, permute, select.int, embedding, cat

7. fallback 机制详解

含义:当 op 不能被 inductor 编译/融合时,wrapper 代码直接生成 torch.ops.aten.X.default(...) 调用,运行时通过 ATen dispatcher 走标准实现。

wrapper 代码示例

# 融合(不 fallback):
dvm_fused_add_mul_0.run(buf1, buf2, buf3, stream=stream0)

# fallback:
buf0 = torch.ops.aten.unsqueeze.default(arg0_1, 0)

fallback 路径成本:每次 ~5-10us host

torch.ops.aten.X.default(...)
   → Python C 扩展
   → ATen dispatcher
   → torch_npu 实现
   → aclnnXxx API
   → CANN 运行时
   → NPU kernel launch

reinterpret_tensor(轻量替代):当节点是 view、下游能吃 strided 输入时,inductor 在 wrapper 里直接生成:

buf6 = reinterpret_tensor(buf5, (s0, 352, 1), (352, 1, 1), 0)  # ~0.1us,C++ inline

比 fallback 便宜 ~50 倍。


8. DCNv2 实例:每个 op 的命运分析

8.1 进 DVM 融合块的 op

来源 op 备注
MLP fc1/fc2 后段 (BN+ReLU 展开) sub, mul, sqrt, reciprocal, expand, add, maximum fp32,全部 elementwise,进 dvm_fused_BN_addmm_relu_*
Cross net 加法链 add, mul dvm_fused_add_mul_*
aten.unsqueeze (fp32) ×2 unsqueeze lowering 改写成 reshape,进 dvm_fused_add_mul_*
Sigmoid 展开 neg, exp, add, reciprocal, mul inductor decomposition 把 sigmoid 拆成这些基础 op

8.2 fallback 的 op(每个 ~5-10us host)

op 数量 fallback 原因
aten.unsqueeze (i64) 1 dtype i64 不通过 common_rule
aten.add.Tensor (i64) 1 同上
aten.embedding.default 1 DVM 没 gather 原语,未在 GRAPH_FUSION_SUPPORT_OP
aten.permute.default 9 没 DVM emitter(DVM 没 transpose 原语),未在 GENERATE_LIST
aten.select.int 4 同上(lowering 转 slice+squeeze,slice 也没 emitter)
aten.mm.default 5 mm_rule 拒 fp32 + 输出维度太小
aten.cat.default 1 没 emitter
合计 22 ~110-220us host 开销

8.3 走 reinterpret_tensor 的 op(启用 reshape 后)

启用 aten.reshape 后,4 个 cross matmul decomposition 边界的 reshape 从 dispatch 调用变成 reinterpret_tensor,host 各 ~5us → ~0.1us。


9. 关键源码位置索引

概念 文件:行号
DVM_OP_REGISTRY 定义 dvm/op_emitter.py:11
@register_dvm_op 装饰器 dvm/op_emitter.py:167
common_rule dvm/op_emitter.py:65
mm_rule dvm/op_emitter.py:106
GENERATE_LIST 赋值 dvm/mlir_fusion.py:40
GENERATE_LIST 读取 ascend_npu_ir/.../npu_lowering.py:46
MFusionPatch.enable mfusion/graph_fusion.py:924
DvmMlirFusionPatch.enable dvm/mlir_fusion.py:357
_patch_lowering_type_checks dvm/mlir_fusion.py:264
DvmOpSupport.is_node_supported dvm/graph_fusion.py:122
DvmCodegenInterpreter.call_function dvm/graph_build.py:133
_codegen_dvm_kernel dvm/mlir_fusion.py:85
_define_dvm_kernel dvm/mlir_fusion.py:112
dvm/decomp.py: patch_decomp dvm/decomp.py:155

10. 调试与验证套路

看 DVM 子图边界

  • torch_compile_debug/*/torchinductor/*/output_code.py 看 wrapper 里有哪些 torch.ops.aten.* 调用(fallback)和 dvm_fused_*.run(...) 调用(融合块)
  • ir_post_fusion.txtExternKernelSchedulerNode(FallbackKernel) 节点

看 DVM kernel 内部 op

  • output_code.py 里每个 def dvm_fused_*_build(k): 函数体即为 DVM IR 调用序列
  • 上方 """class <lambda>:...""" 注释是该子图的 fx graph 源(lowering 后的 target)

修改后必须步骤

  1. 修改 dvm/mlir_fusion.pydvm/op_emitter.py
  2. 清 pyc 缓存:find .../torch_npu/_inductor -name '__pycache__' -exec rm -rf {} +
  3. 清 inductor cache:rm -rf $TORCHINDUCTOR_CACHE_DIR
  4. 重跑测试

验证一个 op 进 DVM 的条件

from torch_npu._inductor.dvm.op_emitter import DVM_OP_REGISTRY
print(aten.X.Y in DVM_OP_REGISTRY)            # 是否注册 emitter
# 再查 GENERATE_LIST 是否包含
# 再看 fx graph 里实际的 target(lowering 可能改写)

11. 已知瓶颈与扩展方向(按 DCNv2 优化收益排序)

方向 难度 预期收益
让 i32 args 通过 common_rule(修 DVM_SUPPORT_FLOAT_TYPE) ~15us host(i64→i32 + unsqueeze/add 进 DVM)
修复 dyn_shape=True reshape 稳定性,让 GENERATE_LIST 启用 reshape 默认开 中(C++ 侧) ~70us host(5 reshape 转 DVM/reinterpret)
给 permute 加 emitter + GENERATE_LIST(DVM IR 需 transpose 原语) ~70us host(9 个 permute)
给 select.int / slice 加 emitter(DVM IR 需 slice 原语) ~20us host(4 select)
Cube-Vector epilogue 融合(mm + add/mul/relu 一个 kernel) ~35us host + ~10us device
GatherV2/embedding 进 DVM ~15us host + ~6us device
Concat 通过 inductor memory planning 消除 ~10us host
结构性天花板:Dynamo dynamic shape guard ~217us preparing time,超出 DVM 后端能解决范围
likedislike
Hhbhu_bin成员
5月27日 修改了issue 的描述
Hhbhu_bin成员
5月27日 修改了issue 的描述