通过交互式可视化,深入理解 Flash Attention 的核心技术:内存瓶颈、Online Softmax、与分块矩阵乘法。

标准 Attention 的内存瓶颈

在深入 Flash Attention 的代码实现之前,我们必须先回答一个底层问题:为什么标准的注意力机制公式 𝑆 𝑜 𝑓 𝑡 𝑚 𝑎 𝑥 ( 𝑄 𝐾 𝑇 ) 𝑉 Softmax(QK T )V 在现代 GPU 上跑得还不够快?

GPU 内存层级:SRAM 与 HBM

首先我们需要建立一个极其重要的概念:在 GPU 架构中,所有的计算(如矩阵加法、乘法)都必须在靠近核心的 SRAM(Shared Memory,共享内存) 中进行。

这意味着:哪怕你的显存(HBM)有 80GB 那么大,数据也必须先被“搬运”到几十 MB 大小的 SRAM 中,才能被计算核心处理。

标准实现的逻辑陷阱

PyTorch 等深度学习框架的“朴素”实现中,Attention 的计算过程被拆分成了多个独立的算子(Op)。这导致了一个严重的效率问题:

  1. 第一步( 𝑄 𝐾 𝑇 QK T ):GPU 把 𝑄 Q 和 𝐾 K 从 HBM 搬到 SRAM,计算出分数矩阵 𝑆 S。
  2. 中间结果过大:由于 𝑆 S 的形状是 ( 𝑁 , 𝑁 ) (N,N),对于长序列来说,这个矩阵大到 SRAM 根本存不下。
  3. 被迫回写:GPU 只能被迫将这个巨型矩阵 𝑆 S 从 SRAM “踢”出去,写回到慢速的 HBM 中。
  4. 反复折腾:到了下一步计算 𝑆 𝑜 𝑓 𝑡 𝑚 𝑎 𝑥 ( 𝑆 ) Softmax(S) 时,GPU 又得重新去 HBM 把刚才存进去的 𝑆 S 再搬回 SRAM。
# 标准 Attention 实现的 IO 噩梦
def standard_attention(Q, K, V):
    # 1. HBM -> SRAM(计算) -> HBM(存 S)
    S = Q @ K.T
    # 2. HBM(读 S) -> SRAM(计算) -> HBM(存 P)
    P = softmax(S)
    # 3. HBM(读 P) -> SRAM(计算) -> HBM(存 O)
    O = P @ V
    return O

这种 “搬进来 -> 算一下 -> 踢出去 -> 再搬回来” 的反复 I/O 往返,就是性能最大的杀手。

SRAM 与 HBM 的带宽差异

你可能会问,存回 HBM 再读回来会有多大影响?

速度差异对比

存储类型 容量示例 (A100) 带宽 速度比喻
SRAM (共享内存) ~20 MB ~19 TB/s F1 赛车 🏎️
HBM (显存) 40~80 GB ~1.5 TB/s 普通轿车 🚗

SRAM 的带宽通常比 HBM 高出约 10 倍以上。

瓶颈本质:IO 受限 (IO-bound)