详解 CUDA 内联汇编与 TMA 指令语法

📅 2026/8/17 18:26:15
详解 CUDA 内联汇编与 TMA 指令语法
详解 CUDA 内联汇编与 TMA 指令语法// 在 kernel 中单个线程触发 TMA 加载 __global__ void tma_kernel(CUtensorMap* desc_ptr){__shared__ __align__(128)half smem[64][64];// 只需要一个线程发起 TMA 操作if(threadIdx.x0){uint64_t smem_addr__cvta_generic_to_shared(smem);// 一条指令加载整个 64x64 tile asm volatile(cp.async.bulk.tensor.2d.shared.global.tile.bulk_group [%0], [%1, {%2, %3}];::l(smem_addr), // 共享内存地址l(desc_ptr), // 张量描述符r(tile_row), // tile 坐标(row)r(tile_col)// tile 坐标(col):memory);}详细解释这里的语法尤其是汇编这段代码涉及几个关键的语法点我会从最基础的开始逐层拆解。第一部分__cvta_generic_to_shared是什么背景CUDA 的地址空间CUDA 有多种内存地址空间Generic通用地址空间 - 一个统一的地址视图 - 可以指向全局内存、共享内存、局部内存 - C 指针默认是 generic 的 Shared共享内存地址空间 - 专门指向共享内存 - 地址范围小每个 SM 只有 228KB - 用 32 位就能表示为什么需要转换__shared__ half smem[64][64];// smem 是一个 generic 指针64 位half*ptrsmem;// 这是 generic 地址// 但 TMA 指令需要的是 shared 地址32 位// 所以需要转换uint64_tsmem_addr__cvta_generic_to_shared(smem);__cvta的含义__cvta ConVerT Address转换地址 generic_to_shared 从通用地址转为共享地址 对应的 PTX 指令cvta.to.shared.u64图解Generic 地址空间64 位 0x00007f0000001000 ← smem 的 generic 地址 ↓ __cvta_generic_to_shared Shared 地址空间相对偏移 0x00001000 ← smem 在共享内存中的偏移为什么 TMA 需要 shared 地址因为 TMA 引擎直接操作共享内存需要知道数据在共享内存中的物理偏移而不是通用视图的地址。第二部分内联汇编Inline Assembly语法GCC/CUDA 内联汇编的基本结构asmvolatile(汇编指令模板// 第 1 部分指令字符串:输出操作数// 第 2 部分输出:输入操作数// 第 3 部分输入:破坏描述// 第 4 部分clobber);用冒号:分隔四个部分。逐部分解析我们的代码asmvolatile(// 第 1 部分指令模板 cp.async.bulk.tensor.2d.shared.global.tile.bulk_group [%0], [%1, {%2, %3}];// 第 2 部分输出操作数空:// 第 3 部分输入操作数 :l(smem_addr),// %0l(desc_ptr),// %1r(tile_row),// %2r(tile_col)// %3// 第 4 部分破坏描述 :memory);第三部分volatile关键字asmvolatile(...);^^^^^^^^作用volatile告诉编译器不要优化这条汇编指令。不加 volatile 的风险// 假设不加 volatileasm(cp.async.bulk.tensor...);// TMA 加载// 编译器可能认为// 这条指令没有输出看起来没用删掉吧// 结果TMA 加载被优化掉了程序出错加了 volatileasmvolatile(cp.async.bulk.tensor...);// 编译器// 有 volatile这条指令有副作用必须保留且不能重排什么时候需要 volatile需要 volatile 的情况 ✅ 有副作用的指令内存操作、I/O ✅ 顺序敏感的指令 ✅ TMA、内存屏障等 可以不加的情况 纯计算指令如 add、mul编译器可以自由优化第四部分操作数约束Constraints约束字母的含义:l(smem_addr),// l 64 位整数寄存器l(desc_ptr),// l 64 位整数寄存器r(tile_row),// r 32 位整数寄存器r(tile_col)// r 32 位整数寄存器CUDA PTX 常用约束约束含义对应类型例子r32 位整数寄存器int,unsigned%r0l64 位整数寄存器long, 指针%rd0f32 位浮点寄存器float%f0d64 位浮点寄存器double%fd0h16 位整数寄存器short,half%rs0n立即数常量编译时常量100为什么地址用l64位坐标用r32位l(smem_addr)// 地址是 64 位虽然共享内存偏移只需 32 位但寄存器是 64 位l(desc_ptr)// 指针是 64 位r(tile_row)// tile 坐标是普通 int32 位足够r(tile_col)// tile 坐标是普通 int图解smem_addr (uint64_t) → 分配到 64 位寄存器 %rd0 desc_ptr (指针) → 分配到 64 位寄存器 %rd1 tile_row (int) → 分配到 32 位寄存器 %r0 tile_col (int) → 分配到 32 位寄存器 %r1第五部分占位符%0, %1, %2, %3编号规则占位符按照从上到下、从左到右的顺序编号asmvolatile(... [%0], [%1, {%2, %3}];:// 输出无:l(smem_addr),// %0 ← 第 0 个操作数l(desc_ptr),// %1 ← 第 1 个操作数r(tile_row),// %2 ← 第 2 个操作数r(tile_col)// %3 ← 第 3 个操作数:memory);如果有输出操作数呢编号会先数输出再数输入asmvolatile(add.u32 %0, %1, %2;:r(result)// %0 ← 输出从 0 开始:r(a),// %1r(b)// %2);注意输出约束前面有号表示写入。编译时的替换过程// 源代码[%0], [%1, {%2, %3}]// 编译器分配寄存器后[%rd0], [%rd1, {%r0, %r1}]// 最终 PTX[%rd0], [%rd1, {%r0, %r1}]第六部分破坏描述memory:memory作用告诉编译器这条指令会修改内存。没有memory的风险// 假设不加 memorysmem[0]100;// 写入共享内存asmvolatile(cp.async.bulk...);// TMA 加载会覆盖 smemintxsmem[0];// 读取// 编译器可能优化// smem[0] 刚写入 100直接用 100不用读内存// int x 100; ← 错误TMA 已经改变了 smem加了memoryasmvolatile(cp.async.bulk...:::memory);// 编译器// 这条指令修改了内存之后所有内存读取都要重新加载// 强制重新读取 smemmemory的完整含义memory 是一个 clobber破坏描述告诉编译器 1. 这条指令可能读写任意内存 2. 不要跨越这条指令缓存内存值 3. 内存操作不能重排到这条指令的另一侧 相当于一个编译器内存屏障第七部分TMA 指令本身的语法现在我们来解析 TMA 指令字符串cp.async.bulk.tensor.2d.shared.global.tile.bulk_group [%0], [%1, {%2, %3}];指令名称分解cp - copy拷贝 .async - 异步执行 .bulk - 批量传输大块数据 .tensor - 张量多维数据 .2d - 2 维 .shared - 目标共享内存 .global - 源全局内存 .tile - tile 模式按块拷贝 .bulk_group - 属于一个 bulk 组用于同步完整含义异步地、批量地、从全局内存拷贝一个 2D 张量 tile 到共享内存。操作数格式[%0], [%1, {%2, %3}] ↓ ↓ ↓ ↓ 目标 描述符 行 列详细解读[%0] → 目标共享内存地址smem_addr 方括号表示这是一个内存地址 [%1, {%2, %3}] → 源从张量描述符 desc_ptr 中 取坐标为 (tile_row, tile_col) 的 tile %1 → 张量描述符指针 {%2, %3} → tile 的坐标大括号表示坐标元组图解 TMA 的工作全局内存中的大矩阵4096 × 4096 ┌─────┬─────┬─────┬─────┐ │ T00 │ T01 │ T02 │ T03 │ ├─────┼─────┼─────┼─────┤ │ T10 │ T11 │ T12 │ T13 │ ← 每个 T 是一个 64×64 的 tile ├─────┼─────┼─────┼─────┤ │ T20 │ T21 │ T22 │ T23 │ └─────┴─────┴─────┴─────┘ 假设 tile_row1, tile_col2 → TMA 拷贝 T12 到共享内存 指令 cp.async.bulk.tensor.2d... [smem], [desc, {1, 2}] ↑ ↑ row col张量描述符 desc 里存了什么desc 包含在主机端用 cuTensorMapEncodeTiled 设置 - 全局内存基地址 - 矩阵总尺寸4096 × 4096 - Tile 大小64 × 64 - 数据类型half - Swizzle 模式128B - 步幅stride4096 TMA 收到坐标 {1, 2} 后自动计算 实际地址 base (1 × 64) × 4096 (2 × 64) base 262144 128第八部分完整流程串联让我们把所有部分串起来看完整的执行流程__global__voidtma_kernel(CUtensorMap*desc_ptr){__shared____align__(128)half smem[64][64];inttile_row1;// 假设inttile_col2;if(threadIdx.x0){// 步骤 1转换地址uint64_tsmem_addr__cvta_generic_to_shared(smem);// smem_addr 现在是共享内存偏移如 0x1000// 步骤 2内联汇编触发 TMAasmvolatile(cp.async.bulk.tensor.2d.shared.global.tile.bulk_group [%0], [%1, {%2, %3}];::l(smem_addr),// %0 → %rd0 (0x1000)l(desc_ptr),// %1 → %rd1 (描述符地址)r(tile_row),// %2 → %r0 (1)r(tile_col)// %3 → %r1 (2):memory);}}编译后的 PTX简化// __cvta_generic_to_shared 变成 cvta.to.shared.u64 %rd0, %rd_smem; // 转换地址 // 内联汇编变成 cp.async.bulk.tensor.2d.shared.global.tile.bulk_group [%rd0], [%rd1, {%r0, %r1}];执行时的硬件行为时刻 T0: 线程 0 执行 cp.async.bulk.tensor ↓ 时刻 T1: TMA 引擎读取描述符 desc_ptr - 得知矩阵 4096×4096, tile 64×64, half 类型 ↓ 时刻 T2: TMA 计算源地址 - 坐标 {1, 2} - 源地址 base 1*64*4096 2*64 ↓ 时刻 T3: TMA 从全局内存读取 64×64 tile - 8192 字节64×64×2 ↓ 时刻 T4: TMA 写入共享内存 [%rd0] - 自动应用 swizzle ↓ 时刻 T5: TMA 完成异步第九部分常见变体与扩展变体 1带屏障的版本实际使用中TMA 通常配合异步屏障asmvolatile(cp.async.bulk.tensor.2d.shared.global.tile.mbarrier::complete_tx::bytes [%0], [%1, {%2, %3}], [%4];// ↑ 新增屏障地址::l(smem_addr),l(desc_ptr),r(tile_row),r(tile_col),l(barrier_addr)// %4 → 屏障地址:memory);区别tile.bulk_group → 用 bulk group 同步较老 tile.mbarrier::... → 用 mbarrier 同步推荐变体 23D 张量asmvolatile(cp.async.bulk.tensor.3d.shared.global.tile.bulk_group [%0], [%1, {%2, %3, %4}];// 三个坐标::l(smem_addr),l(desc_ptr),r(coord_x),// %2r(coord_y),// %3r(coord_z)// %4:memory);变体 3使用 CUDA 内置函数推荐实际上你可以用 CUDA 提供的封装函数避免手写汇编#includecuda/barrier#includecudaTMA.h// 假设的头文件// 使用内置函数更安全cuda::device::experimental::cp_async_bulk_tensor_2d(smem,// 目标desc_ptr,// 描述符tile_row,// 坐标tile_col,barrier// 屏障);为什么还要学汇编✅ 理解底层原理 ✅ 调试时看得懂 PTX ✅ 某些高级优化需要手写 ✅ DeepGEMM 等库大量使用内联汇编第十部分语法速查表内联汇编结构asmvolatile(instruction %0, %1// 模板%n 是占位符:r(out)// 输出 表示写:r(in)// 输入:memory// 破坏描述);约束速查r → 32位整数 l → 64位整数 f → float d → double h → 16位 n → 常量 → 只写输出 → 读写特殊符号%0, %1, ... → 操作数占位符 [...] → 内存地址 {...} → 元组坐标 :: → 分隔可选参数如 mbarrier::complete_tx总结这段代码的核心要点__cvta_generic_to_shared把通用指针转成共享内存地址因为 TMA 需要物理偏移。asm volatile内联汇编volatile防止编译器优化掉。四段式结构指令模板 : 输出 : 输入 : 破坏描述。约束l和r分别表示 64 位和 32 位寄存器。占位符%0-%3按顺序对应操作数。memory告诉编译器这条指令会改内存别乱优化。TMA 指令cp.async.bulk.tensor.2d...一条指令拷贝整个 tile坐标{row, col}指定拷贝哪个 tile。一句话总结这段代码用内联汇编让单个线程发出一条 TMA 指令就能异步地把全局内存中指定坐标的整个 tile 搬到共享内存——这就是 Hopper “硬件搬运工” 的威力。后记2026年8月15日于上海在Claude opus 4.8辅助下完成。