Apache TVM Tirx 指南:smem 变体如何把共享内存上的 elementwise 算子低开销地向量化执行

发布时间:2026/9/23 12:09:11
Apache TVM Tirx 指南:smem 变体如何把共享内存上的 elementwise 算子低开销地向量化执行
模型编译深度学习推理引擎【免费下载链接】tvmOpen Machine Learning Compiler Framework项目地址https://gitcode.com/gh_mirrors/tv/tvm点击查看免费下载本篇技术指南聚焦 Apache TVM当前仓库Tirx 前端中elementwise tile primitive 的smem变体当某个逐元素算子sqrt、exp、add、fma……的所有缓冲区操作数都位于共享内存shared*时smem变体如何从执行作用域“综合”出一个[outer, threads, vec]三段式划分并以向量化方式逐元素应用算子。读完本文你将掌握smem变体的接受条件predicate、分块vec_chunk选择算法、生成的 TIRx IR 与 CUDA 形态以及 unary / binary /fma等不同算子与 dtype 输入如何影响最终调度。背景elementwise 的 reg / smem 双变体 dispatch在 Tirx 中elementwise 是一组以“逐元素应用某个算子”为语义的 tile primitive覆盖cast、fill、unary 算子zero、reciprocal、sqrt、exp、exp2、log2、silu、binary 算子add、sub、mul、fdiv、maximum以及fma。父页面 docs/tirx/tile_primitives/elementwise.rst 明确说明每个算子都会注册两个变体——reg与smem优先级都是 10二者通过“缓冲区操作数的存储作用域”作为互斥判别条件reg变体所有缓冲区操作数都在local作用域寄存器化路径划分由本地缓冲区的布局“诱导”产生见 docs/tirx/tile_primitives/elementwise/reg.rstsmem变体所有缓冲区操作数都在shared*作用域划分从执行作用域的线程数“综合”产生——这正是本文的主题。这里的“双变体模型”直接沿袭了 copy 系列 PR-640 的两变体设计诱导 vs 综合两者共享同一套OpSpec数据模型与 vec 选择辅助函数仅在“划分来源”上分道扬镳。从源码结构看python/tvm/backend/cuda/tile_primitive/elementwise/__init__.py通过from .register import *一次性注册全部算子而 register.py 中_register_smem(spec)为每个ALL_OPS条目调用register_dispatch( spec.name, cuda, variantsmem, priority10, when[predicate(f{spec.name}_smem, is_smem_ewise(spec))], ) def _dispatch(op: TilePrimitiveCall, sctx: DispatchContext, _specspec) - PrimFunc: return emit_smem(op, _spec, sctx)也就是说smem的调度入口是is_smem_ewise(spec)谓词 emit_smem(...)发射函数二者都定义在 smem.py 中。注意父页面同时提醒全局Tx.tile目录中的其他构造器如minimum、memset、select目前并没有这套 CUDAreg/smem变体调用一个算子必须由所选后端注册了对应变体才能成功分发。它接受什么is_smem_ewise谓词与接受条件表is_smem_ewise(spec)构建的谓词在 dispatch 阶段被求值逐条检查执行上下文与操作数任一条件不满足就以(False, reason)拒绝分发。文档给出的谓词骨架如下def check(op_call, sctx): if not sctx.is_target(cuda): return False, non-cuda target if sctx.scope_kind not in (thread, warp, warpgroup, cta): ... ok, reason _all_threads_active(sctx) # full scope plan, msg spec.parse(op_call) # parse the ops operands for br in buffer_regions(plan): if not br.buffer.scope().startswith(shared): # every buffer operand shared* return False, foperand scope {br.buffer.scope()} ! shared* if br.buffer.layout is None: ... # spec.check_extras (dtype rules) and anchor-layout validation对照源码 smem.py 的真实实现检查链比文档骨架更完整依次是目标后端sctx.is_target(cuda)非 CUDA 目标直接拒绝执行作用域scope_kind必须是thread/warp/warpgroup/cta之一线程全活跃_all_threads_active(sctx)——作用域内的线程必须全部参与例如laneid覆盖全部 32 条 lane未被外层if收窄解析算子spec.parse(op_call)把调用解析成(Plan, msg)其中Plan记录目标 regiondst与源列表srcsbuffer region 或标量作用域检查对buffer_regions(plan)即 dst 加所有 buffer 型 src逐一要求br.source.scope().startswith(shared)且布局不能为Nonedtype 规则若该算子定义了spec.check_extras则校验plan.extras与compute_dtype_of(plan)的一致性广播兼容性以plan.dst为锚点anchor每个 buffer src 的 region 形状必须与目标形状满足 NumPy 式右对齐广播对应维度 extent 相等或为 1由shape_broadcast_compat校验。把文档的属性表格与源码实现合并smem变体接受的输入可以归纳为PropertyRequirementtarget / scope / prioritycudathread/warp/warpgroup/cta线程全活跃priority10operands每个缓冲区操作数含输出都在shared*fill、binary 算子与fma允许标量源op父页面列出的任意 CUDA elementwiseOpSpecunarysqrt/exp/zero…binaryadd/mul…fmaspec.check_extras校验 dtype 组合layout操作数必须带布局划分由作用域线程数综合产生dtype、逻辑最内层 region extent、每线程元素数共同约束调度块宽度关于第 7 条广播_common.py的shape_broadcast_compat以目标形状为锚点做右对齐若某 src 形状与锚点同 rank 则逐维要求“相等或为 1”rank 小于锚点时前面补对齐不允许某维既不等也不为 1。这一广播语义在发射阶段由_broadcast_indices落实详见后文“标量回退路径”。演示程序256 线程 CTA 对 32×32 共享 tile 求 sqrt文档给出的演示程序展示了一个完整的最小使用示例一个 CTA 对32×32的float32共享 tile 做 elementwisesqrt改编自test_unary.py这里用 256 线程的 CTA因此划分恰好是一轮s_layout TileLayout(S[(32, 32)]); full (slice(0, 32), slice(0, 32)) Tx.prim_func def unary_op(A_ptr: Tx.handle): A Tx.match_buffer(A_ptr, (32, 32), float32, layouts_layout) Tx.device_entry(); Tx.cta_id([1]); Tx.warp_id([8]); Tx.lane_id([32]); Tx.thread_id([256]) A_smem Tx.alloc_buffer((32, 32), float32, scopeshared, layouts_layout) Tx.tile.cta.copy(A_smem[full], A[full]) Tx.tile.cta.sqrt(A_smem[full], A_smem[full]) # elementwise smem dispatch Tx.tile.cta.copy(A[full], A_smem[full])该程序在 tests/python/tirx/operator/tile_primitive/cuda/elementwise/test_unary.py 中有着对应的真实测试形态test_unary_op_shared以(32, 32)的 global shape、64线程的thread_cnt参数化zero/sqrt两种算子并覆盖float16→float16、float32→float16、float32→bfloat16三组 dtype 组合测试还包含带偏移的场景如 global shape(32, 8, 12)源/结果起点分别为(10, 0, 3)与(20, 0, 2)extent(5, 6, 7)用于验证 region 切片后smem变体仍能正确按块发射。同目录下的test_binary.py、test_fma.py则分别覆盖 binary 算子与fma的 smem 路径。值得注意Tx.tile.cta.sqrt(A_smem[full], A_smem[full])的源与目标是同一个共享缓冲区in-place 使用这正说明smem变体是“先向量化读、逐分量算子、再向量化写回”的读改写模式而不是需要额外临时空间的 copy。在测试的非 in-place 场景src_dtype ! dst_dtype中则声明第二个 global bufferB作为输出sqrt的目标 region 指向B_smem。算法三步走smem变体的降低过程可以拆成三个步骤这与 copy 的gmem_smemvec_auto路径见 docs/tirx/tile_primitives/copy/gmem_smem.rst一脉相承。1. 解析算子并检查操作数spec.parse把Tx.tile.cta.sqrt(dst, src)这样的调用转成Plan记录 dst region、srcs 列表与 extras谓词随后确认每个缓冲区操作数都是shared*。以 unary.py 的_parse_unary为例T.unary(dst, src[, bias, scale])会被解析成src若是TensorRegion则记为 buffer src若是 prim expr 则记为标量 src可选的biasFloatImm常量或 buffer与scale进入extras。binary 的解析binary.py则更讲究两个输入都是常量会被拒绝若常量在左操作数且算子不可交换sub等直接报 “non-commutative op ... cannot have constant lhs”允许时把常量挪到右侧广播场景下若 src1 的元素数更少且算子可交换则交换使大 buffer 留在 src1保持“src1 与 dst 同形”的约定。2. 从线程数综合[outer, threads, vec]划分与gmem_smem一样因为操作数本身不携带线程划分信息smem变体从执行作用域的线程总数综合三段式划分。关键的宽度选择函数是_max_layout_vec(plan, total, thread_cnt)smem.py其逻辑为计算所有 dst/src buffer 中最宽的 dtype 位宽max_bits标量 src 参与其中吗不参与——只统计 buffer region计算per_thread total // thread_cnt若total % thread_cnt ! 0直接返回 1必须整除否则三段式无整数解收集 dst 与各 buffer src 的逻辑最内层 region extent按候选位宽{128, 64, 32, 16, 8}从宽到窄尝试n cand_bits // max_bits必须为正、必须整除per_thread、且必须整除每个操作数的最内层 extent满足则返回该n全部失败则返回 1标量粒度。关键语义_max_layout_vec不检查物理布局的 stride/连续性——正如文档强调的“The current width selection does not inspect physical layout contiguity”。连续性的检查被推迟到向量发射阶段详见“packed 路径”由_emit_vec所依赖的布局最内层 stride 1 条件把关。线程数本身取自get_thread_cnt(sctx)而非 copy 的_thread_cnt——源码注释smem.py解释了原因_thread_cnt计算∏ sctx.intra在 cta 作用域下对 sub-warp 计数会因 warpid extent 向下取整而静默返回 0而get_thread_cnt读取launch_params[threadIdx.x].dom.extent对所有作用域都正确。另外emit_smem断言launch_params中不含threadIdx.y/threadIdx.z即当前 smem 发射假定一维 threadIdx。以演示程序为例32×32 1024个float32256 线程per_thread 4max_bits 32候选128 / 32 4整除per_thread 4也整除最内层 extent 32于是vec 4、outer 1024 / (256 × 4) 1——正好一轮。3. 逐元素应用算子向量化读改写得到vec_chunk后调度在pick_vec_chunk(spec, op_call, sctx, plan, vec_max)_common.py中选择发射方式遍历该算子spec.vec_impls中“按宽度从大到小排序”的 packed 实现找到第一个vec_len整除vec_max且impl.applies(...)通过的实现走_emit_packed否则以vec_max为块宽走标量回退_emit_scalar。_emit_packedsmem.py按outer轮串行循环每轮中计算该线程本轮首个元素的一维融合下标fused0 s * vec_chunk * thread_cnt tid * vec_chunk用if fused0 vec_chunk total跳过尾部的残缺块predicate the call对vec_chunk个 lane 用get_indices(fused0 k, dst_st, dst_ext)从一维融合下标反解多维索引共享缓冲区的多维索引由 buffer 自身布局在 codegen 时解析为物理地址对每个 src标量源直接取值buffer 源则给出对应 lane 的多维索引广播场景经由_src_lane_indices→_broadcast_indices从 dst 索引推导调用_emit_vec(vec_impl, dst_buf, dst_lane_indices, src_args, extras)发射 packed 向量语句循环结束后按作用域插入同步emit_scope_synccta→T.cuda.cta_sync()warpgroup→warpgroup_sync(8)warp→warp_sync()thread无同步。_emit_scalarsmem.py则是逐元素版本for s in serial(n_outer)外层 ×for vec in T.vectorized(vec_chunk)内层每个fused用fetch_src_value取各 src 的值buffer / 标量 / 广播三种来源统一处理再T.cast(compute(src_vals, extras, dst_dtype), dst_dtype)写回 dst。注意这里内层用的是Tx.vectorized——即使没有 packed PTX 实现向量宽度仍然以vectorized标注形式保留给后续 codegen 优化。生成的 TIRx IR 与 CUDA一窥向量化形态文档给出了_emit_scalar对演示程序的 TIRx IR 输出outer 1、vec 4for f in Tx.serial(1): # outer 1 for vec in Tx.vectorized(4): A_smem[tid * 4 vec] Tx.sqrt(A_smem[tid * 4 vec])当算子注册了可用的 packed 实现例如 binary 的add/sub/mul挂有 sm_100 的f32x2packed PTX见 binary.py 与vec_emit/binary_f32x2fdiv、maximum因单条FMNMX/max.f32指令本身精确、无 rounding/ftz 变体可打包而只走标量回退发射会走_emit_packed对应 CUDA 端把vec 4的 4 个float32打包成一个float4一次性向量化读出后逐分量应用算子float4 v_ *(float4*)(A_smem_ptr[tid * 4]); __1.x sqrtf(v_.x); __1.y sqrtf(v_.y); __1.z sqrtf(v_.z); __1.w sqrtf(v_.w);文档注明该结果已在sm_100a上验证——生成的 tile 恰好等于sqrt(A)。这里float4的 16 字节对齐读写在 packed 路径成立的前提是_max_layout_vec选中宽度时 dtype 位宽与 128 位候选的整除关系4 × 32 bit 128 bit以及最内层维度物理 stride 为 1_emit_vec的 packed 发射要求非 swizzle 切片下 lane 物理连续smem.py 的模块注释明确写了这一点。输入如何改变算法文档用一张表格总结了三种输入对算法的影响结合源码可以展开如下inputeffectopunary → 逐分量sqrtf/expf/ …binary → 两个输入组合a bfma→a * b c。具体 CUDA intrinsic 由OpSpec.compute_scalar标量路径或vec_implspacked 路径决定dtype约束候选宽度_max_layout_vec先取所有操作数中最宽的位宽max_bits再用{128,64,32,16,8}位候选求n cand_bits // max_bits每线程元素数与各操作数逻辑最内层 extent 可以进一步收窄它从而改变轮数outer。当前宽度选择不检查物理布局的连续 stridescope决定线程轴名称与线程数warp→laneidcta→threadIdx.x等等进而决定综合出的划分。_tid_exprsmem.py在thread作用域返回常量 0单线程在集合作用域用_axis_decl声明对应轴对比同页的reg变体docs/tirx/tile_primitives/elementwise/reg.rst二者的“输入→算法”映射差异正是判别核心reg的划分由 anchor 布局的线程轴诱导丢弃线程迭代器每线程保留私有 bundle且要求操作数的线程/local/replica 布局签名一致smem的划分则纯粹由作用域线程数综合操作数只需在共享内存中带布局即可。也正因smem不依赖操作数布局携带线程信息它天然适合“共享 tile 上做逐元素变换再写回”的流水线中段——上游 copy 把数据搬进shared*smem就地变换再交给下游 copy/算子消费全程无需经过寄存器搬运。验证路径与运行前提smem变体的行为由测试直接背书tests/python/tirx/operator/tile_primitive/cuda/elementwise/test_unary.py 中的test_unary_op_shared参数化地验证了zero/sqrt在多种 shape含带偏移 region、多种线程数与 dtype 组合下的正确性测试标注了pytest.mark.gpu与skipif(not env.has_cuda())——运行这些测试需要有 CUDA 环境。文档中演示的 TIRx IR 与 CUDA 输出是在sm_100aBlackwell 架构上验证的属于当前仓库后端代码生成覆盖的硬件前提。若要亲手复现文档中的演示程序需要在本仓库已构建的 TVM 环境中使用tvm.script.tirxfrom tvm.script import tirx as Tx与tvm.tirx.layout的S/TileLayout编写 prim_func然后经LowerTIRx等 pass 降级为可运行的内核由于smem变体由register_dispatch以variantsmem注册正常调用Tx.tile.cta.sqrt(...)或其他列出的算子时只要操作数全在shared*作用域、线程全活跃dispatch 就会自动命中该变体无需手动指定变体名。赞分享模型编译深度学习推理引擎【免费下载链接】tvmOpen Machine Learning Compiler Framework项目地址https://gitcode.com/gh_mirrors/tv/tvm点击查看免费下载相关推荐TVM TIRx CUDA Buffers 完全指南从参数缓冲区到共享内存、Tensor Memory 与 Buffer 视图TVM TIRx CUDA Buffers 完全指南从参数缓冲区到共享内存、Tensor Memory 与 Buffer 视图 导读 本文是 TVM 开源仓库模型编译深度学习推理引擎Apache TVM TIRx 安装指南tvm.tirx 编译器与 tirx-kernels 内核库的完整部署方案Apache TVM TIRx 安装指南tvm.tirx 编译器与 tirx kernels 内核库的完整部署方案 TIRx读作 tier ex 是模型编译深度学习推理引擎TVM TIRx 张量布局Layout详解从 S/R/O 规格到 TMEM 与 SMEM 布局实战TVM TIRx 张量布局Layout详解从 S/R/O 规格到 TMEM 与 SMEM 布局实战 导读 TIRx 是 TVM 中面向现代加速器尤其是模型编译深度学习推理引擎创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考