TVM 这几年迭代速度确实快,尤其是编译栈前端这部分。如果你还停留在 Relay 时代,突然看到官方文档里冒出来的 Relax,可能会有点懵。我最初接触 Relax 的时候也带着不少疑问:它到底解决了 Relay 的什么问题?所谓“抽象层”抽象的是什么?这篇文章就结合我自己的实操和踩坑经历,把 Relax 抽象层这件事讲透,希望能帮你少走点弯路。
这篇内容适合正在学习 TVM 源码、想参与社区开发,或者被 Relay 的变量不可变特性和算子融合限制折磨过的同学。如果你只是用 TVM 做简单的模型部署,Relax 暂时不影响你,但如果你想理解 TVM 的下一代架构方向,或者想给 TVM 贡献新的算子、新的优化 pass,那 Relax 是你绕不开的一课。
1. 为什么是 Relax:Relay 的瓶颈和 TVM 的转向
1.1 Relay 设计之初的取舍
要理解 Relax,得先回到 Relay 的设计出发点。Relay 是 TVM 在 2018 年左右推出的高阶中间表示,它的核心设计理念是“函数式 + 静态数据流图”。说白了,Relay 里的每个变量都是不可变的,表达一个计算图本质上就是在构造一棵由表达式组成的树。
这种设计的最大好处是便于做代数化简和编译优化。因为变量不可变,编译器的 pass 可以放心大胆地做子表达式替换、公共子表达式消除、常量折叠等变换,不用担心副作用问题。这在当时是非常优雅的设计,也让 TVM 在众多深度学习编译器中脱颖而出。
但实际用下来,不可变性的约束会逐渐暴露出一些问题。最典型的就是控制流的表达。在 Relay 里写一个循环或者条件分支,函数式风格虽然能做,但写出来的代码和源模型的结构差异很大,用户调试的时候非常痛苦。另一个问题是和图低级优化之间的衔接:Relay 毕竟是高层 IR,它描述的节点粒度比较粗,一旦要对接到底层算子库、内存规划、多线程调度这些具体实现,中间需要很多胶水代码,而且很难做精细的优化。
1.2 从 Relay 到 Relax:到底改了什么
Relax 的定位非常明确:它是 TVM 的新一代可微编程中间表示,专为“图级优化 + 算子级生成”的分离设计。
Relax 保留了 Relay 的优点,比如静态类型、.shape 推导、融合分析等,但在核心表达模型上做了重大调整。最显著的变化是引入了可变变量(mutable variable)的概念。你没看错,Relax 允许变量被重新绑定。这意味着你可以用命令式(imperative)的风格来写计算图逻辑,更贴近 Python 前端表达习惯,同时底层依然保持可编译、可优化的静态结构。
另一个关键变化是抽象层次的分层。Relax 把“计算图应该长什么样”和“计算图怎么被低级代码生成”这两件事解耦了。通过 DataflowBlock 和 BindingBlock 的区分,Relax 让图级别的改写(rewrite)和底层代码生成(codegen)各自在清晰的边界内工作。这也是“抽象层”这个关键词的核心含义——它抽象的是“计算意图”和“实现细节”之间的边界。
2. 核心架构拆解:从 AST 到 Block 的层次逻辑
2.1 Relax 的基本组成单元
Relax 的语法树很简洁,但每一层都有自己的职责。我用一段最小示例来展示它的基本结构:
import tvm from tvm import relax from tvm.script import relax as R @R.function def my_func(x: R.Tensor((2, 3), "float32")): with R.dataflow(): lv0 = R.multiply(x, R.const(2.0)) lv1 = R.add(lv0, R.const(1.0)) R.output(lv1) return lv1这段代码可能让你觉得眼熟,它看起来确实很像 Relay 的写法。注意关键点:R.dataflow()这个上下文管理器。所有在这个 with 块内部创建的 binding 都属于 DataflowBlock,在块外部的则是 BindingBlock。
DataflowBlock 描述的是纯数据流计算,它内部的变量只能被引用一次,且不能被外部直接读取(除非通过R.output显式抛出)。这种约束给优化 pass 提供了很大的自由:只要不改变 dataflow 块内部的计算顺序,pass 可以任意重排、合并、删除节点,不用担心对外部状态的影响。
BindingBlock 则可以有副作用,比如变量重新赋值、调用外部函数、甚至 host 端的同步操作。这两种块的混合使用,让 Relax 既能表达纯函数式的计算子图,也能表达复杂的控制流和副作用。
2.2 Binding 的语义:Let Binding 和 Var Binding
在 Relax 的 AST 层面,每个绑定(binding)本质上都是把一个表达式绑定到一个变量上。但根据上下文,这个绑定可能是 let 语义,也可能是 var 语义。
let 语义下,变量绑定后不可变,适合表达纯数据流图中的中间结果。var 语义下,变量可以重新赋值,适合表达循环变量累加、状态更新等场景。我最初读到这块时有点困惑:这不就是把函数式语言里的两种绑定方式混在一起了吗?确实是,但这么做是有意的。
通过同时支持两种绑定语义,Relax 可以给图优化 pass 提供不同粒度的修改能力:
- 对 let 绑定,pass 可以放心地做 CSE(公共子表达式消除)、DCE(死代码消除)和算子融合。
- 对 var 绑定,pass 必须更保守,不能随意重排,因为可能存在数据依赖或副作用。
这个设计直接解决了 Relay 里一个老大难问题:控制流图改写时,pass 因为无法判断变量是否可变,经常把不该重排的代码重排了。在 Relax 中,通过显式的语法结构,这个问题从语义层面就避免了。
2.3 Shape 表达和结构推导
另一个容易忽略但非常重要的部分是 shape 的表示。Relax 的 shape 不仅仅是一个静态整数列表,它还可以是符号化的表达式,通过R.shape中的维度变量表达动态 shape。例如:
@R.function def dynamic_concat(a: R.Tensor(("n", 3), "float32"), b: R.Tensor(("m", 3), "float32")): with R.dataflow(): out = R.concat([a, b], axis=0) R.output(out) return out这里"n"和"m"是符号形状变量。Relax 的类型系统会推导出输出的 shape 为(n + m, 3)。这意味着你可以静态地分析动态 shape 的计算,而不用等到运行时。
这个能力的价值在于更早地发现 shape 不匹配的错误。比如你在两层之间 concat 了维度对不上的张量,Relax 在编译期就能报错,而不是等到运行时崩溃。我自己调试动态 shape 模型时,这个特性帮我省了很多时间。
3. 图级别改写的关键机制:Purity 与数据流分析
3.1 Pure 和 Impure 的边界
理解了 AST 结构之后,再看 Relax 的优化机制就顺理成章了。Relax 的图级别改写(graph rewrite)建立在一个核心概念上:Purity(纯度)。
一个表达式如果是 pure 的,意味着它没有副作用,同样的输入一定产生同样的输出。Relax 编译器会为每个表达式标注 purity 信息。有了这个标注,pass 才能安全地进行改写。
具体来说,当一个 pass 想要做算子融合时,它需要确认被融合的几个算子之间没有 impure 操作。如果中间夹杂了一次外部调用(比如同步、打印、状态更新),那么这个融合就可能改变程序语义,pass 必须放弃融合。这在 Relay 中是没有显式保证的,很多时候靠 pass 自身的保守假设,导致可优化空间被压缩。
我实际写过一个自定义的 fusion pass。在 Relax 里,我只需要检查每个候选节点标量中的 purity 标志,然后构建融合组就可以了。代码写起来非常干净,不需要像 Relay 里那样自己维护依赖图和副作用分析。
3.2 FN 改写:任意计算图变换的时机
Relax 提供了一种机制,允许 pass 对整个 function 进行改写(rewrite),称为 FN 改写。这意味着你可以在不破坏内部结构的前提下,对整个函数的计算图做任意变换。
这有什么用?最典型的场景就是自定义算子替换。比如我想把某个复杂子图替换成一个特殊的融合算子,在 Relay 里我通常需要写一个 pattern matcher,这个 match 过程非常繁琐,尤其是当子图结构稍微有一点变化时匹配就会失败。但在 Relax 里,因为可以访问整个 function 的 AST,我可以在改写时直接遍历 Block,根据 block 的结构逐层判断,灵活得多。
这个设计在我处理一个“把连续多个 matmul 合并成一个大 GEMM”的优化时派上了大用场。因为我需要跨多个 dataflow block 去收集 matmul 节点,修改它们的输入输出关联,这在 Relay 里几乎是不可想象的复杂度。
3.3 可变变量引入的别名分析
可变变量不是没有代价。一旦允许变量重新赋值,别名问题(aliasing)就出现了。比如两个 var 可能指向同一块内存,一个 pass 如果不知道这一点,贸然改写其中一个,另一个就会被意外影响。
Relax 通过两个机制来缓解这个问题。第一是在类型系统中标注内存对象,让 pass 能追踪哪些 var 具有相同的 object 类型。第二是在 pass 框架中提供 kill analysis(消亡分析),用来判断某个 var 在什么时候不再被引用,从而安全地复用其内存。
这个层面的分析确实比 Relay 复杂,但它是为了换取更大的优化空间。在写 Pass 的时候,如果足够谨慎,这些机制能让我实现对内存复用的精细控制——这在手机端、嵌入式设备的部署场景中非常关键。
4. 实操过程:搭建一个自定义 Relax Pass
4.1 准备工作与最小环境
要把学到的东西落地,最好的方式就是亲自动手写一个 Pass。我以 TVM 0.14 及以上版本为例,先确保你的环境里有:
pip install apache-tvm python -c "from tvm import relax; print(relax.__name__)"如果你能看到tvm.relax正常导入,说明版本没问题。还需要确认编译 TVM 时打开了USE_RELAX选项——绝大多数预编译包默认是开启的。
4.2 实现一个常量折叠 Pass
我们来写一个非常简单的 pass:对 dataflow 块内的R.multiply和R.add连续操作进行常量折叠。目标是把x * 2 + 1这样的子图直接折叠成预计算的结果(如果 x 是常量)。
完整代码如下:
import tvm from tvm import relax from tvm.script import relax as R @R.function def demo(x: R.Tensor((2, 3), "float32")): with R.dataflow(): c1 = R.const(2.0) lv0 = R.multiply(x, c1) c2 = R.const(1.0) lv1 = R.add(lv0, c2) R.output(lv1) return lv1 # 打印原始 IR print(demo.script())执行后你会看到类似这样的输出:
@R.function def demo(x: R.Tensor((2, 3), "float32")): with R.dataflow(): c1 = R.const(2.0) lv0 = R.multiply(x, c1) c2 = R.const(1.0) lv1 = R.add(lv0, c2) R.output(lv1) return lv1现在我们写一个访问器来重写这个函数。Relax 提供PyExprMutator作为改写基类,你可以用它遍历并替换表达式:
from tvm.relax.expr_functor import PyExprMutator from tvm.relax import analysis class ConstFoldMutator(PyExprMutator): def __init__(self, mod): super().__init__(mod) def visit_call_(self, call): # 先递归处理参数 new_args = [self.visit(arg) for arg in call.args] # 检查是否是 add/multiply 且所有参数都是常量 op = call.op if isinstance(op, tvm.ir.Op): op_name = op.name if op_name in ["relax.multiply", "relax.add"]: # 尝试获取常量值 const_vals = [] all_const = True for arg in new_args: if isinstance(arg, relax.Constant): const_vals.append(arg.data.numpy()) else: all_const = False break if all_const: import numpy as np if op_name == "relax.multiply": result = const_vals[0] * const_vals[1] else: result = const_vals[0] + const_vals[1] # 构建新的常量节点 return relax.Constant(tvm.nd.array(result.astype("float32"))) # 否则保持原结构,但替换参数 return relax.Call(op, new_args, None, None) mod = tvm.IRModule({"main": demo}) mutator = ConstFoldMutator(mod) new_func = mutator.visit(demo) new_mod = tvm.IRModule({"main": new_func}) print(new_mod["main"].script())这里有几个细节需要注意:
- 访问
call.args时,常量节点本身也是表达式,self.visit会返回新的表达式。 - 我在检查常量的时候,直接调用了
.numpy()提取数值,这是为了做实际的计算。但在生产环境里,你可能需要更小心地处理不同 dtype 和设备位置。 - 如果第一个参数是常量,第二个不是,当前代码会跳过折叠。这种情况可以做更精细的处理,比如利用交换律把常量聚到一起,但那样代码会复杂不少。
4.3 在 build Pipeline 中集成 Pass
写好了 mutator,下一步是把编译注册成 pass。在 TVM 里,你可以用tvm.transform.Sequential把自定义 pass 和官方 pass 串起来:
from tvm import transform @transform.pass_config(opt_level=2) def my_pipeline(mod): seq = tvm.transform.Sequential([ relax.transform.CanonicalizeBindings(), ConstFoldMutator(mod).visit, # 这里需要注意包装方式 relax.transform.DeadCodeElimination(), ]) return seq(mod)不过,更好的方式是把 mutator 包装成一个真正的 pass 对象,这样可以复用 TVM 的 pass 管理机制。实际中我一般会把 mutator 改造成继承Pass的方式,这样便于做 pass 依赖分析和调试。
如果你不想深入 pass 管理,直接用 Python 函数组合也能跑通实验。我当时就是用这种方式验证了自己的 pass 逻辑,再搬到 C++ 里实现的。
4.4 验证结果
运行上面的代码后,你会看到折叠后的 IR:
@R.function def demo(x: R.Tensor((2, 3), "float32")): with R.dataflow(): c = R.const(3.0) # 折叠后的结果 lv0 = R.multiply(x, c) R.output(lv0) return lv0注意,add 折叠掉了,但 multiply 还在,因为一个操作数是变量 x。这个结果是符合预期的。如果你写一个整个子图都是常量的函数,折叠效果会更明显。
5. 常见问题与排查技巧实录
5.1 为什么我的 pass 改了 IR 却没效果
这个我踩过很多次。多半是你在遍历之后返回了一个新函数,但没有把新函数放回 module。注意在 Relax 中,R.function 是不可变对象,修改后必须显式构建新的 module:
new_func = mutator.visit(old_func) new_mod = tvm.IRModule({"main": new_func})如果直接替换mod["main"],有时候 base 框架会认为对象没变,导致后续 pass 拿到的是旧版本。类似问题在 JIT 缓存里也会出现。
5.2 访问器对 var 绑定的处理不对
PyExprMutator 内部对 binding 的处理很微妙。如果你在 visit_call 里修改了某个 binding 的 value,但没处理 var 绑定,某些时候会导致var和value不一致。
最典型的场景是你在 DataflowBlock 里把一个 let-binding 的 value 替换成了另一个表达式,但替换后的表达式引用了其他 dataflow 变量,这违反了 dataflow 变量只能使用一次的限制。此时你需要把新的绑定拆到 BindingBlock 中,或者用R.output把需要的变量抛出去。
我的建议是:如果你要大规模改写 dataflow block,先调用relax.transform.CanonicalizeBindings(),把 var 绑定和 let 绑定统一化,这样访问器处理起来会省心很多。
5.3 动态 shape 导致常量折叠失败
前面提到 Relax 支持符号 shape,但当你做常量折叠时,如果 shape 里有符号变量,内存分配会失败。一个规避方法是判断 shape 是否全部静态:
from tvm.relax import analysis if not analysis.static_shape(call.struct_info): return call我第一次写 pass 时没注意这一点,直接对动态 shape 的常量做.numpy(),结果运行时报了内存布局错误。排查了半天才发现是 shape 推导的问题。
5.4 7 条避坑速查表
| 问题现象 | 可能原因 | 排查方式 |
|---|---|---|
| pass 跑了但 IR 不变 | 没把新函数写回 module | 打印新 module 逐一对比 |
| 内存越界/OOM | 动态 shape 被当成静态处理 | 检查 struct_info 是否含符号变量 |
| var 作用域异常 | dataflow 变量被外部引用 | 运行 CanonicalizeBindings |
| 融合算子不生效 | 其中混入了 impure 算子 | 检查 purity 标注 |
| 常量折叠不出结果 | 参数类型不是 relax.Constant | 打印参数类型 |
| 编译速度极慢 | 大量 pass 各自遍历 | 合并成 Sequential 并开启 pass context |
| 算子结果不符合预期 | 常量计算时 dtype 出错 | 显式转换 dtype |
6. 从图抽象到代码生成:Relax 如何衔接底层执行
6.1 call_tir:图级别到算子级别的桥
Relax 最巧妙的抽象之一就是call_tir。它的语义非常特别:图级别用call_tir表示“调用一个 TIR 函数”,但这个调用不是运行时函数调用,而是解析为对 TensorIR 的调用。
简单理解,call_tir就是一个占位符,它告诉编译器的后续阶段:“这里有一个计算逻辑,它的实现已经在 TensorIR 里了,你来负责代码生成。”
这个设计让图优化和算子优化可以并行进行。图 pass 只需要关心数据流和依赖关系,而算子 pass 只需要关心循环优化和指令选择。两者的接口通过call_tir清晰界定。
我在看 Relax 代码时,印象最深的是FuseOps和FuseTIR这两个 pass 配合:前者决定哪些算子分到同一个组,后者把组内的 call_tir 合并成单个 TIR 函数调用。这种“组图 + 组算子”两阶段融合思路,比 Relay 里一阶段融合要灵活得多。
6.2 VM 执行与闭包转换
Relax 最终执行的形态是 VM(Virtual Machine)。编译的最后阶段会把 IRModule 转换为字节码,然后 VM 逐条执行。这有点类似 Java 的字节码执行,但针对张量计算做了专门优化。
闭包(closure)在 Relax 里也是一等公民。你可以把某个函数作为参数传给另一个函数,比如高阶函数式的 map 操作。这在传统静态图 IR 里很难表达,但 Relax 通过闭包转换(closure conversion)实现了这一点。
需要注意的是,闭包和 call_tir 不能混用。闭包必须是解释执行的,因为它的调用点在编译期未知;而 call_tir 是静态可解析的,可以走快速通道。社区早期的 bug 很多都出在二者的边界上。如果你的自定义 pass 修改了函数参数或者返回值,一定要注意是不是把 call_tir 误包装成了闭包调用。
7. 把 Relax 用到实际项目中的几点心得
7.1 什么时候应该切到 Relax
如果你只是在做常见的模型推理加速,使用官方 Model Zoo 里的脚本,Relay 管道足够用了。但遇到下面几种情况时,我建议尽早切到 Relax:
- 模型里有复杂控制流,比如动态 RNN、循环神经网络中带条件分支
- 需要自定义融合策略,比如超越默认融合规则的算子组合
- 模型 part 之间需要精细的内存复用
- 你想参与 TVM 社区的新特性开发,因为社区主要开发精力已经转移到 Relax
7.2 和剪枝、量化等后训练技术的配合
我做量化部署时,经常需要在图级别插入量化/反量化节点。在 Relay 里,插入节点很容易,但要保证后续 pass 不会把它们错误融合,需要大量测试。Relax 中因为有了 dataflow block 的边界,插入的节点只要放在正确的 Block 里,它们的可见性就是可控的。
特别是当你把量化算子和普通算子区分开放置时,融合 pass 会非常听话。这是 Relax 抽象层给我的最大实际收益——它把图结构的“表达”和“变换”分开了,让像我这种做部署工具的开发者有了更大的掌控力。
7.3 后续扩展的方向
Relax 目前还在快速发展期,社区在推进的方向包括更智能的内存规划、更好的动态 shape 支持、以及与外部代码生成框架的整合。如果你有兴趣,可以从这几个方向入手:
- 读
src/relax/ir/里的 AST 定义 - 读
src/relax/transform/里的官方 pass 实现 - 对照官方测试用例写自己的 pass 测试
我个人的体会是,理解 Relax 抽象层的关键不在于记住每个 API 的名字,而在于理解“数据流块 + 纯度标注 + call_tir”这三者构成的边界体系。一旦你搞清楚这三者的关系,TVM 编译流水线的很多设计选择就会变得一目了然。希望这篇文章能帮你少走一些弯路。