- 模型编译
- 深度学习
- 推理引擎
【免费下载链接】tvm
Open Machine Learning Compiler Framework
本篇技术指南聚焦 TVM TIRx(TensorIR eXperimental)CUDA 后端copytile primitive 的vec_auto变体,详解其在global(全局)与 shared(共享)内存之间执行同步向量化拷贝的实现路径gmem_smem——包括可接受输入的门控条件、[outer, threads, vec]三维分区合成算法、向量宽度选择规则、生成的 TIRx IR 与 PTX 指令,以及 dtype、执行作用域(exec_scope)和 swizzled 布局对生成代码的影响。读完本文,你将掌握如何在 TIRx 脚本中编写 warp/warpgroup/CTA 级别的 global↔shared 拷贝,理解其底层向量化决策原理,并能据此预测不同输入下的生成代码形态与性能特征。
从copytile primitive 说起:vec_auto变体与gmem_smem路径
在 TIRx 中,copy是一个同步的元素拷贝原语,语义为src → dst,可在 global、shared 与 register(local)三种存储之间搬运数据。CUDA 后端当前注册了八个变体:五个显式固定宽度变体(vec_16b/vec_32b/vec_64b/vec_128b/vec_256b)、ldstmatrix、vec_auto与fallback,优先级与职责见 copy 原语总览:
| 变体 | 存储对 | 优先级 | 降级方式 |
|---|---|---|---|
vec_16b/vec_32b/vec_64b/vec_128b/vec_256b | global ↔ shared/local,或 shared ↔ local | 20 | 显式线程作用域下搬运恰好指定宽度的数据;可带 global-load cache 控制 |
vec_auto:gmem_smem 路径 | global ↔ shared | 10 | 合成[outer, threads, vec]分区,配合直接 PTX 向量 load/store |
vec_auto:reg 路径 | register ↔ shared/global | 10 | 由寄存器布局的线程轴诱导分区 |
ldstmatrix | register ↔ shared | 10 | warp 集体ldmatrix/stmatrix(m8n8 片段) |
fallback | global / shared / local | 0 | 标量单线程拷贝(兜底) |
其中vec_auto变体内部包含两条实现路径:gmem_smem(global↔shared,本文主题)与reg(寄存器参与)。需要特别强调的是,gmem_smem不是可选择的独立 dispatch 名称,而是vec_auto变体在自动选择或显式dispatch="vec_auto"时,根据操作数存储类型路由到的内部实现路径。其 dispatch 注册逻辑位于 vec_auto.py,注册优先级为 10,predicate 依次尝试_is_gmem_smem与_is_reg_copy:
@register_dispatch( "copy", "cuda", variant="vec_auto", priority=10, when=[predicate("vec_auto_applicable", _is_vec_auto_copy)], ) def copy_schedule_vec_auto(op_call, sctx): g_ok, g_reason = _is_gmem_smem(op_call, sctx) if g_ok: return _emit_gmem_smem(op_call, sctx) r_ok, r_reason = _is_reg_copy(op_call, sctx) if r_ok: return _emit_reg(op_call, sctx) fail(f"gmem_smem: {g_reason}; reg: {r_reason}")gmem_smem路径的核心特征是:拷贝两侧都是跨线程存储(global 与 shared),没有寄存器侧可以提供现成的线程分区,因此实现必须从执行作用域(execution scope)中合成一个分区——将目标区域切分为[outer, threads, vec]三维迭代,并发出串行的向量化 load/store 循环。该路径的实现文件为 vec_auto_gmem_smem.py,其布局/分区算法与ldgsts共享 _common.py 中的align_layouts_gs。
接受什么输入:_is_gmem_smem门控条件
vec_auto变体的 gmem_smem 路径由谓词_is_gmem_smem把关(源码见 vec_auto_gmem_smem.py#L79-L93):
def _is_gmem_smem(op_call, sctx): if not sctx.is_target("cuda"): return False, "non-cuda target" if sctx.scope_kind not in ("thread", "warp", "warpgroup", "cta"): return False, f"unsupported exec_scope {sctx.scope_kind}" for check in ( lambda: _all_threads_active(sctx), # full scope, no narrowing lambda: _is_valid_copy(op_call, sctx), # layouts, equal dtype/extents lambda: _scope_allowed(op_call, sctx, allowed_pairs=_GMEM_SMEM_PAIRS), lambda: _divides_thread_cnt(op_call, sctx), ): ok, msg = check() if not ok: return False, msg return True, None门控条件可归纳为下表:
| 属性 | 要求 |
|---|---|
| target | cuda |
| scope(执行作用域) | thread/warp/warpgroup/cta,且所有线程处于激活状态(_all_threads_active——laneid覆盖 32 个线程等,未被外围if收窄) |
| 存储对 | (global, shared*)或(shared*, global)——即_GMEM_SMEM_PAIRS;任一侧都不能是local |
| dtype / shape | 两侧操作数都有 layout、dtype 相等、非单位 extent 相等(_is_valid_copy→validate_copy_op) |
| 整除性 | 区域元素总数可被线程数整除(_divides_thread_cnt)——否则[outer, threads, vec]分区没有整数解,变体拒绝接受 |
其中_divides_thread_cnt的具体逻辑(vec_auto_gmem_smem.py#L51-L76)值得展开:它先通过_thread_cnt(sctx)从sctx.intra推导线程数,若thread_cnt <= 0(作用域为空 intra)则直接拒绝;随后取 global 侧 buffer 的 region,将所有 extent 相乘得到n_elements,若n_elements % thread_cnt != 0则拒绝。这样做的目的是拒绝形状不佳的拷贝(例如 1024 线程的 CTA 搬运一个 64 元素的尾部区域),而不是用慢速标量 emit 来掩盖问题。region extent 必须是常量表达式,否则同样拒绝。
演示程序:warp 往返搬运 32×32 float32 tile
来自 test_gmem_smem.py 的典型用例:一个 warp(32 线程)把32×32的float32tile 从 global 拷入 shared,再拷回 global(往返验证正确性):
from tvm.script import tirx as Tx from tvm.tirx.layout import S, TileLayout shape, dtype = (32, 32), "float32" s_layout = TileLayout(S[shape]) fs = (slice(0, 32), slice(0, 32)) @Tx.prim_func def kernel(A_ptr: Tx.handle, B_ptr: Tx.handle): A = Tx.match_buffer(A_ptr, shape, dtype) B = Tx.match_buffer(B_ptr, shape, dtype) Tx.device_entry() Tx.cta_id([1]); Tx.lane_id([32]); Tx.thread_id([32]) A_smem = Tx.alloc_buffer(shape, dtype, scope="shared", layout=s_layout) Tx.tile.warp.copy(A_smem[fs], A[fs]) # global -> shared (this dispatch) Tx.cuda.cta_sync() Tx.tile.warp.copy(B[fs], A_smem[fs]) # shared -> global (this dispatch)要点解读:
Tx.lane_id([32])声明 32 条 lane,Tx.thread_id([32])声明 32 个 thread,两者结合定义了warp作用域,sctx.intra由此得出thread_cnt = 32;A_smem以scope="shared"分配,并显式给定TileLayout(S[shape])布局;- 两次
Tx.tile.warp.copy分别触发 global→shared 与 shared→global 两个方向的 gmem_smem 路径,中间以Tx.cuda.cta_sync()保证同步。
测试文件还覆盖了warpgroup(128 线程)、cta(256 线程)等作用域以及float16/float32/uint8多种 dtype 的往返用例(TASKS 表),并有针对 swizzle 布局、非对齐 stride、非对齐 region 偏移的算法级回归测试(当前标注 XFAIL 的部分对应已知的align_layouts_gs待修复项)。
算法核心:三步合成向量化拷贝
gmem_smem 路径的 emit 逻辑(vec_auto_gmem_smem.py#L96-L186)由三个步骤构成。
1. 合成三维分区[outer, threads, vec]。以 32 线程、32×32 = 1024元素为例:dispatch 通过align_layouts_gs把两侧布局切片到目标 region,让 global 侧驱动规范(stride 降序)顺序,切出连续的vec尾部和threads块,再把 shared 侧按同样的方式重组以匹配。最终每个线程负责连续的一段融合索引槽位。
2. 由宽到窄选择向量宽度。依次尝试{128, 64, 32, 16, 8}位对应的元素个数,接受满足以下条件的最宽值:(a) 连续尾部能整除该宽度;(b) 两侧所有非 vec 迭代 stride(含线程迭代)以及两个基础偏移量都是它的倍数——这样每线程、每轮的向量指针天然对齐(只有最内层vec迭代被排除在检查之外)。对float32而言vec = 4(4 × 4 B = 16 B = 128 bit),于是outer = 1024 / (32 × 4) = 8。
3. 发出串行循环。emit 刻意使用普通range循环而非Tx.unroll,把最终的展开决策留给 ptxas:
for f in range(total_outer): g_lin = g_p.apply(f, tid, v0, shape=apply_shape)["m"] s_off = s_apply_layout.apply(f, tid, v0, shape=apply_shape)["m"] s_ptr = _ptr_off(s_buf.ptr_to(s_zero), s_off) g_ptr = _ptr_off(g_buf.ptr_to(g_zero), g_lin) if g_is_src: Tx.ptxld_g], g_ptr) Tx.ptxst_s]) else: Tx.ptxld_s], s_ptr) Tx.ptxst_g])每个(f, tid, 0)坐标都由layout.apply以[outer, threads, vec]为 shape 扁平化,因此 emit 代码完全不需要知道分区是如何切分迭代的。ld_g/st_s/ld_s/st_g是按内存方向与向量宽度注册的 direct-PTX 形式——128 位传输使用四个uint32寄存器与v4.u32形式,链名(如ld.global.v4)在 traced body 内以 Python 字符串直接构造(这是 parser 无法携带跨代码块字符串的技术细节,见源码注释)。
生成的 TIRx IR:向量化循环的中间形态
对上述演示程序运行LowerTIRx之后,每个Tx.tile.warp.copy都会被替换为合成后的循环(以 global→shared 方向为例,已精简):
tid: Tx.let = threadIdx_x % 32 A_smem = Tx.alloc_shared((1024,)) tmp = Tx.alloc_local((4,), "uint32") for f in range(8): # outer = 8 s_lin = f * 128 + tid * 4 # 32 threads × vec 4 = 128 / round g_lin = f * 128 + tid * 4 s_ptr = pointer_offset(A_smem, s_lin) g_ptr = pointer_offset(A_1, g_lin) # A_1 = A.view(1024) Tx.ptx.ld.global_.v4.u32(tmp[0], tmp[1], tmp[2], tmp[3], g_ptr) Tx.ptx.st.shared.v4.u32(s_ptr, tmp[0], tmp[1], tmp[2], tmp[3])注意几个实现细节:
tmp是一个(4,)的uint32本地临时 buffer,用于在 load 与 store 之间中转位模式——scratch 只搬运比特,因此按 PTX 容器类型而非元素类型分配(源码注释明确说明这一点);- 两侧地址都通过
pointer_offset计算,A_1 = A.view(1024)表明 global buffer 被扁平化为一维视图后做线性偏移; - 每轮每线程搬运
vec = 4个元素,32 线程一轮共 128 个元素,8 轮恰好覆盖 1024 个元素。
生成的 PTX 指令:每轮一条向量 load + 一条向量 store
CUDA 代码生成器为每一轮发出一对向量指令:
ld.global.v4.u32 {r0, r1, r2, r3}, [g_ptr]; st.shared.v4.u32 [s_ptr], {r0, r1, r2, r3};shared→global 方向则对应ld.shared.v4.u32后接st.global.v4.u32。线程tid每轮处理元素[f·128 + tid·4 .. +4),8 轮 × 32 lane 覆盖全部 1024 个元素,且每个元素恰好以一次 128 位传输完成——这正是向量化拷贝追求的最小指令数与最大带宽利用率。
输入如何改变算法:dtype、scope 与 swizzle
dtype 决定向量宽度与轮数
元素dtype决定向量宽度(取能保持对齐的最宽 128 位传输),进而决定轮数。对同样的32×32tile 与 32 线程:
| dtype | vec | 传输宽度 | outer = 1024 / (32 · vec) |
|---|---|---|---|
float32 | 4 | 16 B(v4.u32) | 8 |
float16 | 8 | 16 B(v4.u32) | 4 |
uint8 | 16 | 16 B(v4.u32) | 2 |
可以看到无论 dtype 如何,只要对齐条件满足,最终都收敛到 128 位传输(v4.u32);差别在于单次向量化覆盖的元素个数与需要的轮数。dtype 位宽越小,单轮搬运元素越多、轮数越少。测试文件 test_gmem_smem.py 还覆盖了int8、float8_e4m3fn、float8_e5m2、bfloat16等 dtype,佐证了这一规律在不同数据宽度下的普适性。
scope 决定线程轴与线程数
执行作用域决定线程 id 的轴名称(warp→laneid,cta→tx,warpgroup→ 对应的 warpgroup 内线程轴等)与线程总数,因而决定分区形态。源码中通过_TID_AXIS_FOR_SCOPE映射作用域到轴名,_thread_cnt(sctx)从sctx.intra推导线程数。当thread_cnt == 1(如thread作用域)时,tid声明退化为常量0,循环退化为单线程的多轮向量搬运。
swizzled shared 布局:向量宽度被 chunk 上限约束
若 shared 侧使用swizzled布局(ComposeLayout),vec被限制为不超过一个 swizzle chunk 的大小,且s_off的计算需经过 swizzle:识别出的 swizzle 每轮只需几条寄存器加法,否则每轮调用swizzle.apply。测试 test_swizzled_smem_vec_len_must_fit_chunk 明确指出:ComposeLayout的底部per_element位不参与 swizzle,vec必须留在该 chunk 内,否则会跨越 XOR 边界读写到错误的物理字节;而test_gmem_smem_swizzle_uses_structured_compose_apply验证了 swizzle 路径生成的是结构化地址形式(P/XOR-low/ADD-high,即^异或加* 256原子对齐加法),且要求每轮 offset 不含完整的除法/取模分解。
对齐约束的兜底:非对齐输入会收窄 vec_len
test_unaligned_strides_must_clamp_vec_len与test_unaligned_region_offset_must_clamp_vec_len两个回归测试揭示了align_layouts_gs的对齐契约:当 global 布局行 stride 非vec_len倍数(例如 fp16 行 stride 20,导致tid=2的基础偏移 40 字节不满足 16 字节对齐),或 region 起点列号非向量对齐(如 fp16 中从第 3 列切片),vec_len必须相应收窄(极端情况退化为 1,即标量),否则 128 位uint4reinterpret 会产生非法内存访问。从源码结构看,这些非对齐场景正是 gmem_smem 路径保证安全性的关键边界。
总结:gmem_smem 路径的适用边界与设计取舍
vec_auto变体的gmem_smem实现路径覆盖了CUDA 上、两侧均非寄存器的 global↔shared 同步拷贝,其设计取舍清晰:
- 合成而非继承分区:两侧都是跨线程存储,分区完全由执行作用域推导,最终统一表达为
[outer, threads, vec]三维坐标,emit 与分区解耦; - 宽优先的向量化:以对齐为硬约束,从 128 位向下搜索最宽向量,保证指令数最小且所有非 vec 迭代保持自然对齐;
- 延迟展开:串行
range循环把最终展开决定交给 ptxas,避免过早展开导致 kernel 膨胀(源码注释明确「keep a serial loop, T.unroll floods the kernel」); - 严格的准入门控:非 CUDA 目标、非全线程激活、存储对不符、dtype/extent 不匹配、元素数不可被线程数整除——任一条件不满足即拒绝并让位给
reg路径或fallback(优先级 0 的标量兜底)。
如需进一步深入,可继续阅读 reg 路径文档(寄存器参与的vec_auto路径,分区由寄存器布局的线程轴诱导)、ldstmatrix 文档(warp 集体矩阵搬运)以及 fallback 文档(标量兜底);实现细节可在 vec_auto_gmem_smem.py 与 _common.py 中追踪,行为验证与回归用例集中在 test_gmem_smem.py。
- 模型编译
- 深度学习
- 推理引擎
【免费下载链接】tvm
Open Machine Learning Compiler Framework
相关推荐
TVM TIRx Tile 原语详解:同步 copy 在 CUDA 上的五种降级路径与向量化实现
TVM TIRx Tile 原语详解:同步 copy 在 CUDA 上的五种降级路径与向量化实现 导读 :本文聚焦 Apache TVM 的 TIRx(TIR
模型编译深度学习推理引擎DORA tensor-pool 内存池传输示例实战:从 CPU↔CPU 到跨机 CUDA 的零拷贝张量搬运
DORA tensor pool 内存池传输示例实战:从 CPU↔CPU 到跨机 CUDA 的零拷贝张量搬运 导读 本文以 dora 仓库中 libraries
机器人人工智能ROS消息路由如何用 create-next-app 创建 Next.js 项目并跑通本地开发服务器
如何用 create next app 创建 Next.js 项目并跑通本地开发服务器 这篇文章解决一个具体的任务:在本地从零创建一个新的 Next.js 项目
模型编译深度学习推理引擎
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考