ARTICLE DETAIL

资讯详情

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

Apache TVM tvm.te 张量表达式(TE)Python API 完全指南:从计算声明到 TensorIR 桥接

Apache TVM tvm.te 张量表达式(TE)Python API 完全指南:从计算声明到 TensorIR 桥接 Apache TVM tvm.te 张量表达式TEPython API 完全指南从计算声明到 TensorIR 桥接【免费下载链接】tvmOpen Machine Learning Compiler Framework项目地址: https://gitcode.com/gh_mirrors/tv/tvm本篇技术指南围绕 Apache TVM 的tvm.teTensor 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 tirxtvm.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计算声明 APIplaceholder、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与Operationtvm.te的一切围绕两个基础对象展开定义于 python/tvm/te/tensor.pyte.TensorTensor表示一个将要被计算的张量注册为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.Operationtensor.py#L307-L334提供output(index)、num_outputs、input_tensors。TE 共有四类 op 子类对应四种计算声明方式类注册名来源 API语义PlaceholderOpte.PlaceholderOpte.placeholder输入占位数据来自外部ComputeOp继承BaseComputeOpte.ComputeOpte.compute在形状域上逐元素计算ScanOpte.ScanOpte.scan沿时间轴递推的扫描ExternOpte.ExternOpte.extern/te.extern_primfunc调用外部函数/内联 TIR PrimFuncTensorSlice切片语法支持TensorSlicetensor.py#L33-L196是辅助数据结构支持A[i, j]、A[0][i]等切片写法并累积索引最终通过asobject()转成真正的TensorLoad表达式。它同样实现了全套运算符重载因此在te.compute的 lambda 中可以直接对切片结果做算术。三、te.placeholder声明输入张量te.placeholder(shape, dtypeNone, nameplaceholder)实现在 operation.py#L37-L57。它构造一个空张量占位符由外部数据在运行时填充shapeTuple of Expr可为符号表达式如te.var也可为具体整数dtype默认float32dtype float32 if dtype is None else dtypename张量名提示便于生成的代码阅读与调试。典型用法来自 tests/python/te/test_te_tensor.pym, n, l te.var(m), te.var(n), te.var(l) A te.placeholder((m, l), nameA) B te.placeholder((n, l), nameB) 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, namecompute, tag, attrsNone, varargs_namesNone)实现在 operation.py#L60-L140是 TE 使用频率最高的 API。其计算规则为result[axis] fcompute(axis)。参数含义shape输出张量形状若传入单个原始表达式会自动包装为元组且 float 形状会被转成 intfcomputeindices - value的 lambda描述每个输出元素的取值name名称提示默认computetag额外标签见第七节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)的返回体包装成ComputeOpoperation.py#L130-L140。fcompute可返回单个表达式或列表/元组多输出compute依据 op 的num_outputs返回单个Tensor或输出元组。示例与测试佐证来自 tests/python/te/test_te_tensor.py 的几个真实用法元素级运算与广播L30-L32T te.compute((m, n, l), lambda i, j, k: A[i, k] * B[j, k])rank-0 输出L55-L58T te.compute((), lambda: te.sum(A[k] * scale(), axisk))——scale是te.placeholder((), names)scale()以零索引调用取得标量切片语法L77-L78B te.compute((n,), lambda i: A[0][i] A[0][i])。五、归约与te.reduce_axis归约是神经网络算子的核心求和、求最大、点积等。reduce_axis实现在 operation.py#L513-L535te.reduce_axis(dom, namerv, thread_tag)dom归约迭代域Rangename变量名默认rvthread_tag可选线程标签。它创建类型为 2reduction的IterVar。归约算子te.sum、te.min、te.max等由tvm.tirx提供并通过 python/tvm/te/init.py 重导出另有comm_reducer可自定义归约算子。多轴归约te.sum的axis参数同时支持元组和列表test_te_tensor.py#L81-L88k1 te.reduce_axis((0, m), namek1) k2 te.reduce_axis((0, n), namek2) 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, dtypet.dtype)) C te.compute((m,), lambda i: mysum(A[i, k], axisk))comm_reducer接收两个函数combiner如何合并两个值和 identity单位元返回可当作归约算子使用的对象。带条件的归约test_tensor_reduce_multiout_with_condtest_te_tensor.py#L123-L135演示了多输出、带if_then_else条件表达式的归约写法其中idx、val均为int32占位符。六、te.scan时序扫描计算scan用于沿时间轴递推的计算如 RNN、cumsum实现在 operation.py#L143-L207te.scan(init, update, state_placeholder, inputsNone, namescan, tag, attrsNone)init前init.shape[0]个时间戳的初始条件Tensor或Tensor列表update给定符号状态张量后的递推更新规则Tensor或列表state_placeholderupdate中使用的状态占位张量inputs扫描的输入列表非必需但有助于编译器更快识别扫描体约束init、update、state_placeholder长度必须一致否则抛ValueError返回单输出返回Tensor多输出返回元组。文档字符串中给出了等价于numpy.cumsum的完整示例operation.py#L175-L186m te.var(m) n te.var(n) X te.placeholder((m, n), nameX) 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)创建类型为 3scan的迭代轴再构造ScanOpoperation.py#L204-L206。七、te.extern与te.extern_primfunc接入外部实现te.extern当某个算子已有高效的外部实现如 BLAS可用extern直接嵌入实现在 operation.py#L210-L351te.extern(shape, inputs, fcompute, nameextern, dtypeNone, in_buffersNone, out_buffersNone, tag, attrsNone)shape输出形状单个元组或多个元组的列表inputs输入Tensor列表fcompute(ins, outs) - stmt其中ins/outs是tvm.tirx.Buffer列表返回值必须是tvm.tirx.Stmt普通表达式会被自动包装为Evaluate否则抛ValueErrordtype输出数据类型默认与输入一致当输入类型不唯一时必须显式给出in_buffers/out_buffers可显式指定输入/输出 buffer数量不匹配会抛RuntimeError。文档字符串中的典型示例是调用tvm.contrib.cblas.matmuloperation.py#L270-L282A te.placeholder((n, l), nameA) B te.placeholder((l, m), nameB) 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), nameC)可以看到当未显式提供 buffer 时源码会为输入自动decl_buffer含elem_offset变量输出则按推断/给定的 dtype 创建对应 bufferoperation.py#L303-L339。te.extern_primfunc更现代的方式是直接把一个 TVMScript 编写的、可调度的 TIR PrimFunc 内联进 TE 计算图operation.py#L354-L435A te.placeholder((128, 128), nameA) B te.placeholder((128, 128), nameB) T.prim_func(s_tirTrue) 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.constte.var(nametindex, dtypeint32, spanNone)创建符号变量operation.py#L438-L457默认int32返回tvm.tirx.Var。它用于构造符号形状、符号索引使计算定义与具体尺寸解耦te.const(value, dtypeint32, spanNone)创建常量表达式operation.py#L460-L479value支持 bool、int、float、numpy 数组、tvm.runtime.Tensor。这两个 API 是 python/tvm/te/init.py 导出的基础构件几乎所有te.compute示例中都用te.var声明符号形状。九、te.thread_axis线程轴声明实现在 operation.py#L482-L510te.thread_axis(domNone, tag, name, spanNone)当dom传字符串时会被解释为 tagtag, dom dom, Nonetag必填否则抛ValueError(tag must be given as Positional or keyword argument)默认name取 tag 值返回类型为 1thread的IterVar。thread_axis在后续调度schedule阶段用于把循环绑定到 GPU/CPU 线程维度是bind操作的关键输入。十、te.tag_scope为算子打标签tag_scopepython/tvm/te/tag.py既可作为上下文管理器也可作为装饰器为作用域内的compute/scan/extern自动附加标签供下游调度与优化参考。实现上基于TagScope单例compute等 API 内部会通过TagScope.get_current()读取当前标签见 operation.py#L91-L94。上下文管理器用法tag.py#L77-L94n, m, l te.var(n), te.var(m), te.var(l) A te.placeholder((n, l), nameA) B te.placeholder((m, l), nameB) k te.reduce_axis((0, l), namek) with tvm.te.tag_scope(tagmatmul): C te.compute((n, m), lambda i, j: te.sum(A[i, k] * B[j, k], axisk))装饰器用法同一文档示例tvm.te.tag_scope(tagconv) 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 通往现代 TensorIRs_tir调度体系的关键桥梁实现在 operation.py#L538-L592te.create_prim_func(ops, index_dtype_overrideNone) - tirx.PrimFuncops源表达式Tensor或tirx.Var的列表index_dtype_override可选覆盖索引数据类型返回可被 TensorIR 调度器如s_tir的 schedule继续变换的PrimFunc。文档字符串中的完整示例operation.py#L548-L583——定义 128×128 的 matmul 并查看生成的 TVMScriptimport tvm from tvm import te from tvm.te import create_prim_func A te.placeholder((128, 128), nameA) B te.placeholder((128, 128), nameB) k te.reduce_axis((0, 128), k) C te.compute((128, 128), lambda x, y: te.sum(A[x, k] * B[y, k], axisk), nameC) func create_prim_func([A, B, C]) print(func.script())生成的等价 TVMScript 为T.prim_func(s_tirTrue) 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,), nameA) B te.compute((n,), lambda i: A[i] 1, nameB) func te.create_prim_func([A, B]) # 1. TE - TensorIR PrimFunc mod tvm.build(func, targetllvm) # 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.pyTensor对象模型、rank-0 计算、切片语法、多轴/自定义归约、带条件归约等tests/python/te/test_te_create_primfunc.pyTE 到 TensorIRPrimFunc的转换与索引类型覆盖tests/python/te/test_te_tag.pytag_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),仅供参考
返回列表