1. 为什么Relay撑不住了:先说清楚Relax出现的背景
如果你用过TVM,大概率接触过Relay。这门基于图(Graph)的中间表示在很长一段时间里是TVM的“门面”——从前端框架导入的模型,首先会转成Relay IR,然后经过图优化、算子融合、布局转换,最后落到TensorIR做算子级优化和代码生成。这套链路到今天依然能用,但它处理的问题越来越吃力。
先说结论:Relax不是Relay的简单升级版,而是TVM社区重新审视机器学习编译器抽象层次后提出的新方案。我没有参与过Relax的设计,但看完设计文档、论文以及和社区里朋友的交流后,一个感受非常强烈:Relax的诞生不是某个团队拍脑袋的决定,而是这几年做编译优化的人被“动态形状”“控制流”“闭包”“模块化”这些东西反复折磨之后,痛定思痛的产物。
Relay最初的设计目标,是解决论文里那种“符号化地描述一个静态计算图”的任务。对于2018年前后的卷积网络、ResNet、BERT这种相对规整的模型,Relay的图IR够用。But——这几个模型都有一个共同点:输入的形状基本固定,控制流极少,整个模型在推理阶段几乎是一条直线。可一旦你开始处理这样的模型:输入序列长度不定的Transformer、训练循环里的梯度累积、动态batch的推荐系统、带有while循环的递归神经网络,Relay就暴露出了结构性的短板。
短板具体在哪?我梳理一下。
第一,Relay的图表示太“平”了。图IR天然适合表达数据流,但它不擅长表达“嵌套”“作用域”“绑定关系”这些编程语言层面的概念。一个训练循环里,同一个算子在不同迭代中要被反复执行,还要处理不同迭代间的依赖关系,用图节点堆叠会非常扭曲,图的规模和可读性都会失控。
第二,Relay对动态形状的支持是打了补丁的。TVM里有一个叫Any的动态形状机制,能让某些维度在编译期未知。但Relay的整个优化管线都是围绕静态形状设计的——形状推断、算子融合、内存规划,一旦遇到编译期不知道的形状,很多优化就直接退化成保守模式。这种“退化”在性能上的代价极其惨痛。
第三,Relay被夹在“表达力”和“可优化性”之间进退两难。如果你把Relay当作一种编程语言,它的表达能力不足(比如对闭包、高阶函数的支持不够);如果你把Relay当作一个优化的载体,它又不够底层——很多图重写和算子级重写逻辑找不到合适的挂载位置。这种“高不成低不就”的状态,才是Relay最尴尬的地方。
Relax就是冲着这几根刺来的。它的设计目标不是做一个“更好的图IR”,而是设计一个能够承载完整机器学习编译流程的抽象层次——从高层次的计算描述,到中层次的形状推导、数据流分析,再到低层次的算子实现和执行规划,每一层各司其职,互相之间又有清晰的接口。这个思路,和编译器经典的分层抽象(前端→中端→后端)在精神上一脉相承。
理解这一点很关键:Relax不是一个“新格式”的问题,而是一个“新抽象层次”的问题。这也是为什么标题里强调的是“抽象层”而不是“新语言”。
2. Relax的核心设计:一个带着“形状史”的中间表示
讲Relax的抽象设计之前,先看一段实际代码,感受一下它长什么样。
@I.ir_module class Module: @R.function def main(x: R.Tensor(("m", "n"), "float32"), w: R.Tensor(("n", "k"), "float32")) -> R.Tensor(("m", "k"), "float32"): with R.dataflow(): lv0 = R.call_dps_packed("matmul", (x, w), out_sinfo=R.Tensor(("m", "k"), "float32")) R.output(lv0) return lv0这段代码来自Relax的测试用例,很能说明问题。注意几个关键点:
- 整个模块用
I.ir_module声明,R.function定义了一个“Relax函数”,返回值是一个带符号形状的张量类型。 - 函数内部的
with R.dataflow()块,把所有纯数据流操作圈定在里面。这是Relax对“图”这一概念的新定义——不是全局一张图,而是函数内部的一个数据流子图。 R.call_dps_packed这个调用很有意思,dps是“destination-passing style”的缩写。这意味着Relax的算子调用采用目标传递风格——提前为输出张量分配存储,再把输出地址传给算子函数。这种方法让内存规划变得直接,也让算子实现更容易对接底层的PackedFunc(TVM里的统一函数接口)。
再看函数签名里的R.Tensor(("m", "n"), "float32")——这里的"m"和"n"是符号形状变量,不是具体数字。这就是Relax处理动态形状的方式:形状可以在编译期是符号化的,在运行期被具体值绑定(binding)。
2.1 StructInfo:让“张量形状”成为一等公民
Relax里最核心的新概念之一是StructInfo(结构化信息)。你可以把它理解为“类型系统 + 形状信息 + 存储信息”的复合体。传统编译器里,每个表达式都有一个类型;Relax里,每个表达式有一个StructInfo,它比类型更精细——不仅告诉你这个值是什么类型(Tensor、Tuple、Function等),还告诉你它的形状、数据类型(dtype)、以及它是否已知的具体张量常量(比如某个权重矩阵的直接值)。
为什么要这么设计?因为机器学习编译器的很多优化,依赖的恰恰是“形状信息”而非“类型信息”。举个例子,算子融合时我要知道两个算子的张量形状是否兼容;内存池规划时我要知道所有中间张量的生存期和大小;算子调度时我要知道维度的顺序和步长。这些信息在Relay的类型系统里是零散地挂在张量上的,没有统一抽象;在Relax里,它们都被收拢到StructInfo这个概念下,成为编译器的统一查询接口。
StructInfo还有一个设计上的巧思:它支持Object类型的StructInfo,里面存的是任意Python对象或用户自定义的结构。这意味着Relax不只能描述张量计算,还能描述高层的数据结构(比如训练状态、数据集迭代器),这为以后接入更复杂的执行模式留了后门。
2.2 DataflowBlock:把“纯计算”和“副作用”分开
学过函数式编程的朋友对“纯函数”这个概念不陌生——同样的输入,永远得到同样的输出,不修改任何外部状态。Relax的R.dataflow()块内包含的都是纯计算,块内的变量是“数据流变量”,只会被使用一次且不可被块外引用(除了通过R.output显式输出的值)。这个设计实现了“纯”和“不纯”的清晰切割。
这样做的好处极其实在:编译器可以放心地对dataflow块内部做各种变换,包括算子融合、重写、并行化、重新调度,只要保证块的输入输出语义不变就行,不用担心副作用导致变换不安全。而R.dataflow块外部允许R.call_packed这种带有副作用的调用(比如打印、随机数生成、惰性求值等),这又给执行层的灵活性留了空间。
我一直觉得,dataflow这个概念是Relax里最值得花时间琢磨的地方。本质上,它把“图中的子图”这个非正式概念给类型化、结构化、模块化了。子图不是一个扁平的节点集合,而是带边界、带作用域的代码块。
2.3 符号形状的跟踪方式:ShapeExpr与match_shape
如果说StructInfo是Relax的“骨架”,那么符号形状的跟踪机制就是它的“神经”。Relax里使用ShapeExpr来表达形状表达式,它是一组符号形状变量的约束。比如R.Tensor(("m", "n"), "float32")里的("m", "n")就是一个ShapeExpr。
动态形状带来的最棘手的问题是:当一个算子的输出形状依赖于输入的具体形状时(比如两个矩阵相乘,输出的行数等于左矩阵的行数),我需要一种机制去“传播”和“绑定”形状信息。Relax给出的答案是R.match_shape——一个运行时绑定操作:
x = R.placeholder(("m", "n"), "float32") m, n = R.match_shape(x, ("m", "n"))这段代码的意思是:在运行时,从张量x的实际形状中解构出m和n两个值,并把它们绑定为后续计算中的形状变量。如果实际形状不匹配声明的形状,会触发一个错误。这就是Relax处理动态形状的“桥接层”——编译时类型检查用符号形状,运行时用match_shape把真实形状绑定进来。
这个机制解决了Relay时代的一个大痛点。以前在Relay里,如果形状是动态的,图重写阶段往往直接放弃优化,因为无法确定重写后的形状是否一致。Relax通过match_shape在运行期“兜底”,让编译期可以继续放心大胆地进行符号化优化,到了运行期再根据真实形状决定是否走特殊化路径。
3. 从Relay到Relax:一个编译器的新架构思路
理解Relax,绕不开它和TVM全家桶的关系。直接放一张我心中的架构图(不画图,用文字描述):TVM现在的前端导入阶段,from_torch、from_tensorflow这些前端可以把模型转换成Relay IR,也可以直接转成Relax IR;Relax IR经过绑定(binding)、形状推断、梯度下降等功能模块的处理后,要么走Legalize把高层的R.call_dps_packed算子目标替换成低层的TensorIR函数,要么直接编译成可执行的VMBuild——Relax有自己的一套虚拟机后端,叫Relax VM。
这套架构的核心思想是:Relax负责“结构”和“策略”,TensorIR负责“算子”和“实现”。这个分界的边界感非常干净——Relax关心的是一个计算图如何被组织、如何被规划、内存如何被复用;TensorIR关心的是一个算子的循环如何分块、如何向量化、如何做缓存优化。两层之间的桥梁,是R.call_tir这个调用——Relax可以调用一个TensorIR函数,告诉它输入输出张量的地址,剩下的交给更底层的优化。
3.1 闭包与回调:Relax的“函数是一等公民”
Relax里函数是一等公民,可以直接作为参数传递和返回。这个设计对现代机器学习框架尤其重要,最典型的场景是——训练循环中的梯度更新。你要更新一个模型的参数,实际上是在回调一个“更新函数”;在数据并行的场景里,你要把不同设备上的梯度合并,这本身就是一个高阶函数。
在Relay时代,表达这种“函数回调”非常痛苦,因为图IR本质上是一个静态结构,函数作为数据流动的载体很难在图里体现。Relax通过R.function和PackedFunc回调机制解决这个问题:你可以把Relax的函数编译成一个可调用的对象,传给某个算子执行。这就让Relax的抽象层次高于单纯的“图”——它具备了一门函数式编程语言的基本表达能力,但又能直接对接编译优化和硬件代码生成。
你可能觉得,这不就是“图里有个函数节点”吗?没那么简单。Relax的函数是一等值,意味着你可以写一个map函数,接受一个算子函数和一个张量列表,批量执行——这在Pytorch里是很自然的写法,但放在图编译器里,它就要求编译器拥有“对函数体做类型推断、对函数调用做内联或惰性求值”的能力。这些都是Relax作为“抽象层”在真正做的事情。
3.2 内存规划为何在Relax里变得顺手
传统图编译器做内存池规划,是在整个图作用域上做liveness分析、图着色,这是很经典的编译技术。但Relay的问题是:图的规模一大,中途一旦动态形状发生变化,内存规划就得把整个分析重做一遍——因为一个算子的输出大小改变了,所有下游张量的大小都可能改变。
Relax通过两层设计简化了这件事。第一,dataflow块提供了明确的边界,块内的liveness分析在小范围内,效率极高;块和块之间的内存交互通过显式的输入输出绑定,编译器不需要做全局分析。第二,前面提到的call_dps_packed目标传递风格,让“先分配输出,再填充计算”成为Compile的默认模式,内存的分配和回收时机变得非常明确。
所以当你问“Relax比Relay快在哪”的时候,很多时候答案不是说Relax生成的代码更快,而是Relax让某些优化在动态形状和复杂控制流的情况下仍然能够成功地进行下去。这一点在静态模型上不明显,一旦进入训练场景、动态batch、多模态模型,两者的差距会急剧放大。
3.3 Relax VM:一个为抽象层服务的执行引擎
有人可能觉得,Relax是一个纯编译期的概念,跟运行期没什么关系。错。Relax有配套的虚拟机Relax VM,它的指令集和运行时充分体现了Relax的抽象设计。
Relax VM的指令不是逐条对应算子的普通虚拟机,而是加载一个编译过的Relax模块后,执行其中的函数调用、数据流计算、形状绑定和闭包调用。它跟传统的DL runtime(比如TVM的GraphRuntime)最大的区别是:它支持递归调用、闭包调用、动态形状绑定这些高级控制流能力,而不是一个简单的前馈图执行器。
用VM而不是端到端编译,这个选择的思路也很有意思。纯粹的AOT编译(要么生成整张图的机器码,要么生成整张图的一个二进制)在动态形状和复杂控制流面前很容易失控——你不可能为每一种可能的分支组合都编译一份机器码。VM把“结构性的控制流”留在解释阶段,把“算子计算”下沉到编译后的TensorIR函数。这个“分层执行”的模式是Relax的一个核心特色。
4. 动态形状:Relax的硬仗和应对策略
前面反复提“动态形状”,为了让你更清楚地理解这个问题的分量,我展开讲一下。在传统推理场景中,模型输入尺寸是编译时就能知道的,比如224x224的图像、64长度的句子。这种静态形状的好处是,编译器可以提前知道每一个中间张量的大小,从而精确地分配内存,并且可以大胆展开循环、做算子融合。但如今的模型越来越“不规矩”:
- NLP里的变长输入,每个样本的padding长度不同
- 推荐系统里动态变化的batch size
- 视频模型里不固定的帧数
- 训练时的梯度累积和参数更新引入了跨step的依赖
这些场景下,静态图编译的“提前规划”优势变成了制约:你无法在编译期确定某个张量有多大,很多算子就没法安全地融合、没法用手写kernel替换、没法进行精确的内存池复用。要么退化为动态调度(每次计算都现场查询形状、现场选kernel),要么用Padding补齐到最大形状(浪费算力、浪费显存)。
Relax给的答案是:尽量在编译期做符号化的规划和决策,把必须推迟到运行期的决定压缩到最小。这个策略具体体现在三个方面:符号形状的静态推断、运行时形状绑定、以及对象的延迟求值(R.call_packed加上闭包机制)。当形状无法在编译期统一时,Relax还能用R.match_shape的“惰性分支”在运行时做出决定——比如检测到batch size是32,就走为32优化的kernel;否则走通用kernel。这种分支决策运行时开销极小,但带来了很大的灵活性。
我必须坦诚:动态形状的优化到现在也没有一个万能的银弹。Relax的贡献在于,把这个问题从“IR根本表达不了”提升到了“IR能表达,且能在多套策略之间做权衡”。这个进步,是图编译器时代做不到的。
5. 实操:亲手编译一个Relax模块(附调试经验)
理论讲了一堆,给你一段能直接跑通的实操流程。我这里使用的是TVM的python包(从源码编译或者pip安装的版本,建议用最新的main分支,Relax部分迭代很快),并假设你已经有了一个TVM环境。
5.1 极简环境准备
pip install apache-tvm # 注意:apache-tvm不是所有版本都有Relax,建议源码编译或从nightly包安装如果是从源码编译,记得打开USE_RELAX相关开关。TVM从2022年底开始,Relax已经并入主线,但不同版本的API变化很快,建议锁定某个release版本或跟踪main分支的commit。
5.2 编写并编译一个Relax程序
直接用TVM的python前端编写Relax IR:
import tvm from tvm import relax from tvm.relax.testing import nn # 定义输入形状为符号变量 m = tvm.tir.Var("m", "int64") n = tvm.tir.Var("n", "int64") k = tvm.tir.Var("k", "int64") # 构造一个简单的模块:y = matmul(x, w) + bias @tvm.script.ir_module class MyModule: @R.function def main(x: R.Tensor((m, n), "float32"), w: R.Tensor((n, k), "float32"), bias: R.Tensor((m, k), "float32")) -> R.Tensor((m, k), "float32"): with R.dataflow(): lv0 = R.call_dps_packed("matmul", (x, w), out_sinfo=R.Tensor((m, k), "float32")) lv1 = R.add(lv0, bias) R.output(lv1) return lv1 # 编译 mod = relax.transform.BindParams("main", {}) (MyModule) ex = relax.build(mod, target="llvm")上述代码里,R.call_dps_packed("matmul", ...)中的"matmul"是注册在TVM runtime里的PackedFunc名字。在真实项目中,你要么用R.call_tir调用一个已经实现好的TensorIR算子,要么用R.call_dps_packed加一个自定义的PackedFunc。
5.3 踩过的坑:Relax API极不稳定
操作阶段最大的阻力不是概念难懂,而是API变动太快。我第一次跑到一个能用的例子,发现relax.transform.BindParams这种接口在三个月后就改了名字。建议:
- 依赖官方仓库的
tests/python/relax目录作为“活文档”,里面有大量可运行示例 - 不要依赖网上两三年前的博客教程,极可能是过时API
- 遇到
AttributeError或ImportError,优先去GitHub搜语法,不要硬猜
第二个容易出问题的点是符号形状变量的作用域。如果你用tvm.tir.Var做形状变量,一定要注意它是否被正确绑定到了对应的R.function签名上。很多时候编译器报“shape mismatch”就是符号变量没有传导到后续算子的out_sinfo,仔细检查每个调用的out_sinfo是否都显式声明。
第三个点是dataflow块的边界。块内定义的变量如果不通过R.output显式输出,出了块就不能引用。新手最容易犯的错误是写了return但忘了在R.output里列出返回值变量,结果在后续访问时直接报错。
5.4 核验编译产物
编译完成后,可以验证生成的VM是否可用:
vm = relax.VMClosedModule(ex, tvm.cpu()) # 或者 vm = relax.vm.VirtualMachine(ex, tvm.cpu()) import numpy as np x = tvm.nd.array(np.random.rand(4, 5).astype("float32")) w = tvm.nd.array(np.random.rand(5, 6).astype("float32")) bias = tvm.nd.array(np.random.rand(4, 6).astype("float32")) res = vm["main"](x, w, bias) print(res)这段代码如果顺利跑通,说明你的Relax环境搭建成功。注意这里的m,n,k在运行时被具体值4,5,6绑定,这正好演示了动态形状从符号到具体的绑定过程。
6. 常见误区与设计取舍:我的一些理解与忠告
讲到这里,我觉得有必要专门聊聊大家对Relax的几个常见误解。这些误解我自己也经历过,踩坑之后才逐渐清晰。
6.1 误区一:认为Relax是替代Relay的“新一代Relay”
Relax的目标不完全是“更好的Relay”,它更像是一个更通用的编译抽象层。Relay仍然可以存在——TVM生态里依然有人用Relay前端做模型导入,再转换成Relax继续后端流程。Relax和Relay不是非此即彼的关系,而是抽象层次上的演进。如果非要类比:Relay是“计算图的专用表示”,Relax是“面向ML编译的通用IR”。通用,意味着它既能表达图,也能表达图论之外更丰富的结构(比如训练循环、闭包回调、动态控制流)。
6.2 误区二:认为Relax VM是一个“慢速的解释器”
我看过有人拿Relax VM和TensorRT的AOT编译产物比性能,然后得出结论“Relax很慢”。这属于典型的“关公战秦琼”。Relax VM的确承担了解释执行的任务,但它解释的是编译后的高层结构指令,算子计算本身是调用了经过TensorIR极致优化的代码。VM层的开销和算子计算的开销根本不在一个数量级——算子动辄微秒到毫秒级,而VM指令的执行是皮秒到纳秒级。所以衡量Relax的端到端性能,核心应该看算子的生成质量和调度策略,而不是VM这条指令i>是不是“解释执行”。
当然,VM层确实是Relax在极致性能诉求下的一个代价。如果某个模型完全静态、形状完全固定、不需要任何动态特性,那么跳过VM层直接做端到端AOT编译(甚至直接调用TensorIR的编译产物)仍然会更快。这个取舍Relax是清楚的:它把自己的职责定位为“结构化和可编程”,而不是“无脑压榨最后一滴性能”。
6.3 误区三:把“图优化”等同于Relax的全部价值
Relax真正让人兴奋的地方,从来不是它把图重写做得更华丽,而是它为训练场景和动态形状场景提供了第一公民的支持。这一点和PyTorch 2.0引入compile的思路有异曲同工之妙——都是在“传统静态图编译”和“动态灵活的执行语义”之间寻找一个更优雅的中间点。Relax的dataflow块、match_shape、闭包回调,都是为了在表达能力上对齐现代深度学习开发者的真实需求,而不是守着图优化那一亩三分地。
6.4 关于抽象层次设计的一点个人体会
现在回头看Relax的设计,我最敬佩的一点不是它发明了多少新概念,而是它敢于把“现状”打碎了重新分层次。以往的DL编译器,要么沉迷于“图就是一切”,要么沉迷于“循环就是一切”。Relax的设计者意识到:真正工程化的编译器需要不同关注点的分离——结构归结构、算子归算子、执行归执行。这种“分层”的思维方式,其实是编译器历史上反复出现的母题,只是Relax把这种母题应用到了“深度学习编译”这个年轻的领域。
如果你做的是编译器相关的工作,或者准备学习编译原理,我建议你把Relax当作一个极好的案例去研究。那些dataflow、闭包、VM的设计,在传统编译器教科书里都有自己的对应物,但Relax把它们应用到一个全新的、快速变化的领域,这个“移植”的过程本身就非常值得琢磨。
7. 下一步的探索方向与实践建议
到这里,Relax的核心抽象已经梳理得差不多了。如果你准备深入学习或直接上手用,我有几条比较实际的建议。
第一,从官方仓库的测试用例入手。在apache/tvm仓库的tests/python/relax目录下,有海量的可运行例子,覆盖了数据类型、算子集合、控制流、融合策略等各个方面。建议挑几个最简单的用例跑一遍,改改形状、改改dtype,观察编译器的行为变化。这是比任何博客教程都可靠的学习路径。
第二,用PyTorch模型做端到端验证。虽然Relax的Python前端可以直接写IR,但绝大多数人不会选择手写IR。目前TVM社区提供了一个叫relax.frontend.torch的模块,可以导入PyTorch模型并转换成Relax IR。你可以拿一个结构稍微复杂的模型(比如带残差块和条件分支的),走一遍from_torch → relax.build → 执行的完整流程,感受一下Relax在真实模型上的表现。这一步能帮你建立“IR设计—编译流程—执行细节”的完整闭环。
第三,关注社区动态,及时调整知识体系。Relax的演进速度非常快,设计文档和RFC在github.com/apache/tvm-rfcs仓库里有非常详细的历史记录。这个仓库是理解Relax设计决策的第一手素材——你可以看到哪些方案被讨论过、为什么被否决、有几个备选方案。我对这个仓库的评价是:它比任何“XX教程”都有价值。
第四,动手写一个自定义算子。理解Relax最好的方法是在它的框架内填一块砖。你可以尝试用TensorIR写一个矩阵乘法,然后通过R.call_tir接入Relax的调用链,再接入一个Python层的PackedFunc。这个过程会把StructInfo、call、TensorIR、PackedFunc这几块核心知识整个串起来。
如果你的目标是研究编译器,我想多说一句:Relax这个项目的价值不止于TVM本身。作为一个“把通用IR思路应用到DL编译器”的完整实践,它提供了一个不太常见的学习样本——你很少能在IR层面,这么清晰地看到一个编译器如何平衡表达力、优化能力和工程可维护性这三者的关系。
最后提一个我自己的体会:如果你第一次接触Relax觉得“怎么这么多概念,太复杂了”,我特别理解。我最初也有这个感觉,觉得StructInfo、dataflow、ShapeExpr、match_shape这些概念像机关枪一样扫过来。但静下心来跑通一两个用例之后会发现,所有这些概念都指向同一个需求——既要让程序写起来灵活,又要让编译器优化起来安全。所有的设计,都在这个坐标轴上展开。想通了这条主线,Relax的面目就会变得亲切许多。