资讯动态

TVM TIRx CUDA 向量化内存拷贝实战:vec_auto 变体的 global↔shared 实现路径(gmem_smem)深度解析

发布时间:2026/9/23 2:58:02 来源:尧图企业网站定制
模型编译深度学习推理引擎【免费下载链接】tvmOpen Machine Learning Compiler Framework项目地址https://gitcode.com/gh_mirrors/tv/tvm点击查看免费下载本篇技术指南聚焦 TVM TIRxTensorIR eXperimentalCUDA 后端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 与 registerlocal三种存储之间搬运数据。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_256bglobal ↔ shared/local或 shared ↔ local20显式线程作用域下搬运恰好指定宽度的数据可带 global-load cache 控制vec_autogmem_smem 路径global ↔ shared10合成[outer, threads, vec]分区配合直接 PTX 向量 load/storevec_autoreg 路径register ↔ shared/global10由寄存器布局的线程轴诱导分区ldstmatrixregister ↔ shared10warp 集体ldmatrix/stmatrixm8n8 片段fallbackglobal / shared / local0标量单线程拷贝兜底其中vec_auto变体内部包含两条实现路径gmem_smemglobal↔shared本文主题与reg寄存器参与。需要特别强调的是gmem_smem不是可选择的独立 dispatch 名称而是vec_auto变体在自动选择或显式dispatchvec_auto时根据操作数存储类型路由到的内部实现路径。其 dispatch 注册逻辑位于 vec_auto.py注册优先级为 10predicate 依次尝试_is_gmem_smem与_is_reg_copyregister_dispatch( copy, cuda, variantvec_auto, priority10, 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(fgmem_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-L93def _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, funsupported 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门控条件可归纳为下表属性要求targetcudascope执行作用域thread/warp/warpgroup/cta且所有线程处于激活状态_all_threads_active——laneid覆盖 32 个线程等未被外围if收窄存储对(global, shared*)或(shared*, global)——即_GMEM_SMEM_PAIRS任一侧都不能是localdtype / 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 的典型用例一个 warp32 线程把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, scopeshared, layouts_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 条 laneTx.thread_id([32])声明 32 个 thread两者结合定义了warp作用域sctx.intra由此得出thread_cnt 32A_smem以scopeshared分配并显式给定TileLayout(S[shape])布局两次Tx.tile.warp.copy分别触发 global→shared 与 shared→global 两个方向的 gmem_smem 路径中间以Tx.cuda.cta_sync()保证同步。测试文件还覆盖了warpgroup128 线程、cta256 线程等作用域以及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 44 × 4 B 16 B 128 bit于是outer 1024 / (32 × 4) 8。3. 发出串行循环。emit 刻意使用普通range循环而非Tx.unroll把最终的展开决策留给 ptxasfor f in range(total_outer): g_lin g_p.apply(f, tid, v0, shapeapply_shape)[m] s_off s_apply_layout.apply(f, tid, v0, shapeapply_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 一条向量 storeCUDA 代码生成器为每一轮发出一对向量指令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 与 swizzledtype 决定向量宽度与轮数元素dtype决定向量宽度取能保持对齐的最宽 128 位传输进而决定轮数。对同样的32×32tile 与 32 线程dtypevec传输宽度outer 1024 / (32 · vec)float32416 Bv4.u328float16816 Bv4.u324uint81616 Bv4.u322可以看到无论 dtype 如何只要对齐条件满足最终都收敛到 128 位传输v4.u32差别在于单次向量化覆盖的元素个数与需要的轮数。dtype 位宽越小单轮搬运元素越多、轮数越少。测试文件 test_gmem_smem.py 还覆盖了int8、float8_e4m3fn、float8_e5m2、bfloat16等 dtype佐证了这一规律在不同数据宽度下的普适性。scope 决定线程轴与线程数执行作用域决定线程 id 的轴名称warp→laneidcta→txwarpgroup→ 对应的 warpgroup 内线程轴等与线程总数因而决定分区形态。源码中通过_TID_AXIS_FOR_SCOPE映射作用域到轴名_thread_cnt(sctx)从sctx.intra推导线程数。当thread_cnt 1如thread作用域时tid声明退化为常量0循环退化为单线程的多轮向量搬运。swizzled shared 布局向量宽度被 chunk 上限约束若 shared 侧使用swizzled布局ComposeLayoutvec被限制为不超过一个 swizzle chunk 的大小且s_off的计算需经过 swizzle识别出的 swizzle 每轮只需几条寄存器加法否则每轮调用swizzle.apply。测试 test_swizzled_smem_vec_len_must_fit_chunk 明确指出ComposeLayout的底部per_element位不参与 swizzlevec必须留在该 chunk 内否则会跨越 XOR 边界读写到错误的物理字节而test_gmem_smem_swizzle_uses_structured_compose_apply验证了 swizzle 路径生成的是结构化地址形式P/XOR-low/ADD-high即^异或加* 256原子对齐加法且要求每轮 offset 不含完整的除法/取模分解。对齐约束的兜底非对齐输入会收窄 vec_lentest_unaligned_strides_must_clamp_vec_len与test_unaligned_region_offset_must_clamp_vec_len两个回归测试揭示了align_layouts_gs的对齐契约当 global 布局行 stride 非vec_len倍数例如 fp16 行 stride 20导致tid2的基础偏移 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。赞分享模型编译深度学习推理引擎【免费下载链接】tvmOpen Machine Learning Compiler Framework项目地址https://gitcode.com/gh_mirrors/tv/tvm点击查看免费下载相关推荐TVM TIRx Tile 原语详解同步 copy 在 CUDA 上的五种降级路径与向量化实现TVM TIRx Tile 原语详解同步 copy 在 CUDA 上的五种降级路径与向量化实现 导读 本文聚焦 Apache TVM 的 TIRxTIR模型编译深度学习推理引擎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),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价