一句话:FlashAttention 是一种「省着用显存带宽」的注意力(attention)计算方法——它算的东西和标准注意力一模一样,只是坚决不把中间那张巨大的矩阵搬进搬出显存。
原理:慢的原因往往不是算得慢,而是搬得慢
注意力是这样算的:Q 乘 K 得到分数矩阵,做 softmax,再乘 V 得到输出。麻烦出在中间那两张 N×N 的矩阵——句子长度翻倍,它们就大四倍。GPU 的显存(HBM)容量有限,带宽更有限;而真正干活的算力单元在贴着它的片上缓存(SRAM)里,又快又小。
打个比方:你在厨房做菜,灶台(算力)出菜极快,但冰箱在另一个房间(HBM)。标准做法是每加一味调料都跑一趟冰箱,灶台大部分时间在等人。FlashAttention 的做法是把食材提前切好分份,一次性搬到灶台边的小台面上,在小台面上把这道菜做完再端走——中间过程从不回冰箱。
技术上这叫分块(tiling)加在线 softmax(online softmax):把 Q、K、V 切成小块放进 SRAM,一边算一边维护「当前见过的最大值」和「当前归一化因子」,逐块修正已有结果。于是那张 N×N 矩阵根本不需要被完整写出来。
和相邻概念的区别
一句话说清:FlashAttention 不减少计算量(FLOPs),只减少显存读写(IO)。这是它最容易被误解的地方——它不是近似算法。
| 维度 | 普通注意力实现 | FlashAttention | 稀疏/低秩等近似注意力 |
|---|---|---|---|
| 数学结果 | 基准 | 完全一致(精确) | 改变,属近似 |
| 是否写出 N×N 中间矩阵 | 是 | 否 | 通常否 |
| FLOPs | 基准 | 基本不变 | 一般减少 |
| 额外显存 | O(N²) | O(N) | 视方案而定 |
它和 KV cache、PagedAttention 也不是一回事:后两者管的是推理时缓存怎么存,FlashAttention 管的是每次计算怎么读写。训练做反向传播时,它不保存 N×N 矩阵,而是重算一遍,用算力换显存。
对从业者的意义
- 直接好处是长上下文变得可行:显存开销随序列长度线性增长,而不是平方,更长的文本、更高分辨率的图像、更长的音频都能塞进去。
- 它通常被封装进框架底层的注意力算子,你调用常规接口时可能已经在用;是否真正生效取决于硬件、数据类型和具体实现,以官方页面为准。
- 更普适的启发:GPU 利用率上不去时,先怀疑数据搬运,而不是先怀疑算法。很多性能优化空间就藏在「少搬一次」里。
