从零编写现代编译器 29 - MLIR 入门:定义 Sprout 方言
到第 14 篇为止,Sprout 的编译管线把 Typed HIR 降低到 LLVM IR,再交给 LLVM 的 pass pipeline 做优化和代码生成。这条路径能用,但有一个结构性的缺陷:LLVM IR 太低了。循环变成了跳转和 phi,数组访问变成了 getelementptr,函数调用的高层语义全部摊平成了调用约定。一旦进入 LLVM IR,编译器就很难再回答"这段代码原来是一个模式匹配"或者"这个数组创建应该做边界检查"这类高层问题。
MLIR(Multi-Level Intermediate Representation)正是为了解决这个问题而设计的。它允许在同一个框架内定义多个抽象层级的 IR——称为方言(dialect),每个方言保留特定层级的语义信息,再逐层降低到最终的机器表示。这一篇我们为 Sprout 定义一个自己的 MLIR 方言,并搭建从 Sprout 方言到 LLVM IR 的降低管线。
MLIR 的核心结构
MLIR 的设计围绕几个递归嵌套的概念展开。
操作(Operation) 是最基本的单元。每个操作有一个名字(例如 arith.addi)、零到多个操作数、零到多个结果、一组属性(编译期常量)和零到多个区域。操作名由方言前缀和操作名两部分组成,中间用点分隔。
区域(Region) 是基本块的容器。一个操作可以包含区域,区域内部又包含操作,形成递归结构。func.func 操作包含一个区域来容纳函数体,scf.for 操作包含一个区域来容纳循环体。
块(Block) 是操作的线性序列,必须以终结操作结尾。块可以接收参数,用于替代 LLVM IR 中的 phi 节点。这是 MLIR 与 LLVM IR 的一个重要区别——块参数比 phi 更容易分析和变换。
类型(Type) 和属性(Attribute) 可以由方言自行定义。memref<4x4xf64> 是 memref 方言定义的类型,表示一个 4×4 的双精度浮点内存引用;#sprout.source_loc<"main.spr":12:5> 可以是 Sprout 方言定义的属性,在 IR 层级保留源码位置。
方言(Dialect) 是操作、类型和属性的命名空间。MLIR 自带的方言各有分工:arith 提供整数和浮点算术;scf 提供结构化控制流(for、while、if);memref 提供内存引用和读写;func 提供函数定义和调用;llvm 是通向 LLVM IR 的出口。这些方言可以在同一段 IR 中混合使用——一个函数体内既可以有 arith.addi,也可以有 sprout.bounds_check。这种混合正是逐层降低的基础:每一步只替换一部分操作,不需要一次性翻译全部。
参考 MLIR Language Reference 获取完整的语法规范。
定义 Sprout 方言
Sprout 方言的目的是在 MLIR 层级保留 Sprout 语言的特有语义。我们为它定义以下操作:
| 操作 | 语义 | 操作数与结果 |
|---|---|---|
sprout.call |
调用 Sprout 函数,保留 Sprout 级别的调用信息 | 函数符号 + 实参 → 返回值 |
sprout.array_new |
创建 Sprout 数组,携带元素类型和长度 | 长度(i64)→ memref |
sprout.bounds_check |
数组边界检查,检查失败时触发运行时错误 | 索引(i64)+ 长度(i64)→ 无结果 |
sprout.match |
模式匹配,包含多个分支区域 | 被匹配值 → 区域内结果 |
以 MLIR 文本格式写出一个具体的 Sprout 函数。假设源码是:
1 | |
对应的 Sprout 方言 IR:
1 | |
这段 IR 混合了三个方言:func 定义函数,arith 做比较,scf 表达结构化的 if-else。比较和分支仍然保持着源程序的嵌套结构,而不是被打散成基本块和跳转。
再看一个用到 Sprout 自定义操作的例子——数组创建与边界检查:
1 | |
sprout.bounds_check 是 Sprout 方言独有的操作。它的语义是:如果 %idx 不在 [0, %len) 范围内,触发运行时错误并终止。在后续降低中,它会变成一个比较加条件分支加错误调用的序列。但在 Sprout 方言层级,它是一个原子操作,优化 pass 可以识别它、移动它或消除冗余的检查。
降低管线
从 Sprout 方言到可执行代码,经过三级降低:
1 | |
第一步:Sprout → 结构化中层。 这一步消除所有 sprout.* 操作。sprout.bounds_check 展开为 arith.cmpi 加 scf.if 加对运行时错误函数的 func.call。sprout.array_new 展开为 memref.alloc 加长度记录。sprout.match 展开为嵌套的 scf.if 序列。(sprout.match 的降低规则留作练习 2:将每个 arm 映射为 scf.if 的嵌套链。)降低完成后,IR 中不再有 sprout. 前缀的操作。
sprout.bounds_check 的降低过程值得展开:
1 | |
第二步:结构化中层 → LLVM 方言。 scf.for 变成 llvm.br 和 llvm.cond_br 构成的循环结构;memref.load 变成 llvm.load;arith.addi 变成 llvm.add。MLIR 自带这一步的标准转换 pass。
第三步:LLVM 方言 → LLVM IR。 mlir-translate 工具把 LLVM 方言的 .mlir 文件转换为 .ll 文本文件,之后进入我们已经熟悉的 opt → llc → 链接流程。
每一步降低只消除一个层级的抽象。这意味着在每一层都可以运行只关心该层级的优化。
每个层级的优化机会
分层带来的好处是优化可以在最合适的抽象层级进行。
Sprout 方言层级。 编译器知道哪些操作是 Sprout 特有的,可以做语言级优化。例如:连续两个 sprout.bounds_check 检查同一个数组的相邻索引,第二个检查可以弱化为只比较上界。sprout.match 如果某个分支的条件是常量 false,可以直接消除。sprout.call 调用的 Sprout 函数如果足够小,可以在降低之前内联——此时内联决策可以利用 Sprout 的类型信息,而不是在 LLVM 层级靠启发式猜测。
结构化中层。 scf.for 保留了循环的结构,循环变换(交换、分块、展开)可以直接在这个层级表达。如果我们后续引入张量运算(第 30 篇的内容),这一层是做 tiling 和向量化决策的位置。memref 层级的别名分析也比 LLVM 的 ptr 更精确,因为 memref 携带了形状和步幅信息。
LLVM 层级。 这是我们已经用了十几篇的老朋友。寄存器分配、指令选择、窥孔优化、尾调用优化——这些 LLVM 做得很好的事情仍然交给 LLVM。
关键认识是:不存在一个"最好的"IR 层级。高层看得见结构但看不见寄存器,低层看得见机器但看不见意图。多级表示不是增加复杂度,而是让每个层级只回答自己擅长的问题。
在 Rust 中实现
melior 是 MLIR 的 Rust 绑定,基于 MLIR 的 C API 构建。它的 API 围绕 context、module、block、operation 展开,与 MLIR 的概念一一对应。
定义 Sprout 方言的骨架:
1 | |
ODS 是 MLIR 推荐的操作定义方式——用 TableGen 描述操作的签名、约束和文档,自动生成解析、打印和验证代码。对于教学目的,手动构造操作更容易理解内部结构,但生产级方言通常会使用 ODS 减少重复代码。
melior 要求系统安装了 MLIR 的共享库。本篇使用 MLIR/LLVM 19。melior 的版本需要与系统安装的 LLVM 版本匹配——melior 0.20.x 对应 LLVM 19。
方言验证
MLIR 的一个核心设计是每个方言定义自己的验证器(verifier)。验证器在 IR 构造之后、降低之前运行,检查操作是否满足方言定义的约束。
Sprout 方言的验证规则:
sprout.bounds_check的两个操作数必须是i64类型。如果传入了f64或i1,验证失败。sprout.array_new的长度操作数必须是i64,结果类型必须是memref<?xT>。sprout.match必须至少有一个分支区域,每个分支区域必须以sprout.yield终结。sprout.call引用的函数符号必须在模块中存在,参数数量和类型必须匹配签名。
验证器的价值在于把 IR 构造错误提前暴露。不经验证就降低,错误的 IR 可能产生难以追踪的后端崩溃。这与第 06 篇类型检查的思路一样:越早拒绝非法输入,调试成本越低。
运行验证的方式:
1 | |
MLIR 的诊断系统支持附带源码位置的错误消息。如果我们在构造操作时传入了 Sprout 源码的 Location,验证失败时错误消息会指向原始的 .spr 文件位置——这与第 20 篇的源码映射衔接。
降低管线的测试策略
降低管线的正确性需要逐层验证。
最直接的办法是在每一步降低之后转储 IR 文本,检查输出是否符合预期。MLIR 支持 --mlir-print-ir-after-all 选项,在每个 pass 之后打印 IR。
1 | |
对于 sprout.bounds_check 的降低,测试需要覆盖:正常索引不触发错误路径;负索引和越界索引各触发一次运行时错误;长度为零的数组对任意索引都越界。这些测试可以直接复用第 15 篇的受检数组测试用例。
练习
-
给 Sprout 方言增加一个
sprout.print操作,用于打印一个i64值。定义它的操作数(一个i64)、结果(无)和验证规则。实现它到中层的降低:先把i64值存入一个memref<1xi64>,再生成一个func.call @__sprout_print_i64调用。考虑一下:为什么不直接生成func.call而要经过sprout.print?因为在 Sprout 方言层级,优化 pass 可以识别连续的sprout.print并合并为批量输出,或者在确认值是常量时直接折叠。 -
画出
sprout.match降低到scf.if嵌套的过程。假设有三个分支,被匹配值是一个枚举标签(i64)。每个分支先比较标签,匹配则执行分支体。注意最后一个分支可以省略比较——如果前两个都不匹配,必然是第三个(穷尽性已由第 17 篇的类型检查保证)。
