MLIR 与多层次 IR

MLIR 用方言把「一种 IR 走到底」拆成多层次的渐进降级。本文讲解方言与操作的组织方式、渐进式降级的时机与收益、pass pipeline 的调度,以及最终如何衔接 LLVM 后端。

1. 为什么需要多层次 IR

一句话总结: 一种 IR 无法同时表达高层语义与低层细节,MLIR 的答案是让多种 IR 共存,并规定它们之间如何渐进转换。

传统编译器通常只有一两种 IR:前端产出高层 IR,中端把它降低到低层 IR,后端再降到机器码。这个结构在通用语言上工作得很好,但遇到领域专用场景就暴露了问题。

传统两层 IR 的困境:
  高层 IR 需要表达:张量、卷积、矩阵乘、循环分块、内存层级
  低层 IR 需要表达:寄存器、指令选择、栈帧、调用约定
  中间缺失的部分(循环结构、访存模式、并行映射)无处安放

如果把这些都塞进高层 IR,它会变成一个大杂烩,每条优化 pass 都要处理所有可能的形式;如果都塞进低层 IR,那么高层的结构信息(比如「这是一个卷积」)在降低时就丢失了,优化器只能看到一堆循环。

MLIR 的思路是把这条链路切成多个明确层次,每一层用自己合适的抽象:

层次抽象典型方言
领域层张量、算子、图tosa、stablehlo、linalg
结构层循环、仿射映射、分块affine、scf、linalg
通用层函数、内存、控制流func、memref、cf
低层指针、算术、内建llvm、arith、builtin
// 一个高层次操作:linalg 卷积,语义完整、可被领域优化识别
%0 = linalg.conv_2d ins(%input, %filter : tensor<1x28x28x1xf32, tensor<3x3x1x8xf32>)
                      outs(%init : tensor<1x26x26x8xf32>) -> tensor<1x26x26x8xf32>

这个操作在 MLIR 里是一个单一操作,优化器可以整体地做算子融合、布局选择;而在传统 IR 里它已经被展开成六层嵌套循环,那些高层机会全都消失了。

2. 方言体系

一句话总结: 方言是一组操作、类型与属性的命名空间,它让不同层次的抽象可以共存于同一份 IR 中,并各自定义自己的验证与优化规则。

2.1 操作与方言

一句话总结: 操作是 MLIR 的基本单位,方言是操作的命名空间;一份 IR 里可以同时出现多种方言的操作,这被称为「混合 IR」。

// 混合 IR:scf、arith、memref 三个方言共存
func.func @add(%a: memref<16xf32>, %b: memref<16xf32>) {
  %c0 = arith.constant 0 : index
  %c16 = arith.constant 16 : index
  %c1 = arith.constant 1 : index
  scf.for %i = %c0 to %c16 step %c1 {
    %x = memref.load %a[%i] : memref<16xf32>
    %y = memref.load %b[%i] : memref<16xf32>
    %s = arith.addf %x, %y : f32
    memref.store %s, %a[%i] : memref<16xf32>
  }
  return
}

定义一个新方言需要三部分:

// 用 TableGen 声明一个操作(简化)
def MyAddOp : MyDialect_Op<"add", [Pure]> {
  let arguments = (ins AnyType:$lhs, AnyType:$rhs);
  let results = (outs AnyType:$result);
  let assemblyFormat = "$lhs `,` $rhs attr-dict `:` type($result)";
  let hasVerifier = 1;      // 自定义验证器
}
组成作用例子
操作计算的基本单位arith.addf、linalg.matmul
类型值所属的集合tensor<4x4xf32>、memref<?xi8>
属性编译期常量信息affine_map<(i,j)->(i,j)>
接口跨方言的通用能力内存效果接口、形状推断接口
# 用 mlir-opt 查看与操作一份 IR
mlir-opt input.mlir --canonicalize --cse -o out.mlir
mlir-opt input.mlir --pass-pipeline='builtin.module(func.func(canonicalize,cse))'

