我标注"执行了什么 Pass "以及"IR 树发生了什么变化 "。
假设我们的模型是:
Y = ReLU(MatMul(A, B))。
"一个静态计算图被转换为 MLIR 后,初始是一堆高层的 Ops(如 torch dialect)。编译器运行一系列 Pass,将这些高层 Ops 逐步转换为中层 Ops(如 linalg, tensor),在这个过程中进行图优化(融合、折叠)。接着,这些中层 Ops 被进一步转换为底层 Ops(如 scf, memref, arith),在这个过程中进行循环优化和内存规划。最后,所有剩余的 Ops 都被统一转换为 llvm dialect 的 Ops,导出为 LLVM IR。"
第 0 步:初始状态(刚从前端进来,纯高层抽象)
这是我们的起点。所有东西都是高层的、声明式的。没有内存概念,没有循环概念。
mlir
// 【状态:纯高层】
module { // builtin dialect
func.func @main(%A: tensor<4x8xf32>, %B: tensor<8x4xf32>) -> tensor<4x4xf32> {
%0 = tensor.empty() : tensor<4x4xf32> // tensor dialect
%1 = linalg.matmul ins(%A, %B) outs(%0) : ... // linalg dialect (矩阵乘法)
%2 = linalg.generic { ... } ins(%1) outs(...) // linalg dialect (代表 ReLU)
return %2 : tensor<4x4xf32>
}
}
第 1 步:运行 Pass convert-linalg-to-loops(只针对 matmul)
编译器决定先把最耗时的矩阵乘法降级。它运行了一个 Pass,只把 linalg.matmul 变成了循环和内存操作 。注意,这个 Pass 不认识 tensor.empty,也不认识 ReLU 的 linalg.generic,所以它们原封不动。
变化: %1 从 linalg 变成了 scf + memref + arith。引入了内存分配 memref.alloc。
mlir
// 【状态:混合状态 1】
module {
func.func @main(%A: tensor<4x8xf32>, %B: tensor<8x4xf32>) -> tensor<4x4xf32> {
%0 = tensor.empty() : tensor<4x4xf32> // tensor dialect (没变)
// --- 原来的 linalg.matmul 消失了,变成了下面这坨 ---
%1_mem = memref.alloc() : memref<4x4xf32> // memref dialect (分配内存)
scf.for %i = 0 to 4 step 1 { // scf dialect (循环)
scf.for %j = 0 to 4 step 1 {
scf.for %k = 0 to 8 step 1 {
%a = memref.load %A[%i, %k] : ... // memref dialect
%b = memref.load %B[%k, %j] : ... // memref dialect
%mul = arith.mulf %a, %b : f32 // arith dialect (标量乘)
%add = arith.addf %sum, %mul : f32 // arith dialect (标量加)
memref.store %add, %1_mem[%i, %j] : ... // memref dialect
}
}
}
%1 = tensor.cast %1_mem : memref<4x4xf32> to tensor<4x4xf32> // 强行转回 tensor 以适配下游
// --- ReLU 还在,因为 Pass 只处理了 matmul ---
%2 = linalg.generic { ... } ins(%1) outs(...) // linalg dialect (没变)
return %2 : tensor<4x4xf32>
}
}
(注:为了简化,我省略了 tensor.cast 的细节,实际中 MLIR 会用 bufferization 接口来处理这种类型不匹配,但核心意思是 %1 已经变成内存了。)
第 2 步:运行 Pass one-shot-bufferize(内存化)
现在编译器发现,IR 里既有 tensor(抽象值)又有 memref(内存),太混乱了。它运行 Bufferization Pass,把所有 tensor 都变成 memref。
变化: tensor.empty 变成了 memref.alloc。ReLU 的 linalg.generic 的输入输出也从 tensor 变成了 memref。
mlir
// 【状态:混合状态 2】
module {
func.func @main(%A: memref<4x8xf32>, %B: memref<8x4xf32>) -> memref<4x4xf32> {
// tensor.empty 变成了 memref.alloc
%0 = memref.alloc() : memref<4x4xf32> // memref dialect
// --- 第一步转换好的 matmul 循环 (没变) ---
%1_mem = memref.alloc() : memref<4x4xf32>
scf.for %i = 0 to 4 step 1 { ... } // scf + memref + arith
// --- ReLU 的 linalg.generic 还在,但类型变了 ---
// 它现在接收 memref,输出 memref
%2 = linalg.generic { ... } ins(%1_mem) outs(%0) // linalg dialect (依然是高层壳子)
return %0 : memref<4x4xf32>
}
}
第 3 步:运行 Pass convert-linalg-to-loops(这次针对 ReLU)
现在轮到 ReLU 了。编译器再次运行 convert-linalg-to-loops,把剩下的 linalg.generic 也变成循环。
变化: %2 从 linalg 变成了 scf + memref + arith。
mlir
// 【状态:纯底层混合 (全是 scf + memref + arith)】
module {
func.func @main(%A: memref<4x8xf32>, %B: memref<8x4xf32>) -> memref<4x4xf32> {
%0 = memref.alloc() : memref<4x4xf32> // memref dialect
// --- 矩阵乘法循环 (没变) ---
%1_mem = memref.alloc() : memref<4x4xf32>
scf.for %i = 0 to 4 step 1 { ... } // scf + memref + arith
// --- 原来的 ReLU (linalg.generic) 消失了,变成了下面这坨 ---
scf.for %i = 0 to 4 step 1 { // scf dialect (循环)
scf.for %j = 0 to 4 step 1 {
%val = memref.load %1_mem[%i, %j] : ... // memref dialect (读)
%zero = arith.constant 0.0 : f32 // arith dialect
%relu = arith.maxf %val, %zero : f32 // arith dialect (ReLU 计算)
memref.store %relu, %0[%i, %j] : ... // memref dialect (写)
}
}
return %0 : memref<4x4xf32>
}
}
注意: 此时,所有的高级 Dialect(tensor, linalg)已经全部消失了!IR 树变成了纯粹的 scf + memref + arith。这时的代码已经非常像 C 语言了。
第 4 步:运行 Pass convert-scf-to-cf 和 convert-memref-to-llvm
现在编译器要准备交给 LLVM 了。LLVM 不认识 scf.for,也不认识 memref。它只认识"基本块(Basic Block)"和"裸指针(Pointer)"。
编译器运行两个 Pass:
convert-scf-to-cf:把scf.for变成cf.br(跳转指令)。convert-memref-to-llvm:把memref变成llvm.ptr和 GEP(指针计算)。
变化: scf 和 memref 消失了,变成了 cf 和 llvm。
mlir
// 【状态:即将进入 LLVM】
module {
// 函数签名变成了裸指针
llvm.func @main(%A: !llvm.ptr, %B: !llvm.ptr) -> !llvm.ptr {
// memref.alloc 变成了 llvm 的堆分配调用
%0 = llvm.call @malloc(...) : ... // llvm dialect
// --- 矩阵乘法循环变成了基本块和跳转 ---
// (这里代码极长,简化表示)
cf.br ^bb1(...) // cf dialect (跳转)
^bb1:
// 指针计算和加载
%ptr = llvm.getelementptr ... // llvm dialect
%val = llvm.load %ptr : ... // llvm dialect
...
// --- ReLU 循环也变成了基本块和跳转 ---
cf.br ^bb_reLU(...)
^bb_reLU:
%relu_ptr = llvm.getelementptr ...
%relu_val = llvm.load %relu_ptr ...
%zero = llvm.mlir.constant(0.0) : f32 // llvm dialect
%relu = llvm.intr.maxnum(%relu_val, %zero) // llvm dialect (调用 LLVM 内置 max)
llvm.store %relu, ...
cf.br ^bb_end(...)
^bb_end:
llvm.return %0 : !llvm.ptr
}
}
第 5 步:运行 Pass convert-to-llvm(最终统一)
最后,编译器运行一个"大扫除"Pass,把剩下的 cf(跳转)和 arith(算术)也全部翻译成 llvm dialect。
变化: 所有非 llvm 的东西全部消失。
mlir
// 【状态:纯 LLVM Dialect】
module {
llvm.func @main(%A: !llvm.ptr, %B: !llvm.ptr) -> !llvm.ptr {
%0 = llvm.call @malloc(...) : ...
// cf.br 变成了 llvm.br
llvm.br ^bb1(...)
^bb1:
%ptr = llvm.getelementptr ...
%val = llvm.load %ptr : ...
// arith.addf 变成了 llvm.fadd
%add = llvm.fadd %val, %val : f32 // llvm dialect
...
llvm.br ^bb_reLU(...)
^bb_reLU:
// arith.maxf 变成了 llvm.intr.maxnum
%relu = llvm.intr.maxnum(%relu_val, %zero) // llvm dialect
...
llvm.return %0 : !llvm.ptr
}
}
第 6 步:导出为 LLVM IR
最后,MLIR 调用 mlir-translate --mlir-to-llvmir,把这棵纯 llvm dialect 的树,翻译成标准的 LLVM IR(.ll 文件)。
llvm
; 【最终的 LLVM IR】
define ptr @main(ptr %A, ptr %B) {
entry:
%0 = call ptr @malloc(...)
br label %bb1
bb1:
%ptr = getelementptr ...
%val = load float, ptr %ptr
%add = fadd float %val, %val
br label %bb_relu
bb_relu:
%relu_val = load float, ptr %relu_ptr
%relu = call float @llvm.maxnum.f32(float %relu_val, float 0.0)
store float %relu, ptr ...
br label %bb_end
bb_end:
ret ptr %0
}
总结
现在可以清晰地看到:
- 第 0 步: 纯
tensor+linalg(声明式,像数学公式)。 - 第 1-2 步: 混搭状态(
tensor和memref共存,linalg和scf共存)。 - 第 3 步: 纯
scf+memref+arith(命令式,像 C 语言)。 - 第 4-5 步: 纯
llvmdialect(底层,像汇编的抽象)。 - 第 6 步: 标准 LLVM IR。
这就是 MLIR "渐进式降级"和"混合方言"的完整慢动作回放。 每一个 Pass 只做自己负责的那一小块转换,逐步把高层抽象"融化"成底层机器码。