CUDA 算子工程:手写 FlashAttention v2 之路
第 20 章 PTX 与 SASS:编译器到底干了什么
PTX 是你写下来的东西,SASS 才是真正在 SM 上跑的东西。 这一章讲清楚 ptxas 在这两者之间做了什么,以及什么时候值得绕过它自己写。
20.1 CUDA 编译流程
CUDA C++ 到 GPU 跑代码经过三层抽象:
flowchart LR CXX[CUDA C++ source] -->|nvcc 前端| PTX[PTX · 虚拟 ISA] PTX -->|nvcc 后端 ptxas| CUBIN[CUBIN · SASS 机器码] CUBIN -->|GPU 加载| EXEC[GPU 执行] PTX -.JIT 编译.-> CUBIN
- CUDA C++ source:你写的
.cu文件。 - PTX (Parallel Thread Execution):NVIDIA 的虚拟 ISA。前向兼容——按较老虚拟架构(比如
compute_70)生成的 PTX,可以在 Volta 及之后的 Ampere/Hopper 上跑(运行时再 JIT 成具体 SASS)。例外是带a后缀的架构专用目标(sm_90a等),它们的代码不保证前向兼容。 - SASS (Streaming ASsembler):真正在硬件上跑的机器码。每代架构的 SASS 不一样(Hopper 的 SASS 和 Ampere 的不同)。
编译时:
nvcc -arch=sm_90 my_kernel.cu -o my_kernel
-arch=sm_90 告诉编译器目标是 Hopper。nvcc 会生成同时包含 PTX 和 sm_90 SASS 的 fatbin。注意:kernel 里只要用到 WGMMA、setmaxnreg 这类 Hopper 专用指令(20.6.1 节),就必须写 -arch=sm_90a,只写 sm_90 会被 ptxas 直接拒绝(nvcc 13.4 实测报 not supported on .target 'sm_90')。(nvcc 13.4 实测还有一层:-arch=sm_90a 直接编目标文件或可执行文件时,会顺带生成一份通用的 compute_90 PTX,里面的WGMMA、setmaxnreg照样被 ptxas 拒绝。解法是改写成 -gencode arch=compute_90a,code=sm_90a,或者像 CUTLASS 那样把这段 asm 包进 #if defined(__CUDA_ARCH_FEAT_SM90_ALL),见 cutlass-4.7.0/include/cutlass/arch/config.h:48;只出 -cubin 时没有这个问题。不同 CUDA 版本可能不同。)
可以用 --keep 保留中间产物:
nvcc -arch=sm_90 --keep my_kernel.cu
# 生成 my_kernel.ptx / my_kernel.sm_90.cubin / my_kernel.fatbin / my_kernel.cudafe1.cpp 等中间文件
# 注意没有 .sass —— SASS 要用 cuobjdump 从 .cubin 反汇编出来
20.2 PTX 的语法
PTX 是一种汇编语言,但不是 GPU 真正的机器码——它是"假装的汇编",给编译器后端处理用的。
PTX 的几个关键语法元素:
// 加载一个浮点数到寄存器
ld.global.f32 %f1, [%rd1]; // %f1 = *(float*)%rd1
// 浮点加法
add.f32 %f3, %f1, %f2; // %f3 = %f1 + %f2
// 浮点 fma (融合乘加)
fma.rn.f32 %f4, %f1, %f2, %f3; // %f4 = %f1 * %f2 + %f3
// 存储
st.global.f32 [%rd2], %f4;
// 控制流
@%p1 bra TARGET; // if (%p1) goto TARGET
PTX 寄存器:
%r0..%rN:32-bit 通用寄存器%rd0..%rdN:64-bit 寄存器(指针)%fX:浮点寄存器%pX:谓词(条件)寄存器
这些前缀是 NVVM 生成 PTX 时的命名习惯,寄存器的类型由 .reg 声明决定。PTX 是无限寄存器的——程序员(或编译器)可以声明任意多个 %r0..%rN,最终 ptxas 后端会把它们映射到物理寄存器。
20.3 看 PTX 输出
编译时加 -ptx 让 nvcc 只生成 PTX:
nvcc -arch=sm_90 -ptx my_kernel.cu -o my_kernel.ptx
或者 -keep 保留。对下面这个 kernel:
__global__ void my_kernel(float* a, float* b, int n) {
int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i < n) b[i] = a[i] + 1.0f;
}
打开生成的 PTX 看(nvcc 13.4、-arch=sm_90 所得,略去文件头注释;不同 CUDA 版本的寄存器编号、标签名可能不同):
.version 9.4
.target sm_90
.address_size 64
.visible .entry _Z9my_kernelPfS_i(
.param .u64 _Z9my_kernelPfS_i_param_0, // float* a
.param .u64 _Z9my_kernelPfS_i_param_1, // float* b
.param .u32 _Z9my_kernelPfS_i_param_2 // int n
)
{
.reg .pred %p<2>;
.reg .f32 %f<3>;
.reg .b32 %r<6>;
.reg .b64 %rd<8>;
ld.param.u64 %rd1, [_Z9my_kernelPfS_i_param_0];
ld.param.u64 %rd2, [_Z9my_kernelPfS_i_param_1];
ld.param.u32 %r2, [_Z9my_kernelPfS_i_param_2];
mov.u32 %r3, %ctaid.x;
mov.u32 %r4, %ntid.x;
mov.u32 %r5, %tid.x;
mad.lo.s32 %r1, %r3, %r4, %r5; // r1 = blockIdx.x * blockDim.x + threadIdx.x
setp.ge.s32 %p1, %r1, %r2;
@%p1 bra $L__BB0_2;
cvta.to.global.u64 %rd3, %rd1;
mul.wide.s32 %rd4, %r1, 4;
add.s64 %rd5, %rd3, %rd4;
ld.global.f32 %f1, [%rd5];
add.f32 %f2, %f1, 0f3F800000; // 0f3F800000 即 1.0f 的十六进制写法
cvta.to.global.u64 %rd6, %rd2;
add.s64 %rd7, %rd6, %rd4;
st.global.f32 [%rd7], %f2;
$L__BB0_2:
ret;
}
读 PTX 的几个要点:
.reg声明寄存器。.b32是 32-bit,.f32是浮点。%ctaid/%ntid/%tid是 blockIdx / blockDim / threadIdx 的内置寄存器。mad.lo.s32是 32-bit 整数 multiply-add(a*b+c)。.lo表示取低 32 位。@%p1 bra是有条件跳转。cvta.to.global.u64是 generic 指针转 global 指针(CUDA 有 generic/global/shared 多种地址空间)。
20.4 看 SASS 输出
SASS 是真正的机器码。用 cuobjdump 或 nvdisasm 反汇编:
cuobjdump --dump-sass my_kernel.cubin
# 或者直接对 .o 文件:
cuobjdump --dump-sass my_kernel.o
输出(nvcc 13.4、-arch=sm_90 所得;每条指令后面的两行十六进制编码和函数末尾补齐用的 NOP 已略去,不同 CUDA 版本会有出入):
Function : _Z9my_kernelPfS_i
.headerflags @"EF_CUDA_SM90 EF_CUDA_VIRTUAL_SM(EF_CUDA_SM90)"
/*0000*/ LDC R1, c[0x0][0x28] ;
/*0010*/ S2R R0, SR_TID.X ;
/*0020*/ S2UR UR4, SR_CTAID.X ;
/*0030*/ LDC R7, c[0x0][RZ] ;
/*0040*/ IMAD R7, R7, UR4, R0 ;
/*0050*/ ULDC UR4, c[0x0][0x220] ;
/*0060*/ ISETP.GE.AND P0, PT, R7, UR4, PT ;
/*0070*/ @P0 EXIT ;
/*0080*/ LDC.64 R2, c[0x0][0x210] ;
/*0090*/ ULDC.64 UR4, c[0x0][0x208] ;
/*00a0*/ LDC.64 R4, c[0x0][0x218] ;
/*00b0*/ IMAD.WIDE R2, R7, 0x4, R2 ;
/*00c0*/ LDG.E R2, desc[UR4][R2.64] ;
/*00d0*/ IMAD.WIDE R4, R7, 0x4, R4 ;
/*00e0*/ FADD R7, R2, 1 ;
/*00f0*/ STG.E desc[UR4][R4.64], R7 ;
/*0100*/ EXIT ;
/*0110*/ BRA 0x110;
读 SASS 的关键点(SASS 没有 NVIDIA 官方的完整语义文档,官方的 CUDA Binary Utilities 文档只列了各代的指令助记符表;下面对寄存器约定、修饰符的解读有一部分来自社区逆向资料):
c[0x0][0x28]、c[0x0][0x210]是 constant memory 访问,c[bank][offset]。kernel 参数就放在 bank 0 里:这里0x210、0x218、0x220依次是a、b、n,c[0x0][RZ](偏移 0)是blockDim.x。LDC把常量读进普通寄存器,ULDC读进 uniform 寄存器(UR,整个 warp 共享一份)。开头那句LDC R1, c[0x0][0x28]几乎每个 kernel 都有,一般认为是把栈指针初值装进 R1——后面 spill 的LDL/STL就是以[R1+偏移]寻址的。S2R R0, SR_TID.X把 special register 读到 R0;S2UR读进 uniform 寄存器——blockIdx 对整个 warp 相同,所以放UR里。LDG.E是 Load Global,.E通常解读为 64 位(extended)地址;后面还可以跟宽度后缀(LDG.E.128一次读 128 bit)和 cache 策略后缀(如 20.6.2 节的.EF)。desc[UR4]是 Hopper 新出现的写法,UR4取自c[0x0][0x208],逆向资料一般称其为全局访存的内存描述符;Ampere 上同一段代码只写[R2.64]。@P0 EXIT是有条件退出(基于 predicate P0)。IMAD.WIDE是 32×32→64 乘加,常用于地址计算。- 别的 kernel 里还经常看到
IMAD.MOV.U32 Rx, RZ, RZ, ...这类写法:RZ是恒零寄存器,所以它等价于一条 MOV。社区的一般解释是编译器借整数乘加流水分担 MOV 的发射压力,NVIDIA 没有公开说明。
注意 SASS 比 PTX 少了几条指令——编译器把 PTX 中的 cvta 等"虚拟操作"消掉了,mul.wide + add 也合并成了一条 IMAD.WIDE。
20.5 看 SASS 找性能问题
ncu 告诉你哪里慢(第 19 章讲过,它的 Source 页本身也能逐行显示 SASS),读 SASS 则能解释为什么慢:
20.5.1 寄存器 spill 检测
如果 SASS 里看到大量 LDL/STL(local memory load/store),多半是编译器把寄存器溢出到了 local memory(物理上在显存里,经 L1/L2 缓存)。这是性能毒药。编译时加 -Xptxas -v,ptxas 会直接报出 bytes spill stores / bytes spill loads,比数指令更快。另外,按运行时下标访问的局部数组也会被放进 local memory,同样表现为 LDL/STL。
STL [R1+0x110], R0 ; // ← 溢出到 local memory
LDL.LU R10, [R1+0xe0] ; // ← 从 local memory 读回, 慢!
(两行摘自 nvcc 13.4、-arch=sm_90a 编译的一个故意溢出的 kernel;地址都以栈指针 R1 为基址。)
修复:减少局部变量、降低 unroll 程度、用 __launch_bounds__ 提示编译器降低寄存器使用。
20.5.2 Bank Conflict 实证
SMEM 访问指令是 LDS 和 STS。如果 SASS 里看到这些指令,可以用 ncu 测 bank conflict 数。SASS 本身能告诉你的是访问宽度和编译器排出来的偏移,但要当心一个常见误读:
LDS.128 R8, [R0] ; // 一次读 128 bit (4 个 fp32)
LDS.128 R12, [R0+0x80] ;
LDS.128 R4, [R0+0x100] ;
(nvcc 13.4、-arch=sm_90a 所得,源码是 float4 数组的 s[tid]、s[tid + 8]、s[tid + 16]。)
三条指令的地址差是 128 bytes(32 banks × 4 bytes),看上去都落在"同一组 bank"上,但这不构成 conflict:bank conflict 只发生在同一条指令内部、同一 warp 的不同 lane 之间,不同指令本来就是分开执行的。这个例子里每条指令的各 lane 地址是连续的 float4,没有冲突。真正决定冲突的是 R0 在各 lane 上的取值,也就是要顺着 SASS 往回追 R0 是怎么由 threadIdx 算出来的;最终仍以 ncu 的 bank conflict 计数为准,SASS 只是让你更快地猜到该去查哪里。
20.5.3 控制依赖与延迟
Volta 之后每条 SASS 指令(128 bit)里都编码了一段调度信息(control codes):固定 stall 周期数、yield 提示、读写 scoreboard barrier 的分配,以及要等待哪些 barrier 的 wait mask。NVIDIA 没有公开这部分格式,目前的认识来自逆向,如 Jia et al. 2018《Dissecting the NVIDIA Volta GPU Architecture via Microbenchmarking》。cuobjdump --dump-sass 把每条指令打成两行十六进制,控制位在第二行的高位里(形如 /* 0x000fe20000000800 */),并不直接解码;社区工具会把它翻成可读形式(Maxwell 时代的 maxas 写作 --:-:-:-:1,Volta 之后的 turingas、CuAssembler 一类写作 [B------:R-:W-:Y:S01])。
按上述逆向格式解 20.4 节的输出:LDG.E 分配了写 barrier SB2,隔了一条指令的 FADD 的 wait mask 恰好等 SB2——这就是"等 global load 回来"在机器码里的样子。所以要找依赖链上的等待点,看的是哪条指令在等 barrier;stall 字段只管固定延迟指令之间那几个周期。
编译器通常做得不错,但偶尔会看到次优的调度。要注意,手写 PTX 一般改不了这一层:ptxas 会重新做指令调度和寄存器分配,PTX 里的指令顺序不会原样保留到 SASS。真要逐条控制调度,只能用上面那些非官方的 SASS 汇编器,工程上很少这么做。
20.6 Inline PTX:何时用、怎么用
CUDA C++ 大部分情况下足够。但以下情况需要 inline PTX:
20.6.1 用还没暴露 C++ API 的硬件特性
Hopper 引入的很多新指令(TMA、WGMMA、setmaxnreg)在 C++ 层面只有部分包装,最完整的接口是 PTX。例如 wgmma.mma_async(示意:寄存器列表中间用 ... 省略了;这类代码须用 -arch=sm_90a 编译,前后还要配 wgmma.fence、wgmma.commit_group、wgmma.wait_group):
asm volatile(
"wgmma.mma_async.sync.aligned.m64n128k16.f32.f16.f16 "
"{%0, %1, ..., %63}, %64, %65, %66, 1, 1, 0, 0;\n"
// D 是 64 个 f32 累加器; %64/%65 是 A/B 的 SMEM 描述符;
// %66 = scale-D (0 覆盖 / 1 累加); 后面 4 个立即数是 scale-A/B 与 trans-A/B
: "+f"(c[0]), "+f"(c[1]), /* ... */ "+f"(c[63])
: "l"(desc_a), "l"(desc_b), "n"(scale_d) // "n" 要求编译期常量
);
把省略处补全后,nvcc 13.4、-arch=sm_90a 编出来的核心是 WARPGROUP.ARRIVE、HGMMA.64x128x16.F32 R24, gdesc[UR4], R24, gsb0、WARPGROUP.DEPBAR.LE gsb0, 0x0 三条。scale-D 在 PTX 里是谓词操作数;要在运行时决定覆盖还是累加,就得像 CUTLASS 那样先 setp.ne.b32 p, %66, 0; 再把谓词 p 填进去(此时 SASS 里多出一个 UP0 操作数)。
CUDA 12.x 起的 libcu++(CCCL)提供了 cuda::ptx:: 这一层薄包装(mbarrier_*、cp_async_bulk_*、tensormap_* 等),但覆盖并不完整:以 CUDA 13.4 附带的 CCCL 为例,setmaxnreg、ldmatrix 和 Blackwell 的 tcgen05_* 都有了,Hopper 的 wgmma 却没有,只能自己写。FA3 和 CUTLASS 里都是大量手写 inline PTX——CUTLASS 干脆把它们封成了 cute::SM90_64x128x16_F32F16F16_SS 这类 atom(别名在 cutlass-4.7.0/include/cute/atom/mma_traits_sm90_gmma.hpp:1244,真正那段 inline PTX 在 cutlass-4.7.0/include/cute/arch/mma_sm90_gmma.hpp:1632)。
20.6.2 控制编译器无法表达的优化
某些 PTX 指令的"flag"在 C++ 层没有现成写法。以 ld.global.cs(cache streaming)为例,它告诉硬件这份数据大概只读一次,按 evict-first 策略进缓存,以减少对 L1/L2 的污染:
// 用 streaming load 减少缓存污染
asm("ld.global.cs.v4.b32 {%0,%1,%2,%3}, [%4];\n"
: "=r"(v0), "=r"(v1), "=r"(v2), "=r"(v3)
: "l"(addr));
nvcc 13.4、-arch=sm_90a 编出来是一条 LDG.E.EF.128(.EF 即 evict-first)。四个 32 位输出要写 .v4.b32 向量形式;写成 .b128 配四个寄存器,ptxas 会报 Argument vector size mismatch。单独的 .cs 其实 CUDA 已有 intrinsic __ldcs()(sm_32_intrinsics.hpp 里就是这条 inline PTX),__ldcs((const uint4*)addr) 编出来同样是 LDG.E.EF.128;要叠加别的修饰时才轮到手写,比如 apex 的 apex-25.09/apex/contrib/csrc/groupbn/nhwc_batch_norm_kernel.h:110 直接写了 ld.global.cs.nc.s32。
20.6.3 强制特定指令
编译器有时会"自作聪明"——把你的 int4 解码成多条指令而不是一条 lop3.b32。比如 C++ 写 (q & 0x000F000F) | 0x64006400,nvcc 13.4、-arch=sm_90a 编出来是两条 LOP3.LUT(先与再或)。这时手写 PTX 强制用 lop3 能省指令带宽(第 9 章 Marlin 的例子,vLLM 源码注释原话是 "Guarantee that the (a & b) | c operations are LOP3s.")。
20.6.4 让 fragment 直接进 mma
mma.sync 对 fragment 有固定的布局要求:warp 里每个 lane 的每个寄存器装矩阵的哪几个元素,PTX 文档都规定死了。C++ 的 WMMA API 把 fragment 布局藏了起来,没法直接接上 ldmatrix 读出的寄存器;直接用 inline PTX 写 mma.sync,就能自己掌控每个寄存器装什么(物理寄存器编号仍由 ptxas 分配)。
20.7 一组高频 PTX/SASS Pattern
LLM kernel 中高频出现的指令(下面的 SASS 均按 nvcc 13.4 实际编译输出核对过,寄存器编号随上下文而变,不同 CUDA 版本也可能不同):
A. mma.sync (Tensor Core 矩阵乘)
mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32
{%0, %1, %2, %3},
{%4, %5, %6, %7},
{%8, %9},
{%0, %1, %2, %3};
SASS:
HMMA.16816.F32 R4, R8, R12, R4 ;
B. cp.async (Ampere+ 异步拷贝)
cp.async.cg.shared.global [%0], [%1], 16;
cp.async.commit_group;
cp.async.wait_group 0;
SASS:
LDGSTS.E.BYPASS.128 [R7], desc[UR6][R2.64] ;
LDGDEPBAR ;
DEPBAR.LE SB0, 0x0 ;
(.cg 绕过 L1,对应 SASS 的 .BYPASS。)
C. cp.async.bulk.tensor (Hopper TMA)
cp.async.bulk.tensor.2d.shared::cluster.global.tile.mbarrier::complete_tx::bytes
[%0], [%1, {%2, %3}], [%4];
SASS(Hopper):一条 PTX 对应一条 UTMALDG 系列指令(global→shared 是 UTMALDG,shared→global 是 UTMASTG)。nvcc 13.4 所得为 UTMALDG.2D [UR4], [UR8]:维度写在后缀上,操作数全在 uniform 寄存器(UR)里。具体助记符随 CUDA 版本会有出入,以自己机器上 cuobjdump --dump-sass 的输出为准。
D. ldmatrix (从 SMEM 加载 fragment)
ldmatrix.sync.aligned.m8n8.x4.shared.b16
{%0, %1, %2, %3}, [%4];
SASS(.M88 是不转置,加了 .trans 的那一版对应 .MT88;末尾的 .4 是 .x4):
LDSM.16.M88.4 R4, [R20] ;
E. lop3 (三输入逻辑运算)
lop3.b32 %0, %1, 0x000F000F, 0x64006400, 0xea;
SASS:
HFMA2.MMA R9, -RZ, RZ, 1024, 1024 ; // 把 0x64006400 装进 R9
LOP3.LUT R7, R2, 0xf000f, R9, 0xea, !PT ;
(nvcc 13.4、-arch=sm_90a 所得。)一条 SASS 指令只能带一个 32 位立即数,所以第二个常数先被装进寄存器——0x6400 恰好是 FP16 的 1024,编译器顺手用一条 HFMA2.MMA 把它"算"了出来(换个上下文也可能是一条 IMAD.MOV.U32)。
熟悉这些 pattern 后,读 SASS 会变成"翻译"——每条 SASS 都对应一个清晰的工程意图。
20.8 一个有趣的实战:让编译器生成正确的指令
经常见到这种情况:你写了一段看起来高效的 C++ 代码,但 SASS 显示它生成了次优指令。
例:把一个字节转成 float。直接从内存读一个 uint8_t 再转换:
float f = (float)p[i]; // p 是 const uint8_t*
SASS(nvcc 13.4、-arch=sm_90a):
LDG.E.U8 R4, desc[UR4][R4.64] ; // 读字节时硬件顺手零扩展
I2FP.F32.U32 R7, R4 ; // 一条转换
但如果这个字节是从一个已经读进寄存器的 32 位打包字里取出来的:
uint32_t w = ...;
float f = (float)(uint8_t)(w >> 8);
SASS(同上编译条件):
LOP3.LUT R0, R2, 0xffff, RZ, 0xc0, !PT ;
SHF.R.U32.HI R0, RZ, 0x8, R0 ;
LOP3.LUT R0, R0, 0xffff, RZ, 0xc0, !PT ;
I2FP.F32.U32 R7, R0 ;
同样是"字节转 float",这里要 4 条指令而不是 1 条;而同一个 kernel 里取有符号字节的 (float)(int8_t)(w >> 16) 只要 SHF + I2F.S8 两条。这里的重点不是"哪种类型一定快"——具体生成什么强烈依赖 nvcc 版本、这个字节是怎么从内存/打包字里取出来的、以及周围的上下文;同一段代码换个编译器版本可能就变了。重点是这种事会发生,而且从 C++ 源码上完全看不出来。
量化 kernel 里这类"多出来的 PRMT / SHF / LOP3"最容易成规模(第 9 章对比 INT4 解码的朴素写法和 lop3 写法,比的正是 SASS 里的指令条数)。读 SASS 是发现它们的唯一方法。
20.9 这一章的小结与下一章
PTX/SASS 是 CUDA 工程师的"反汇编技能":
- PTX 是虚拟 ISA,SASS 是真实机器码:ptxas(编译期)或驱动 JIT(运行期)把 PTX 转成 SASS,每代架构 SASS 不同。
-keep保留 PTX,cuobjdump --dump-sass看 SASS:这是日常工具。- 读 SASS 找寄存器 spill / bank conflict 线索 / 次优指令:ncu 告诉你哪里慢,SASS 解释为什么慢。
- Inline PTX 是高级优化的最后一招:当 C++ 编译器不够聪明时,手写 PTX 强制特定指令。
- 熟悉高频指令 pattern:mma、cp.async、ldmatrix、lop3 等 LLM kernel 的关键指令。
第 21 章我们结束第五篇(也是本专栏的核心内容)——讲性能陷阱与反模式。一些常见的"看起来对、实际上慢"的写法,把它们罗列出来作为读者的避坑指南。读完第 21 章读者就完成了从理论到实战到诊断的完整训练。
本章动手练习:
- 编译你之前写的某个 kernel,用
cuobjdump --dump-sass看 SASS。找一行 SASS,对应回 CUDA C++ 源码。- 写一段故意有寄存器 spill 的 kernel(比如几百个局部变量),看 SASS 中的
LDL/STL指令。- 用 inline PTX 写一条
lop3.b32,对比让编译器自动生成的版本。