到第 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
2
3
4
5
6
7
8
9
fn clamp(x: i64, lo: i64, hi: i64) -> i64 {
if x < lo {
return lo;
}
if x > hi {
return hi;
}
return x;
}

对应的 Sprout 方言 IR:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
module {
func.func @clamp(%x: i64, %lo: i64, %hi: i64) -> i64 {
%lt = arith.cmpi slt, %x, %lo : i64
%r0 = scf.if %lt -> i64 {
scf.yield %lo : i64
} else {
%gt = arith.cmpi sgt, %x, %hi : i64
%r1 = scf.if %gt -> i64 {
scf.yield %hi : i64
} else {
scf.yield %x : i64
}
scf.yield %r1 : i64
}
return %r0 : i64
}
}

这段 IR 混合了三个方言:func 定义函数,arith 做比较,scf 表达结构化的 if-else。比较和分支仍然保持着源程序的嵌套结构,而不是被打散成基本块和跳转。

再看一个用到 Sprout 自定义操作的例子——数组创建与边界检查:

1
2
3
4
5
6
func.func @get_element(%arr: memref<?xi64>, %idx: i64) -> i64 {
%len = memref.dim %arr, %c0 : memref<?xi64>
sprout.bounds_check %idx, %len : i64
%val = memref.load %arr[%idx] : memref<?xi64>
return %val : i64
}

sprout.bounds_check 是 Sprout 方言独有的操作。它的语义是:如果 %idx 不在 [0, %len) 范围内,触发运行时错误并终止。在后续降低中,它会变成一个比较加条件分支加错误调用的序列。但在 Sprout 方言层级,它是一个原子操作,优化 pass 可以识别它、移动它或消除冗余的检查。

降低管线

从 Sprout 方言到可执行代码,经过三级降低:

1
2
3
4
5
6
7
8
9
10
11
12
13
Sprout dialect + scf + arith

│ (1) Sprout → 结构化中层

scf + arith + memref + func

│ (2) 结构化中层 → LLVM 方言

llvm dialect

│ (3) LLVM 方言 → LLVM IR 文本

LLVM IR (.ll)

第一步:Sprout → 结构化中层。 这一步消除所有 sprout.* 操作。sprout.bounds_check 展开为 arith.cmpiscf.if 加对运行时错误函数的 func.callsprout.array_new 展开为 memref.alloc 加长度记录。sprout.match 展开为嵌套的 scf.if 序列。(sprout.match 的降低规则留作练习 2:将每个 arm 映射为 scf.if 的嵌套链。)降低完成后,IR 中不再有 sprout. 前缀的操作。

sprout.bounds_check 的降低过程值得展开:

1
2
3
4
5
6
7
8
9
10
11
12
// 降低前
sprout.bounds_check %idx, %len : i64

// 降低后
%c0 = arith.constant 0 : i64
%below = arith.cmpi slt, %idx, %c0 : i64
%above = arith.cmpi sge, %idx, %len : i64
%oob = arith.ori %below, %above : i1
scf.if %oob {
func.call @__sprout_bounds_error(%idx, %len) : (i64, i64) -> ()
// 该函数标记为 noreturn
}

第二步:结构化中层 → LLVM 方言。 scf.for 变成 llvm.brllvm.cond_br 构成的循环结构;memref.load 变成 llvm.loadarith.addi 变成 llvm.add。MLIR 自带这一步的标准转换 pass。

第三步:LLVM 方言 → LLVM IR。 mlir-translate 工具把 LLVM 方言的 .mlir 文件转换为 .ll 文本文件,之后进入我们已经熟悉的 optllc → 链接流程。

每一步降低只消除一个层级的抽象。这意味着在每一层都可以运行只关心该层级的优化。

每个层级的优化机会

分层带来的好处是优化可以在最合适的抽象层级进行。

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
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
use melior::{
Context, dialect::DialectHandle,
ir::{Module, Block, Location, Operation, Region, Type, Value},
};

