- 编译器
- 深度学习
- 模型优化
【免费下载链接】tvm
Open deep learning compiler stack for cpu, gpu and specialized accelerators
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_key为DFPatternNode,所有模式节点(如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):调用运算符,构造CallPattern。args为子模式列表;传入None表示匹配任意参数。注意__call__对args的处理:CallPattern(self, args)中的args会被包装成数组,用于匹配 Call 节点的参数列表。pattern | other:AltPattern,两个模式二选一。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是二者之间"模糊路径"(通常匹配逐元素算子)。若path为None,默认使用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四、约束模式:类型、数据格式、形状与属性
模式可以叠加约束,让匹配更加精确。所有这些约束方法都返回新的模式对象(TypePattern、DataTypePattern、ShapePattern、AttrPattern),因此可以链式调用。
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)约束张量形状,shape为List[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_map为tvm.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)是其实例方法形式。
参数说明:
attrs:Optional[Dict[str, Object]],添加到被划分函数上的属性字典;check:Callable[[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()) |
| 卷积后必有 ReLU | is_op("nn.relu")(is_op("nn.conv2d")(wildcard(), wildcard())) |
| 卷积后可有可无 ReLU | is_op("nn.conv2d")(wildcard(), wildcard()).optional(is_op("nn.relu")) |
| ReLU 或 LeakyReLU | is_op("nn.relu")(wildcard()) \| is_op("nn.leaky_relu")(wildcard()) |
| 权重为常量的 Dense | is_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 实践建议
- 先用
match验证模式:在写重写/划分逻辑前,用pattern.match(expr)快速验证模式是否符合预期拓扑。 - 善用
is_var而非wildcard划分参数:partition以模式中的VarPattern为依据生成被划分函数的参数,需要"哪些输入应成为函数参数"时用is_var精确控制。 - 属性约束注意类型:
has_attr只支持 Op 属性;数组属性(如kernel_size)在 C++ 端会做递归展开比较,可放心使用嵌套结构。 - 重写时注意
rewrite_once与require_type:需要类型信息做判断时设require_type=True;只想改第一处匹配时设rewrite_once=True。 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
相关推荐
TVM Relay DataFlow Pattern 模式匹配语言:从图重写到算子融合的声明式框架
TVM Relay DataFlow Pattern 模式匹配语言:从图重写到算子融合的声明式框架 本文围绕 TVM 中 Relay 数据流图的模式匹配语言(D
编译器深度学习模型优化Apache TVM Relax 数据流模式语言(DPL)完全指南:从图匹配到算子融合与后端分发
Apache TVM Relax 数据流模式语言(DPL)完全指南:从图匹配到算子融合与后端分发 本指南围绕 Apache TVM 中 Relax 前端内置的数
模型编译深度学习推理引擎Apache TVM Relax 数据流模式语言(DPL)实战指南:图模式匹配与自动改写
Apache TVM Relax 数据流模式语言(DPL)实战指南:图模式匹配与自动改写 导读 tvm.relax.dpl (Dataflow Pattern
模型编译深度学习推理引擎
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考