资讯中心

从Relay到Relax:TVM新架构下从零构建Relax模块实战指南

📅 2026/9/29 13:41:50
从Relay到Relax:TVM新架构下从零构建Relax模块实战指南
去年做边缘设备部署我还在用 Relay 写 TVM 的部署脚本碰到动态 shape、复杂控制流和自定义算子拼接每次都折腾得够呛。直到 2023 年之后 Relax 以官方教程主角的身份正式进入 TVM 主分支我才花了一个周末把原来的推理链路全部重写了一遍。这个决定很值——不仅是换一种写法而是整个建图、优化、编译的思考方式都变了。这篇文章我想从一个实际使用者的角度聊聊 Relax重点放在“怎么从零创建一个 Relax 模块”这件事上。我不会把官方文档复述一遍而是把一年多来踩过的坑、验证过的写法、以及真正能跑通的代码给出来。无论你是第一次听说 Relax还是已经从 Relay 转过来但总觉得别扭这篇文章应该都能帮你省下不少试错时间。1. 为什么 TVM 要把 Relay 留给新架构Relax 解决的表达瓶颈1.1 Relay 时代最难写的几类模型先说结论Relax 不是对 Relay 的小修小补而是 TVM 在计算图表示层面的重新设计。搞清楚“为什么”你后面写代码时才不会带着 Relay 的惯性去写 Relax。我自己在 Relay 里最痛苦的三类场景是动态 shape。RNN、Transformer 这类序列长度不固定的模型在 Relay 里要做Any维度的标注加上shape_func才能推导形状很多 pass 遇到动态维度直接罢工。多设备异构执行。想把图的一部分放在 GPU、一部分放在 CPU 跑在 Relay 里需要手动插 annotation pass而且这些注解在后续优化中很容易被改写掉。组合高层算子。比如把nn.conv2d、nn.batch_norm、nn.relu组合成一个自定义融合算子在 Relay 里你得深入 pass 层改 pattern门槛很高。Relay 设计时把“计算图表达”和“算子实现”绑得太紧。每个子图节点必须对齐到已有的 TIR 算子高层语义一旦没法降到 TIR后续优化就很难做。1.2 Relax 换了什么思路把“表达”和“实现”彻底分层Relax 的设计核心是计算图 IR 只负责描述计算结构和数据流不强制要求每个节点都能直接映射到某个 TIR 算子。高层算子可以先挂在图上由后续 pass 决定怎么 lowering、怎么融合、怎么布局转换。这就带来几个直观变化动态 shape 变成了一等公民。Relax 里维度可以是一个符号变量shape 本身也是运行时对象很多分析在编译期做不了就放到运行时做。dataflow block 概念。显式地把无副作用的纯计算区域包起来优化器能非常安全地做公共子表达式消除、算子融合和内存复用。自定义算子友好。你可以把一个 Python 函数直接声明为一个R.function的一部分只要给它写出 TIR 内核Relax 就能接入而 Relay 里做同样的事要动 pass 管线。举一个最直观的例子。Relay 中如果你想表达“两个张量逐元素相乘然后把结果累加”你需要将操作转换成 Relay 的 Call 节点每个节点都要在 pass 管线里被逐个识别。Relax 里直接用R.multiply和R.add这类内建算子写出来在 dataflow block 内任何 pass 都可以基于这个更松散的表示做激进优化。建议已经熟悉 Relay 的朋友第一件事是忘掉relay.Function的嵌套结构。Relax 的模块是扁平的 binding 序列看起来更像你平时写的 Python 函数体。2. 从源码编译启用 Relax环境准备与依赖坑2.1 版本选择主分支才是 Relax 的主场如果你用的是 pip 直接安装的 TVM 稳定版大概率就有 Relax但 API 可能已经变动过好几轮。Relax 本身经历了一个从 RFC 到进入主干的过程很多早期 API 在正式版本里已经被替换掉了。我个人的建议是直接拉 GitHub 主分支源码编译。不必害怕主分支不稳定TVM 的主分支目前已经相当可靠而且 Relax 相关的示例、测试、文档都是按主分支代码维护的你搜索到的大多数代码片段在主分支上能直接跑通。如果你手头是 0.15 之前的版本遇到tvm.relax模块不存在或者接口对不上不要怀疑自己多半是版本太旧。2.2 CMake 配置与编译选项编 TVM 最烦的是 LLVM 那一环。Relax 在生成可执行文件时需要 LLVM 后端支持所以编译期最好把 LLVM 打开。我的编译步骤大致如下git clone --recursive https://github.com/apache/tvm tvm cd tvm mkdir build cp cmake/config.cmake build/然后修改build/config.cmake重点开这几个选项set(USE_LLVM ON) set(USE_OPENMP ON) set(USE_CUDA OFF) # 如果你用 GPU 就写 ON并指定 CUDA 路径 set(USE_RELAY_DEBUG ON)USE_LLVM ON会尝试自动探测 llvm-config。如果你机器上有多个 LLVM 版本最好直接写完整路径set(USE_LLVM /usr/bin/llvm-config-15)我踩过的一个坑是 conda 环境里的 LLVM 和系统 LLVM 版本混在一起cmake 探测到了错误的llvm-config结果生成的 TVM 在运行时找不到某些 LLVM 符号。解法也很简单编译前用which llvm-config看清楚或者直接在 config.cmake 里写死路径。然后就是标准构建cd build cmake .. make -j$(nproc)编译时间取决于机器8 核机器差不多 20 到 40 分钟别急着关终端。编译完成后把tvm/python加入 Python 路径推荐用软链接方式这样以后 git pull 更新代码后 Python 包也是新的cd tvm export TVM_HOME$(pwd) export PYTHONPATH$TVM_HOME/python:${PYTHONPATH}2.3 验证 Relax 模块是否可用编译完先不要急着跑模型先确认 Relax 能正常导入python -c import tvm; print(tvm.__version__); print(tvm.relax)如果看到类似module tvm.relax from ...的输出说明环境 OK。我这里遇到的一个经典错误是ModuleNotFoundError: No module named tvm.relax排查后发现是 Python 路径没指到编译出来的 python 目录而是指向了 pip 装的旧版本。用python -c import tvm; print(tvm.__file__)先看看到底导入的是谁这一步能解决 80% 的环境困惑。3. 第一个 Relax 脚本跑通最小 IRModule3.1 用 TVMScript 写一个可运行的模块Relax 最友好的地方是支持 TVMScript也就是直接用 Python 语法描述 IRModule。下面这个最小例子包含了“主函数 TIR 内核”两层结构也是后面所有改造的起点。import tvm from tvm.script import ir as I from tvm.script import relax as R from tvm.script import tir as T I.ir_module class MyModule: T.prim_func def tir_mul( A: T.Buffer((4, 4), float32), B: T.Buffer((4, 4), float32), C: T.Buffer((4, 4), float32), ): for i, j in T.grid(4, 4): C[i, j] A[i, j] * B[i, j] R.function def main( x: R.Tensor((4, 4), float32), w: R.Tensor((4, 4), float32), ) - R.Tensor((4, 4), float32): with R.dataflow(): gv R.call_tir(tir_mul, (x, w), R.Tensor((4, 4), float32)) R.output(gv) return gv这段代码做了什么tir_mul是底层 TIR 内核负责真正执行 4x4 矩阵逐元素乘法main是 Relx 层的入口函数它把输入x和权重w打包传给tir_mul通过R.call_tir调用底层内核。R.dataflow()声明了一个无副作用计算区域R.output(gv)标记这个区域对外输出的变量。把这个模块跑起来的代码非常简单ex tvm.compile(MyModule, targetllvm) vm tvm.relax.VirtualMachine(ex, tvm.cpu()) import numpy as np x np.ones((4, 4), dtypefloat32) w np.full((4, 4), 2.0, dtypefloat32) out vm[main](x, w) print(out)如果没有意外你应该看到一个全是 2.0 的 4x4 矩阵。不要小看这个小例子它把 Relax 最重要的调用约定演示清楚了高层R.function负责组织调度低层T.prim_func负责实际数学运算两者通过call_tir连接。3.2 打印 IR 结构理解高层表示跑通之后你一定想看看 Relax 到底把这段代码表示成了什么样子。用MyModule.script()可以打印规范化后的 IRprint(MyModule.script())你会发现脚本和原始输入的 TVMScript 几乎一致这正是 TVMScript 设计的巧妙之处——打印出来的 IR 还能再解析回去。当你需要调试 pass 优化后的结果时这个能力非常关键你可以把中间产物 dump 出来人工检查。此外如果你只想看函数级别的结构不关心具体运算实现用tvm.relax.analysis里的一些 API 也可以做 AST 级别的查看不过日常调试中script()已经够用了。提示tvm.compile在不同版本里可能写作relax.build。如果你手上的版本里找不到tvm.compile试试tvm.relax.build(MyModule, targetllvm)功能一致。4. 用 Python API 从零构建 Relax 函数4.1 变量、结构和函数签名TVMScript 适合手写和阅读但如果你要动态生成计算图比如后端动态解析用户配置来组装模型就必须掌握用 Python API 构建的方法。这就像写 SQL 可以用查询工具也可以直接写 JDBC 代码后者更灵活但细节更多。构建一个 Relax 模块的核心组件是三样Var变量、StructInfo结构信息和BlockBuilder图构建器。变量就像计算图上的“导线”它本身没有具体数据只携带类型和形状信息。结构信息StructInfo描述变量或者函数签名张量维度、数据类型等。BlockBuilder是最关键的它负责把你在 Python 里调用的操作逐步记录成 IR 节点。上节那个乘法模块用 Python API 构建长这样import tvm from tvm import relax from tvm.script import tir as T # 第一步定义输入变量 x relax.Var(x, relax.TensorStructInfo([4, 4], float32)) w relax.Var(w, relax.TensorStructInfo([4, 4], float32)) # 第二步创建 BlockBuilder并开启一个名为 main 的函数 bb relax.BlockBuilder() with bb.function(main, [x, w]): # 第三步在 dataflow block 中 emit 一个乘法操作 with bb.dataflow(): y bb.emit(relax.multiply(x, w)) # 标记 dataflow 输出 bb.emit_output(y) # 标记函数返回值 bb.emit_func_output(y) mod bb.get() print(mod.script())打印出来的 IR 会和上面R.function写的几乎一样只不过tir_mul不存在因为这里用的是内建算子relax.multiply。Relax 自带了一批内建的高层算子它们既可以直接放到图上也可以在后续 lowering 阶段被转换到 TIR 内核。4.2 张量运算与 call_tir 的分工这里有个核心问题什么时候用内建算子什么时候用call_tir以我的经验规则很简单——如果底层已经有一段写好的 TIR 内核比如手写的 Conv、Attention用call_tir把它挂进来如果只是代数的组合、拼接、切片这类高层操作直接用内建算子让 Relax 后续自己去 lower。call_tir的参数很有讲究gv R.call_tir(tir_mul, (x, w), R.Tensor((4, 4), float32))第三个参数是输出结构信息out_sinfo它告诉编译器这个内核会产生什么样的输出。很多新手在这里随便填一个 shape结果后续 pass 在类型推导时直接报错。out_sinfo 必须和 TIR 内核中输出的 buffer shape 完全一致这是 Relax 静态类型安全的一个基本保障。如果你要在某个循环结构内动态决定调用哪个内核call_tir还支持闭包形式的参数这在 Relay 里实现起来很麻烦Relax 里可以直接写 Python 逻辑来构造不同分支。4.3 模块构建与调用验证构建完成后和 TVMScript 版本一样用tvm.compile编译再调用即可。这里我想强调一个经验用 Python API 构建模块时最好在每一步都打印一次bb.get().script()不一定要等到最后。因为 emit 的过程中如果存在结构信息不一致一些错误会在后面的 pass 中延迟暴露定位起来很困难。养成随手打印的好习惯能省掉至少一半的调试时间。5. 打造一个带训练场景的端到端示例从 ONNX 到 Relax5.1 转换链路的选型直接转还是经 Relay 中转上面的例子都是从零建图但真实项目里更多是加载现成模型。先讲转换路径TVM 生态里有两个常见的入口一个是从 ONNX 直接到 Relax另一个是 ONNX 先转 Relay 再转 Relax。两个选择我都试过。我的结论是如果模型结构规整想快速预览效果直接走tvm.relax.frontend.onnx.from_onnx一步到位。如果模型结构复杂或者你后续要做 Relay 的算子改写、debug先转 Relay 再转 Relax 更稳。Relax 的很多新的优化 pass 是针对 Relax IR 实现的但 Relay 的算子库和 pass 生态成熟度高于 Relax。经 Relay 中转你可以借助 Relay 的算子融合先把图做一遍规模化简化再交给 Relax 做后阶段的优化。转换代码大致如下import onnx import tvm from tvm import relax, relay onnx_model onnx.load(model.onnx) shape_dict {input: (1, 3, 224, 224)} # 先转到 Relay relay_mod, params relay.frontend.from_onnx(onnx_model, shape_dict) # 再从 Relay 转到 Relax from tvm.relax.frontend.relay import from_relay relax_mod from_relay(relay_mod[main]) print(relax_mod.script())如果你的 TVM 版本较新也可以试试直接转换from tvm.relax.frontend.onnx import from_onnx relax_mod from_onnx(onnx_model, shape_dict) print(relax_mod.script())两路转换最后得到的 Relax 模块结构不完全一样。经 Relay 中转的模块通常已经带上了融合分组直接转换过来的则更“原始”可以留给 Relax 的 pass 去处理。没有绝对优劣取决于你要做哪一层的研究。5.2 端到端推理流程编译、加载参数、执行转换完成后的推理流程和前面的最小例子是统一的一套 APIex tvm.compile(relax_mod, targetllvm) vm tvm.relax.VirtualMachine(ex, tvm.cpu()) # 参数从转换时返回的 dict 获取 input_data np.random.rand(1, 3, 224, 224).astype(float32) out vm[main](input_data, **params)注意这里的**paramsRelax 的入口函数签名里如果包含了参数名字VM 调用时可以按关键字传入。很多 Relay 来的老代码习惯把参数 dict 再塞进输入列表里在 Relax 里不这么做。我自己常用的调试小技巧是把转换后的relax_mod.script()保存成文本文件先看一眼函数签名是不是符合预期特别是有没有意外的全局变量或未绑定参数。这一步能提前发现 shape_dict 写错的问题。5.3 从 Relay 迁移旧代码到 Relax 的实操套路如果你已经有一批 Relay 时代写的部署代码迁移起来其实没有想象中可怕。我自己的套路是先把 Relay 模型中main函数里的计算逻辑梳理出来搞清楚所有子图的输入输出。在 Relax 中定义一个同名R.function把输入和输出的StructInfo照抄。逐层替换relay.nn.conv2d这类调用为relax.nn.conv2d、R.nn.relu语法相似度很高。用 dataflow block 把原本互不关联的 Call 节点包在一起。编译后和 Relay 版本对拍输出几乎能对齐到个位小数点的差异。有几个 API 变化需要格外注意。Relay 里relay.var创建变量Relax 里是relax.VarRelay 里调用算子直接relay.nn.relu(data)Relax 里很多时候要套R.call_tir或者bb.emit。最容易被坑的是shape参数现在放在StructInfo里而不是算子调用参数里。6. 新手最容易踩的八个坑调试手记6.1 TVMScript 解析报错先查缩进和类型注解I.ir_module这类装饰器实际上是一个解析器它读的不是 Python 的 AST 语义而是从源码字符串中还原 IR。所以如果你在写T.prim_func时括号不齐、类型注解少了个引号、或者R.Tensor((4, 4), float32)写成了R.Tensor((4, 4), float32)解析器会抛出一堆让人困惑的语法错误。这种报错的第一个检查步骤是把代码复制到一个干净的 Python 文件里跑排除 notebook 环境导致的 AST 源码获取问题。TVMScript 解析器和 Python 的 import 系统耦合很深Jupyter Notebook 偶尔会出现source拿不到导致解析失败。6.2 call_tir 的 out_sinfo 不匹配call_tir的 out_sinfo 必须和 TIR 内核的输出 buffer 一一对应。包括维度顺序、dtype任何一个不匹配都会在后继类型推导中报错。比如(4, 4)写了(4, 4, 1)报错信息通常很抽象像什么Cannot prove: 16 4看起来一头雾水。排查思路先用print(mod.script())把 TIR 内核的输出 buffer 定义打出来直接逐字核对不要凭记忆填。6.3 动态 shape 的符号维度处理Relax 主推符号形状但这带来的问题就是符号推断复杂。把维度写成字符串时R.Tensor((n, 4), float32)可能因为符号名冲突导致无法统一。最稳妥的做法是先用静态 shape 把流程跑通再替换成动态维度。直接上动态 shape你会同时面对编译期和运行期的两类错误新手很难分清问题出在哪一层。6.4 运行时调用接口不匹配VirtualMachine调用时输入必须是 numpy 数组或者 DLDeviceType 对应的 dltensor不能是 Python list。我遇到最多的是用户传了个 list 进去VM 直接报AttributeError: list object has no attribute dtype。养成习惯入口统一转 numpy 或 tvm.nd.array。6.5 表格式速查我把常见问题整理成一个速查表方便你定位问题现象最可能原因处理建议ModuleNotFoundError: tvm.relaxPython 路径指向 pip 旧版检查tvm.__file__切到源码编译路径TVMScript 解析报语法错误类型注解或括号格式不对逐行对照R.Tensor(shape, dtype)写法Cannot prove ...类型错误out_sinfo 和 TIR 输出不一致打印脚本逐字段核对 shape 和 dtype动态 shape 运行时错误符号维度约束未满足先静态跑通再逐步动态化VM 调用时 list 参数报错传参类型不是 ndarray统一np.array(...)后传入编译时间莫名卡住编译选项打开了太多后端关闭不用的 CUDA/METAL仅保留 LLVM还有几个小坑因为篇幅关系我合并说一是多设备要提前统一 target二是shape_dict里的名字必须是 ONNX 输入节点的名字三是 dataflow block 里不要写打印或者副作用操作。这些都满足之后Relax 的日常开发体验会顺畅很多。最后说一个我个人的体会。Relax 的script()可逆性是真的好用调试的时候我经常把 pass 前和 pass 后的 IR 都 dump 出来用 diff 工具直接看变化很多优化问题一眼就能定位。如果你也准备在 TVM 上做编辑器层面的工作我建议先照着第三章的最小例子跑通再试着用 Python API 重写一遍最后才碰真实模型。这个顺序看起来慢但实际上是避开混乱的最佳路径。

看完文章,想为自己的企业也做一次专业网站诊断?

尧图顾问免费为您评估现有网站,并给出建站/改版建议与报价方案。

免费获取方案