第 3 课:FlashAttention-2 Loop Order 与并行度

为什么把 query tile 放到外层循环能让 FlashAttention 快近一倍

1. 回顾 FlashAttention-1 的循环结构

在 FA-1 中,外层循环遍历 KV tile,内层循环遍历 query tile:

# FlashAttention-1 风格:外层 KV tile,内层 query tile
for j in range(T_c):            # 外层:遍历 K_j, V_j
    K_j = load(K[j*B_c:(j+1)*B_c, :])
    V_j = load(V[j*B_c:(j+1)*B_c, :])

    for i in range(T_r):        # 内层:遍历 Q_i
        Q_i = load(Q[i*B_r:(i+1)*B_r, :])
        # 计算 S_ij,更新 running max/sum,累加 O_i
        # ...
        store(O_i)              # 每个 query tile 被反复加载/写入

这种结构的问题:

2. FlashAttention-2:把 query tile 放到外层

FA-2 把循环顺序反过来:

# FlashAttention-2 风格:外层 query tile,内层 KV tile
for i in range(T_r):            # 外层:遍历 Q_i
    Q_i = load(Q[i*B_r:(i+1)*B_r, :])
    m_i = full(B_r, -inf)
    l_i = zeros(B_r)
    O_i = zeros(B_r, d)

    for j in range(T_c):        # 内层:遍历 K_j, V_j
        K_j = load(K[j*B_c:(j+1)*B_c, :])
        V_j = load(V[j*B_c:(j+1)*B_c, :])
        # 计算 S_ij,更新 m_i, l_i, O_i
        # ...

    O_i = O_i / l_i[:, None]
    store(O[i*B_r:(i+1)*B_r, :], O_i)   # 每个 Q_i 只写回一次

这个改动的直接收益:

关键直觉

FA-2 不是改了算法正确性,而是改了 数据复用和并行粒度。把 "output accumulator 的生命周期" 和 "query tile 的生命周期" 对齐,减少了 HBM 往返。

3. 并行度:从 batch/head 到 sequence length

FA-1 的并行主要来自 batch 和 head 维度。当 batch size 小、head 数少时,一个 batch-head 内只有一个 block 在工作,SM 利用率不高。

FA-2 引入 sequence-length parallelism:同一个 batch-head 内的不同 query tile 可以分配到不同 block。这样即使 batch/head 很小,长序列也能产生足够多的 block 填满 GPU。

\(\text{FA-1 并行度} \approx \text{batch} \times \text{heads}\)
\(\text{FA-2 并行度} \approx \text{batch} \times \text{heads} \times T_r\)

其中 \(T_r = \lceil N / B_r \rceil\) 是 query tile 的数量。序列越长,\(T_r\) 越大,并行度越高。

4. Work Partitioning:一个 block 负责一个 query tile

FA-2 的工作划分更清晰:

这比 FA-1 减少了 block 间的同步开销。FA-1 中不同 query tile 共享同一个 KV tile,需要在更新全局 running statistics 时做更多同步。

5. 减少非 matmul 开销

FlashAttention 的计算包括两部分:

FA-2 通过更好的 warp-level work partition 减少了 non-matmul 对 Shared Memory 的读写次数。例如,running statistics 可以更多地保留在寄存器中,减少 shared memory 的往返。

性能数字

FA-2 在 A100 上达到约 225 TFLOP/s,相当于峰值算力的 50–73%。这主要归功于并行度提升和非 matmul 开销降低。

6. 完整伪代码对比

FA-1(外层 KV)

for j in range(T_c):
    K_j, V_j = load(K_j), load(V_j)
    for i in range(T_r):
        Q_i = load(Q_i)                 # 每次 j 都重新加载
        m_i, l_i, O_i = load_stats()    # 全局 running stats
        # ... compute, update, store
        store_stats()
        store(O_i)                      # 多次写回

FA-2(外层 query)

for i in range(T_r):
    Q_i = load(Q_i)                     # 只加载一次
    m_i, l_i, O_i = init_local()        # 局部 accumulator
    for j in range(T_c):
        K_j, V_j = load(K_j), load(V_j)
        # ... compute, update local accumulators
    O_i = O_i / l_i[:, None]
    store(O_i)                          # 只写回一次

7. 一句话总结

FlashAttention-2 的核心优化是 调整循环顺序增大并行度:把 query tile 放到外层,让每个 query tile 的计算独立并行,减少 \(Q_i\) 和 \(O_i\) 的 HBM 往返,并把并行维度从 batch×head 扩展到 batch×head×sequence。

8. 快速测验

1. FlashAttention-2 相比 FA-1,主要调整了哪一点?
2. FA-2 把 query tile 放到外层后,并行度增加了哪个维度?
3. FA-2 中一个 CUDA block 通常负责什么?

9. 延伸阅读

本课推荐阅读 FlashAttention-2 原论文 Faster Attention with Better Parallelism and Work Partitioning 的 Section 3 和 Section 4。

下节课预告:FlashAttention-3:Hopper 异步执行与 FP8

有疑问?

FA-2 的优化点其实不复杂,但它是从"算法正确"到"硬件高效"的关键一跃。如果有任何不清楚的地方,随时问我。