2.2 类型与属性

一句话总结: MLIR 的类型系统把「值的形状与布局」也纳入类型,属性则承载编译期已知的映射关系,这两者共同让高层结构信息在降级前一直可见。

// 同一份数据的三种类型表达,代表三个不同的抽象层次
tensor<4x8xf32>              // 值语义:不可变的张量值
memref<4x8xf32>              // 内存语义:有地址的缓冲区
memref<4x8xf32, affine_map<(i,j)->(j,i)>>  // 带布局的缓冲区(转置)
// 属性承载编译期信息:仿射映射描述索引变换
#map = affine_map<(d0, d1) -> (d0 * 4 + d1)>
%0 = affine.load %buf[%i, %j] : memref<16xf32>   // 索引经 #map 变换

这个设计的意义在于:布局选择可以作为属性在 IR 中显式存在,而不是隐式地体现在地址计算里。优化器可以在不改动循环结构的前提下尝试不同的布局,只需换一个属性值。

3. 渐进式降级

一句话总结: 渐进降级不是一次把高层 IR 变成低层 IR,而是分多次、每次只跨一个抽象层次,让每一层都有机会做该层特有的优化。

// 阶段一:领域操作 -> 结构层(卷积变成循环)
linalg.conv_2d ins(%in, %filt) outs(%out) -> ...
// 降级后:
scf.for %n = ... { scf.for %h = ... { scf.for %w = ... {
  scf.for %kh = ... { scf.for %kw = ... {
    // 累加计算
  }}}}}}
// 阶段二:结构层 -> 通用层(循环变成 cf 分支)
// 阶段三:通用层 -> 低层(memref 变成裸指针运算)
%ptr = llvm.getelementptr %base[%i] : (!llvm.ptr, i64) -> !llvm.ptr, f32
%v = llvm.load %ptr : !llvm.ptr -> f32
降级阶段输入抽象输出抽象保留的机会
领域到结构算子循环与分块循环变换、分块、向量化
结构到通用仿射与张量memref 与控制流内存提升、缓存优化
通用到低层内存与函数指针与内建指令选择、寄存器分配
低层到目标LLVM IR机器码由 LLVM 后端完成

渐进降级的核心收益是每层只做自己能做的事。循环分块在结构层做,因为那时循环还可见;向量化也在结构层做,因为那时还有张量与仿射映射的信息;一旦降级到 LLVM IR,这些结构信息已经展开成地址计算,再想分块就困难得多。

# 分步降级,每步之间可以插入自己的 pass
mlir-opt --linalg-generalize-named-ops input.mlir -o s1.mlir
mlir-opt --convert-linalg-to-loops s1.mlir -o s2.mlir
mlir-opt --convert-scf-to-cf s2.mlir -o s3.mlir
mlir-opt --convert-to-llvm s3.mlir -o s4.mlir

4. pass pipeline 与调度

一句话总结: MLIR 的 pass 可以运行在任意操作上(而非只在模块上),这带来了极大的灵活性,也带来了「何时运行、运行在哪个层级」的调度复杂度。

// pass 的嵌套运行:外层在模块上,内层在函数上
builtin.module(
  func.func(canonicalize, cse, loop-invariant-code-motion),
  convert-scf-to-cf,
  convert-to-llvm
)
// 用 C++ API 构建 pipeline
void buildPipeline(OpPassManager &pm) {
  pm.addNestedPass<func::FuncOp>(createCanonicalizerPass());
  pm.addNestedPass<func::FuncOp>(createCSEPass());
  pm.addPass(createConvertSCFToCFPass());
  pm.addPass(createConvertToLLVMPass());
  // 降级到 LLVM 之后交给 LLVM 的优化管线
  pm.addPass(createReconcileUnrealizedCastsPass());
}
调度问题表现应对
pass 顺序顺序不同结果不同明确依赖关系,写成固定序列
层级选择在模块还是函数上跑按作用域选择,粒度越细越易并行
迭代到不动点需要反复运行用 -pass-pipeline 中的 repeat 或自定义循环
合法性检查降级后残留不合法操作用 legality 声明与动态合法性
调试难度pipeline 长,出错难定位每步 -o 落盘或 --mlir-print-ir-after-all
# 定位 pipeline 中哪一步出了问题
mlir-opt --pass-pipeline='...' --mlir-print-ir-after-all input.mlir 2> ir_trace.log
# 或者用 crash reproducer 保存失败现场
mlir-opt --verify-diagnostics input.mlir