fn register_sprout_dialect(context: &Context) {
// 方言注册在 melior 中通过 DialectRegistry 完成,具体实现见练习 1
}

fn build_bounds_check(
context: &Context,
block: &Block,
index: Value,
length: Value,
loc: Location,
) {
// 比较 index >= length
let cmp = block.append_operation(
arith::cmpi(context, arith::CmpiPredicate::Sge, index, length, loc)
);
let cond = cmp.result(0).unwrap().into();

// 条件分支:越界时调用 panic 函数
let then_region = Region::new();
let then_block = Block::new(&[]);
// 调用运行时的 bounds_check_failed
then_block.append_operation(func::call(
context, "sprout_bounds_check_failed", &[index, length], &[], loc,
));
then_block.append_operation(scf::r#yield(&[], loc));
then_region.append_block(then_block);

let else_region = Region::new();
let else_block = Block::new(&[]);
else_block.append_operation(scf::r#yield(&[], loc));
else_region.append_block(else_block);

block.append_operation(scf::r#if(cond, &[], then_region, else_region, loc));
}

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 类型。如果传入了 f64i1,验证失败。
  • sprout.array_new 的长度操作数必须是 i64,结果类型必须是 memref<?xT>
  • sprout.match 必须至少有一个分支区域,每个分支区域必须以 sprout.yield 终结。
  • sprout.call 引用的函数符号必须在模块中存在,参数数量和类型必须匹配签名。

验证器的价值在于把 IR 构造错误提前暴露。不经验证就降低,错误的 IR 可能产生难以追踪的后端崩溃。这与第 06 篇类型检查的思路一样:越早拒绝非法输入,调试成本越低。

运行验证的方式:

1
2
3
4
let module: Module = /* ... */;
// module.as_operation().verify() 检查整个模块
// 失败时返回包含位置信息的诊断
assert!(module.as_operation().verify());

MLIR 的诊断系统支持附带源码位置的错误消息。如果我们在构造操作时传入了 Sprout 源码的 Location,验证失败时错误消息会指向原始的 .spr 文件位置——这与第 20 篇的源码映射衔接。

降低管线的测试策略

降低管线的正确性需要逐层验证。

最直接的办法是在每一步降低之后转储 IR 文本,检查输出是否符合预期。MLIR 支持 --mlir-print-ir-after-all 选项,在每个 pass 之后打印 IR。

1
2
3
4
sproutc --emit=mlir-sprout  programs/clamp.spr  # Sprout 方言
sproutc --emit=mlir-mid programs/clamp.spr # 降低到 scf+arith+memref
sproutc --emit=mlir-llvm programs/clamp.spr # 降低到 llvm 方言
sproutc --emit=llvm-ir programs/clamp.spr # 最终 LLVM IR

对于 sprout.bounds_check 的降低,测试需要覆盖:正常索引不触发错误路径;负索引和越界索引各触发一次运行时错误;长度为零的数组对任意索引都越界。这些测试可以直接复用第 15 篇的受检数组测试用例。

练习

  1. 给 Sprout 方言增加一个 sprout.print 操作,用于打印一个 i64 值。定义它的操作数(一个 i64)、结果(无)和验证规则。实现它到中层的降低:先把 i64 值存入一个 memref<1xi64>,再生成一个 func.call @__sprout_print_i64 调用。考虑一下:为什么不直接生成 func.call 而要经过 sprout.print?因为在 Sprout 方言层级,优化 pass 可以识别连续的 sprout.print 并合并为批量输出,或者在确认值是常量时直接折叠。

  2. 画出 sprout.match 降低到 scf.if 嵌套的过程。假设有三个分支,被匹配值是一个枚举标签(i64)。每个分支先比较标签,匹配则执行分支体。注意最后一个分支可以省略比较——如果前两个都不匹配,必然是第三个(穷尽性已由第 17 篇的类型检查保证)。

上一篇:28 - 用 WIT 封装组件接口
下一篇:30 - 张量方言与 Lowering 到 LLVM