拓展阅读:GPU 内存层次与 CUDA 执行模型基础

理解 FlashAttention 每代优化背后的硬件直觉

1. 为什么要学这一课

FlashAttention 的优化不是纯算法游戏,而是 算法与硬件的协同设计。想要理解 FA-2 为什么调 loop order、FA-3 为什么用 TMA/WGMMA、FA-4 为什么重架构 Blackwell,需要先建立 GPU 的内存层次和 CUDA 执行模型直觉。

这一课是拓展阅读,不需要记住所有数字,但要理解三个核心问题:

2. GPU 内存层次:离计算单元越近越快

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。

3. CUDA 执行模型:Grid → Block → Warp → Thread

一个 CUDA kernel 启动时,会创建一个由线程组成的网格(grid):

一个 GPU 由多个 Streaming Multiprocessor(SM) 组成。一个 block 会被调度到一个 SM 上执行,一个 SM 可以同时驻留多个 block。SM 上有自己的寄存器文件、Shared Memory、Tensor Core、CUDA Core。

为什么 Warp 是 32 个线程?

NVIDIA GPU 以 warp 为单位取指令、发射指令。一个 warp 内的 32 个线程执行同一条指令,但处理不同的数据(SIMT)。如果 warp 内线程走不同分支(branch divergence),硬件会串行执行各分支,降低效率。

4. CUDA Core vs Tensor Core

CUDA CoreTensor 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)。

5. Memory-bound vs Compute-bound

一个 kernel 的瓶颈通常分两类:

标准 attention 是 memory-bound:构造 S、P 两个大矩阵并写回 HBM,读写量太大。FlashAttention 通过 tiling 把中间矩阵留在 SRAM,把同一个 KV tile 复用多次,从而把问题推向 compute-bound。

算术强度(Arithmetic Intensity)

\(\text{算术强度} = \frac{\text{FLOPs}}{\text{Bytes moved}}\)。算术强度越高,越接近 compute-bound。FlashAttention 通过减少 HBM 流量(分母)提高了算术强度。

6. 几个影响性能的关键概念

6.1 Coalesced Memory Access(合并内存访问)

一个 warp 的 32 个线程同时访问全局内存时,如果访问的地址连续,硬件可以合并成一次宽事务(如 128 byte),效率最高。如果访问分散,就需要多次事务,带宽利用率下降。

6.2 Shared Memory Bank Conflict

Shared Memory 被分成多个 bank。一个 warp 内多个线程同时访问同一个 bank 的不同地址时会发生 bank conflict,访问会串行化。FlashAttention 的 tile layout 会尽量避免这种情况。

6.3 Occupancy(占用率)

一个 SM 上同时活跃的 warp 数占理论最大值的比率。Occupancy 高可以更好地隐藏延迟(当一个 warp 在等内存时,调度器可以执行另一个 warp)。但高 occupancy 不是唯一目标——有时候寄存器太多会限制 occupancy,需要权衡。

6.4 Latency Hiding(延迟隐藏)

GPU 通过大量并发线程来隐藏内存延迟。当一个 warp 发起内存请求后进入等待,SM 会切换到其他就绪的 warp。FlashAttention-3 的异步化(TMA + WGMMA)把这一点推到了极致:内存拷贝、矩阵计算、softmax 可以并行进行。

7. 这些概念如何串起 FlashAttention

概念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

8. 快速测验

1. GPU 上哪一级内存的延迟最低、带宽最高?
2. FlashAttention 把 S、P 矩阵留在 SRAM 而不是写回 HBM,主要解决了什么问题?
3. CUDA 中,一个 warp 包含多少个线程?

9. 延伸阅读

NVIDIA CUDA C++ Programming Guide 的 Memory HierarchySIMT Architecture 章节是最权威的入门资料。

下节课:FlashAttention-2:Loop Order 与并行度

有疑问?

随时可以问我:occupancy 和 latency hiding 的关系、bank conflict 具体怎么避免、或者某个概念和 FlashAttention 的联系。