一个实用的经验是:pass 之间的耦合越少,pipeline 越容易维护。MLIR 的规范做法是每个 pass 只声明自己需要什么(比如「我要求输入是 Linalg 方言」),不假设上一个 pass 做了什么。这样 pass 可以被自由组合与重排,也让调试时的二分定位成为可能。

5. 与 LLVM 后端衔接

一句话总结: MLIR 的终点通常不是机器码,而是 LLVM IR:通过 convert-to-llvm 把所有方言降级为 llvm 方言,再交给 LLVM 完成指令选择与寄存器分配。

// 降级完成后的 llvm 方言 IR(接近 LLVM IR)
llvm.func @add(%arg0: !llvm.ptr, %arg1: !llvm.ptr) {
  %0 = llvm.mlir.constant(0 : i64) : i64
  %1 = llvm.getelementptr %arg0[%0] : (!llvm.ptr, i64) -> !llvm.ptr, f32
  %2 = llvm.load %1 : !llvm.ptr -> f32
  llvm.return
}
# 生成 LLVM IR 并交给 llc 或 clang
mlir-opt --convert-to-llvm input.mlir | mlir-translate --mlir-to-llvmir -o out.ll
llc out.ll -o out.s
clang out.ll -o out
衔接方式用途命令
翻译成 LLVM IR复用 LLVM 全部后端mlir-translate –mlir-to-llvmir
直接生成目标代码绕过 IR 文本mlir-translate 加 ExecutionEngine
混合编译部分函数用 MLIR通过外部函数声明链接
JIT 执行交互式与调试ExecutionEngine 的 JIT 模式
// 在进程内 JIT 执行 MLIR
mlir::ExecutionEngineOptions opts;
opts.transformer = ...;                 // 挂载降级 pipeline
auto engine = mlir::ExecutionEngine::create(module, opts);
auto fn = engine->lookupPacked("add");

unrealized conversion cast 是衔接时最常见的报错来源。降级过程中若某个操作的类型还没被转换,MLIR 会插入一个占位 cast 保证 IR 合法。如果最后这些 cast 仍然存在,说明降级不完整,此时 reconcile-unrealized-casts 会失败并报错——它其实是一个很有用的完整性检查。

6. 工程实践

一句话总结: 使用 MLIR 的工程决策集中在三处:选择哪个上游方言作为入口、降级到哪个层次后交给 LLVM、以及如何组织自己的方言与 pass。

一个典型的领域编译流程:
  领域前端(自定义语法/框架图)
     |
     v
  高层方言(自定义或 stablehlo)    <- 领域优化:算子融合、常量折叠
     |
     v
  linalg 方言                       <- 分块、向量化、布局选择
     |
     v
  scf 与 memref                     <- 循环变换、内存提升
     |
     v
  llvm 方言 -> LLVM IR -> 机器码    <- 指令选择、寄存器分配
决策选项取舍
入口方言自定义 / stablehlo / tosa自定义灵活但生态少;上游方言有现成 pass
降级终点LLVM 方言 / 目标方言前者复用 LLVM,后者可控但工作量大
pass 组织单一大 pass / 多个小 pass小 pass 可组合可调试,大 pass 性能好
验证策略每步验证 / 只在末端验证每步验证慢但定位快
# 用 mlir-opt 快速试验一段 IR 的降级效果
mlir-opt --linalg-tile='tile-sizes=32,32' --convert-linalg-to-loops test.mlir
mlir-opt --test-vectorization test.mlir        # 观察向量化是否触发

