- 模型编译
- 深度学习
- 推理引擎
【免费下载链接】tvm
Open Machine Learning Compiler Framework
导读
本文聚焦 TVM TIRx CUDA 后端归约算子(sum/max/min)的local变体——即当源(src)与目标(dst)缓冲区都位于local(寄存器/线程私有)存储作用域时,编译器如何完成算子分发、布局校验与代码生成。你将掌握该变体的三层递进算法路径(线程级顺序归约、warp 级 laneid shard→replica 专用 shuffle 归约、通用 warp/warpgroup 视图归约)、每个输入的语义影响,以及对应源码位置与测试用例的验证方式,可直接用于阅读或扩展 TIRx 归约管线。
变体定位:归约算子家族中的 local 路径
在 TIRx 的归约算子体系中,sum、max、min三种操作都通过register_dispatch在 CUDA 目标上注册了多个变体,按操作数存储作用域与硬件能力区分,见 reduction 索引页:
| 变体 | 优先级 | 降级策略 |
|---|---|---|
reduction/local | 10 | local 缓冲区 src/dst;线程级顺序归约(可选 warp shuffle) |
reduction/shared | 10 | shared 缓冲区 src/dst;自适应分组__shfl_xor树 |
reduction/sm100_packed(packed_add_sum/3input_maxmin) | 20 | CUDA SM100+ 线程级 fp32 ≥8 元素:打包add.f32x2/max3/min3 |
本文讨论的local变体位于 python/tvm/backend/cuda/tile_primitive/reduction/local.py,是三个变体中唯一要求 src 与 dst 均为local作用域的实现。
接受条件:predicate 双层校验
注册声明
local变体在local.py末尾通过循环注册,三个算子共用同一套校验逻辑(见 local.py#L471-L489):
for op_name, op_type in [ ("sum", ReduceOpType.SUM), ("max", ReduceOpType.MAX), ("min", ReduceOpType.MIN), ]: @register_dispatch( op_name, "cuda", variant="local", priority=10, when=[ predicate("storage_scope", _match_reduction_storage_scope, expected_scope=["local"]), predicate("local_valid", validate_reduction_local), ], ) def _local_dispatch(op: TilePrimitiveCall, sctx: DispatchContext, _op_type=op_type) -> PrimFunc: op = TilePrimitiveCall.downcast(op) return reduction_local_impl(op, _op_type, sctx)第一层_match_reduction_storage_scope(定义于 reduction/utils.py#L61-L71)要求src 与 dst 的作用域同时匹配local;第二层validate_reduction_local再深入检查执行作用域、布局与轴信息。
详细要求
validate_reduction_local(local.py#L152-L224)逐条落实以下约束:
| 属性 | 要求 |
|---|---|
| target / priority | cuda;优先级10 |
| 操作数作用域 | src 与 dst 均为local,且 dtype 相等(src.dtype != dst.dtype直接拒绝) |
| 执行作用域 | thread(顺序、线程本地);warp/warpgroup需要合法且非 swizzle 的TileLayout。warp 作用域下若 src/dst 呈 laneid shard→replica 模式,则自动选择专用 shuffle 路径;否则thread_reduce=True可为通用路径附加 shuffle 步骤。warpgroup 作用域拒绝thread_reduce=True(报错"thread_reduce=True is only supported in warp scope; warpgroup local reduction is thread-local only") |
| 形状 | thread 作用域下校验器不检查轴、不比较 src/dst 尺寸,轴分析与循环构造推迟到 emit 阶段;宽作用域(warp/warpgroup)视图归约要求空间维的布局结构匹配(线程/局部 extent 一致),被归约维在 dst 中的局部 extent 为 1 |
从源码结构看,warp 作用域的校验顺序有一个值得注意的细节:先尝试_analyze_shuffle_reduce识别 laneid shard→replica 模式并提前放行,因为该模式中 laneid 出现在 dst 的 replica(广播)里,会被后续通用布局校验拒绝(见 local.py#L188-L194 的注释)。
演示程序:线程级 4 元素向量归约
原文档给出的演示程序(出自 test_reduction.py)在单个线程内把一个 4 元素float32local 向量归约为标量:
@Tx.prim_func def test_func(A_ptr: Tx.handle, B_ptr: Tx.handle): A = Tx.match_buffer(A_ptr, [4], "float32", layout=TileLayout(S[(4,)])) B = Tx.match_buffer(B_ptr, [1], "float32", layout=TileLayout(S[(1,)])) Tx.device_entry(); Tx.cta_id([1]); Tx.thread_id([1]) A_local = Tx.alloc_buffer([4], "float32", scope="local") B_local = Tx.alloc_buffer([1], "float32", scope="local") for i in Tx.serial(4): A_local[i] = A[i] Tx.tile.sum(B_local, A_local, accum=False) # reduction local dispatch B[0] = B_local[0]由于归约长度只有 4(< 8),此用例停留在local路径而非跳到 SM100 的packed_add_sum/3input_maxmin快速路径(参见 reduction/sm100_packed 文档与 reduction/sm100_packed.py)。归约算子其余参数(reduce_axes、accum、config)由_reduction_args统一解析,见 reduction/utils.py#L48-L58。
算法:三条路径的选择逻辑
reduction_local_impl(local.py#L403-L464)根据执行作用域与布局形态选择路径:
路径一:线程级顺序归约(_emit_reduction_local_thread_wise)
当执行作用域为thread时走此路径(local.py#L227-L283):对输出位置做空间循环,每个位置先初始化为算子单位元(sum→0.0、max→T.min_value(dtype)、min→T.max_value(dtype),见 reduction/utils.py#L40-L45),除非accum=True;随后内层归约循环逐元素累积源数据。该路径完全没有跨线程通信:
for spa in Tx.serial(spatial_len): if not accum: dst[spa] = identity for red in Tx.serial(reduction_len): dst[spa] = op(dst[spa], src[spa, red])其中spatial_len与reduction_len分别是空间维与归约维 extent 的乘积,负轴先经_analyze_axes归一化为非负下标(reduction/utils.py#L74-L82)。_emit_reduction_local_thread_wise的 docstring 给出了一个 2D 示例:Tx.sum(B_local[0:2, 0:3], A_local[0:2, 0:3, 0:4], [-1], False)展开为spatial_len=6、reduction_len=4的双层循环(local.py#L24-L34)。
路径二:专用 laneid shard→replica shuffle 归约(_gen_warp_shuffle_reduce)
当 src 具有全跨度(32 lanes)的 laneid shard、dst 具有 2 的幂次 laneid replica 时,_analyze_shuffle_reduce(local.py#L90-L127)返回(reduce_width, local_elems),其中:
reduce_width:参与每组归约的 lane 数(dst replica 中 laneid 迭代子 extent 的乘积),必须为正、≤32 且为 2 的幂;local_elems:src 中非 laneid shard extent 的乘积,即每个线程持有的元素个数;- swizzle 布局直接返回
None(不支持)。
命中后由_gen_warp_shuffle_reduce(local.py#L130-L149)生成实现:每个 lane 先把本线程对应的局部值复制到 dst 视图(src/dst为同一缓冲区时跳过复制),再对每个元素执行T.cuda.warp_reduce(dst_local[k], op_str, reduce_width)。该路径由布局自动触发,与thread_reduce无关,也不先运行通用局部轴循环;同时它不分支处理accum参数,因此总是用 shuffle 结果覆盖 dst。
路径三:通用 warp/warpgroup 视图归约(_emit_reduction_local_view)
未命中专用模式时,warp/warpgroup 走此路径(local.py#L286-L400):通过get_local_region从TileLayout中分解出每个线程的局部区域,把 src 的局部归约维逐一归约进 dst 的每个位置。在 warp 作用域下,thread_reduce=True会额外发出显式的tvm_warp_shuffle_xor步骤,掩码来自_compute_shuffle_masks(reduction/utils.py#L111-L128)——对归约维中每个线程迭代子按其 stride 生成stride * 2^i(i 从 0 到 log2(extent))的 XOR 掩码并升序排列,同时用T.tvm_warp_activemask()取活跃掩码。warpgroup 作用域只支持纯 local 部分,不生成 shuffle。
该路径还有一个关键组合:accum=True且需要 shuffle 时,实现会先在归约前把旧 dst 值保存到临时old_val局部缓冲区,完成局部归约与 shuffle 后再与新结果结合,保证累积语义正确(local.py#L355-L378)。
生成的 IR 与 CUDA 代码
对演示中的 4 元素线程级归约,生成的 TIRx IR 为:
for spa in Tx.serial(1): dst[...] = Tx.float32(0) for red in Tx.serial(4): dst[...] = dst[...] + src[...] # op = sum对应的 CUDA C++ 为:
for (int red = 0; red < 4; ++red) B_local_ptr[0] = B_local_ptr[0] + A_local_ptr[red];该用例已在sm_100a上验证,B == sum(A)成立(测试由 test_reduction.py 中的 GPU 测试覆盖)。注意 IR 与 CUDA 均不含跨线程通信,完全符合路径一的顺序语义。
输入如何改变算法行为
| 输入 | 影响 |
|---|---|
| op | sum→+,max→max,min→min(以及对应单位元);算子到字符串的映射表_REDUCE_OP_TO_STR见 reduction/utils.py#L216 |
| exec scope | thread→顺序路径;warp 下匹配 shard→replica 布局→专用warp_reduce;否则 warp/warpgroup 走通用局部视图路径,仅thread_reduce=True时附加 warp shuffle(warpgroup 不允许) |
| axes / shape | 决定空间循环与归约循环的 extent;_analyze_axes将负轴归一化 |
| accum | 线程级与通用视图路径按上文方式复用旧 dst 值(线程级为跳过初始化,视图+shuffle 为保存旧值后合并);专用 shard→replica 路径忽略该标志,总是覆盖 dst |
测试与验证
local变体的行为由 tests/python/tirx/operator/tile_primitive/cuda/reduction/test_reduction.py 系统性验证,覆盖了三类场景:
test_reduction_local_thread_wise:对多种 shape/axes 组合(1D→1D、4D→2D、3D→2D、非 2 的幂、带 offset 切片)参数化测试线程级顺序归约,并同时参数化op_type(sum/max/min)、dtype(float32/float16)与accum;test_reduction_local_view_basic/test_reduction_local_view_complex:前者验证纯 local 布局的视图归约(含切片输入),后者使用 WGMMA 风格的多层 tile 布局并参数化thread_reduce(shuffle 开关)与accum,还演示了“先thread_reduce=False归约、再用Tx.warp.sum(red_view, red_view, thread_reduce=True)补一步 shuffle”的用法;test_reduction_op_warp_shuffle系列:覆盖 laneid shard→replica 专用路径——32 lane 全 warp 归约到 1 值并广播、每线程多元素(4 元素×32 lane)、稀疏存储槽(storage span 7 但仅 4 元素)等边界情形,另有test_reduction_warp_shuffle_multi_warp_loop验证循环内 thread→warp 作用域交替的跨 warp 归约组合。
这些测试在tir_pipeline="tirx"下编译,并在真实 CUDA 设备上用tvm.testing.assert_allclose与 NumPy 参考结果比对。此外 test_dispatcher.py 覆盖分发框架本身的注册与 predicate 判定逻辑。
小结
local变体是 TIRx 归约体系中最贴近硬件线程模型的实现:线程级路径完全避免通信,专用 shard→replica 路径把归约折叠进单条warp_reduce,通用视图路径则以布局分解为代价换取 warp/warpgroup 下对任意 local 视图的归约能力。阅读 local.py 及其配套的 reduction/utils.py 和 test_reduction.py,即可完整掌握该变体的分发条件、三条算法路径与验证手段,为后续在 TIRx 中编写或扩展归约类 tile primitive 提供直接参考。
- 模型编译
- 深度学习
- 推理引擎
【免费下载链接】tvm
Open Machine Learning Compiler Framework
相关推荐
Apache TVM TIRx CUDA 归约原语(Reduction Tile Primitive)完全指南:local / shared / SM100 packed 三变体的调度、算法与代码生成
Apache TVM TIRx CUDA 归约原语(Reduction Tile Primitive)完全指南:local / shared / SM100 p
模型编译深度学习推理引擎TVM TIRx tcgen05 张量内存与寄存器异步拷贝:copy_async tmem<->local 变体(tcgen05.ld/st)深度解析
TVM TIRx tcgen05 张量内存与寄存器异步拷贝:copy_async tmem< local 变体(tcgen05.ld/st)深度解析 本篇技术指
模型编译深度学习推理引擎TVM TIRx CUDA 分布式共享内存拷贝:copy_async 的 dsmem 变体深度解析
TVM TIRx CUDA 分布式共享内存拷贝:copy_async 的 dsmem 变体深度解析 导读 在 TVM TIRx(Tile IR eXtensio
模型编译深度学习推理引擎
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考