FlashAttention:注意力计算的 IO 优化
- Published on
FlashAttention:注意力计算的 IO 优化
很多人第一次看到 FlashAttention,会以为它是一个更快的注意力公式。其实它并没有改变标准注意力的数学结果,真正改变的是计算过程:尽量少把中间结果在 GPU 的不同存储层级之间来回搬运。
这件事听起来不像模型创新,却非常关键。大模型运行时,GPU 不只是“算得快”,还要不断从显存取数据、把中间结果写回显存。计算单元在等数据时,即使理论算力很高,也只能闲着。FlashAttention 研究的就是这段经常被忽略的等待。
我以前优化推理速度时,第一反应也是换更大的 GPU,后来才发现,有些场景不是算力不够,而是内存访问太浪费。把数据少搬几次,可能比单纯增加计算单元更有效。
一、标准注意力为什么占内存
对输入计算 Query、Key、Value 后,标准注意力大致是:
S = Q × Kᵀ
P = softmax(S)
O = P × V
其中 S 和 P 的形状通常包含序列长度的平方。序列从 1K 增加到 4K,注意力矩阵的元素数量理论上增加 16 倍。即使最终输出不大,中间矩阵也可能迅速占满显存。
type AttentionShape = {
batch: number
heads: number
sequenceLength: number
headDimension: number
}
function attentionMemory(shape: AttentionShape, bytesPerElement: number) {
return (
shape.batch *
shape.heads *
shape.sequenceLength ** 2 *
bytesPerElement
)
}
问题不只是“矩阵很大”。标准实现通常会先把 S 写入显存,再读回来做 Softmax,再写入 P,最后读 P 和 V 做乘法。中间结果在高带宽显存和片上高速存储之间来回移动,IO 成本很高。
二、GPU 有很多层存储
可以把 GPU 存储粗略理解成一座仓库:
寄存器:最快,容量最小
共享内存 / SRAM:很快,容量有限
显存 HBM:容量大,但访问相对慢
主机内存:更大,但跨总线访问更慢
高效 Kernel 的目标,是把当前正在计算的那一小块数据放进更快的存储里,尽量避免反复访问大显存。FlashAttention 的核心思路就是 Tiling:不构造完整的 S 和 P,而是把 Q、K、V 分块,边读边算、边归约、边写出结果。
三、分块计算怎么保持正确结果
如果不保存完整的 Softmax 矩阵,就必须在线计算 Softmax。Softmax 需要全局最大值和归一化分母,分块后可以使用在线归约的方式维护中间状态:
当前块最大值 m
当前块归一化统计 l
当前块输出累积 o
当读到下一块时,先更新最大值,再按比例修正之前的统计和输出。最终得到的结果与一次性计算完整矩阵等价或在浮点误差范围内一致。
这也是 FlashAttention 难的地方:不是把矩阵切成小块这么简单,还要正确处理 Softmax 的数值稳定性、边界块、因果 Mask 和不同序列长度。
四、因果 Mask 和变长序列
语言模型通常只能看到当前位置之前的 Token,需要使用因果 Mask:
第 1 个 Token:只能看自己
第 2 个 Token:能看前两个
第 3 个 Token:能看前三个
高效 Kernel 会在分块时跳过不可能参与计算的上三角区域,减少无效工作。变长序列还需要处理 Padding 和不同样本的有效长度,不能为了一个长样本让整个 Batch 都按最大长度浪费计算。
因此,FlashAttention 的收益和 Batch 形状、序列长度、Mask 类型、数据类型及硬件有关。不能看到论文中的加速倍数,就认为所有任务都会得到同样收益。
五、它优化的是 IO,不是把复杂度变成线性
标准全注意力仍然有序列长度平方的计算关系。FlashAttention 主要减少中间矩阵的显存读写和峰值占用,并通过更好的 Kernel 提高实际吞吐。
这意味着它不能单独解决无限长上下文问题。序列特别长时,计算量本身仍然很大;如果业务真正需要超长上下文,还需要稀疏注意力、滑动窗口、分层摘要、检索或其他结构。
把 FlashAttention 说成“让注意力变成线性”是不准确的。它让同样的注意力计算更接近硬件的高效路径,但没有改变所有数学复杂度。
六、为什么实际速度可能没有论文那么快
常见原因包括:
- GPU 架构不匹配,Kernel 没有走最佳路径;
- 输入太短,Kernel 启动和调度开销占主要部分;
- Batch 太小,计算单元没有充分利用;
- 数据类型、Mask 或布局触发了回退实现;
- 端到端瓶颈其实在 Tokenizer、通信或采样;
- 测试只比较了 Kernel 时间,没有计算模型加载和数据传输。
性能测试要分层:单独测注意力 Kernel,再测单层,再测完整模型,最后测真实服务的首 Token、生成速度、P95 和显存。一个局部 Kernel 快 2 倍,不代表用户看到的响应也快 2 倍。
七、使用时先确认版本和回退路径
不同框架和版本对 FlashAttention 的支持方式不同。启用后要确认:
type AttentionRuntime = {
backend: 'flash' | 'math' | 'memory_efficient'
dtype: 'fp16' | 'bf16' | 'fp32'
causal: boolean
fallbackReason?: string
}
生产环境要记录实际使用的后端。如果某些输入触发回退,团队应该知道,而不是只看配置文件里写着“flash”。回退是正常能力,但回退比例过高就需要重新评估。
数值差异也要进入测试。不同 Kernel、精度和硬件会有微小浮点差异,对生成结果可能产生放大影响。关注任务指标和输出稳定性,不要要求每个中间浮点数逐位相同。
八、和量化、KV Cache 一起看
FlashAttention 主要优化注意力计算,量化主要减少权重和部分 Cache 的内存,KV Cache 则影响自回归生成阶段的历史状态占用。三者解决的问题不同,却共同影响显存和吞吐。
权重内存 → 量化
Prefill 注意力 IO → FlashAttention
Decode 历史状态 → KV Cache 优化
部署时要分开测 Prefill 和 Decode。长 Prompt 的 Prefill 可能受注意力 IO 影响更明显;逐 Token Decode 可能更受 KV Cache 读取和内存带宽影响。只测一种阶段,容易把优化收益看错。
九、从工程角度如何判断值得使用
先回答三个问题:
- 当前瓶颈确实是注意力 Kernel 和显存 IO 吗?
- 目标硬件和框架能稳定支持吗?
- 端到端质量、延迟和显存是否都改善?
如果模型很小、序列很短,切换复杂 Kernel 可能不值得;如果服务使用长上下文、高并发和现代 GPU,收益通常更有价值。任何性能优化都应该有基线、实验条件和回滚方案。
总结:性能优化常常是少搬数据
FlashAttention 给我的最大启发,是性能问题不能只用“算力不够”解释。现代 GPU 计算很快,数据移动反而可能成为瓶颈。标准注意力把巨大的中间矩阵反复写入和读出,FlashAttention 通过分块、在线 Softmax 和片上存储复用,减少了这部分浪费。
它没有改变注意力的数学定义,也没有让任意长度上下文突然变成线性;它做的是让算法更贴近硬件真实的存储层级。理解这一点,比记住一个加速倍数更重要。
我现在做性能优化,会先问数据去了哪里、搬了几次、哪一层存储在等待,而不是马上购买更大的机器。量化、FlashAttention、KV Cache 和批处理各自解决不同瓶颈,只有把它们放到真实模型、真实输入和真实硬件上测量,优化才不是漂亮的实验室数字,而是用户确实能感受到的更快、更稳和更省。