一个常见的误区是「用了 MLIR 就自动获得高性能」。实际上 MLIR 只提供表示与转换的基础设施:它让「分块」「向量化」这类变换变得容易实现与组合,但变换本身的质量仍取决于代价模型与调优。上游提供的 linalg 系列 pass 在常见形状上表现良好,遇到特殊形状与硬件时,仍需要自己写针对性的 pass。

7. 实现要点与陷阱

一句话总结: MLIR 的坑主要来自「抽象层次切换时的信息丢失」与「pass 顺序的隐含依赖」,两者都表现为「IR 合法但结果不对」。

陷阱表现应对
过早降级高层优化机会丢失尽量在高层完成领域与结构优化
残留 unrealized cast降级不完整,后端报错用 reconcile-unrealized-casts 检查
pass 顺序隐含依赖换顺序结果不同显式声明前置条件,不依赖副作用
自定义验证器缺失非法 IR 静默通过为每个操作实现 verifier
混合 IR 类型不匹配跨方言传值失败用 unrealized cast 或统一类型
调试信息未保留降级后无法定位源码用 location 传播,保留源位置
// 陷阱:过早把张量降级成 memref,丢失了值语义,后续无法做算子融合
// 差:进入 pipeline 就是 memref 操作
// 好:保持在 tensor 上做融合,最后一步再 bufferize
// 陷阱:location 丢失导致报错无法定位
// 每个操作都应带 location,用 TableGen 的 `let hasVerifier` 之外还要注意
// 生成 IR 时传 location:builder.create<arith::AddFOp>(loc, a, b)
# 用 location 信息定位问题
mlir-opt --mlir-print-debuginfo input.mlir | grep "loc("

bufferization 是 MLIR 里最值得单独提的一个环节:它把值语义的 tensor 转成有地址的 memref,同时决定在哪里分配缓冲区、能否复用。这个决策直接影响内存占用与拷贝次数,也是许多性能问题的根源。one-shot bufferize 用冲突分析决定缓冲区的原地复用,但它对别名与生命周期的假设需要程序员理解,否则容易出现「本该拷贝却被复用」的错误。

8. 总结

环节要点
动机单一 IR 无法兼顾高层语义与低层细节
方言操作的命名空间,允许混合 IR 共存
类型与属性布局作为属性显式存在,可独立替换
渐进降级分多次跨层,每层保留该层特有的优化机会
pass 调度pass 可运行在任意操作上,需显式管理顺序与层级
与 LLVM 衔接降级到 llvm 方言后翻译成 LLVM IR
完整性检查unrealized cast 残留说明降级不完整
bufferization值语义到内存语义,决定分配与复用
常见误区MLIR 提供基础设施,不自动带来性能

MLIR 真正的贡献不在于某个具体的优化,而在于它把「编译器的抽象层次」变成了可编程的一等公民:方言可以自由定义,层次之间的转换可以自由组合,每个层次都能承载适合它的信息。这解决了长期以来「领域专用编译器要么重复造轮子、要么硬塞进通用 IR」的两难。代价是引入了一整套新的概念与调试负担——pass 顺序、合法性声明、降级完整性,每一项都需要工程经验。至此,本批六篇从目标格式、循环变换、异常机制、类型求解、构建确定性一路走到多层次 IR,它们共同勾勒出编译器工程的一个基本事实:每一个抽象层次的引入,都是用复杂度换取某种能力,而工程判断的价值就在于知道何时该付这个代价。

延伸阅读

继续阅读

探索更多技术文章

浏览归档,发现更多关于系统设计、工具链和工程实践的内容。

全部文章 返回首页

「compiler」更多文章

  1. 可复现构建与确定性输出
  2. 约束求解与类型类
  3. 异常处理编译与栈展开