ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

TVM Relax抽象层深度解析:从动态形状到控制流的IR设计

TVM Relax抽象层深度解析:从动态形状到控制流的IR设计 刚接触TVM的时候我对Relax这个抽象层其实是一头雾水的。因为当时市面上大多数入门教程都停留在Relay IR层面讲到如何用Relay写一个网络、如何做算子融合和代码生成再往后就没了。直到后来要在自己的项目里处理一些带有动态形状和控制流的模型看了官方文档也磕磕绊绊才认真把这个IR设计吃透。这篇文章就从一个实际使用者的角度把我对Relax抽象层的理解、动手验证的过程、以及在真实项目里踩过的坑系统性地梳理一遍。Relax并不是Relay的简单改名它是TVM在迈向“编译动态形状模型、支持函数式表达、更好对接高层框架”这三个目标时重新设计的可解析中间表示。简单说Relay擅长的是静态图优化而Relax更贴近“带类型、带作用域、可被Python端直接构造和调试的脚本化IR”。搞清楚这一点很多设计选择就顺理成章了。如果你正在学TVM源码、打算在新项目里用TVM做推理优化或者被Relay派生的控制流写法折磨过这篇文章应该能帮你省下我当初反复翻源码和文档花的几周时间。1. Relay的边界在哪里为什么需要再造一个抽象层要理解Relax你得先承认Relay有它自己的舒适区。Relay是函数式的、静态形状优先的IR它的计算图优化做得非常好——算子融合、布局转换、常量折叠这些都是Relay的看家本领。但它有几个骨子里的问题在一开始用的时候可能不觉得等你开始写复杂模型时就压不住了。1.1 控制流表达与形状推导的割裂Relay里表达条件分支用relay.If循环用while从功能上讲它们都存在。但问题在于编译器对这两种结构的静态分析能力非常弱。典型的场景是“形状依赖数据”——比如模型要处理一批变长序列而序列长度是运行时才知道的数值后续的层要依据这个数值动态分配计算图分支。在Relay里你通常得把这类逻辑挪到上层Python代码里做分包或者用where之类的算子把多个分支计算结果都算一遍再选导致上层的控制流被“压平”成一个巨大的静态图。这是一个很尴尬的妥协逻辑上它是动态的物理上你却不得不静态展开。Relax重新做了一个很关键的决定把If、While、函数调用这些高层控制流作为IR里的“一等公民”并且明确允许“形状还未确定”的Tensor存在。你可以在一个节点上写一个张量的shape是变量n或m而不必在构造图的当下就限定成具体整数。这对于序列模型、多模态输入、或者那些依赖输入尺寸来决定内部计算策略的网络是质的改变。1.2 Dataflow Block与函数级作用域的纠缠Relay里也有let绑定、局部变量和函数参数这些作用域概念本身没问题。但当你真的想写一个“模块化”的优化Pass时你会发现Dataflow Block和函数体之间的边界经常非常模糊。举个例子Relay里一个Function的参数可以在DataflowBlock中被任意读取而DataflowBlock内部也可调用外部函数——这给了优化器很大的灵活性但也让“这个数据流到底在哪个基本块里结束”变得很难判断。Relax在这方面做了一个减法它把“数据流块”定义为纯粹的、单赋值的、只包含计算节点的局部区域块内是纯函数式语义块边界就是数据依赖的边界。这样一来任何数据流分析都只需要关注块本身而控制流分支、循环、调用被明确放在数据流块之外两者职责清晰不会再混在一起。1.3 构建时与运行时的“两张皮”问题这是我个人最痛的体会。用Relay写一个动态形状的模型时最恼火的是你在Python端构造的是“抽象语法树形态的计算图”但这个图在编译期被“冻结”了一次到运行期你想要根据数据动态决定某些行为会发现根本无从下手。更准确地说Relay在本质上更多的是一种“描述计算图”的DSL它把“构建计算”和“执行计算”之间的边界画得比较死。Relax的核心设计之一就是推迟执行deferring execution。Relax的脚本式IR允许你定义“尚未执行的调用”这些调用只有到运行时被调度到具体的PackedFunc之后才真正发生计算。换句话说Relax把“我想调一个函数”和“这个函数真的去跑”明确地区分开来。这个设计带来的直接好处是你可以写出“很动态的”程序结构同时依然能对它们做后期静态分析和优化。2. Parsability 三级抽象Relax 到底在解析什么如果你直接翻开Relax的论文或者源码注释大概率会被“Parsability”这个词砸到。不少中文技术博客会把Parsability翻译成“可解析性”但看完一脸问号什么可解析性谁解析谁为什么要解析出三个级别我自己的理解是Relax把“一段IR程序里哪些部分是编译器能立刻看透的”分层了。从高到低分为三个层次可解析的函数式表达Parsable Functional通常简写为RX、可解析的数据流表达Parsable Dataflow简称RX注意和第一个缩写雷同只是来源不同、可解析的张量表达Parsable Tensor简写为RX。这里的“可解析”指的是形状推导、数据依赖分析、以及后续优化Pass能否在不动这笔代码的情况下可靠地分析和变换它。2.1 函数式层让控制流拥有基于符号形状的推理能力函数式层是最接近普通函数式语言的层次它包含Function、Call、If、While、Tuple、MatchCast等结构同时还引入了DataflowVar与局部变量绑定。这些结构背后实际执行的语义依赖一批内置的PackedFunc例如R.shape、R.prim_value等而你在编写Relax脚本时可以直接像调用普通Python函数那样调用它们。这一层的核心价值在于控制流是显式的形状可以是符号化的。在Relay里面如果你写if x.shape[0] 3:编译器在构建阶段大概率直接报错形状还不是常量。但Relax允许构造阶段对符号形状的表达式进行“推理”比如把if的条件绑定成n 3其中n是一个符号变量然后分支内继续使用n参与张量的shape表示。这个特性的实际意义很大它让“形状相关的逻辑”成为IR计算图的一部分而不是只能靠Python外挂脚本去动态处理。实际操作中我经常用这层能力处理“padding策略随输入长度变化”的模型。以前用Relay的时候这种动态策略得在Python端预先判断形状再为每一种情况生成一个静态图现在用Relax直接把符号形状写进IR让TVM在运行时观察具体数值并调度到对应分支代码量直接下降一个数量级而且可读性高得多。2.2 数据流层数据依赖清晰、作用域边界严格的“纯计算块”数据流层就是我前面提到的那个减法设计。Relax的数据流块是DataflowBlock用R.dataflow()开始、R.output()结束。块内只能做纯数据流计算而所有控制流语句比如R.if、调用外部函数必须放在块外面。这个约束看上去是限制实际上给了优化器一个极大的便利块内没有副作用、没有循环、没有条件跳转你可以放心地在这个区域内做公共子表达式消除、算子合并、布局变换、常量折叠而不用操心“跨块”的依赖关系。一个值得了解的细节是DataflowBlock内部变量是单一赋值single static assignment的而且块边界不允许“漏数据”。所有块内的局部变量只能通过R.output()输出到外部作用域外部也不能反向只读取块内部的值。这有点像函数式编程里的“纯函数”约束——输入输出清晰外部只能拿到你明确导出的结果。我在自定义Pass的时候特别喜欢这个性质。比如实现一个“基于数据流块内算子语义的预算分配算法”因为在DataflowBlock中我可以直接枚举所有的CallNode而且它们的输入输出关系已经通过变量名建立好不需要额外去维护use-def链。这在Relay里做同样的事情要费很多功夫因为那里的Dataflow作用和函数边界会纠缠。2.3 张量层与TIR无缝对接的原子计算单元张量层是离硬件最近的一层它对应的是TVM里“最终会落到TIRTensor IR去做的实际循环计算”。在Relax中一个张量算子调用通常是R.call_tir(global_var, args, out_sinfo)这里的global_var指向一个TIR函数或者一个打包好的PackedFunc。换句话说这层让你的“高级抽象”和“低级循环”第一次发生联系。这个设计的好处是你可以为一个算子在TIR里手写一个高性能实现然后在Relax里通过call_tir直接引用它而不需要走“算子注册”再“匹配模式”的繁琐路径。类型的匹配是通过out_sinfo显式声明的。这个sinfo参数经常被新手忽略但它是Relax类型系统里非常核心的纽带——它告诉编译器这个调用返回的张量应该是什么形状、什么dtype从而让后续的shape传播变得可能。我自己在实现一个自定义卷积时就是先用TVMScript写了一个TIR函数然后在Relax脚本里用call_tir绑定调用。这套流程比在Relay里手动注册算子然后还要写一个TOpPattern要直观得多因为你在同一段脚本里就能看到高级调用和底层实现的对应关系。2.4 三级抽象之间的协作机制这三个抽象层级并不是独立存在的它们之间有严格的嵌套关系函数式层包含数据流层数据流层包含张量层。这个嵌套关系每次都让我想起DOM树的结构——父节点可以访问子节点但子节点并不会反向影响父节点的结构。当一个Relax函数被构造时最外层通常是Function节点函数体内依次包含若干语句其中可能有一个DataflowBlock块内都是call_tir或R.call_packed等张量级调用。控制流节点If、While则出现在函数体的顶层区域中它们的body分支体/循环体各自又是一段可包含数据的调用序列。这种设计让优化Pass可以按层去访问。例如“在张量层做算子融合”的Pass只需要遍历DataflowBlock内部的调用序列而不会误伤控制流“做控制流扁平化”的Pass则主要扫描函数层。分层不只是为了清晰更是为了给编译器内部的模块化设计提供一个稳定的结构基础。3. 从零写一个可运行的 Relax 示例TVMScript 视角我从来不建议一上来就看IR的JSON格式太劝退了。TVMScript提供了一种Python风格的、可读性很好的Relax程序描述语法你写出来的脚本可以直接被解析成一个Relax IR模块所以非常推荐作为学习入口。下面这段是我实际跑通过的代码用来展示Relax的书写方式以及“解析级别”在生成IR中的体现。3.1 准备环境与最小依赖在跑下面的代码之前你需要一个完整可用的TVM Python包。我建议直接用官方预编译的发布版本或者从源码自行编译。Relax相关的API在TVM 0.14之后的版本里趋于稳定我使用的是TVM 0.16.dev0的源码构建版本。因为Relax内部API还在演进每个小版本接口可能略有调整如果你发现某些函数名对不上优先在源码里搜索同类名称找出最新的替代接口——这在开源社区里属于常态不用惊慌。pip install apache-tvm0.16.dev0 # 或者是支持Relax的新版本3.2 TVMScript写一个带符号形状的函数考虑一个这样的场景输入一个形状为(n, m)的矩阵和一个标量threshold当m大于threshold时对矩阵执行exp变换后的行求和反之则执行tanh变换后的列求和。这里n和m都是运行时才能确定的符号量核心是想展示Relax如何描述动态形状和控制流。import tvm from tvm import relax tvm.script.ir_module class MyModule: R.function def main( x: R.Tensor((n, m), float32), threshold: R.Tensor((), float32), ) - R.Tensor((n,), float32): n T.int64() m T.int64() with R.dataflow(): gv R.call_tir(my_exp, (x,), out_sinfoR.Tensor((n, m), float32)) cond R.call_packed(my_compare, m, threshold, sinfo_args(R.Tensor((), bool))) R.output(gv, cond) if cond: out R.call_tir(my_sum_rows, (gv,), out_sinfoR.Tensor((n,), float32)) else: h R.call_tir(my_tanh, (x,), out_sinfoR.Tensor((n, m), float32)) out R.call_tir(my_sum_cols, (h,), out_sinfoR.Tensor((n,), float32)) return out乍一看这段代码既像Python又像一种DSL。实际上它就是一个合法的Relax模块可以被编译、执行。你可以直接运行并得到一个输出张量。但为了更深入地学习我们要逐行拆解它的含义。第一件值得注意的事是顶部R.function装饰器。它告诉TVM这是一个Relax函数参数里有R.Tensor((n, m), float32)。这里的(n, m)就是符号形状n和m并不是Python变量而是形状变量——这一点和Relay里只能用Any或者具体的整数完全不同。你可以直接在这个函数体里引用n、m它们会被替换成具体的形状表达式。第二件事是with R.dataflow()块的存在。这里的call_tir和call_packed都是纯计算它们没有任何副作用也没有分支。cond的构造依赖m和threshold它本身是一个布尔张量。这个条件张量虽然在数据流块内完成计算但它的值要到运行时才知道所以整个if cond其实是一种“运行时分支”——Relax允许这种分支与数据流块同时存在只要数据流块本身语义保持干净。第三件事是R.call_packed(my_compare, ...)。这个调用会去运行时环境里查找名为my_compare的PackedFunc。如果你没有注册这个PackedFunc这个模块可以构造成功但在执行时会报错。这提醒了我们一个很重要的机制Relax脚本中的高层调用不一定编译成TIR它可能直接映射到运行时PackedFunc。这也是Relax“延迟执行”的体现——脚本在构造期只是登记了“我要调用my_compare”真正执行要等到运行时。3.3 编译与执行处理符号形状与具体数据的动态绑定写完模块之后接下来的步骤就是把它编译成可执行形式。这里有一个非常核心的概念match_cast和形状绑定。在调用模块之前你需要把符号形状n和m“具体化”。TVM里最常用的方式是在运行时动态绑定输入张量的实际形状并执行一次“模式匹配”来确认类型是否一致。ex relax.build(MyModule, targetllvm) vm relax.VirtualMachine(ex, tvm.cpu()) x tvm.nd.array(np.random.rand(4, 5).astype(float32)) threshold tvm.nd.array(np.array(3.0, dtypefloat32)) result vm[main](x, threshold) print(result)整个过程看起来和调用一个Python函数差不多但背后发生了几层关键操作。首先relax.build会把整个IR模块编译成可执行的机器码和运行时函数表。对于call_tir指向的TIR函数这里我并没有定义真正的TIR函数体实际跑的时候需要补全或改用内置的relax.op算子接口会走TVM的代码生成管线生成底层kernel。然后vm[main]返回一个编译后的可调用函数。当你传入实际形状为(4,5)的x时VM会执行模式匹配将符号形状绑定成n4, m5并沿着IR的数据流做具体的形状推导再执行条件判断。我在这里还想强调一个新手容易迷糊的点符号形状并不等于动态形状必需绑定到Python的symbolic变量。在Relax的VM执行阶段符号形状只是在模式匹配时用于确认类型一致性一旦绑定成功所有内部计算都是基于具体数值的。它不比静态图慢多少但多了运行时类型检查和形状推导的开销。对于批量推理场景这种开销通常在微秒到几十微秒级别可以通过批量大小固定来显著摊薄。3.4 深入控制流If 分支的编译行为我们回过头再仔细看看if cond这个结构在Relax编译中的实际行为。它会被编译为一个IfNode带有true_branch和false_branch两个子函数不是闭包而是编译后的函数体运行时VM会先计算cond的值然后直接执行对应分支。这个做法的精妙之处在于两个分支内的算子并不会在另一个分支执行时被调度。这意味着你可以把两个分支都描述成合法的计算即使其中一个分支的数据类型或临时张量形状在另一个分支不存在也不会引起错误因为另一个分支根本不会被调用。这解决了Relay那种静态展开导致的分支计算浪费问题。另外Relax还支持多分支选择节点R.switch和循环结构WhileNode但核心原理一致。一旦你理解了If的处理方式再理解While就不难它同样是运行时循环循环体内的变量每次迭代重新绑定循环次数可以是动态的只要你能保证类型一致。不过While的实际执行开销比If高一些因为每次迭代都需要重新做一次变量遮蔽和类型检查在写高性能推理模型时要注意避免无谓的循环嵌套。3.5 对“match_cast”的进一步解释我在前面的示例里无意间用到了R.call_packed但我还没有正式介绍match_cast。这个节点在Relax里扮演着“类型再绑定”的角色专门用来处理“形状从动态变成半动态/静态”的情况。一个典型的场景是某个算子的输出形状依赖前一个算子的数值但在数据流块中这个输出形状会被声明为一个符号表达式比如(n, m * 2)。问题在于当m真的是一个运行时读取的值时n*m*2这个表达式无法在构造阶段解析成确定的整数值。match_cast的职责就是专门处理这种类型兼容性验证并生成新的类型注解。match_cast的实现方式不算复杂它本质上是一个运行时检查节点的包装。在编译时它会对输入的结构信息与目标结构信息做模式匹配如果匹配成功则返回一个带新类型标注的变量失败则报错。这种“安全网”机制让我这种喜欢写复杂形状表达式的用户感到非常安心——它不会默默接受错误的形状推导而是明确地告诉你哪一个环节的类型对不上省去了反复打印中间张量形状的调试痛苦。4. 工程落地中真正需要注意的几个细节讲完示例我想从工程实践的角度聊几个容易被文档一笔带过、但真实项目里会反复踩到的坑。这些内容大多不是官方教程会写进去的但如果你打算把Relax用到自己的项目里几乎一定会碰到。4.1 PackedFunc你函数的“跨语言边界”我第一次看到R.call_packed时困惑了很久它调用的到底是一个Python函数还是一个编译后的C算子答案是它调用的是一个PackedFunc——这是TVM里所有跨语言函数调用的统一方式可以是Python函数注册的也可以是C编译期注册的。当你希望把自定义的Python逻辑嵌入到Relax脚本中时你确实可以这样做先把一个Python函数注册成PackedFunc然后在Relax的call_packed里引用。但你要小心性能问题——因为每次调用都会跨过Python/C的边界如果这笔调用频率很高性能开销会很明显。tvm.register_func(my_compare) def my_compare(m, threshold): return m threshold执行时期Relax VM会把这个函数当作一个普通的PackedFunc来调用。它返回的结果需要显式转换成一个R.Tensor或R.Prim类型的值否则类型匹配阶段会失败。我自己在开发中经常用这种方法在IR层实现一些“用于调试的输出打印”逻辑很方便但正式部署时一定会把它替换成真正的TIR实现。4.2 数据流块内的隐式约束别在里面搞事情DataflowBlock是Relax的优化重心但它有两条硬核规则违反了在编译时会报错或者产生不可预知的行为。规则一块内不允许有任何副作用操作。这意味着你不能在数据流块内调用会修改外部状态的PackedFunc。所有块内调用都必须是纯函数式的否则编译器无法做重排或删除优化。也就是说你不能在数据流块里写“除了输出变量还要打印日志”这种操作。规则二块内变量必须是单赋值的。你可以对一个Var做一次绑定但后面不能改变它的值。如果你有类似“在循环中累积更新”的需求需要把它移到数据流块之外或者重新设计算法否则编译器会认为你在试图创建不纯的计算。这两条规则在当前阶段看起来像约束但当你开始调用TVM优化管线中的Pass时才会体会到它们的好处。允许Pass假设块内无副作用意味着可自由地做公共子表达式消除、死代码消除和算子融合。我就是因为在自定义Pass中经常依赖这些规则才真正理解为什么Relax要如此“刻意”地划分这些边界。4.3 形状推断的“符号传播”并不是免费的很多初学者会觉得Relax既然支持符号形状那么我在IR里写n和m编译器就会自动精确地推导出所有中间张量形状。事实上符号传播在大多数情况下是能工作的但并非总是精确。当你在数据流块内部计算时形状通常是由算子本身的结构信息推导出来的。比如add两个(n, m)形状的张量结果当然是(n, m)。但当你执行call_tir调用一个自定义TIR算子时TVM无法自动推理出它的输出的精确形状你必须提供out_sinfo——这相当于你向编译器做出的显式承诺。我在开发中已经形成了这样一个习惯每写一个自定义的TIR算子我就会在Relax脚本里手工写上它的out_sinfo并且在单元测试中验证实际输出的形状是否与sinfo一致。这个习惯极大地减少了由于“传播失败”导致的隐藏Bug因为一旦sinfo写错运行时模式匹配会立刻暴露问题。4.4 从Relay迁移到Relax时最容易出现的思维惯性最后说一个迁移层面的问题。很多Relay老用户在初写Relax时会下意识地照搬Relay的“算子即节点”思维把每个Relax调用都写成“图的node”然后试图用图匹配的方式去做优化。这在Relax里虽然也能做但会错失它最核心的设计优势以“数据流块符号形状显式控制流”为单位做分析和变换。比如在Relay中如果你想做“专用算子融合”需要分析整个函数的计算图、匹配到特定的子图模式、然后把它替换成一个新的算子调用。在Relax中你通常只需要在DataflowBlock里遍历调用序列检查相邻算子是否满足融合条件然后把它们合并成一个新的call_tir即可。因为块内数据依赖天然有序、无副作用、且形状信息已经在每个sinfo里声明好融合的安全性和可行性检查都简化了。另一个常见的思维惯性是在Python层“过度抽取子图”。Relay时代因为控制流支持孱弱很多人会写很多relay.Function做子图划分然后在外部负责调用。Relax不需要这样它的Function天然有清晰的作用域和类型接口子图划分可以直接在IR层做并且可以携带符号形状和动态控制流。这会改变你设计编译流程的方式——不再是为每种输入配置硬编码静态图而是把动态变化的部分留在IR中让运行时调度负责执行。5. 扩展思路Relax在真实项目里的几种应用方向讲完基础我想稍微发散一下聊聊Relax在真实项目中比较有潜力的几个方向。这些方向并不一定都对应完整的开源实现但它们都是我判断未来几年里TVM生态里会比较活跃的领域适合自己去摸索验证。5.1 动态形状推理引擎的建模动态形状在 NLP 推理里特别常见比如可变长输入。以往的做法是paddings到固定长度不可避免有计算浪费。Relax允许符号形状与运行时控制流结合意味着你可以设计一个更细粒度的调度策略当序列长度没有超过某个上限时走一个快速分支超过上限时走另一个更保守的分支。这种“基于实际数据的运行时分支”很难用Relay优雅地表达但Relax里就是现有能力。我目前在自己项目的BERT推理中做预填充和解码的两段式调度正是利用Relax的符号形状和控制流把两种模式写进同一个IR函数并让运行时根据输入长度自动切换省去了在C层手写调度逻辑的维护成本。5.2 形状相关的算子选择与自动调度符号形状和运行时绑定的另一个非常直接的应用是算子的自动调度选择。同一个算子在形状小时可能直接使用通用kernel更划算形状大时则切换到分块kernel甚至不同的并行策略。Relax能够让你在IR层面表达这个决策逻辑避免在设备端做多次形状判断。比如可以写一个函数先匹配输入形状是否为(1, max_len)再决定是否调用高度特化的kernel。这个逻辑如果是写在Python前端会引入一次额外的Python解释层开销但放在Relax IR里VM直接一次性编译完成决策在执行期瞬间完成开销几乎可以忽略。5.3 面向硬件后端的分层编译Relax天然把编译流程分成了“高层IR优化”和“底层TIR代码生成”两个阶段。这意味着你可以把前端模型描述和后端硬件优化解耦。前端用Relax把网络逻辑表达清楚后端用TVM的TIR机制做循环变换和代码生成。我试过把同一个Relax模块分别编译到CPU和GPU后端主流程几乎没有变化只有底层TIR的Pass序列和target不同。这个抽象带来的灵活性在自己维护多个硬件后端时会非常舒服——你不必为一个新后端重写整个前端逻辑只需要关注后端算子实现和调度策略。5.4 可定制的Pass管线与调试利器最后提一点工程上的体会Relax对自定义Pass的支持比Relay更友好。因为解析级别分层你在写DataflowBlock级别的优化时可以精准地拿到所有需要的结构信息而在写函数级的变换时又能通过访问者模式遍历所有控制流节点不需要额外处理各种查询分支。我经常在调试的时候用TVM内置的IRModule.show()方法把所有Pass应用前后的Relax脚本打印下来做对比并通过tvm.script.parser.relax.parse把自己写的测试片段快速变成模块验证Pass是否按预期工作。这套调试流程在开发自定义编译流程时极其高效也适合作为学习Relax内部机制的入口。我个人在实际操作中的体会是理解Relax抽象层不能只靠看文档必须亲手写一个带控制流和符号形状的小模块用TVMScript把它描述出来、编译、跑通然后再试着改动它的某个分支看看IR和运行结果如何变化。这种“代码在手、语义我有”的正反馈循环是搞懂任何IR系统最快的方式。
返回列表