本质上不管是静态图还是动态图都是由3个部分组成的Tensor、算子(Ops)、依赖关系Edge。这是我们在图捕获得到静态/动态图之后需要这个图Lowering到相应的中间层IR做相关的优化。
我们假设你有一个简单的 PyTorch 模型,比如一个包含 Linear -> ReLU -> Linear 的简单全连接网络。下面我为你详细拆解它进入 MLIR 后,经历各种 Pass 的"奇幻漂流"。
整个过程可以分为 5 个主要阶段。
阶段一:前端接入与高层 IR 生成 (Frontend & High-Level IR)
PyTorch 本身是一个动态图框架,MLIR 需要静态结构。所以第一步是捕获计算图。
- TorchScript / FX Graph 捕获:
PyTorch 先将模型转换为TorchScript或者FX Graph(一种静态的 Python 级别的计算图)。 - Torch Dialect 生成 (Torch-MLIR):
通过torch-mlir项目,将 PyTorch 的图转换为 MLIR 的torchdialect 。
此时 IR 看起来非常像 PyTorch 代码:torch.operator表示具体的算子(如torch.aten.linear,torch.aten.relu)。torch.tensor表示张量类型(带有 shape 和 dtype)。- 这时的 IR 是最高层的,完全保留了 PyTorch 的语义。
阶段二:高层图优化与 Dialect 转换 (Graph Optimization)
现在 IR 在 MLIR 手里了。第一步是"去 PyTorch 化",将其转换为更通用的 MLIR 方言。
- Torch -> Linalg / TOSA 转换:
调用convert-torch-to-linalg或convert-torch-to-tosaPass。torch.aten.linear被分解为linalg.matmul+linalg.add(加上 bias)。torch.aten.relu被转换为linalg.generic(一个逐元素的最大值操作)或者arith.maxf。torch.tensor变成了tensordialect。
- 常量折叠与死代码消除 (Canonicalization):
此时会运行canonicalizePass。比如如果权重是常量,可能会在此时进行计算;如果某个算子对最终结果没影响,会被删掉。
此时的 IR 状态: 混合了 tensor(数据)和 linalg(计算)方言。这是 MLIR 最适合做高层图优化的层级。
阶段三:中层循环优化与 Bufferization (Loop Optimization & Memory Planning)
这是 MLIR 区别于传统编译器最精彩的部分。linalg 只是说"做个矩阵乘法",但没说怎么做。
- Tiling (循环分块):
调用linalg-tilePass。考虑到 Cache 大小,编译器决定把大矩阵切成小块(比如 32x32)。linalg.matmul被转化为scf.for循环嵌套,内部包含小的linalg.matmul。
- Fusion (算子融合):
这是 AI 编译器的关键。linalg.matmul后面跟着linalg.generic(ReLU)。
通过linalg-fusePass,编译器把 ReLU 的计算融合进了矩阵乘法的循环里。这样就不需要把矩阵乘法的结果写回内存再读出来做 ReLU,大大节省了带宽。 - Bufferization (内存分配):
之前我们用的是tensor(抽象值,不可变)。现在必须变成memref(内存缓冲区,可变)。
调用iree-bufferize或one-shot-bufferizePass。tensor类型变成memref类型。- 编译器在此刻进行内存规划 :分析哪些内存可以复用,插入
memref.alloc和memref.dealloc,或者尽量提升到栈上分配。
- Lowering to SCF / Affine:
linalg结构被完全打散,变成scf.for/affine.for循环和memref.load/memref.store操作。
此时的 IR 状态: memref(内存) + scf/affine(循环) + arith(标量计算)。
阶段四:底层硬件相关优化 (Low-Level Optimization)
现在代码已经很像 C 语言了,接下来要针对硬件(CPU/GPU)做优化。
- Vectorization (向量化):
调用scf-vectorize或linalg-vectorizePass。
将scf.for循环中的标量操作(一次算一个 float)转换为vectordialect 操作(一次算 8 个 float,利用 AVX2/AVX512 指令)。arith.addf->vector.add(SIMD 指令)。
- 循环展开与不变量提取:
运行affine-loop-unroll等 Pass,进一步优化指令流水线。 - Memref 降级为裸指针:
调用memref-to-llvmPass。memref的多维结构被展平,变成llvm.ptr(裸指针)和 GEP 计算。
阶段五:LLVM IR 生成与后端编译 (LLVM Lowering)
这是最后一步,把所有剩余的 MLIR 方言全部翻译成 LLVM IR。
- 转换为 LLVM Dialect:
运行convert-scf-to-cf(结构化控制流变跳转),convert-arith-to-llvm,convert-func-to-llvm等 Pass。
此时,所有的scf.for变成了cf.br(基本块跳转),所有的arith.addf变成了llvm.add。 - 导出 LLVM IR:
调用mlir-translate --mlir-to-llvmir。
生成标准的.ll文件。 - LLVM 后端:
LLVM 接手,进行指令调度、寄存器分配,最终生成.o或.s汇编文件,链接成可执行程序。
总结:一个 Pass Pipeline 的伪代码
在实际工程(如 IREE 或 Torch-MLIR)中,这一长串过程是被封装在一个 Pass Pipeline 里的。看起来大概是这样:
text
// 1. 接入与高层转换
pass pipeline:
convert-torch-to-linalg
canonicalize
inline
// 2. 中层优化 (核心魔法发生的地方)
linalg-tile (tile sizes: 32, 32, 32)
linalg-fuse
one-shot-bufferize
convert-linalg-to-loops
// 3. 底层与硬件优化
scf-vectorize
affine-loop-unroll
convert-vector-to-llvm
convert-memref-to-llvm
// 4. 收尾
convert-scf-to-cf
convert-arith-to-llvm
convert-func-to-llvm
reconcile-unrealized-casts
在这个过程中,你可以清晰地看到:
- Dialect 是逐渐混合又逐渐统一的: 一开始全是
torch,然后变成tensor+linalg,再加入scf,再加入memref,最后全部统一为llvm。 - 不是所有 Dialect 都走完全程: 比如
torchdialect 在阶段二就消失了,linalg在阶段三消失了。它们只在最适合自己的层级存在,做完优化就"功成身退"。 - 优化是分层进行的:
- 在
linalg层面做算子融合(图优化)。 - 在
scf层面做循环分块(访存优化)。 - 在
vector层面做SIMD 指令生成(指令级优化)。
- 在
这就是 MLIR 设计哲学的魅力:让最适合的 Dialect 在最合适的阶段做最擅长的事。