cute(11)hopper之TMA_copy
前置阅读
11 TMA Copy
对应代码:
11_tma_copy.cu需要 GPU(SM90 / Hopper)。
核心概念
TMA(Tensor Memory Accelerator) 是 Hopper 架构的硬件加速搬运引擎:
- 单线程即可发起整个 tile 的 Global → Shared Memory 搬运
- 硬件自动处理地址计算、Swizzle、多维索引
- 搬运完成后自动通过 mbarrier 通知
- 与 cp.async(09 课)的根本区别:cp.async 每线程搬一小块,TMA 一个线程搬整个 tile
代码实现
1 | /** |
打印信息如下
1 | [yuyanqi@set-hldy-mlp-codelab-pc75 build]$ ./11_tma_copy |
1. TMA 的工作流程
1 | Host 端(一次性准备): |
2. Host 端:构造 TMA Descriptor
2.1 GMEM Tensor
1 | auto gmem_tensor = make_tensor( |
2.2 SMEM Layout(带 Swizzle)
1 | using SmemLayoutAtom = GMMA::Layout_K_SW128_Atom<half_t>; |
为什么用 Layout_K_SW128_Atom 而不是 Layout_MN_SW128_Atom?
| 属性 | Layout_K | Layout_MN |
|---|---|---|
| stride-1 方向 | 第二维 (K/N) | 第一维 (M/N) |
| 适用场景 | gmem row-major: stride-1 沿 N | gmem col-major: stride-1 沿 M |
我们的 gmem 是 (M, N) stride (N, 1) → stride-1 沿 N(第二维)→ smem 也应 stride-1 沿第二维 → Layout_K。
TMA 对 Swizzle 的要求:
1 | Swizzle<B, M, S> 中 M 必须 ≥ 4 |
TMA 的 swizzle base 需要 16B 对齐(M=4)、32B(M=5)或 64B(M=6),对应:
Layout_K_INTER_Atom: Swizzle<0,4,3> — 无 swizzleLayout_K_SW32_Atom: Swizzle<1,4,3> — 32B swizzleLayout_K_SW64_Atom: Swizzle<2,4,3> — 64B swizzleLayout_K_SW128_Atom: Swizzle<3,4,3> — 128B swizzle
2.3 make_tma_copy
1 | auto tma_load = make_tma_copy(SM90_TMA_LOAD{}, gmem_tensor, SmemLayout{}); |
这一步在 Host 端执行,创建一个包含 TmaDescriptor(128 字节对齐结构体)的 Copy_Traits 对象。Descriptor 编码了:
- gmem 的基地址、shape、stride
- smem 的 swizzle 模式
- 搬运的元素类型和大小
3. Kernel 端:发起 TMA Copy
3.1 grid_constant
1 | __global__ void tma_copy_kernel( |
__grid_constant__ 是 Hopper 引入的修饰符:
- 让 kernel 参数存在于 constant memory 中(而非被拷贝到每个线程的 local memory)
- TMA descriptor 是 128 字节的大结构体,必须用此修饰
- 所有线程共享同一份,减少 register 压力
3.2 坐标 Tensor 和 Partition
1 | // 1. 获取坐标 tensor(arithTuple 格式) |
为什么 TMA 用"坐标 tensor"而不是数据 tensor?
cp.async 需要每线程自己算地址→搬一小块。TMA 不同:
- 线程不直接访问 gmem 地址
- 而是告诉 TMA 引擎"请搬坐标 (row, col) 处的 tile"
- 地址计算由硬件完成
3.3 发起搬运
1 | if (tid == 0) { |
关键步骤:
- initialize_barrier: 设 arrive_count=1(
set_barrier_transaction_bytes内部会 arrive) - set_barrier_transaction_bytes: 设置期望搬运量 + 隐式 arrive
- copy(tma.with(mbar)):
.with(mbar)绑定 mbarrier,TMA 搬运完成后自动 complete_tx - 所有线程
wait_barrier(mbar, 0)→ 等待搬运完成
1 | 时间线: |
3.4 同步代码详解:两个 __syncthreads 和 tid==0 wait
1 | __syncthreads(); // ① 确保 mbarrier 初始化完成 |
为什么是 tid==0 做 wait_barrier,而不是所有线程?
wait_barrier 内部是 spin-wait 忙等循环:
1 | LAB_WAIT: |
128 个线程全部 spin-wait → 128 个线程同时轮询同一个 smem 地址 → 浪费算力、增加 smem bank 争用。只让 1 个线程探测就够了,然后用 __syncthreads 广播"数据已就绪"。
三个同步点各自的作用:
1 | tid==0: init_barrier → set_tx_bytes → copy(TMA) |
能不能让所有线程都 wait_barrier?
可以,但没必要。wait_barrier 是纯轮询,N 个线程同时轮询不会比 1 个更快通过——翻转时机由 TMA 完成决定,与轮询线程数无关。1 个线程 wait + 1 个 __syncthreads 是更高效的做法。
对比 12 课 Pipeline: 12 课中 pipeline.consumer_wait() 是所有 consumer 线程都调用的,因为 PipelineTmaAsync 内部做了优化(分散到不同 warp、支持 cluster),不存在简单 spin-wait 的问题。
4. 验证方式
搬运完成后,每个线程从 smem 读数据写回另一块 gmem,然后 host 端对比原始数据和写回数据:
1 | // 通过 smem tensor sA 访问(自动处理 swizzle 解码) |
注意:直接用裸指针读 smem 数据会因为 swizzle 而得到错误的位置。必须通过 sA(row, col) 访问,CuTe 会自动应用 swizzle 映射。
5. TMA vs cp.async 对比
| 特性 | cp.async (SM80/Ampere) | TMA (SM90/Hopper) |
|---|---|---|
| 发起线程数 | 所有线程各搬一部分 | 单线程发起整个 tile |
| 地址计算 | 每线程自己算 | 硬件自动(通过 Descriptor) |
| Swizzle | 需要手动 Layout | Descriptor 编码,硬件处理 |
| 同步机制 | cp_async_fence / wait | mbarrier (phase 翻转) |
| Descriptor | 不需要 | Host 端构造,Kernel 端传入 |
| 多维支持 | 只能 1D(逐行搬) | 最多 5D tensor 直接搬 |
| Warp 专用化 | 困难(所有线程都搬运) | 自然支持(单线程搬运) |
TMA 的关键优势:搬运只需 1 个线程 → 其余线程可以做计算 → Warp Specialization 的基础。
6. TMA 与 LSU 的代理可见性(Proxy Visibility)
6.1 两条独立的硬件通路
Hopper 上共享内存有两条独立的访问通路:
1 | LSU (Load/Store Unit): 普通线程的 ld.shared / st.shared 指令 |
它们是不同的硬件代理(proxy),互相看不到对方的 in-flight 写入。
6.2 TMA 写 → LSU 读:无需额外操作
本课使用的场景:
1 | TMA 引擎写入 smem → mbarrier wait → LSU 读取 smem |
这是安全的,因为:
mbarrier.wait保证 TMA 搬运已完成(transaction bytes 归零)- TMA 写入完成后,数据对 LSU 可见(TMA → LSU 方向有硬件保证)
6.3 LSU 写 → TMA 读:需要 fence!
如果反过来,线程用普通指令写 smem,然后让 TMA 引擎读 smem(如 TMA_STORE:Smem → Global):
1 | // 危险!LSU 写入可能对 TMA 不可见 |
问题:__syncthreads 保证所有线程执行到同一点,但 LSU 的 st.shared 可能还在流水线上,数据尚未真正写入 smem。如果是 LSU 自己读(走同一条流水线),顺序性保证数据可见;但 TMA 引擎走的是另一条通路,看不到 LSU 流水线中尚未提交的写入。
修复:在 __syncthreads 前加 fence.proxy.async.shared::cta:
1 | sdata[idx] = value; |
fence.proxy.async.shared::cta 的作用:
- 等待 fence 之前所有 LSU 对 smem 的写入完成
- 使这些写入对异步代理(TMA)可见
- TMA 后续的读取会等待 fence 完成后才开始
6.4 CUTLASS 中的封装
CUTLASS 在 cute/arch/copy_sm90_tma.hpp 中提供了封装函数:
1 | // cute::tma_store_fence() — 在 LSU 写 smem 和 TMA Store 之间调用 |
典型的 TMA Store 流程:
1 | // 1. 线程通过 LSU 写入 smem(例如 MMA 累加器 → smem) |
6.5 总结:四种组合
| 写入方 | 读取方 | 是否需要 fence | 原因 |
|---|---|---|---|
| TMA | LSU | 不需要 | mbarrier.wait 已保证可见性 |
| LSU | LSU | 不需要 | 同一流水线,顺序保证 |
| LSU | TMA | 需要 fence.proxy.async.shared::cta |
跨代理,LSU 写入对 TMA 不可见 |
| TMA | TMA | 不需要 | 同一代理,顺序保证 |
本课只用了 TMA_LOAD(TMA 写 → LSU 读),所以不需要 fence。但在实际 GEMM kernel 的 Epilogue 中(累加器 → smem → TMA Store → global),必须加 fence。
7. API 总结
Host 端
| API | 作用 |
|---|---|
make_tma_copy(SM90_TMA_LOAD{}, gmem, smem_layout) |
构造 TMA descriptor |
SM90_TMA_LOAD{} |
TMA Load 操作类型(Global→Smem) |
Kernel 端
| API | 作用 |
|---|---|
tma_load.get_tma_tensor(shape) |
获取坐标 tensor(arithTuple) |
local_tile(coord, tile_shape, coord) |
切出当前 block 的 tile 坐标 |
tma_load.get_slice(0) |
获取线程 0 的搬运 partition |
thr.partition_S(coord_tile) |
Source partition(坐标) |
thr.partition_D(smem_tensor) |
Destination partition(smem) |
tma_load.with(mbar) |
绑定 mbarrier,搬运完成自动 arrive |
copy(tma.with(mbar), src, dst) |
发起 TMA 搬运 |
mbarrier(第 10 课 API)
| API | 作用 |
|---|---|
initialize_barrier(mbar, 1) |
初始化(1 个 arrive) |
set_barrier_transaction_bytes(mbar, bytes) |
设置期望搬运字节 + 隐式 arrive |
wait_barrier(mbar, 0) |
等待 phase 翻转(搬运完成) |
- 标题: cute(11)hopper之TMA_copy
- 作者: 鱿鱼圈
- 创建于 : 2026-06-17 12:13:32
- 更新于 : 2026-06-14 22:17:07
- 链接: https://yuyanqi.com/2026/06/17/cute(11)tma_copy/
- 版权声明: 本文章采用 CC BY-NC-SA 4.0 进行许可。