news 2026/9/19 7:25:19

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
从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.conv2dnn.batch_normnn.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.multiplyR.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 脚本:跑通最小 IRModule

3.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, target="llvm") vm = tvm.relax.VirtualMachine(ex, tvm.cpu()) import numpy as np x = np.ones((4, 4), dtype="float32") w = np.full((4, 4), 2.0, dtype="float32") out = vm["main"](x, w) print(out)

如果没有意外,你应该看到一个全是 2.0 的 4x4 矩阵。不要小看这个小例子,它把 Relax 最重要的调用约定演示清楚了:高层R.function负责组织调度,低层T.prim_func负责实际数学运算,两者通过call_tir连接。

3.2 打印 IR 结构,理解高层表示

跑通之后,你一定想看看 Relax 到底把这段代码表示成了什么样子。用MyModule.script()可以打印规范化后的 IR:

print(MyModule.script())

你会发现脚本和原始输入的 TVMScript 几乎一致,这正是 TVMScript 设计的巧妙之处——打印出来的 IR 还能再解析回去。当你需要调试 pass 优化后的结果时,这个能力非常关键,你可以把中间产物 dump 出来人工检查。

此外,如果你只想看函数级别的结构,不关心具体运算实现,用tvm.relax.analysis里的一些 API 也可以做 AST 级别的查看,不过日常调试中script()已经够用了。

提示:tvm.compile在不同版本里可能写作relax.build。如果你手上的版本里找不到tvm.compile,试试tvm.relax.build(MyModule, target="llvm"),功能一致。

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 到 Relax

5.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 端到端推理流程:编译、加载参数、执行

转换完成后的推理流程和前面的最小例子是统一的一套 API:

ex = tvm.compile(relax_mod, target="llvm") 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)

注意这里的**params,Relax 的入口函数签名里如果包含了参数名字,VM 调用时可以按关键字传入。很多 Relay 来的老代码习惯把参数 dict 再塞进输入列表里,在 Relax 里不这么做。

我自己常用的调试小技巧是把转换后的relax_mod.script()保存成文本文件,先看一眼函数签名是不是符合预期,特别是有没有意外的全局变量或未绑定参数。这一步能提前发现 shape_dict 写错的问题。

5.3 从 Relay 迁移旧代码到 Relax 的实操套路

如果你已经有一批 Relay 时代写的部署代码,迁移起来其实没有想象中可怕。我自己的套路是:

  1. 先把 Relay 模型中main函数里的计算逻辑梳理出来,搞清楚所有子图的输入输出。
  2. 在 Relax 中定义一个同名R.function,把输入和输出的StructInfo照抄。
  3. 逐层替换relay.nn.conv2d这类调用为relax.nn.conv2dR.nn.relu,语法相似度很高。
  4. 用 dataflow block 把原本互不关联的 Call 节点包在一起。
  5. 编译后和 Relay 版本对拍输出,几乎能对齐到个位小数点的差异。

有几个 API 变化需要格外注意。Relay 里relay.var创建变量,Relax 里是relax.Var;Relay 里调用算子直接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 重写一遍,最后才碰真实模型。这个顺序看起来慢,但实际上是避开混乱的最佳路径。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/19 7:24:39

SAP Fiori内容模型配置与业务角色权限管理

1. 项目概述&#xff1a;SAP Fiori内容模型的业务价值解析在SAP S/4HANA实施过程中&#xff0c;Fiori作为新一代用户体验框架&#xff0c;其内容模型的配置质量直接决定了最终用户的系统使用效率。Business Role&#xff08;业务角色&#xff09;与Target Mapping&#xff08;目…

作者头像 李华
网站建设 2026/9/19 7:24:37

SpringBoot整合MyBatis分页插件PageHelper全解析

1. 分页处理的必要性与应用场景在开发企业级应用时&#xff0c;数据分页几乎是每个项目都会遇到的刚需。想象一下&#xff0c;当数据库中有10万条用户记录时&#xff0c;如果一次性全部加载到内存中&#xff0c;不仅会造成服务器内存压力&#xff0c;前端渲染也会变得极其缓慢。…

作者头像 李华
网站建设 2026/9/19 7:23:10

数据结构从理论到代码:手写链表、二叉树、哈希表与调试实战

简介&#xff1a;这份PDF是山东大学《数据结构》课程内容整理&#xff0c;面向计算机专业本&#xff08;专&#xff09;科生、考研与期末复习者&#xff0c;帮助快速建立从数据组织到算法分析的知识框架。资源共1个文件&#xff0c;为PDF格式&#xff0c;压缩包大小仅324KB&…

作者头像 李华
网站建设 2026/9/19 7:19:55

聚氨酯一体板vs铝单板:建筑外围护选型全维度对比与决策指南

1. 建筑外围护选型&#xff1a;聚氨酯一体板vs铝单板1.1 核心需求解析建筑外围护选型这件事&#xff0c;说大不大&#xff0c;说小也绝对不小。往小了说&#xff0c;它决定了建筑外立面好不好看、耐不耐用&#xff1b;往大了说&#xff0c;它直接关系到项目的综合造价、施工周期…

作者头像 李华