FlashAttention 的基本原理
注意力机制无疑是现代深度学习架构中最重要的基石之一。它被应用于从自然语言处理到计算机视觉等各种任务的大多数先进模型中。然而,注意力机制也是这些模型中成本最高的操作之一。因此,自然而然地,人们进行了大量研究,致力于使其运行速度更快、内存利用率更高。其中大多数研究都基于对注意力机制的近似,这可能会导致准确性的损失。
我们将深入了解 FlashAttention 的细节,看看它如何实现 7.6 倍的大幅加速,以及如何在计算精确的注意力分数的同时实现O(N)内存复杂度。
1. 第一部分的内容
FlashAttention 的第一原则
在本系列的第一部分中FlashAttention 的第一原则,我们为理解 FlashAttention 论文奠定了基础。
我们对注意力机制有了基本的了解,包括它的工作原理以及在模型中如何找到它。我们还简要讨论了现代 GPU、CUDA 编程模型以及 GPU 的内存层次结构。
然后,我们转换思路,研究如何在 GPU 上将两个矩阵相乘,以及如何通过利用 GPU 的内存层次结构(共享内存)、在单个线程中计算多个结果(块平铺)以及将多个内核融合为一个内核(内核融合)来优化此过程。
然后我们得到了下面的图片,它说明了注意力机制的主要问题:我们读取和写入巨大的NxN矩阵来获取中间结果,这不仅花费大量时间,而且还需要大量内存。
我们最终弄清楚了是什么阻止我们将相同的技术应用于优化矩阵乘法时使用的注意机制:
SoftMax 操作阻止我们融合内核,因为它对整个向量 x 进行操作,而不是对其中的部分进行操作, 我们需要反向传递中的这些中间结果来计算梯度。
理解这两点需要付出不少努力,所以如果你能坚持到这里,恭喜你!🎉
现在和我一起进入旅程的下一部分:使用 FlashAttention 解决这些问题!
2. FlashAttention
论文标题“ FlashAttention:FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness(快速且内存高效的精确注意力机制与 IO 感知)”非常具体,并且已经告诉了我们很多它所解决的问题:
快速:执行时间大幅缩短。注意力层本身缩短了 7.6 倍,从头开始训练 GPT-2 缩短了 3.5 倍。 内存高效:内存占用大幅降低,使我们能够在更大的上下文和更大的批次上训练更大的模型。这是因为标准注意力机制的内存复杂度为*O(N²) ,而 FlashAttention 的内存复杂度为**O(N)*。 精确注意力:与Linformer、Performers或Reformer等其他方法相比,FlashAttention 通过近似注意力机制的各个部分来实现加速和减少内存需求,而 FlashAttention 则在计算精确注意力分数的同时实现了减少内存需求。 IO感知:FlashAttention 利用现代 GPU 的内存层次结构来优化不同内存级别之间的数据传输。具体方法是使用共享内存存储中间结果,并仅将最终结果写入全局内存。
将这些收益可视化如下:
虽然对于大多数人来说,这一切仍然显得相当抽象,但我们可以明确地说明它可以做什么:
更快地训练相同的模型, 在更大的批次上训练相同的模型, 以相同的成本训练更大(可能性能更好)的模型, 训练具有更大上下文窗口的模型, 并在较小的 GPU 上训练模型。
例如,在 8x A100 上训练 GPT-2 small 花费了 2.7 天而不是 9.5 天,在 OpenWebText 数据集上训练 GPT-2 medium 花费了 6.9 天而不是 21.0 天。
❓那么我们究竟如何达到这一点以及我们需要解决哪些问题?
2.1. 我们需要解决什么问题?
正如我们在第一部分FlashAttention 的第一原则中了解到的,我们的主要问题是需要在中间结果中读写巨大的NxN矩阵,这不仅耗费大量时间,也占用大量内存。理想情况下,如果我们可以将 SoftMax 核与矩阵乘法核融合,并仅将最终结果写入全局内存,就可以避免这种情况。
但我们遇到了两个阻碍我们这样做的问题:
从第一个问题开始:SoftMax 操作。
2.2. 修复 SoftMax
回想一下第一部分FlashAttention 的第一原则,为了将两个矩阵相乘,我们利用了共享内存。输入被分割成更小的块,并将结果累加起来得到最终结果:
如果我们想将 SoftMax 核与矩阵乘法核融合,我们还必须对输入进行分块,并以某种方式将它们组合起来以获得最终结果。这就是 SoftMax 运算分解发挥作用的地方。
首先假设我们可以访问计算 SoftMax 所需的所有值,并将 SoftMax 定义转换为与论文中使用的更类似的定义:
到目前为止,这只是对标准 SoftMax 定义的重新表述。
❗ 现在可能是本系列最重要的部分:我们将看到如何分解 SoftMax 操作,以便它可以一次应用于各个块,这将允许我们将 SoftMax 核与矩阵乘法核融合,因为它们将对相同的数据进行操作。
我们想要计算 SoftMax 的并不是单个向量x ,而是共享内存中可用的x的一半。
要非常清楚:我们不能简单地计算 x 的前半部分和 x 的后半部分的 SoftMax,因为我们需要知道整个向量 x 的最大值才能计算数值稳定的 SoftMax。
相反,我们将执行以下操作:我们跟踪一些额外的指标,然后对于每次迭代,我们计算x的当前块的 SoftMax ,我们将当前迭代的结果与之前的迭代相结合,更新指标,然后相应地更新输出。
如果这还不能告诉你什么,别担心!我们会详细介绍。
首先,让我向你展示当将 SoftMax 运算应用于向量x的两个块时的新方程:
我知道这看起来很复杂。但如果你算一下,并代入数值,你会发现它等价于我们之前看到的标准 SoftMax 公式。为了说明这一点,我深入研究了l(x) ,并为你做了以下练习:
❓但是,我们从中得到了什么呢?
关键在于,我们现在可以分别计算x的第一部分和第二部分的SoftMax(这两个部分相对于整个向量x来说都是错误的),重新调整中间结果,然后将它们组合起来,得到x的最终有效 SoftMax 。仔细观察就会发现,函数m(x)和l(x)需要知道x的所有值。因此,我们在全局内存中跟踪它们并迭代更新。我们这样做有两个原因:
我们需要它们在后续迭代中更新和重新调整中间结果 在计算梯度时,我们不想为后向传递重新计算它们。
好吧,这仍然相当抽象。而且,我们讨论了很多关于迭代的内容……所以,现在是时候更具体一点,向你展示 FlashAttention 中前向传播的实际算法了。
2.3 前向传播
在深入细节之前,我认为有必要先看一下前向传播的高层执行过程,并说明如何迭代数据。这将帮助我们理解算法的大部分内容。
动图有些粗糙,我们关注:
我们对输入进行两次循环: 外循环对以j为索引的K和V块进行迭代 内循环对以i为索引的Q块进行迭代 我们对输出矩阵O进行多次迭代,在外循环的每次迭代中更新其值。 我们还跟踪并迭代更新SoftMax 计算中的指标m(x)和l(x) 我们永远不需要将NxN矩阵写入全局内存,而是将它们的小块保存在共享内存中,并且只将最终结果写入全局内存。 每个块包含与多个标记对应的多行,因此我们并行计算多个 SoftMax 操作。
理解这个动画意味着你理解了算法的核心,并且可能开始理解它如何节省时间和内存。虽然有些人可能已经了解了这些,但我认为更深入地探究其中的原理是值得的。
继续仔细看看这些分块的中间结果,以了解如何为它们计算 SoftMax。
内循环迭代索引为i的行,而外循环迭代索引为j的列。回想一下,SoftMax 是针对每一行(代表输入序列中的一个标记)计算的。因此,如果我们处理一个数据块,则需要同时计算多个 SoftMax 操作。此外,随着外循环的推进,我们需要用前一个数据块x(1)的 SoftMax 更新当前数据块x(2) 的SoftMax 。
有了这些知识,我们终于可以看看 FlashAttention 论文中提出的实际算法了(我们将在稍后介绍初始化和迭代更新的细节):
尝试将算法与上面的动画进行匹配。另外,值得一提的是,该算法融合了两个矩阵乘法核和 SoftMax 核,并且没有将NxN矩阵写入全局内存(又称 HBM)。
你现在可能会问自己:如果没有中间结果,我们该如何在反向传播中计算梯度?我们稍后会讨论这个问题。但首先,我们先完成正向传播。
仔细看看步骤 6-11,其中我们计算 x 的当前块的 SoftMax 并相应地更新指标m(x) 和 *l(x)*。
此时,我们已经计算了Q和K的矩阵乘法、当前块所需的度量 m(x) 和 *l(x)*,并将它们与前一个块的度量相结合以获得更新的度量。
现在,我们可以用当前的x块和更新后的指标来更新输出矩阵O。这在步骤 12 中完成。不过,这步有点棘手,因为它已经将 SoftMax 核与P和V的矩阵乘法核融合在一起了。所以,详细看看这一步:
diag(vector) 构造一个矩阵,对角线上为给定向量,其他位置为零。这是一个技巧,将矩阵的每一行与向量中的对应值相乘。我们这样做是因为我们分别计算了块的每一行的 SoftMax,因为它们对应于单个 token。
说实话,这种递归更新加上融合核的结合,几乎让人无法理解😅。不过,如果你代入数值,进行计算,就会发现一切都完美地融合在一起了。
我们尚未涉及的是算法第 2 步中度量m(x)和l(x)以及输出矩阵O的初始化。
由于O、m(x)和l(x)的更新是递归执行的(这意味着我们会重用之前迭代的结果),我们需要考虑一个基准情况,以防没有之前的迭代。因此,我们将O和l(x)初始化为零,将m(x)初始化为负无穷大。如果你进行计算并将初始值代入分解后的 SoftMax 方程,你会发现我们得到了应用于x第一个块的标准 SoftMax 运算。
剩下的就是运行整个算法并输出最终的Nxd矩阵O。
这就是前向传播!
现在,来讨论一下大多数博客文章中经常被忽略的内容:反向传播。在这里,我们计算梯度并更新模型参数。
2.4 反向传播
通常情况下,反向传播并不是什么大问题,那么为什么我们需要在这里讨论它呢?原因有二:
我们缺少计算梯度所需的前向传递的中间结果,因为我们省略了保存它们以节省内存和时间。 在反向传播中,我们遇到的问题与前向传播相同:我们不想实现大型的NxN矩阵。因此,我们需要一个在共享内存上运行的融合核。
不过,好消息是,该算法与前向传递非常相似,我们分别在内循环和外循环中迭代相同的数据块。
那么,那些缺失的中间结果怎么办呢?其实,我们可以在反向传播中即时计算它们。记住,我们的 GPU 核心比从全局内存读取和写入数据所需的时间要快得多。所以,即使我们像在正向传播中那样低效地重新计算相同的中间结果,也仍然比将它们写入全局内存要快。
实际上,我们不需要重新计算所有内容:还记得我们在全局内存中跟踪了指标 m(x) 和 l(x) 吗?现在我们可以用它们来计算当前x块的 SoftMax ,而且这次它从一开始就是正确的!
在开始算法之前,先尝试了解我们需要计算什么。
我们需要计算损失函数相对于模型参数的梯度。这可以通过应用链式法则来实现,每次计算一个梯度,直到得到损失函数相对于输入矩阵Q、K和V 的梯度。所以,正如我的教授常说的:“这只是一种练习,而不是惩罚”,把这些方程式写下来。
我们马上就会看到,这些等式中有一些棘手的地方,我想明确指出:
计算 SoftMax 操作的梯度时需要小心,因为 SoftMax 是按行计算的,而不是针对整个矩阵。这意味着我们需要分别计算输出矩阵O的每一行的梯度! 这个术语可能有点令人困惑,因为损失L相对于某个其他值(假设V)的梯度由𝜕L/𝜕V给出,缩写为dV。 最后,我们再次需要迭代阻塞的图块,所以我们也需要考虑这一点。
但是,先从“简单”的链式法则开始,看看我们需要计算哪些项。回想一下,我们感兴趣的是dV、dK和dQ。
现在,逐层地处理这些梯度。我们将从O=PV的矩阵乘法开始。
这样,我们已经计算出了第一个感兴趣的梯度:dV。可能最复杂的部分是理解如何对矩阵求梯度,尤其是转置从何而来,以及为什么dO是从左乘还是从右乘。如果你扩展了表达式,并计算了矩阵每个元素的导数,你就会明白转置从何而来。根据经验法则,如果我们在乘法中对第一个矩阵求梯度,则上游梯度dO(这是数学术语,指的是损失函数对输出矩阵O的梯度)从左乘以,如果我们在乘法中对第二个矩阵求梯度,则上游梯度 dO(这是数学术语,指的是损失函数对输出矩阵 O 的梯度)从右乘以。
计算dK和dQ需要我们通过 SoftMax 以及Q与K的矩阵乘法进行反向传播。我们继续讨论 SoftMax。回想一下,SoftMax 是针对每一行i计算的,而不是针对整个矩阵。索引表示,i:表示第i行对应所有列j。
利用一些代数技巧,我们现在得到了单行梯度dS,然后通过计算所有行的梯度并将它们叠加在一起,我们最终得到了dS。计算出dS后,我们现在可以计算 Q和K矩阵乘法的梯度,从而得到我们感兴趣的缺失梯度:dK和dQ。
这样,我们就计算出了注意力机制的所有梯度dQ、dK和dV,然后我们可以进一步向后传播并更新模型参数。
在查看论文中的具体算法之前,还有最后一件事需要说明:我们需要考虑共享内存的平铺和迭代,因为在每个时间点我们只能访问一小部分数据。因此,在内循环的每次迭代中,我们都需要重新计算这些中间结果。好消息是,我们将度量m(x)和l(x)存储在全局内存中,因此我们可以用它们来计算当前x块的 SoftMax ,而无需再次迭代更新它们。
现在,我们终于可以看看反向传播的算法了。
请注意,它还涵盖了掩蔽和丢弃操作,为了简单起见,我们省略了这些操作,以便我们可以真正专注于 SoftMax 部分。
这样,我们就讲完了反向传播的整个算法。🎉
快速结束一下并总结一下我们在本系列的这一部分中学到的内容。
我很好奇其他人的想法——如果有什么突出的地方或者你有不同的看法,请随时发表评论。
如果这篇文章有用,点个赞能让更多人看到。感谢阅读!
3. 总结
FlashAttention 的基本思想很简单:我们将注意力层中的所有核融合在一起,以避免将大型NxN矩阵写入全局内存。为了实现这一点,我们必须解决两个问题:
分解 SoftMax 操作,以便将其应用于输入矩阵的各个块,并 重新计算反向传递中的中间结果,而无需实现大型NxN矩阵。
我们花时间深入研究算法的细节并了解其工作原理。