TVM Relay 数据流模式语言(DataFlow Pattern)完全指南:从模式匹配到子图重写与划分
2026/9/23 22:38:58 网站建设 项目流程
  • 编译器
  • 深度学习
  • 模型优化

【免费下载链接】tvm

Open deep learning compiler stack for cpu, gpu and specialized accelerators

项目地址:https://gitcode.com/gh_mirrors/tvm7/tvm
点击查看免费下载

TVM Relay 的数据流模式(DataFlow Pattern,简称 DFPattern)是一套面向数据流图的模式匹配语言,它让开发者可以用接近 Python 表达式的语法描述"卷积 + 激活"这类算子组合,从而在 Relay IR 中定位、重写、划分子图。本文以 dataflow_pattern.rst 所对应的 Python API 为主体,结合 python/tvm/relay/dataflow_pattern/init.py 的完整源码实现、C++ 侧的模式匹配器以及 test_dataflow_pattern.py 测试用例,系统讲解模式语言的全部构件、组合语法、约束条件,以及rewrite(重写)和partition(子图划分)两大核心工作流,帮助你在自定义 Pass、算子融合和代码生成中熟练运用这套工具。

一、什么是 Relay 数据流模式语言

在 Relay 中做图优化(如算子融合、算子替换)时,传统做法是用 ExprVisitor 手工遍历表达式树,用大量 if/else 判断节点类型与拓扑结构,代码冗长且容易出错。数据流模式语言提供了一种声明式的替代方案:先用模式(Pattern)描述目标子图结构,再让引擎在图上寻找所有匹配

模式与普通 Relay 表达式(relay.Expr)的区别在于:模式是"结构模板",其中的某些位置可以留空(通配符),也可以附加数据类型、形状、属性等约束;而表达式是具体的图。DFPattern 的基类定义在 include/tvm/relay/dataflow_pattern.h 中,其_type_keyDFPatternNode,所有模式节点(如relay.dataflow_pattern.ExprPattern)均继承自它。

Python 侧的所有 API 都通过tvm._ffi._init_api("relay.dataflow_pattern", __name__)绑定到 C++ 实现(见 python/tvm/relay/dataflow_pattern/_ffi.py),而模式匹配的核心引擎位于 src/relay/ir/dataflow_matcher.cc,其中DFPatternMatcher::Match负责把模式与表达式逐一比对,DFPatternVisitor则在 src/relay/ir/dataflow_pattern_functor.cc 中实现对模式树的递归遍历。

二、模式的基础构件:语法糖函数

tvm.relay.dataflow_pattern提供了一组以is_开头的语法糖函数,用于快速构造各种基础模式。它们大多只是具体模式类的构造函数包装。

