FlashAttention -1 原理
在介绍FlashAttention之前,先介绍 LLM 推理的两个阶段,一个是 prefill 阶段,一个是 decode 阶段。
假设输入为[Batch_size=1,句子长度len,模型隐藏维度d],数据类型为FP16
标准 Attention 在做什么
首先,计算计算
将结果写入 HBM ,然后应用一个 softmax
从 HBM 中读取 S ,再将结果 P 写入 HBM
然后从 HBM 读取 P ,再将 O 写入 P。整个过程中读入了[Q,K,S,P,V] 写入了 [S,P,O]。每次矩阵乘法是一次加乘,即两次浮点运算,一共执行了
prefill与decode阶段的区别
在 prefill 阶段,我们会把整段 prompt 交给 LLM ,此时模型得到 KV cache,输入是一个
以A100为例,显存宽带为 2 TB/s ,浮点性能为 312 TFLOPS ,当模型计强度超过
在 prefill 阶段,每一层都输入一个完整的句子,可以得到计算强度为
以往的博客告诉我们, prefill 是 computed-bound 的,而 decode 是memory-bound 的,然后会告诉我们下面这些东西。
当
而在 decode 阶段,输入长度变为
实际的具体情况还是要具体分析。
假设
如果我们使用了 chunked_prefill 技术,则也需要考虑这上面的带宽问题。
FlashAttention 问题的背景
看它解决了什么问题,就得以当时的视角看问题是怎么样的。
如果 attention 是一个标注的矩阵乘法,则上面的问题可以忽视
实际上,它还有 softmax ,normalization 以及 scaling 等,同时期还有其他的attention计算方式,例如 sparse attention,linear attention,kernel-based attention 等。它们尝试将算法的时间复杂度降低,但是实际运行起来也没比普通 attention 快多少。
于是尝试在带宽上优化attention。
显然,以上访问出现了中间结果S与P,它被计算并存储,然后重新调用,如果我们一直放在SRAM里面,就能减小访存次数。
Tiling
如果我们想让中间结果一直在SRAM上面,那么我们就需要SRAM一直去存储它们,假设这些数据类型都一样,可以存储
回到上面的三次计算,我们看看每个时间段需要哪些变量参与
| 步骤 | Q | K | V | S | P | O |
|---|---|---|---|---|---|---|
| 1 | 1 | 1 | 1 | |||
| 2 | 1 | 1 | ||||
| 3 | 1 | 1 | 1 |
每个时刻至少需要3份,与此同时,QK其中一个可以一直放在内存块里面,所以简单来看我们可以分成四份。
假设我们取
为了确保能分成四份,首先令
这就是 FlashAttention 分块的思路,将 SRAM 分成四份,然后中间变量不写入 HBM。
此时继续计算,会遇到几个问题:
- softmax公式为
,在计算attention的时候,需要一整行的元素,但是我们已经将矩阵分块了,S得到的仅仅是一部分结果。 - 如果仅做推理,则不需要用到 P这个中间结果,但在训练及反向传播的时候,还是需要这些中间变量。
下面我们来解决这些问题
online Softmax
普通的softmax公式为
在 attention 中,这个是逐行计算的,为了防止过大数造成精度损失或者其他情况,根据平移不变性,我们要减去这一行最大值
问题就出现在这里,我们无法一下取得这一行的最大值。
假设我们这一行的attn被分成两块计算。
并且已经得到两块的最大值
先合并最大值
然后合并指数和
这样,在计算 softmax 的时候,公式为
现在完成了理论基础,这里还没完,对于前面已经产生的输出
假设在之前计算中,我们得到了这一行每个块的
首先会更新全局最大值
然后更新 softmax 分母
接下来更新
当前贡献是
于是,我们就解决了这些问题。不考虑drop与mask,我们得到了它的伪代码

Recomputation
在反向传播的时候,我们有两个中间变量




