Apache TVM tvm.te 张量表达式(TE)Python API 完全指南:从计算声明到 TensorIR 桥接
【免费下载链接】tvmOpen Machine Learning Compiler Framework项目地址: https://gitcode.com/gh_mirrors/tv/tvm
本篇技术指南围绕 Apache TVM 的tvm.te(Tensor Expression Language,张量表达式语言)Python 命名空间展开,它是 Python API 参考文档 所对应的核心模块。读者将掌握tvm.te中张量声明、算子定义、归约、扫描、外部函数接入等全部核心 API 的签名与用法,理解其底层实现(te.Operation/te.Tensor对象模型),并学会用create_prim_func将 TE 计算无缝桥接到可调度的 TensorIR,最终产出可在 TVM 运行时上执行的编译模块。
一、tvm.te的定位与模块结构
在 Apache TVM 中,tvm.te承担"计算声明"(compute declaration)的职责:开发者用 Python 语法描述"计算什么"(形状、数据依赖、算术规则),而不关心"怎么算"。后续由调度(schedule)、编译与代码生成负责"如何高效地算"。
tvm.te 的 API 参考页 通过 Sphinxautomodule指令自动生成模块文档:
tvm.te ------ .. Exclude the ops imported from tirx. .. automodule:: tvm.te :members: :imported-members: :autosummary:该页面被收纳在 Python API 总览 的tvm.te目录下(与topi并列)。注意页面中的注释"Exclude the ops imported from tirx":tvm.te会重新导出大量来自tvm.tirx的算子(exp、tanh、sigmoid、if_then_else、sum、min、max等),文档生成时会排除这些第三方来源的算子,只聚焦 TE 自身定义的 API。
从源码看,tvm.te的实际实现分布在 python/tvm/te/ 下的四个文件中:
| 文件 | 职责 |
|---|---|
| python/tvm/te/init.py | 命名空间入口,统一导出 TE 核心 API 与 tirx 算子 |
| python/tvm/te/operation.py | 计算声明 API:placeholder、compute、scan、extern等 |
| python/tvm/te/tensor.py | 对象模型:Tensor、TensorSlice、Operation及其子类 |
| python/tvm/te/tag.py | 算子标签机制:tag_scope/TagScope |
其中 python/tvm/te/init.py 的模块文档字符串明确写道:"""Namespace for Tensor Expression Language"""。它负责:
- 从
tvm.tirx重导出算术与逻辑算子(exp、log、power、floordiv、isnan等),保证向后兼容的tvm.te.xxx写法; - 从
tvm.tirx重导出归约算子与基础设施(comm_reducer、min、max、sum、CommReducer、Reduce); - 从本目录导出 TE 自身的
TensorSlice、Tensor、tag_scope、placeholder、compute、scan、extern、var、const、thread_axis、reduce_axis、create_prim_func、extern_primfunc以及PlaceholderOp、ComputeOp、ScanOp、ExternOp等类型。
二、核心对象模型:Tensor与Operation
tvm.te的一切围绕两个基础对象展开,定义于 python/tvm/te/tensor.py:
te.Tensor
Tensor表示一个将要被计算的张量,注册为"te.Tensor"对象(tensor.py#L251-L303)。其关键属性/方法:
shape:张量形状;ndim:维度数,即len(self.shape);dtype:元素数据类型(通过 FFI 调用TensorDType获得);op:产生该张量的操作节点;name:单个输出的 op 直接取op.name,多输出 op 则为{op.name}.v{value_index};__call__(*indices):按索引取元素,返回加载表达式(TensorLoad),索引数量必须等于ndim,否则抛出ValueError;__getitem__:返回TensorSlice,支持切片语法A[0][i]与A[0, i]两种写法;- 丰富的运算符重载(
+、-、*、/、%、位运算、比较运算、astype、equal等),使 TE 表达式可以像普通 Python 数值一样书写。
te.Operation及其子类
Operation表示"产生张量的操作",注册为"te.Operation"(tensor.py#L307-L334),提供output(index)、num_outputs、input_tensors。TE 共有四类 op 子类,对应四种计算声明方式:
| 类 | 注册名 | 来源 API | 语义 |
|---|---|---|---|
PlaceholderOp | te.PlaceholderOp | te.placeholder | 输入占位,数据来自外部 |
ComputeOp(继承BaseComputeOp) | te.ComputeOp | te.compute | 在形状域上逐元素计算 |
ScanOp | te.ScanOp | te.scan | 沿时间轴递推的扫描 |
ExternOp | te.ExternOp | te.extern/te.extern_primfunc | 调用外部函数/内联 TIR PrimFunc |
TensorSlice:切片语法支持
TensorSlice(tensor.py#L33-L196)是辅助数据结构,支持A[i, j]、A[0][i]等切片写法并累积索引,最终通过asobject()转成真正的TensorLoad表达式。它同样实现了全套运算符重载,因此在te.compute的 lambda 中可以直接对切片结果做算术。
三、te.placeholder:声明输入张量
te.placeholder(shape, dtype=None, name="placeholder")实现在 operation.py#L37-L57。它构造一个空张量(占位符),由外部数据在运行时填充:
shape:Tuple of Expr,可为符号表达式(如te.var),也可为具体整数;dtype:默认"float32"(dtype = "float32" if dtype is None else dtype);name:张量名提示,便于生成的代码阅读与调试。
典型用法(来自 tests/python/te/test_te_tensor.py):
m, n, l = te.var("m"), te.var("n"), te.var("l") A = te.placeholder((m, l), name="A") B = te.placeholder((n, l), name="B") T = te.compute((m, n, l), lambda i, j, k: A[i, k] * B[j, k])placeholder底层通过_ffi_api.Placeholder创建PlaceholderOp,其输出即Tensor。当形状中使用te.var时,生成的 IR 会保留符号维度,使同一个计算定义可适配不同尺寸的输入。
四、te.compute:定义逐元素计算
te.compute(shape, fcompute, name="compute", tag="", attrs=None, varargs_names=None)实现在 operation.py#L60-L140,是 TE 使用频率最高的 API。其计算规则为result[axis] = fcompute(axis)。参数含义:
shape:输出张量形状;若传入单个原始表达式会自动包装为元组,且 float 形状会被转成 int;fcompute:indices -> value的 lambda,描述每个输出元素的取值;name:名称提示,默认"compute";tag:额外标签(见第七节tag_scope);若当前处于tag_scope上下文内再传tag会抛出ValueError("nested tag is not allowed for now");attrs:附加辅助属性字典;varargs_names:为fcompute的可变参数(*args)指定名字,默认按i1, i2, ...自动命名。
fcompute的签名规则
源码用inspect.getfullargspec解析fcompute的签名,规则如下(operation.py#L100-L128):
- 无参数:
lambda: expr,自动获得i0, i1, ...命名,适用于 rank-0 计算; - 有
*varargs:可变参数吃掉剩余维度,可用varargs_names自定义名称,数量不匹配会抛出RuntimeError; - 固定参数少于输出维度:只保留
len(args)个维度,剩余维度被隐式广播(implicit broadcast); - 禁用
**kwargs、默认参数、仅关键字参数:fcompute中不支持,源码中分别以assert拒绝。
最终通过tvm.tirx.IterVar((0, s), x, 0)为每个输出维度创建迭代变量,把fcompute(*vars)的返回体包装成ComputeOp(operation.py#L130-L140)。fcompute可返回单个表达式或列表/元组(多输出),compute依据 op 的num_outputs返回单个Tensor或输出元组。
示例与测试佐证
来自 tests/python/te/test_te_tensor.py 的几个真实用法:
- 元素级运算与广播(L30-L32):
T = te.compute((m, n, l), lambda i, j, k: A[i, k] * B[j, k]) - rank-0 输出(L55-L58):
T = te.compute((), lambda: te.sum(A[k] * scale(), axis=k))——scale是te.placeholder((), name="s"),scale()以零索引调用取得标量; - 切片语法(L77-L78):
B = te.compute((n,), lambda i: A[0][i] + A[0][i])。
五、归约与te.reduce_axis
归约是神经网络算子的核心(求和、求最大、点积等)。reduce_axis实现在 operation.py#L513-L535:
te.reduce_axis(dom, name="rv", thread_tag="")dom:归约迭代域(Range);name:变量名,默认"rv";thread_tag:可选线程标签。
它创建类型为 2(reduction)的IterVar。归约算子te.sum、te.min、te.max等由tvm.tirx提供并通过 python/tvm/te/init.py 重导出,另有comm_reducer可自定义归约算子。
多轴归约
te.sum的axis参数同时支持元组和列表(test_te_tensor.py#L81-L88):
k1 = te.reduce_axis((0, m), name="k1") k2 = te.reduce_axis((0, n), name="k2") C = te.compute((1,), lambda _: te.sum(A[k1, k2], axis=(k1, k2))) C = te.compute((1,), lambda _: te.sum(A[k1, k2], axis=[k1, k2]))自定义归约算子comm_reducer
用comm_reducer可定义自己的归约语义。测试中的例子(test_te_tensor.py#L91-L97)展示了自定义归约与内建sum的等价性:
mysum = te.comm_reducer(lambda x, y: x + y, lambda t: tvm.tirx.const(0, dtype=t.dtype)) C = te.compute((m,), lambda i: mysum(A[i, k], axis=k))comm_reducer接收两个函数:combiner(如何合并两个值)和 identity(单位元),返回可当作归约算子使用的对象。
带条件的归约
test_tensor_reduce_multiout_with_cond(test_te_tensor.py#L123-L135)演示了多输出、带if_then_else条件表达式的归约写法,其中idx、val均为int32占位符。
六、te.scan:时序扫描计算
scan用于沿时间轴递推的计算(如 RNN、cumsum),实现在 operation.py#L143-L207:
te.scan(init, update, state_placeholder, inputs=None, name="scan", tag="", attrs=None)init:前init.shape[0]个时间戳的初始条件,Tensor或Tensor列表;update:给定符号状态张量后的递推更新规则,Tensor或列表;state_placeholder:update中使用的状态占位张量;inputs:扫描的输入列表,非必需,但有助于编译器更快识别扫描体;- 约束:
init、update、state_placeholder长度必须一致,否则抛ValueError; - 返回:单输出返回
Tensor,多输出返回元组。
文档字符串中给出了等价于numpy.cumsum的完整示例(operation.py#L175-L186):
m = te.var("m") n = te.var("n") X = te.placeholder((m, n), name="X") s_state = te.placeholder((m, n)) s_init = te.compute((1, n), lambda _, i: X[0, i]) s_update = te.compute((m, n), lambda t, i: s_state[t-1, i] + X[t, i]) res = tvm.te.scan(s_init, s_update, s_state, X)底层实现通过tvm.tirx.IterVar((init[0].shape[0], update[0].shape[0]), f"{name}.idx", 3)创建类型为 3(scan)的迭代轴,再构造ScanOp(operation.py#L204-L206)。
七、te.extern与te.extern_primfunc:接入外部实现
te.extern
当某个算子已有高效的外部实现(如 BLAS),可用extern直接嵌入,实现在 operation.py#L210-L351:
te.extern(shape, inputs, fcompute, name="extern", dtype=None, in_buffers=None, out_buffers=None, tag="", attrs=None)shape:输出形状(单个元组或多个元组的列表);inputs:输入Tensor列表;fcompute:(ins, outs) -> stmt,其中ins/outs是tvm.tirx.Buffer列表,返回值必须是tvm.tirx.Stmt(普通表达式会被自动包装为Evaluate,否则抛ValueError);dtype:输出数据类型,默认与输入一致;当输入类型不唯一时必须显式给出;in_buffers/out_buffers:可显式指定输入/输出 buffer,数量不匹配会抛RuntimeError。
文档字符串中的典型示例是调用tvm.contrib.cblas.matmul(operation.py#L270-L282):
A = te.placeholder((n, l), name="A") B = te.placeholder((l, m), name="B") C = te.extern((n, m), [A, B], lambda ins, outs: tvm.tirx.call_packed( "tvm.contrib.cblas.matmul", ins[0], ins[1], outs[0], 0, 0), name="C")可以看到:当未显式提供 buffer 时,源码会为输入自动decl_buffer(含elem_offset变量),输出则按推断/给定的 dtype 创建对应 buffer(operation.py#L303-L339)。
te.extern_primfunc
更现代的方式是直接把一个 TVMScript 编写的、可调度的 TIR PrimFunc 内联进 TE 计算图(operation.py#L354-L435):
A = te.placeholder((128, 128), name="A") B = te.placeholder((128, 128), name="B") @T.prim_func(s_tir=True) def before_split(a: T.handle, b: T.handle) -> None: A = T.match_buffer(a, (128, 128)) B = T.match_buffer(b, (128, 128)) for i, j in T.grid(128, 128): with T.sblock("B"): vi, vj = T.axis.remap("SS", [i, j]) B[vi, vj] = A[vi, vj] * 2.0 C = te.extern_primfunc([A, B], func)其底层逻辑是:通过DomainTouchedAccessMap分析 PrimFunc 参数的读写访问,自动区分输入/输出 buffer,支持原地(inplace)输出,并逐一校验传入input_tensors与 PrimFunc 输入 buffer 的形状一致性(operation.py#L392-L426),最终复用extern构造ExternOp。
八、符号变量与常量:te.var/te.const
te.var(name="tindex", dtype="int32", span=None):创建符号变量(operation.py#L438-L457),默认int32,返回tvm.tirx.Var。它用于构造符号形状、符号索引,使计算定义与具体尺寸解耦;te.const(value, dtype="int32", span=None):创建常量表达式(operation.py#L460-L479),value支持 bool、int、float、numpy 数组、tvm.runtime.Tensor。
这两个 API 是 python/tvm/te/init.py 导出的基础构件,几乎所有te.compute示例中都用te.var声明符号形状。
九、te.thread_axis:线程轴声明
实现在 operation.py#L482-L510:
te.thread_axis(dom=None, tag="", name="", span=None)- 当
dom传字符串时,会被解释为 tag(tag, dom = dom, None); tag必填,否则抛ValueError("tag must be given as Positional or keyword argument");- 默认
name取 tag 值; - 返回类型为 1(thread)的
IterVar。
thread_axis在后续调度(schedule)阶段用于把循环绑定到 GPU/CPU 线程维度,是bind操作的关键输入。
十、te.tag_scope:为算子打标签
tag_scope(python/tvm/te/tag.py)既可作为上下文管理器,也可作为装饰器,为作用域内的compute/scan/extern自动附加标签,供下游调度与优化参考。实现上基于TagScope单例:compute等 API 内部会通过TagScope.get_current()读取当前标签(见 operation.py#L91-L94)。
上下文管理器用法(tag.py#L77-L94):
n, m, l = te.var('n'), te.var('m'), te.var('l') A = te.placeholder((n, l), name='A') B = te.placeholder((m, l), name='B') k = te.reduce_axis((0, l), name='k') with tvm.te.tag_scope(tag='matmul'): C = te.compute((n, m), lambda i, j: te.sum(A[i, k] * B[j, k], axis=k))装饰器用法(同一文档示例):
@tvm.te.tag_scope(tag="conv") def compute_relu(data): return te.compute(data.shape, lambda *i: tvm.tirx.Select(data(*i) < 0, 0.0, data(*i)))注意TagScope的实现约束:不允许嵌套(__enter__中已有当前作用域时抛ValueError("nested op_tag is not allowed for now"));作用域退出时若标签从未被任何算子使用,会发出UserWarning。仓库中的 tests/python/te/test_te_tag.py 专门覆盖了这些行为。
十一、te.create_prim_func:桥接 TensorIR 调度
create_prim_func是 TE 通往现代 TensorIR(s_tir)调度体系的关键桥梁,实现在 operation.py#L538-L592:
te.create_prim_func(ops, index_dtype_override=None) -> tirx.PrimFuncops:源表达式(Tensor或tirx.Var的列表);index_dtype_override:可选,覆盖索引数据类型;- 返回可被 TensorIR 调度器(如
s_tir的 schedule)继续变换的PrimFunc。
文档字符串中的完整示例(operation.py#L548-L583)——定义 128×128 的 matmul 并查看生成的 TVMScript:
import tvm from tvm import te from tvm.te import create_prim_func A = te.placeholder((128, 128), name="A") B = te.placeholder((128, 128), name="B") k = te.reduce_axis((0, 128), "k") C = te.compute((128, 128), lambda x, y: te.sum(A[x, k] * B[y, k], axis=k), name="C") func = create_prim_func([A, B, C]) print(func.script())生成的等价 TVMScript 为:
@T.prim_func(s_tir=True) def tir_matmul(a: T.handle, b: T.handle, c: T.handle) -> None: A = T.match_buffer(a, (128, 128)) B = T.match_buffer(b, (128, 128)) C = T.match_buffer(c, (128, 128)) for i, j, k in T.grid(128, 128, 128): with T.sblock(): vi, vj, vk = T.axis.remap("SSR", [i, j, k]) with T.init(): C[vi, vj] = 0.0 C[vi, vj] += A[vi, vk] * B[vj, vk]其中S/R分别表示空间轴与归约轴,T.init()给出累加初值。该函数可直接交给s_tir调度器做分块、向量化、并行等变换。仓库中的 tests/python/te/test_te_create_primfunc.py 对该桥接路径做了系统验证。
十二、从 TE 到可运行模块的完整链路
将 TE 计算编译为可执行模块的经典路径(可在 docs/get_started/tutorials/quick_start.py 找到端到端示例):
import tvm from tvm import te n = te.var("n") A = te.placeholder((n,), name="A") B = te.compute((n,), lambda i: A[i] + 1, name="B") func = te.create_prim_func([A, B]) # 1. TE -> TensorIR PrimFunc mod = tvm.build(func, target="llvm") # 2. 编译为可执行模块要点:
- 先用
placeholder/compute/scan/extern声明计算图; - 用
create_prim_func(或用tvm.lower等工具)把 TE 计算转换成 IR;若需深度性能优化,可在 TensorIR 阶段接入调度; - 用
tvm.build按目标平台(llvm、cuda、opencl等)生成模块,之后即可分配ndarray输入并调用模块执行。
tvm.te本身不负责调度与代码生成,它聚焦"计算是什么"的声明;而driver/build_module(见 python/tvm/driver/build_module.py)负责把声明变成可运行产物。
十三、测试与进一步探索
tvm.te的功能正确性由仓库中的专项测试保障,可作为深入学习与验证的入口:
- tests/python/te/test_te_tensor.py:
Tensor对象模型、rank-0 计算、切片语法、多轴/自定义归约、带条件归约等; - tests/python/te/test_te_create_primfunc.py:TE 到 TensorIR
PrimFunc的转换与索引类型覆盖; - tests/python/te/test_te_tag.py:
tag_scope上下文管理器/装饰器语义; - tests/python/te/test_te_verify_compute.py:计算声明合法性校验。
从源码结构可以推断,tvm.te作为 TVM 中最稳定、最底层的 Python 计算声明层,其设计目标始终如一:用 Python 写出可验证、可调度、可跨后端部署的计算描述。无论是手写算子、接入外部库(te.extern),还是与 TensorIR 调度体系衔接(te.create_prim_func/te.extern_primfunc),它都是进入 TVM 编译栈的第一站。结合 Python API 参考 中tvm.te与tvm.topi的目录关系,可将 TE 视为高层算子库(topi)的底层语言基础。
【免费下载链接】tvmOpen Machine Learning Compiler Framework项目地址: https://gitcode.com/gh_mirrors/tv/tvm
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考