Apache TVM tvm.te 张量表达式(TE)Python API 完全指南:从计算声明到 TensorIR 桥接
2026/9/23 4:19:46 网站建设 项目流程

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的算子(exptanhsigmoidif_then_elsesumminmax等),文档生成时会排除这些第三方来源的算子,只聚焦 TE 自身定义的 API。

从源码看,tvm.te的实际实现分布在 python/tvm/te/ 下的四个文件中:

文件职责
python/tvm/te/init.py命名空间入口,统一导出 TE 核心 API 与 tirx 算子
python/tvm/te/operation.py计算声明 API:placeholdercomputescanextern
python/tvm/te/tensor.py对象模型:TensorTensorSliceOperation及其子类
python/tvm/te/tag.py算子标签机制:tag_scope/TagScope

其中 python/tvm/te/init.py 的模块文档字符串明确写道:"""Namespace for Tensor Expression Language"""。它负责:

  • tvm.tirx重导出算术与逻辑算子(explogpowerfloordivisnan等),保证向后兼容的tvm.te.xxx写法;
  • tvm.tirx重导出归约算子与基础设施(comm_reducerminmaxsumCommReducerReduce);
  • 从本目录导出 TE 自身的TensorSliceTensortag_scopeplaceholdercomputescanexternvarconstthread_axisreduce_axiscreate_prim_funcextern_primfunc以及PlaceholderOpComputeOpScanOpExternOp等类型。

二、核心对象模型:TensorOperation

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]两种写法;
  • 丰富的运算符重载(+-*/%、位运算、比较运算、astypeequal等),使 TE 表达式可以像普通 Python 数值一样书写。

te.Operation及其子类

Operation表示"产生张量的操作",注册为"te.Operation"(tensor.py#L307-L334),提供output(index)num_outputsinput_tensors。TE 共有四类 op 子类,对应四种计算声明方式:

注册名来源 API语义
PlaceholderOpte.PlaceholderOpte.placeholder输入占位,数据来自外部
ComputeOp(继承BaseComputeOpte.ComputeOpte.compute在形状域上逐元素计算
ScanOpte.ScanOpte.scan沿时间轴递推的扫描
ExternOpte.ExternOpte.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。它构造一个空张量(占位符),由外部数据在运行时填充:

  • shapeTuple 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;
  • fcomputeindices -> 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):

  1. 无参数lambda: expr,自动获得i0, i1, ...命名,适用于 rank-0 计算;
  2. *varargs:可变参数吃掉剩余维度,可用varargs_names自定义名称,数量不匹配会抛出RuntimeError
  3. 固定参数少于输出维度:只保留len(args)个维度,剩余维度被隐式广播(implicit broadcast);
  4. 禁用**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))——scalete.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.sumte.minte.max等由tvm.tirx提供并通过 python/tvm/te/init.py 重导出,另有comm_reducer可自定义归约算子。

多轴归约

te.sumaxis参数同时支持元组和列表(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条件表达式的归约写法,其中idxval均为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]个时间戳的初始条件,TensorTensor列表;
  • update:给定符号状态张量后的递推更新规则,Tensor或列表;
  • state_placeholderupdate中使用的状态占位张量;
  • inputs:扫描的输入列表,非必需,但有助于编译器更快识别扫描体;
  • 约束:initupdatestate_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.externte.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/outstvm.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.PrimFunc
  • ops:源表达式(Tensortirx.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按目标平台(llvmcudaopencl等)生成模块,之后即可分配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 到 TensorIRPrimFunc的转换与索引类型覆盖;
  • 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.tetvm.topi的目录关系,可将 TE 视为高层算子库(topi)的底层语言基础。

【免费下载链接】tvmOpen Machine Learning Compiler Framework项目地址: https://gitcode.com/gh_mirrors/tv/tvm

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

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

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

立即咨询