Flash Attention 原理详解
会员专享通过交互式可视化,深入理解 Flash Attention 的核心技术:内存瓶颈、Online Softmax、与分块矩阵乘法。
标准 Attention 的内存瓶颈
在深入 Flash Attention 的代码实现之前,我们必须先回答一个底层问题:为什么标准的注意力机制公式 在现代 GPU 上跑得还不够快?
登录以继续阅读
这是一篇付费内容,请登录您的账户以访问完整内容。
通过交互式可视化,深入理解 Flash Attention 的核心技术:内存瓶颈、Online Softmax、与分块矩阵乘法。
在深入 Flash Attention 的代码实现之前,我们必须先回答一个底层问题:为什么标准的注意力机制公式 在现代 GPU 上跑得还不够快?
这是一篇付费内容,请登录您的账户以访问完整内容。
首先我们需要建立一个极其重要的概念:在 GPU 架构中,所有的计算(如矩阵加法、乘法)都必须在靠近核心的 SRAM(Shared Memory,共享内存) 中进行。
这意味着:哪怕你的显存(HBM)有 80GB 那么大,数据也必须先被“搬运”到几十 MB 大小的 SRAM 中,才能被计算核心处理。
在 PyTorch 等深度学习框架的“朴素”实现中,Attention 的计算过程被拆分成了多个独立的算子(Op)。这导致了一个严重的效率问题:
GPU 把 和 从 HBM 搬到 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 =
这种 “搬进来 -> 算一下 -> 踢出去 -> 再搬回来” 的反复 I/O 往返,就是性能最大的杀手。
你可能会问,存回 HBM 再读回来会有多大影响?
| 存储类型 | 容量示例 (A100) | 带宽 | 速度比喻 |
|---|---|---|---|
SRAM (共享内存) | ~20 MB | ~19 TB/s | F1 赛车 🏎️ |
HBM (显存) | 40~80 GB | ~1.5 TB/s | 普通轿车 🚗 |
SRAM 的带宽通常比 HBM 高出约 10 倍以上。
由于这个速度鸿沟的存在,如果算法不断在 HBM 和 SRAM 之间搬运中间数据,就会出现一种尴尬的局面:
GPU 强大的计算核心大部分时间都在“空转”,苦苦等待数据从缓慢的 HBM 运送过来。
这种状态被称为 I/O 受限(IO-bound),即计算能力被内存传输速度拖了后腿。
量化感受一下:在 standard Attention 中,读写 的中间矩阵 和 所花费的时间,远远超过了实际进行矩阵乘法计算的时间。
既然 SRAM 这么快,为什么不把整个 Attention 矩阵都放在 SRAM 里算完再走?这里涉及到了物理与经济的刚性制约:
SRAM 的存储密度极低,导致成本极高。根据相关资料(如 FlashAttention 论文引用的背景):
80GB 容量的 SRAM 存储器,成本可能高达 13,000 美元(估算数量级)。HBM 仅需 2,000 美元。在实际硬件中,A100 的 HBM 可以达到 80GB,但 SRAM 通常只有 192 KB / SM(每个流式多处理器)。由于这种容量限制,你无法一次性把整个 的注意力矩阵塞进 SRAM。当序列长度 增加时,中间矩阵的大小呈 爆炸式增长。
Flash Attention 的核心逻辑就在于:既然 SRAM 贵且小但快,HBM 大且便宜但慢,那么我们就必须放弃“全量读写”的幻想。
我们需要引入两个核心思想:
SRAM 的“小方块”。SRAM 内部,一气呵成完成 、Softmax 和 的乘法,只在最后一步将最终结果 写回 HBM。通过这种方式,我们根本不生成(也不写入 HBM)那个巨大的 中间矩阵。
| 方法 | HBM 读写量 | 复杂度 |
|---|---|---|
Standard Attention | 随着序列变长,IO 爆炸 | |
Flash Attention | 线性增长,极大节省带宽 |
SRAMGPU 只能被迫将这个巨型矩阵 从 SRAM “踢”出去,写回到慢速的 HBM 中。GPU 又得重新去 HBM 把刚才存进去的 再搬回 SRAM。