ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

TVM Blackwell `tcgen05.cp` 全解析:shared→tmem 异步拷贝的形状选择、矩阵描述符与调度算法

2026/9/23 19:12:01 拓冰建站 浏览量
TVM Blackwell `tcgen05.cp` 全解析:shared→tmem 异步拷贝的形状选择、矩阵描述符与调度算法 TVM Blackwelltcgen05.cp全解析shared→tmem 异步拷贝的形状选择、矩阵描述符与调度算法【免费下载链接】tvmOpen Machine Learning Compiler Framework项目地址: https://gitcode.com/gh_mirrors/tv/tvmTVM 的 CUDA 后端在 Blackwellsm_100a上为copy_async注册了smem-tmem变体通过tcgen05_cp规划器把共享内存到张量内存tensor memorytmem的异步拷贝降级为覆盖所有合法形状的tcgen05.cpPTX 指令。本篇以 docs/tirx/tile_primitives/copy_async/tcgen05_cp.rst 为骨架结合 tcgen05_cp.py 的实现与 test_tcgen05_cp.py 的端到端验证讲解形状选择规则、shape/multicast配置、矩阵描述符ldo/sdo/swizzle的推导、单线程发出与 mbarrier 完成协议以及读者如何在自家 kernel 中复现这套调用模式。copy_async在 TVM 中的定位一个算子、六个变体在深入tcgen05.cp之前先看它所在的家族。docs/tirx/tile_primitives/copy_async.rst 说明copy_async是一个异步拷贝算子每个变体只发出传输的 issue 指令完成协议由调用方提供。CUDA 后端按源/目的存储对与作用域选择变体目前注册了六个变体名变体存储对优先级issue 指令ldgstsglobal → shared20cp.asyncLDGSTS每线程向量化tma_auto/tma_explicitglobal ↔ shared10cp.async.bulk.tensorTMA描述符驱动单线程dsmemshared → shared跨 CTA10cp.async.bulkshared::clustermapa远端地址smem-tmemshared → tmem10已注册的tcgen05.cp形状矩阵描述符驱动tmem-localtmem ↔ 寄存器10tcgen05.ld/tcgen05.stwarpgroup原子形状匹配从 python/tvm/backend/cuda/tile_primitive/copy_async/init.py 可以看到copy_async的 dispatch 由dsmem、ldgsts、tcgen05_cp、tcgen05_ldst、tma五个子模块共同注册tcgen05_cp对应 shared → tmem 这一路。本篇文章的主角——smem-tmem变体——通过tcgen05_cp规划器把copy_async从共享内存复制到张量内存。其关键设计决策有三形状全量覆盖规划器不局限于单一形状而是覆盖所有tcgen05.cp形状128x256b、4x256b、128x128b、64x128b.warpx2::02_13、64x128b.warpx2::01_23、32x128b.warpx4。描述符派生一个共享内存矩阵描述符matrix descriptor命名源 tile其全部字段ldo/sdo/swizzle与 cp issue 序列都由两个 buffer 的布局推导得出。纯异步发出dispatch 只发出拷贝调用方用tcgen05.commit对 mbarrier 信号完成源码模块 docstring 明确指出 tcgen05.cp is inherently async; this dispatch emits the cp loop only and leaves completion signaling to the caller。形状选择显式配置与布局推断tcgen05.cp的每个形状都是「行数 × 每行位数」的形式对应 PTX ISA 8.8 §9.7.16.9.2部分形状强制要求multicast 限定符。Multicast 会把复制的数据在目的 tmem 的 warp lane slab 上复制多份。shape显式配置通过shape配置可以强制某个 PTX 形状当该形状有多个合法限定符时即64x128b的两个warpx2变体还必须给出multicast。这一逻辑在 tcgen05_cp.py 的_CP_SHAPE_MULTICASTS第 126-132 行和_resolve_cp_shape第 196-220 行中实现_CP_SHAPE_MULTICASTS { 128x256b: (,), 4x256b: (,), 128x128b: (,), 64x128b: (warpx2::02_13, warpx2::01_23), 32x128b: (warpx4,), }要点128x256b、4x256b、128x128b不需要 multicast合法列表只有空串。32x128b强制warpx4。64x128b有两个合法 multicast 限定符必须显式二选一若省略_resolve_cp_shape会抛出requires an explicit multicast config的ValueError。未知形状如96x128b会被拒绝错误信息列出全部合法形状。无shape时的布局推断没有shape配置时规划器按候选列表从最宽的 atom 开始逐一尝试取第一个通过布局校验的计划顺序为128x256b → 4x256b → 128x128b → 64x128b.warpx2::02_13 → 64x128b.warpx2::01_23 → 32x128b.warpx4这对应源码中的_CP_SHAPE_CANDIDATES第 143-150 行。推断逻辑的关键洞察除了 256b/128b 这一对即128x256b与128x128b外其余候选按 tmem 的 (lane, replica) 模式互斥——一个布局只能匹配其中一种模式的形状所以一个裸的warpx4拷贝最终仍解析为32x128b.warpx4与历史 dispatch 完全一致128x128b兼容的源布局同时满足更宽的128x256batom指令数减半推断会优先选择更宽的 atom——test_tcgen05_cp.py 中test_cp_128x128b_layout_infers_wider_256b_atom专门断言推断出的指令是128x256b而非128x128b。形状 → 行/复制模式映射表每个形状钉死一个 tmem 行→lane 映射和 replicamulticast模式。原文档给出的表格在 B200 上由 test_tcgen05_cp.py 逐位验证shapemulticastt lane patternt replica128x256b无(128, 1TLane)—4x256b无(4, 32TLane)—128x128b无(128, 1TLane)—64x128bwarpx2::02_13(64, 1TLane)(2, 64TLane)64x128bwarpx2::01_23(2,32):(64,1)TLane(2, 32TLane)32x128bwarpx4(32, 1TLane)(4, 32TLane)表中模式由_cp_lane_replica_patterntcgen05_cp.py 第 153-193 行实现其注释给出了 PTX ISA 8.8 p675 的权威语义warpx2「数据被 multicast 到 warp pairpair 中的每个 warp 收到一半数据」——pair 内两个 warp 互为镜像持有相同一半pair 之间按列出的顺序瓜分 64 行warpx2::02_13pair {0,2}、{1,3}行 0-31 到 warp 0/2行 32-63 到 warp 1/3。一次拷贝占据 lane 0-63按行序在 lane 偏移 64 处复制warp2warp3→ lane 模式(64, 1TLane)、replica(2, 64TLane)。这恰好就是 Layout E / Layout B 的 datapath 组织PTX ISA 8.8 Figure 213/207M 0-31 在 warp-rank 0 和 2M 32-63 在 warp-rank 1 和 3。warpx2::01_23pair {0,1}、{2,3}行 0-31 到 warp 0/1行 32-63 到 warp 2/3 → 一次拷贝在 lane 0-31行 0-31与 lane 64-95行 32-63在 32 处复制 → lane 模式(2, 32):(64, 1)TLane、replica(2, 32TLane)。warpx4行 0-31 映射到 lane 0-31replica 偏移 0/32/64/96即(4, 32TLane)。4x256b每个 warp 象限 datapath 写一行——行 0..3 落在 lane 0/32/64/96 → lane 模式(4, 32TLane)无 replica。它接受什么两个谓词 完整约束表该变体通过register_dispatch注册tcgen05_cp.py 第 798-809 行priority10variantsmem-tmem携带两个谓词# register_dispatch(..., variantsmem-tmem, priority10, when[ predicate(validate_smem_tmem_copy, _validate_smem_tmem_copy), predicate(exec_scope, _single_thread_exec), # exec_scope thread # ])内存作用域信封_validate_smem_tmem_copy第 701-716 行只做快速检查源shared*、目的tmem、双方都带 layout、位宽一致允许等宽重解释如 nvfp4 把 sf 存成 uint8 按 fp8 读、目的 buffer 设置了allocated_addr由先前的tcgen05.alloc提供。所有形状/布局的详细校验都在规划器里以可读错误抛出。单线程执行作用域_single_thread_exec在 python/tvm/backend/cuda/tile_primitive/copy/utils.py第 27-31 行要求exec_scope thread——即单线程发出拷贝区别于tmem-local变体的 warpgroup 集体操作。原文档给出的完整约束表属性要求target / prioritycudatarget 且支持tcgen05以sm_100a测试priority10scope单线程发出拷贝CTA groupcta_group为1默认或2直接转发给 PTX 指令memory pair源shared*→ 目的tmemallocated_addr由先前的tcgen05.alloc设置两个 buffer 都携带 layout元素位宽一致允许等宽重解释tmem layout必须能切片为上表中某个形状的 (lane, replica) 模式smem layout行按 8 行一组组成描述符 core-matrix groupatom 行宽推导 swizzle 模式K-byte ∈ {16, 32, 64, 128} → sw 0..3且必须与 buffer 自身 swizzle若有一致注意 smem 布局这行还有一个隐藏前提描述符模板以 base_offset0 编码因此 smem buffer 基址必须对齐到 swizzle 周期8 * atom_K字节——copy_smem_tmem_impl中的注释说明align1024即可满足这一要求这也是所有测试里alloc_buffer都带align1024的原因。演示程序一次完整的 smem→tmem 异步拷贝原文档给出了取自 test_tcgen05_cp.py 的演示程序一个 warpgroup 分配 16 个 tmem 列填充一个32×16的uint8共享 tile再拷入 tmem——不指定shape规划器从布局推断出32x128b.warpx4读回与释放的尾部省略from tvm.tirx.layout import R, S, TCol, TileLayout, TLane A_smem Tx.alloc_buffer([32, 16], uint8, scopeshared, layoutTileLayout(S[(32, 16) : (16, 1)]), align1024) tmem_addr Tx.alloc_shared([1], uint32) cp_mbar Tx.alloc_shared([1], uint64) if warp_id 0: Tx.ptxtcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32, Tx.uint32(16)) # ... mbarrier.init, fence, cta_sync, fill A_smem from global ... tmem Tx.decl_buffer([32, 16], uint8, scopetmem, allocated_addrtmem_addr[0], layoutTileLayout(S[(32, 16) : (1 TLane, 1 TCol)] R[4 : 32 TLane])) if tid_in_wg 0: Tx.tile.copy_async(tmem[0:32, 0:16], A_smem[0:32, 0:16], cta_group1) # smem - tmem # caller signals Tx.ptx.tcgen05.commit.cta_group__1.mbarrier__arrive__one.shared__cluster.b64( cp_mbar.ptr_to([0])) Tx.cuda.mbarrier_wait(cp_mbar.ptr_to([0]), 0) # ... readback via tcgen05.ld, then tcgen05.dealloc ...对照源码可以拆出这套调用模式的五个组成部分tmem 分配warp 0 发出tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32(tmem_addr, 16)把 tmem 句柄写到共享变量tmem_addr。目的 tmem bufferdecl_buffer声明scopetmem的视图allocated_addrtmem_addr[0]挂上句柄其布局S[(32,16):(1TLane,1TCol)] R[4:32TLane]正是32x128b.warpx4的 (lane, replica) 模式——R[4:32TLane]声明了 warpx4 的 4 个 lane slab 复制。单线程发出tid_in_wg 0下调用Tx.tile.copy_async(tmem[...], A_smem[...], cta_group1)——这与exec_scope thread谓词吻合。调用方完成协议dispatch 不发出tcgen05.commit调用方在同一线程里对cp_mbar执行tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64。等待与清理所有线程mbarrier_wait随后tcgen05.ld读回、tcgen05.dealloc释放演示中省略。调度算法四步从布局到指令_build_plantcgen05_cp.py 第 315-663 行实现了文档描述的 A-I 步骤可归纳为四步第 1 步解析形状显式shape必要时multicast或按布局推断见上文。解析结果用_shape_dims从形状名rowsxbitsb提取(atom_rows, atom_bits)第 135-138 行。第 2 步校验计划对两个布局在给定区域做 slice canonicalize校验t.replica与形状表条目一致非 multicast 形状要求 replica 为空计算把 TLane 放首位、TCol 按 stride 降序的置换分别作用于 t 与 ss 用grouppermute_by_groups隔离 broadcast按 stride-0 切分两侧要求切分序列一致丢弃 stride-0 的迭代子三段式切分把展平的元素空间按atom_rows × n_mid × elem_per_atom分组elem_per_atom atom_bits / dtype_bits*_lane恰好覆盖一条指令的行数乘积 atom_rows*_col一条指令自身的列数*_middle两者之间的一切——外层列决定哪条指令以及当区域行数超过一个 atom 时例如4x256b为每个 warp 前 16 行做 tile 化——即 M64 Layout-F 的散布多出的行迭代子。逐项校验t_lane行→TMEM lane匹配形状表 lane 模式t_col一行的元素落在连续的 TMEM 列s_col一行 smem 中连续的 16B 单元对 256b atom未 swizzle 时两个单元可处于非连续 stride→ LDO 字段s_lane行在 smem 中按(atom_rows/8, 8)分组排列stride 为(SDO_stride, atom_K_stride)4x256b是单个 4 行组SDO0atom_K_byte ∈ {16, 32, 64, 128}→ swizzle 模式 0..3必须与s_buf.layout的 swizzle若有一致。swizzle 校验还包含两道「方向/族」检查描述符的地址置换假定标准 mma atom 族8 行 128-bit atomper_element/atom_len不符即拒绝swizzle family mismatch且硬件遍历实现的是swizzle_innerTrue的置换x ^ ((x outer_mask) atom_len)若布局是swizzle_innerFalse镜像置换会静默拷错必须显式拒绝。对齐检查第 G 步t 的 TCol 偏移必须 128-bit 对齐t 的 TLane 偏移要能折进 taddr 的 lane 半字s 的 m 偏移在 sw0 时 16B 对齐、sw0 时按 atom 大小对齐middle 迭代子 stride 16B 对齐。最后是 middle 的一一对应simple-modet_middle 与 s_middle 迭代子数量相同、逐位置 extent 匹配。测试test_align_middle_2_to_1_nvfp4_sfb覆盖了双方 middle 规范化后结构不对称TMEM 侧合并为单迭代子、SMEM 侧因步长顺序无法合并时_align_middles的 union-cut 算法。第 3 步一次性编码矩阵描述符64-bit 共享描述符在 smem buffer 基址处、紧接其分配之后编码一次按(smem_buf, ldo, sdo, swizzle)缓存_get_or_create_desc第 671-685 行提升hoist到 smem 分配之后add_post_buffer_def_stmt。每次 cp 只修补 14-bit 地址字段def _desc_set_addr(desc_val, addr_ptr): Patch a SMEM matrix descriptors 14-bit address field with cvta(addr)4 — matches the hand-rolled replace_smem_desc_addr (descriptor encoded at 0). start_addr T.cast( T.bitwise_and( T.shift_right(T.cuda.cvta_generic_to_shared(addr_ptr), T.uint32(4)), T.uint32(0x3FFF), ), uint64, ) return T.bitwise_or(T.bitwise_and(desc_val, T.bitwise_not(T.uint64(0x3FFF))), start_addr)即cvta(addr) 4 0x3FFF取 14-bit 共享地址与模板做按位「清地址位 或入新地址」。编译级测试test_multi_cp_encodes_descriptor_once_and_patches_addr与test_cp_shape_config_routes_to_generic_planner断言了「一个encode_matrix_descriptor、多次cp_desc_ptr[0] 修补」的生成形态前者验证 4-tile 拷贝只有一次编码 4 次地址修补后者验证128x256b的 4 次 cp 各带0x3FFF掩码修补。第 4 步发出拷贝middle 维展平成一个 unrolled 循环每一步同时推进描述符的 16 字节共享偏移和 tmem 地址列位外加 lane-tiled atom 的 lane 半字for flat in Tx.unroll(total): t_off, s_off Tx.meta_var(compute_offsets(flat)) Tx.ptxftcgen05.cp.cta_group::{cta_group}.{shape}{multicast_seg}, smem_desc_add_16B_offset(desc_buf[0], init_off_16B s_off))循环计数total由所有 middle 迭代子的 extent 乘积得到total 1时特判为单条指令避免退化的T.unroll(1)。compute_offsets对每个(n, s_step, t_step)计算(flat // div) % n并累加进 t/s 偏移。指令串按tcgen05.cp.cta_group::{cta_group}.{shape}拼装multicast 写在 shape 之后如tcgen05.cp.cta_group::1.32x128b.warpx4。关键语义dispatch 不发出tcgen05.commit/wait——调用方需像演示程序那样对 mbarrier 提交tcgen05.commit。模块 docstring 明确提示需要同步语义的调用方应在拷贝后自行发出tcgen05.commit。生成的 CUDA在sm_100a上编译后一次 warpx4 拷贝生成的 CUDA 内联 PTX 为原文档示例// one warpx4 copy: shared (named by the matrix descriptor) - tensor memory tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1; // [%0]tmem addr, %1descriptor[%0]是 tmem 地址t_addr[0] t_addr_off t_off的 uint32 形式%1是共享描述符携带被修补的 14-bit 地址。编译级测试test_cp_4x256b_compile_emits_shape_and_count进一步验证指令形状与数量4x256b、2 个 middle 迭代 → 恰好 2 条tcgen05.cp.cta_group::1.4x256b指令。端到端正确性含 tmem 读回由 test_tcgen05_cp.py 的 GPU round-trip 测试覆盖。输入如何改变算法原文档给出的「输入 → 效果」表输入效果tmem layout推断时通过 (lane, replica) 模式选择形状R[4:32TLane]的 replica 解析为32x128b.warpx4dtype设定elem_per_atom atom_bits / dtype_bits并结合 smem 行 stride 决定描述符 swizzle 模式region sizemiddle 维一个 atom → 单条tcgen05.cp更宽区域 → 对列做 unrolled 循环行数超过一个 atom → lane-tiled cp步进 tmem lane 半字shared swizzle layout改变编码的swizzle模式必须匹配推导出的 atom K-bytecta_group转发给指令cta_group::2是 pair-collective——偶数 CTA 单次发出使每个 CTA 从自己的 smem 拷入自己的 tmem由 cta_group2 的 round-trip 测试在 B200 上钉死对照源码补充细节dtypeelem_per_128b 128 // dtype_bits、elem_per_32b 32 // dtype_bits全程参与三段切分与地址计算等宽重解释如 uint8 存 nvfp4是合法的但真正的位宽改变需要解压缩decompress规划器不支持——copy_smem_tmem_impl开头对decompress配置直接抛出tcgen05.cp planner does not support decompress测试test_cp_rejects_decompress_on_generic_path锁定该行为。region size行数超过一个 atom 时如4x256b填充每个 warp 前 16 行t 侧的 middle 迭代子走 TLane 轴t_step stride 16折进 taddr 的 lane 半字第 630-633 行test_cp_4x256b_lane_tiled_layout_f_scatter验证 16 条4x256bcp 各落 16 个不同行M64 Layout-F 散布test_cp_4x256b_lane_tiled_partial_window则验证区域行偏移折进 lane 半字后精确落到{32q 4..11}。cta_group2test_cp_cta_group_2_128x256b_pair_collective生产 small_topk 的 q-cp 形态与test_cp_cta_group_2_64x128b_warpx2_pair_collective生产 head128 的 q-cp 形态验证偶数 CTA 发出cta_group2的 cp 与commit cta_mask3后两个 CTA 各自从自己的 smem 拷入自己的 tmem描述符 smem 地址按 CTA rank 解析奇数 CTA 保持不被写入。负向校验可读的错误信息规划器把形状/布局校验集中在_build_plan用可读的ValueError拒绝非法输入。测试文件里有一组专门的负向测试正好是排查问题的清单测试触发错误test_cp_rejects_wrong_replica_for_multicast02_13 声明的 TMEM 布局配 01_23反之亦然→replica mismatchtest_cp_rejects_illegal_shape_multicast_combo128x256b配warpx4→illegal multicasttest_cp_rejects_missing_multicast_for_64x128b64x128b省略 multicast →requires an explicit multicasttest_cp_rejects_non_128b_tcol_offsetTCol 偏移不足 128-bit 对齐 →not provably 128b-alignedtest_cp_rejects_unknown_shape非法形状串 →unknown tcgen05.cp shapetest_cp_rejects_nonzero_tmem_lane_offsetlane 偏移超出 128 lane 空间 →overflows the 128-lane spacetest_cp_rejects_decompress_on_generic_pathdecompress配置 →does not support decompresstest_cp_rejects_non_16b_aligned_row_group_stride8 行组 stride 非 16B 倍数SBO 以 16B 为单位编码→not 16B-alignedtest_cp_rejects_non_canonical_swizzle_familyswizzle 布局不在标准 mma atom 族 →swizzle family mismatchtest_cp_rejects_flipped_swizzle_innerswizzle_innerFalse镜像置换→ 拒绝防止静默拷错test_dispatch_rejects_bad_inputs若干32x128b无法读取的 sub-tile 配置 → 编译期ValueError测试矩阵B200 上的逐位验证test_tcgen05_cp.py 的核心测试方法是 GPU round-trip对每个(shape, multicast, swizzle, dtype, tile size)组合——填充主机 bufferA→ 以 MMA 风格可选 swizzle布局存入 smem → 发出通用Tx.copy_async(tmem_region, smem_region, shape..., multicast...)不带任何desc_*字段迫使通用规划器推导描述符与 cp 循环→ 每个 warp 用tcgen05.ld.32x32b读回自己的 32-lane slab → 在主机端纯由目的 tmem 布局重建期望值含 multicast 复制偏移逐位比对。覆盖维度test_cp_shape_roundtrip_swizzled全部 6 种 (shape, multicast) × swizzle {1,2,3} × dtype {bfloat16, float32} × middle {1, 4}单 cp / 多 cptest_cp_shape_roundtrip_nonswizzledsw0 的源128b 行 单个 16B 单元256b 行携带非平凡 LDOtest_cp_shape_roundtrip_offsets非零 smem 行偏移 / tmem 列偏移的 sub-region 拷贝test_cp_shape_inferred_from_layouts_matches_explicit裸copy_async的布局推断与显式配置逐字节同构test_cp_cta_group_2_*cta_group2的 pair-collective 形态B200 钉死。运行前提cuda compute 10.0所有 GPU 测试都有pytest.mark.skipif(not env.has_cuda_compute(10), ...)即 Blackwell 硬件。与其他异步变体的协作关系smem-tmem变体是 TVMcopy_async异步家族中面向 Blackwell tmem 的两条路径之一tcgen05_cpshared → tmem单线程、描述符驱动、只发 issue调用方tcgen05.committcgen05_ldsttmem ↔ local见 docs/tirx/tile_primitives/copy_async/tcgen05_ldst.rstwarpgroup 集体、原子形状匹配.16x64b/.16x128b/.16x256b/.32x32b方向按存储对推断tmem→local降为tcgen05.ldlocal→tmem降为tcgen05.st完成协议用tcgen05.wait.ld/st。两者都不发出完成等待——「谁 issue、谁 commit/wait」是 TVM 异步原语的一致设计哲学。配合 TMAglobal ↔ shared与 dsmemshared → sharedcopy_async覆盖了 Blackwell 上数据搬运的完整链路。小结一条可复用的降级范式tcgen05.cp规划器展示的是一条通用降级范式可直接迁移到其他张量内存指令信封谓词只做粗筛scope/位宽/allocated_addr把详细校验留给规划器错误信息可读形状由布局推断而非硬编码候选按「最宽 atom 优先」排序天然得到最少指令数描述符模板化ldo/sdo/swizzle 从布局推导64-bit 描述符编码一次、每次 cp 只补 14-bit 地址生成代码极小middle 维展平为 unrolled 循环t/s 偏移用meta_var在编译期算好严格异步issue 与完成分离调用方用 mbarrier tcgen05.commit自定义同步粒度。实际使用中掌握三个要点即可快速上手目的 tmem buffer 的 layout 必须声明对应形状的 (lane, replica) 模式如 warpx4 要写R[4:32TLane]64x128b必须显式给出multicast二选一cta_group2时只让偶数 CTA 发出 cp 与 commitcta_mask3两个 CTA 各自拷自己的 smem→tmem。【免费下载链接】tvmOpen Machine Learning Compiler Framework项目地址: https://gitcode.com/gh_mirrors/tv/tvm创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考