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