语法糖函数底层模式类语义
wildcard()WildcardPattern匹配任意表达式(通配符)
is_var(name="")VarPattern匹配 Relay 变量;name为空时匹配任意 Var
is_constant()ConstantPattern匹配 Relay 常量
is_expr(expr)ExprPattern按结构等价匹配指定表达式
is_op(op_name)ExprPattern匹配指定名字的算子(如"nn.conv2d"
is_tuple(fields)TuplePattern匹配元组,fields为各字段模式
is_tuple_get_item(tuple, index=None)TupleGetItemPattern匹配取元组元素;index=None表示任意下标
is_if(cond, true_branch, false_branch)IfPattern匹配If表达式三分支
is_let(var, value, body)LetPattern匹配Let绑定的三个部分
has_type(ttype, pattern=None)TypePattern为模式附加类型约束
has_dtype(dtype, pattern=None)DataTypePattern为模式附加数据类型约束
has_shape(shape, pattern=None)ShapePattern为模式附加形状约束
has_attr(attrs, pattern=None)AttrPattern为模式附加属性约束
dominates(parent, path, child)DominatorPattern匹配"支配"结构

2.1 通配符与变量

wildcard()是最常用的构件,表示"这里可以是任何表达式"。is_var(name)除了能匹配特定名字的变量外,还在partition流程中承担重要角色——它决定了被划分出的函数的参数如何生成。

2.2 算子与表达式

is_op通过tvm.relay.op.get(op_name)取得算子后包装成ExprPattern,因此它匹配的是"以该算子为头的 Call 节点":

from tvm.relay.dataflow_pattern import is_op, wildcard # 匹配任意输入的 relu 算子 relu_pat = is_op("nn.relu")(wildcard())

is_expr使用**结构等价(structural equality)**判断,源码注释明确说明ExprPattern的匹配依赖表达式结构等价(见 include/tvm/relay/dataflow_pattern.h)。

三、模式组合语法:运算符重载

基础模式通过DFPattern上的运算符重载进行组合,这是模式语言最直观的部分(Python 实现见 python/tvm/relay/dataflow_pattern/init.py,C++ 侧声明见 include/tvm/relay/dataflow_pattern.h):

  • pattern(*args):调用运算符,构造CallPatternargs为子模式列表;传入None表示匹配任意参数。注意__call__args的处理:CallPattern(self, args)中的args会被包装成数组,用于匹配 Call 节点的参数列表。
  • pattern | otherAltPattern,两个模式二选一。
  • pattern + other/pattern - other/pattern * other/pattern / other:分别等价于is_op("add")(self, other)is_op("subtract")(self, other)is_op("multiply")(self, other)is_op("divide")(self, other)

例如,测试用例 test_dataflow_pattern.py 中的模式:

add_pattern = is_op("add")(wildcard(), wildcard())

等价写法是wildcard() + wildcard()。C++ 侧还额外提供了operator||(对应 Python 的|)与Optional方法。

3.1 可选模式:optional

DFPattern.optional(option_constructor)AltPattern的快捷方式,等价于self | option_constructor(self)。它用于描述"可有可无"的后续算子,例如卷积后可能跟 bias 再跟 relu:

pattern = is_op("nn.conv2d")(wildcard(), wildcard()).optional( lambda x: is_op("nn.bias_add")(x, wildcard()) )

测试 test_dataflow_pattern.py 展示了链式optional的写法:pattern.optional(is_op("nn.relu")).optional(is_op("tanh")),即卷积后面可以有 relu,relu 后面还可以有 tanh,两者都可缺省。

3.2 支配模式:dominates

DFPattern.dominates(parent, path=None)构造DominatorPattern,用于匹配"单生产者、单最终消费者"的链式结构:parent是产出数据的节点,child(即self)是该数据链上所有节点的最终使用者,path是二者之间"模糊路径"(通常匹配逐元素算子)。若pathNone,默认使用wildcard()

测试 test_dataflow_pattern.py 中:

P = is_op("nn.conv2d")(wildcard(), wildcard()) # 'parent' I = is_op("nn.relu")(wildcard()) # 'intermediate' ('path') pattern = I.dominates(P) # 卷积支配 relu

四、约束模式:类型、数据格式、形状与属性

模式可以叠加约束,让匹配更加精确。所有这些约束方法都返回新的模式对象(TypePatternDataTypePatternShapePatternAttrPattern),因此可以链式调用。

4.1 has_type:类型约束

DFPattern.has_type(ttype)要求被匹配的表达式具有指定的 Relay 类型(tvm.ir.type.Type)。底层构造TypePattern(pattern, ttype)

4.2 has_dtype:数据类型约束

DFPattern.has_dtype(dtype)约束张量数据类型,dtype为字符串,如"float32"。例如只匹配 float32 的卷积:

conv_fp32 = is_op("nn.conv2d")(wildcard(), wildcard()).has_dtype("float32")

4.3 has_shape:形状约束

DFPattern.has_shape(shape)约束张量形状,shapeList[tvm.ir.PrimExpr]。例如:

conv_3x3 = is_op("nn.conv2d")(wildcard(), wildcard()).has_shape([1, 3, 224, 224])

4.4 has_attr:算子属性约束

DFPattern.has_attr(attrs)接受Dict[str, Object],内部通过make_node("DictAttrs", **attrs)构造属性字典,然后包装成AttrPattern注意源码注释特别说明:目前只支持 Op 属性匹配,不支持 Call 属性(见 python/tvm/relay/dataflow_pattern/init.py)。

测试用例 test_dataflow_pattern.py 展示了多种属性约束:

# 匹配 NCHW 布局的卷积 is_conv2d = is_op("nn.conv2d")(wildcard(), wildcard()).has_attr({"data_layout": "NCHW"}) # 匹配 3x3 卷积 is_conv2d = is_op("nn.conv2d")(wildcard(), wildcard()).has_attr({"kernel_size": [3, 3]})

属性值支持标量(字符串、数字)与数组(如[3, 3])。在 C++ 匹配端,MatchRetValue会递归展开数组,并对运行时类型(如runtime::Int)与编译期 IR 类型做自动转换,保证((0,0),(0,0))这类嵌套结构也能正确比较(见 src/relay/ir/dataflow_matcher.cc)。

五、核心操作一:match 模式匹配

match(pattern, expr)返回布尔值,判断表达式是否命中模式。DFPattern.match(expr)是它的实例方法形式,最终都调用 FFI 函数ffi.match(见 python/tvm/relay/dataflow_pattern/init.py)。

C++ 侧 src/relay/ir/dataflow_matcher.cc 的Match实现要点:

  • 每次匹配前清空memo_(记忆化缓存)与matched_nodes_记录;
  • VisitDFPattern采用带回溯的记忆化匹配:命中则把 (pattern, expr) 记入 memo,未命中则通过ClearMap回滚到匹配前的"水位线",避免分支匹配失败污染后续结果;
  • 匹配成功后,matched_nodes_中保存了 pattern 与 expr 的对应关系,供后续重写/划分使用。
from tvm.relay.dataflow_pattern import is_op, wildcard, match pat = is_op("add")(wildcard(), wildcard()) x, y = relay.var("x"), relay.var("y") assert match(pat, x + y) # True assert not match(pat, x * y) # False

六、核心操作二:rewrite 模式重写

6.1 DFPatternCallback 回调

模式重写通过继承DFPatternCallback实现(python/tvm/relay/dataflow_pattern/init.py):

  • require_type:回调前是否需要先执行类型推断(InferType);
  • rewrite_once:为True时只执行一次回调;
  • pattern:子类必须提供要匹配的模式;
  • callback(self, pre, post, node_map):命中时被调用,三个参数分别是——pre为原始图中命中的表达式;post为输入已被改写后的表达式;node_maptvm.ir.container.Map[DFPattern, List[Expr]],保存模式节点到命中的表达式列表的映射。回调返回值将替换匹配到的子图。

测试 test_dataflow_pattern.py 中的经典示例——把加法改写成减法:

class TestRewrite(DFPatternCallback): def __init__(self): super(TestRewrite, self).__init__() self.pattern = add_pattern def callback(self, pre, post, node_map): return post.args[0] - post.args[1] out = rewrite(TestRewrite(), x + y) assert sub_pattern.match(out)

6.2 rewrite 函数

rewrite(callbacks, expr, mod=None)接受单个回调或回调列表,对表达式执行重写(python/tvm/relay/dataflow_pattern/init.py)。实现细节:

  • mod参数可选,用于关联 IRModule;若未提供则内部创建空IRModule()
  • 回调在传入 C++ 前会被包装为_DFPatternCallback(即 C++ 侧DFPatternCallback对象),并断言pattern非空;
  • 返回重写后的表达式。

test_rewrite_func(test_dataflow_pattern.py)还展示了重写发生在函数调用内部的场景——对func(x, w) + y执行重写后,函数体中的加法同样被改写。

七、核心操作三:partition 子图划分

partition(pattern, expr, attrs=None, check=None)是数据流模式最强大的能力:把表达式中所有命中模式的子图提取为独立的 Relay 函数,并用对该函数的调用替换原子图(python/tvm/relay/dataflow_pattern/init.py)。DFPattern.partition(expr, attrs, check)是其实例方法形式。

参数说明:

  • attrsOptional[Dict[str, Object]],添加到被划分函数上的属性字典;
  • checkCallable[[Expr], bool],对命中的表达式做更复杂检查的回调,返回True才继续划分,默认恒为True

7.1 实战:划分 BatchNorm

测试 test_partition_batchnorm 展示了完整的 BatchNorm 划分流程:

BN = gamma * (x - mean) / relay.op.sqrt(var + eps) + beta # 定义匹配 BN 计算链的模式 class BatchnormCallback(DFPatternCallback): def __init__(self): super(BatchnormCallback, self).__init__() self.pattern = ( is_op("add")( is_op("divide")( is_op("multiply")(wildcard(), is_op("subtract")(wildcard(), wildcard())), is_op("sqrt")(is_op("add")(wildcard(), is_constant())), ), wildcard(), ) ) partitioned = BatchnormCallback().pattern.partition(BN)

划分结果是一个带PartitionedFromPattern属性的函数调用,属性值形如"subtract_multiply_add_sqrt_divide_add_"(记录模式的操作序列),这正是下游代码生成器(如外部代码生成后端)识别子图来源的依据。test_partition_double_batchnorm(test_dataflow_pattern.py)验证了嵌套 BatchNorm 会被划分成两个嵌套的函数调用;test_partition_check(第 1545 行附近)则验证check回调可以阻止不符合额外条件的匹配被划分。

7.2 划分与融合的关系

partition是 Relay 中算子融合(如FuseOps)和外部代码生成(BYOC)的基础设施:先定义"哪些算子组合可以融合/下放",再通过模式划分把它们聚合成函数,最后交给后端。测试 test_partition_overused 与test_partition_fuzzy_tuple(第 1489 行)等用例覆盖了变量被多处复用、元组/函数参数等复杂拓扑下的划分行为。

八、模式树遍历与 C++ 实现原理

理解了 Python API 之后,再看底层实现能更好地把握模式语言的能力边界。

8.1 模式节点体系

所有模式类在 Python 侧通过register_df_node注册到 C++ 对象系统(python/tvm/relay/dataflow_pattern/init.py),其 C++ 类型 key 均为relay.dataflow_pattern.*。完整的模式节点家族定义在 include/tvm/relay/dataflow_pattern.h 中,与 Python 一一对应:

  • ExprPattern(匹配字面表达式,结构等价比较)
  • VarPattern(匹配变量,可选名字)
  • ConstantPattern(匹配常量)
  • CallPattern(匹配 Call 节点:算子 + 参数列表,args可为空表示任意参数)
  • FunctionPattern(匹配函数:参数列表 + 函数体)
  • IfPattern/LetPattern(匹配控制流结构)
  • TuplePattern/TupleGetItemPattern(匹配元组及其取元素操作)
  • AltPattern(二选一)
  • WildcardPattern(通配)
  • TypePattern/DataTypePattern/ShapePattern/AttrPattern(四类约束)
  • DominatorPattern(支配结构)

8.2 匹配器与访问器

  • DFPatternMatcher(src/relay/ir/dataflow_matcher.cc):负责模式与表达式的比对。核心是带 memo 的回溯匹配,并维护node_map供重写/划分使用。AltPattern的匹配即left || right短路求值(第 71-73 行)。
  • DFPatternVisitor(src/relay/ir/dataflow_pattern_functor.cc):对模式树本身做递归遍历,用visited_集合去重防止环。例如CallPattern会依次访问其算子模式与每个参数模式(第 46-53 行)。
  • 匹配结果中"模式节点 → 命中的表达式列表"的映射通过matched_nodes_累积,最终暴露给 Python 回调的node_map参数。

九、速查与最佳实践

9.1 常用组合速查

需求写法
匹配任意卷积is_op("nn.conv2d")(wildcard(), wildcard())
卷积后必有 ReLUis_op("nn.relu")(is_op("nn.conv2d")(wildcard(), wildcard()))
卷积后可有可无 ReLUis_op("nn.conv2d")(wildcard(), wildcard()).optional(is_op("nn.relu"))
ReLU 或 LeakyReLUis_op("nn.relu")(wildcard()) \| is_op("nn.leaky_relu")(wildcard())
权重为常量的 Denseis_op("nn.dense")(wildcard(), is_constant())
限定数据布局is_op("nn.conv2d")(wildcard(), wildcard()).has_attr({"data_layout": "NCHW"})
限定算子类型标记is_op("nn.dense").has_attr({"TOpPattern": K_ELEMWISE})

9.2 实践建议

  1. 先用match验证模式:在写重写/划分逻辑前,用pattern.match(expr)快速验证模式是否符合预期拓扑。
  2. 善用is_var而非wildcard划分参数partition以模式中的VarPattern为依据生成被划分函数的参数,需要"哪些输入应成为函数参数"时用is_var精确控制。
  3. 属性约束注意类型has_attr只支持 Op 属性;数组属性(如kernel_size)在 C++ 端会做递归展开比较,可放心使用嵌套结构。
  4. 重写时注意rewrite_oncerequire_type:需要类型信息做判断时设require_type=True;只想改第一处匹配时设rewrite_once=True
  5. node_map用于获取匹配子表达式:回调中通过node_map[pattern]取出命中的原始表达式,例如拿到卷积的输入/权重做进一步分析。

十、参考资料与延伸阅读

  • API 参考入口:docs/reference/api/python/relay/dataflow_pattern.rst
  • Python 实现(本文主体):python/tvm/relay/dataflow_pattern/init.py
  • FFI 绑定:python/tvm/relay/dataflow_pattern/_ffi.py
  • C++ 模式节点定义:include/tvm/relay/dataflow_pattern.h
  • 匹配引擎:src/relay/ir/dataflow_matcher.cc 与 src/relay/ir/dataflow_matcher_impl.h
  • 模式树访问器:src/relay/ir/dataflow_pattern_functor.cc
  • 单元测试:tests/python/relay/test_dataflow_pattern.py(覆盖匹配、约束、optional、dominates、rewrite、partition 全场景)
  • C++ 测试:tests/cpp/dataflow_pattern_test.cc
  • Relax 侧同类实现(供对照):include/tvm/relax/dataflow_pattern.h 与 tests/python/relax/test_dataflow_pattern.py
  • 编译器
  • 深度学习
  • 模型优化

【免费下载链接】tvm

Open deep learning compiler stack for cpu, gpu and specialized accelerators

项目地址:https://gitcode.com/gh_mirrors/tvm7/tvm
点击查看免费下载

相关推荐

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询