通过交互式可视化,深入理解 Flash Attention 的核心技术:内存瓶颈、Online Softmax、与分块矩阵乘法。
在深入 Flash Attention 的代码实现之前,我们必须先回答一个底层问题:为什么标准的注意力机制公式 𝑆
𝑜
𝑓
𝑡
𝑚
𝑎
𝑥
(
𝑄
𝐾
𝑇
)
𝑉
Softmax(QK
T
)V 在现代 GPU 上跑得还不够快?
首先我们需要建立一个极其重要的概念:在 GPU 架构中,所有的计算(如矩阵加法、乘法)都必须在靠近核心的 SRAM(Shared Memory,共享内存) 中进行。
这意味着:哪怕你的显存(HBM)有 80GB 那么大,数据也必须先被“搬运”到几十 MB 大小的 SRAM 中,才能被计算核心处理。
在 PyTorch 等深度学习框架的“朴素”实现中,Attention 的计算过程被拆分成了多个独立的算子(Op)。这导致了一个严重的效率问题:
# 标准 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 往返,就是性能最大的杀手。
你可能会问,存回 HBM 再读回来会有多大影响?
| 存储类型 | 容量示例 (A100) | 带宽 | 速度比喻 |
|---|---|---|---|
| SRAM (共享内存) | ~20 MB | ~19 TB/s | F1 赛车 🏎️ |
| HBM (显存) | 40~80 GB | ~1.5 TB/s | 普通轿车 🚗 |
SRAM 的带宽通常比 HBM 高出约 10 倍以上。