理解 FlashAttention 每代优化背后的硬件直觉
FlashAttention 的优化不是纯算法游戏,而是 算法与硬件的协同设计。想要理解 FA-2 为什么调 loop order、FA-3 为什么用 TMA/WGMMA、FA-4 为什么重架构 Blackwell,需要先建立 GPU 的内存层次和 CUDA 执行模型直觉。
这一课是拓展阅读,不需要记住所有数字,但要理解三个核心问题:
GPU 的内存是一个层次结构,每一级在容量、带宽、延迟上相差巨大:
| 层级 | 容量(单 SM / 单卡) | 典型带宽 / 延迟 | 谁管理 |
|---|---|---|---|
| Register | ~256 KB / SM(数万个 32-bit 寄存器) | ~10–20 TB/s,~1 cycle | 编译器自动分配 |
| Shared Memory / L1(SRAM) | ~100–200 KB / SM | ~10–20 TB/s,~20–30 cycles | 程序员显式分配(`__shared__`) |
| L2 Cache | 数 MB 到数十 MB(全卡共享) | ~1–3 TB/s,~100 cycles | 硬件自动缓存 |
| HBM(全局显存) | 40–80 GB+ | ~1–2 TB/s,~300–500 cycles | 程序员显式读写 |
从 HBM 读一次数据的时间,足够从 SRAM 读十几次。FlashAttention 的核心策略就是:尽量少访问 HBM,把能留在 SRAM 的计算都留在 SRAM。
一个 CUDA kernel 启动时,会创建一个由线程组成的网格(grid):
一个 GPU 由多个 Streaming Multiprocessor(SM) 组成。一个 block 会被调度到一个 SM 上执行,一个 SM 可以同时驻留多个 block。SM 上有自己的寄存器文件、Shared Memory、Tensor Core、CUDA Core。
NVIDIA GPU 以 warp 为单位取指令、发射指令。一个 warp 内的 32 个线程执行同一条指令,但处理不同的数据(SIMT)。如果 warp 内线程走不同分支(branch divergence),硬件会串行执行各分支,降低效率。
| CUDA Core | Tensor Core | |
|---|---|---|
| 运算类型 | 标量/向量运算(FP32、INT32 等) | 矩阵乘加(GEMM)运算 |
| 典型吞吐 | 较低 | 高很多(A100 FP16 ~312 TFLOP/s,H100 FP8 ~1978 TFLOP/s) |
| attention 中作用 | softmax 的 exp、max、rescale | QK^T 和 PV 的矩阵乘法 |
FlashAttention 的大部分 FLOPs 都在 QK^T 和 PV 两个矩阵乘法上,所以 让 Tensor Core 满负荷运转 是性能的关键。FA-3/4 的很多优化都是在消除 Tensor Core 周围的"气泡"(bubbles)。
一个 kernel 的瓶颈通常分两类:
标准 attention 是 memory-bound:构造 S、P 两个大矩阵并写回 HBM,读写量太大。FlashAttention 通过 tiling 把中间矩阵留在 SRAM,把同一个 KV tile 复用多次,从而把问题推向 compute-bound。
\(\text{算术强度} = \frac{\text{FLOPs}}{\text{Bytes moved}}\)。算术强度越高,越接近 compute-bound。FlashAttention 通过减少 HBM 流量(分母)提高了算术强度。
一个 warp 的 32 个线程同时访问全局内存时,如果访问的地址连续,硬件可以合并成一次宽事务(如 128 byte),效率最高。如果访问分散,就需要多次事务,带宽利用率下降。
Shared Memory 被分成多个 bank。一个 warp 内多个线程同时访问同一个 bank 的不同地址时会发生 bank conflict,访问会串行化。FlashAttention 的 tile layout 会尽量避免这种情况。
一个 SM 上同时活跃的 warp 数占理论最大值的比率。Occupancy 高可以更好地隐藏延迟(当一个 warp 在等内存时,调度器可以执行另一个 warp)。但高 occupancy 不是唯一目标——有时候寄存器太多会限制 occupancy,需要权衡。
GPU 通过大量并发线程来隐藏内存延迟。当一个 warp 发起内存请求后进入等待,SM 会切换到其他就绪的 warp。FlashAttention-3 的异步化(TMA + WGMMA)把这一点推到了极致:内存拷贝、矩阵计算、softmax 可以并行进行。
| 概念 | FlashAttention 中的应用 |
|---|---|
| SRAM 高速但小 | 把 Q/K/V/O 切成 tile,只把当前需要的 tile 放进 SRAM |
| HBM 慢但大 | 避免写回完整的 S、P 矩阵,只从 HBM 加载原始输入 tile |
| Block / Warp 并行 | 一个 query tile 分配给一个 block,block 内多个 warp 协作完成 tile 内计算 |
| Tensor Core 高吞吐 | QK^T 和 PV 用 Tensor Core 加速,softmax 用 CUDA Core 做 exp/rescale |
| 延迟隐藏 | FA-3/4 用 TMA/WGMMA 异步重叠 load、MMA、softmax |
NVIDIA CUDA C++ Programming Guide 的 Memory Hierarchy 和 SIMT Architecture 章节是最权威的入门资料。
下节课:FlashAttention-2:Loop Order 与并行度。
随时可以问我:occupancy 和 latency hiding 的关系、bank conflict 具体怎么避免、或者某个概念和 FlashAttention 的联系。