为什么把 query tile 放到外层循环能让 FlashAttention 快近一倍
在 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 被反复加载/写入
这种结构的问题:
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 往返。
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\) 越大,并行度越高。
FA-2 的工作划分更清晰:
这比 FA-1 减少了 block 间的同步开销。FA-1 中不同 query tile 共享同一个 KV tile,需要在更新全局 running statistics 时做更多同步。
FlashAttention 的计算包括两部分:
FA-2 通过更好的 warp-level work partition 减少了 non-matmul 对 Shared Memory 的读写次数。例如,running statistics 可以更多地保留在寄存器中,减少 shared memory 的往返。
FA-2 在 A100 上达到约 225 TFLOP/s,相当于峰值算力的 50–73%。这主要归功于并行度提升和非 matmul 开销降低。
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) # 多次写回
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) # 只写回一次
FlashAttention-2 的核心优化是 调整循环顺序 和 增大并行度:把 query tile 放到外层,让每个 query tile 的计算独立并行,减少 \(Q_i\) 和 \(O_i\) 的 HBM 往返,并把并行维度从 batch×head 扩展到 batch×head×sequence。
本课推荐阅读 FlashAttention-2 原论文 Faster Attention with Better Parallelism and Work Partitioning 的 Section 3 和 Section 4。
下节课预告:FlashAttention-3:Hopper 异步执行与 FP8。
FA-2 的优化点其实不复杂,但它是从"算法正确"到"硬件高效"的关键一跃。如果有任何不清楚的地方,随时问我。