大家好,我们现在开始上课。今天的课程旨在讲解如何加快语言模型(Language Model: 用于理解和生成人类语言的AI模型)的生成速度。本课预设大家已非常清楚语言模型内部的运作原理,例如Transformer(Transformer: 一种基于自注意力机制的神经网络架构)的工作方式。在已知Transformer运作原理的前提下,我们将探讨如何加速其生成过程,而不是训练过程。换言之,本课假设一个语言模型已被训练好,我们将聚焦于如何加速其推理阶段,也就是我们常说的推论(Inference: 语言模型训练完成后,用于生成新内容或预测的过程)。
首先,我们用两页投影片快速回顾一下语言模型的内部运作。众所周知,语言模型本质上在做“文字接龙”,你给它一个未完成的句子,它预测接下来应该产生哪个token(Token: 文本在模型中处理的最小单元,可以是词、字或字符)。当前主流语言模型的神经网络架构是Transformer。Transformer由多层组成,每层都包含一个名为自注意力机制(Self-Attention: Transformer 中用于处理序列中不同位置信息的核心机制)的机制。自注意力机制使得Transformer能够考虑整个输入Sequence(序列: 输入到模型中的一系列数据,如文本中的词语顺序)中的所有信息。
自注意力机制的工作原理是:输入一排向量(例如 X1 到 X5),它将输出另一排向量(O1 到 O5)。具体运作方式是,首先将 X1 到 X5 各自乘上三个不同的Transformation Matrix,将 X1 转换为 V1、K1、Q1,X2 转换为 V2、K2、Q2,以此类推。在大多数课程中,我们通常会省略这个从 X 转换为 QKV 的过程,直接假设你已了解输入的 X 会变成 QKV 三个向量。
接下来,以生成第四个位置的输出为例进行计算(其他位置的计算方式相同)。我们将第四个位置的查询向量(Query: 在自注意力机制中,用于查询其他所有键的向量)q4 与前面所有的键向量(Key: 在自注意力机制中,用于与查询向量进行匹配的向量)k1 到 k4 进行点积(Dot Product: 向量之间的一种乘法运算),计算它们的内积。这些内积计算结果我们记作 a1 到 a4。由于这些内积值可能从负无穷大到正无穷大,我们会通过Softmax函数(Softmax Function: 将任意实数向量压缩为概率分布的函数)对其进行归一化(Normalization: 将数据缩放到特定范围,使其符合某种分布),使其数值介于 0 到 1 之间且总和为 1。经过Softmax后的注意力权重(Attention Weight: 表示输入序列中各部分对当前输出的贡献程度)我们记作 A1 hat 到 A4 hat。我们将 A 视为点积后得到的注意力权重,而 A hat 表示经过归一化处理的注意力权重。有了 A1 Hat 到 A4 Hat 后,每个 A Hat 都会乘上其对应的值向量(Value: 在自注意力机制中,承载实际信息的向量),然后进行加权和(Weighted Sum: 各项乘以其权重再求和),这就是注意力层在第四个位置的最终输出。以上这些都是我们预设你已经了解的Transformer基础知识。
加速Transformer:代价与策略
今天,我们将探讨一系列加速Transformer运算的方法。在评估任何加速方法时,我们需要问:“其代价是什么?”任何号称能加速的方法,背后往往都付出了某种代价,即用你可能不那么关心的资源,来换取Transformer的加速。常见的代价包括:
- 改变自注意力计算:某些方法可能并非真正计算自注意力,而是一种近似(Approximation),导致计算结果与原始自注意力有所不同。
- 模型绑定:有些方法是模型绑定的,要求你必须训练特定的模型或对模型进行定制化(Customization)才能使用,并非即插即用(Plug-and-Play)的方法。
- 其他代价:如果以上两点都不是代价,那么它可能付出了其他形式的代价来实现加速。
我们将通过一个表格来总结这些方法。今天课程的重点在于详细解释Flash Attention这项技术。事实上,现在大多数人在使用语言模型时,即便不自知,也已经在利用Flash Attention。
此外,第二部分会介绍一系列与KV Cache(Key-Value Cache: 存储Transformer层中Key和Value向量以加速推理的技术)相关的方法,虽然方法众多,但每个方法都会快速带过。有一个今天不会讲的有效加速推理方法是推测解码(Speculative Decoding: 一种通过草稿模型预生成序列,再用大模型验证以加速推理的方法),因为它已在之前的课程(2024年生成式AI导论第16讲)中讲解过,并包含在作业中。
Flash Attention核心机制:GPU存储层级与计算瓶颈
接下来,我们将深入探讨Flash Attention。这项技术最早来源于2022年的一篇论文,可谓是“上古时代”的技艺,但其威力巨大。Flash Attention不会改变注意力机制的计算结果,其输出与传统方法完全一致,并非近似。更重要的是,它是一个即插即用的方法,可以直接应用于任何使用自注意力机制的Transformer,不与特定模型绑定。它的代价非常小。
Flash Attention之所以能实现加速,其核心思想在于它考虑了GPU(Graphics Processing Unit: 图形处理器,擅长并行计算)运算的底层逻辑。我们需要先了解GPU运算时发生了什么,以及需要关注的关键部分。需要强调的是,以下描述并非极其精确,而是为了便于理解而进行的简化,且并非只针对GPU,许多高速运算都有类似概念。
在GPU中,存在大量负责运算的执行单元(Execution Unit: GPU内部负责实际计算的部件)。这些执行单元并非单一存在,而是数量众多,可以想象成拥有多个分身,能同时执行任务。它们强大之处在于数量多,运算速度极快。然而,它们的弱点是工作台(即SRAM,Static Random-Access Memory: 静态随机存取存储器,GPU的工作台,速度快容量小)太小,能存放的数据有限。如果你有一个非常大的矩阵或向量,无法直接放在工作台上,每次只能放置少量数值进行运算。
大量的数据则存放在一个仓库(即HBM,High Bandwidth Memory: 高带宽内存,GPU仓库,容量大速度相对慢)中。虽然仓库并非无限大,但相较于工作台,仓库的容量要大得多。因此,在实际运算时,执行单元需要从仓库中将数据搬到工作台上,每次只能搬运正好能放入工作台的数据。一旦数据进入工作台,运算速度极快,可在瞬间完成各种运算,然后将结果搬回仓库。所以,只要数据上了工作台,就能进行高速运算,工作台上的一切运算几乎瞬间完成。然而,搬运数据是耗时的,是拖慢运算的瓶颈。
Flash Attention的目标正是改变自注意力机制的计算算法,在不改变计算结果的前提下,通过调整数值计算顺序,减少搬运数据的次数,从而缓解这个较慢的瓶颈,使整体运算更快。
朴素Attention的计算流程与瓶颈
在讲解Flash Attention之前,我们先来看看传统注意力机制(朴素Attention)是如何运算的。我们知道查询向量(Query)、键向量(Key)等大量数值都存放在仓库中,而工作台(用青色区域表示)非常小。现在我们要计算注意力。
假设我们有一个查询向量需要与一长串键向量进行点积计算。实际上,GPU可以同时处理多个查询向量,但为简化理解,我们只考虑一个查询向量的情况(多查询向量的情况可以类推)。我们无法将所有键向量一次性放入工作台,因为键向量数量太多了。你需要想象,键向量的数量就是当前要处理的序列长度(Sequence Length)。一个语言模型作为AI代理使用时,其输入序列往往非常长,可能是上万、十万甚至百万个token。因此,我们必须将整排键向量(长度记作大L)切分成一个个块(Chunk),每个块包含N个键向量。每个块的大小取决于GPU****工作台的大小,工作台大则块可大,工作台小则块亦小。
当包含 N 个键向量的块被放入工作台,查询向量也被放入工作台后,查询向量与所有键向量的点积计算瞬间完成。这些点积结果(灰色框表示 Q 与 K 的点积)计算完成后,会被放回仓库,无法长期停留在工作台上,因为需要清出空间进行其他操作。第一个块读取进来,计算出查询向量与每个键向量的点积;接着读取第二个块,再读取第三个块,以此类推。每次可以计算出一个块中所有键向量与查询向量的点积。
Attention权重的归一化:Amax的作用
我们将每个灰色的框框记作 A_i,它代表注意力权重。但我们真正需要用于后续运算的是经过归一化(Normalized)的注意力权重,我们记作 A Hat。归一化的公式是将每个 A_i 取指数(Exponential),作为分子,分母则是所有 A_i 取指数后的数值总和。这样就能得到归一化后的注意力权重 A Hat。
在实际实现中,通常会额外加入一个 Amax 项。在取指数之前,会将 A_i 的数值减去 Amax。Amax 是在点积计算过程中,所有计算出的点积数值(A_i)中的最大值。将 A_i 减去 Amax 的目的是确保指数项的最大值是 0。因为 A_i 的数值可能非常大,如果在指数项中放入过大的数值,计算时可能会导致溢出(Overflow)或其他问题。所以,通过减去 Amax,可以确保指数项的最大值是 0。接下来,我们假设使用这个带有 Amax 的公式将 A_i 转换为 A hat。
那么,如何将 A_i 转换为 A hat 呢?你可能会想,我们能否将整个序列中所有的 A_i 一次性放入工作台,然后工作台非常强大,瞬间就能算出 A hat 并将其读出。然而,这种操作是不允许的。虽然在这个例子中只有 9 个数值(看起来像标量),但你需要注意,工作台不能存放任何与序列长度 L 相关的东西。你之所以觉得这些数值可以放入工作台,是因为在这个例子中只展示了 9 个数值。但在实际语言模型运算中,序列长度可能达到 10 万甚至 100 万。即使只是 100 万个标量,也无法放入工作台。
既然不能存放与序列长度 L 相关的东西,那么从 A_i 转换为 A hat 的过程就变得非常繁琐。首先,我们必须找到 Amax。如果我们只加载了一个块中的数值,我们根本不知道哪个是 Amax,因为我们没有遍历所有数字。所以,寻找 Amax 的方法是:一次先加载一个块中的 A_i,找出这个块中最大的 A_i,假设为 D1。然后将 D1 这个数值存放在工作台的一个区域,我们称之为 D。D 中存储着 D1,它是当前块中最大的 A_i,而非整个序列中最大的 A_i。
第一个块处理完后,保留 D1 的值,接着从第二个块中读取注意力权重。然后比较这些新读取的注意力权重与 D 中当前存储的值(D1)哪个更大。也就是说,比较上一个块中最大的值 D1 与当前读取的注意力权重 a_{n+1} 到 a_{2n}(假设每个块有 n 个 a_i)哪个更大。将最大的那个值保存为 D2。现在 D 中存储的是 D2。这个过程反复进行,总共执行大 B 次,其中大 B 等于序列长度 L 除以块的大小 N。
直到最后第 B 次操作时,比较 D 中存储的数值 D_{b-1} 与最后一个块中读取的注意力权重,将最大的存入 D 中,这个最大的数值就是 D_B。此时,D 中存储的值就是 Amax。因为我们每次都将当前看到的最大 A_i 放入 D 中,所以当所有块都遍历一遍后,D 中存储的便是 Amax。
Flash Attention简化版:一次性找出Amax与分母
找到 Amax 后,我们将 Amax 的值仍保留在工作台上。然后从第一个块开始读取,将读取到的 A_i 减去 D(即 Amax),再取指数。这样我们就得到了分子项。对每个块都执行相同的操作。我们将 Exponential(A_i - D)(其中 D 是 Amax)记作 A_i prime。
现在我们已经计算出了分子项,接下来要计算分母项。分母项需要将所有分子项,即所有 A_i prime 加起来。寻找总和的方法与之前寻找最大值的方法非常相似:每次读取一个块,将块中的所有数值加起来存入 S 中。第一次存入的数值是 S_1,即第一个块中所有 A_i prime 的总和。然后保留 S 的值,再读取第二个块中的内容。将 S 中原有的 S_1 加上当前读取的第 N+1 到 2N 个 A_i prime 的值,得到 S_2。将 S_2 存放在工作台上。这个过程反复进行,直到读取到最后一个块时,S 中存储的数值是 S_{B-1}。再加上最后一个块中所有 A_i prime 的总和,得到 S_B。这个 S_B 就是所有 A_i prime 的总和,即整个序列中所有 A_i prime 的总和。这个 A_i prime 就是 Exponential(A_i - Amax),也就是我们所需 A hat 的分母。
现在我们已经求出了分母。有了分母项和 A_i prime(分子项),接下来计算 A hat 似乎是轻而易举。我们将分子除以分母,然后读出结果,就得到了 A hat。经过这一番折腾,我们终于从 A_i 转换到了 A hat,中间已经多次读写仓库,将工作台上的数据存入仓库,这是一个非常耗时的过程。最终算出 A hat 后,任务仍未结束。此时,你需要从仓库中搬运值向量 V。将 A hat 与 V 进行加权和。每次同样只能读取一个块的数据到工作台上。将这批 V 读取到工作台,将 A hat 读取到工作台,然后计算 V 与 A hat 的加权和,得到 O。O 是我们最重要的输出,也就是 V 对 A hat 的加权和。
一开始我们只能计算第一个块的 O,所以它并非最终答案。有了第一个块的 O 后,接着将 O_1 加上第二个块中 A hat 乘以 V 的加权和。这里总共有 N 个数值,N 个加权和的结果,加上 O_1 得到 O_2。这个过程反复进行,直到最后 O_{B-1} 加上当前块中 A hat 乘以 V_i 的加权和,得到 O_B。这个 O_B 就是注意力层在某个位置的最终输出,即 Q、K、V 进行注意力计算后最终得到的结果。
在整个过程中,我们对仓库进行了多次读写。Flash Attention提出的问题是:真的需要这么多次读写吗?有没有减少读写次数的方法?我们先介绍一个Flash Attention的简化版,稍后再讲真正的Flash Attention。
简化版的想法是:从 A_i 到 A hat 的多次读取,能否减少读取次数?这个 Amax 和分母的总和项(Summation)可以一次性找出。也就是说,将 A_i 读取到内存一次,就能同时算出 Amax 和这个总和项。之前我们必须将这两项分开计算,因为总和项依赖于 Amax 的值,没有 Amax 就无法计算总和项。但神奇的是,即使总和项依赖于 Amax,也有办法同时找出 Amax 和总和项的结果。
具体做法是:首先读取一个块中的数值,找出块中最大的数值,记作 D1。我们暂时假设 D1 就是真正的 Amax。然后计算 Exponential(A_i - D1) 并将其总和记作 S1。这里我们假设 D1 就是 Amax。如果 D1 确实是 Amax,那么计算结果就是正确的。但实际上并非如此,我们来看看Flash Attention如何解决这个问题。
我们现在计算出了 D,它是这堆数值中最大的。我们计算出了 S,它是假设 Amax 被替换为 D1 时的总和。接着进入下一个块,其中包含上一个块在工作台上留下的 D 和 S 的值。我们读取第二个块的数值,并与 D 中原有的值 D1 进行比较。你可能会找到更大的数值,记作 D2。即使没有找到更大的数值,即 D1 等于 D2,也不会影响后续计算结果。这里我们假设找到了一个不同的值 D2,即当前块中有比第一个块更大的数值。
显然,我们应该将 D2 作为 Amax,而之前将 D1 作为 Amax 是错误的,因为 D2 至少比 D1 大,更接近 Amax。所以我们将 D2 作为 Amax,并计算 Exponential(A_i - D2)。然后,我们将 Exponential(A_i - D2) 在第 N+1 到 2N 个块中的每个数值的总和计算出来。接下来如何处理呢?直观的方法可能是将其与之前计算出的总和 S1 直接相加,得到 S2。但问题是,S1 是假设 D1 是 Amax 时计算出的,而这里我们假设 D2 是 Amax。我们现在已经知道 D1 是 Amax 是错误的,因为 D2 比 D1 大。所以 S1 显然是错误的。前一个块中计算出的数值是错误的。
但怎么办呢?你不能将前一个块的这些数值重新读取进来重新计算,这会浪费读写时间。所以需要想个办法,直接根据现有的 S1 数值进行调整,让错误变为正确,仿佛错误从未发生过。这里的调整是一个公式:S1 乘以 Exponential(D1 - D2)。这样就好像什么都没发生过,直接抹去了之前将 D1 作为 Amax 的痕迹,转变为将 D2 作为 Amax。
为什么将 S1 乘以 Exponential(D1 - D2) 就能移除 D1 的影响,并将其视为以 D2 为 Amax 的计算呢?如果你将 S1 的表达式代入,很容易就能看出:S1 是 summation over i 等于 1 到 N 的 Exponential(A_i - D1)。再乘以 Exponential(D1 - D2),其中 Exponential(-D1) 和 Exponential(D1) 可以抵消,结果变为 summation over i 等于 1 到 N 的 Exponential(A_i - D2)。这与我们在第 N+1 到 2N 的总和计算中(Exponential(A_i - D2))一致。所以最终,S2 变成了 summation over i 等于 1 到 2N 的 Exponential(A_i - D2)。这就好像我们以 D2 为 Amax 计算了 i 等于 1 到 2N 的总和。
因此,我们可以在不回溯过去数值的情况下,直接对 S1 进行调整,就好像我们以 D2 为 Amax 进行计算一样。这就是Flash Attention的核心技巧,稍后会反复利用此技巧减少读写次数。
Flash Attention完整版:直接计算输出O
现在我们已经了解在一个块中如何运作。我们以第一个和第二个块为例,后续块的算法是完全相同的。我们只是将刚才讲述的内容重新叙述一遍。在第 k 个块时,我们做的事情是:读取第 k 个块中的 A_i,与 D 中存储的值进行比较,找出最大的那个,记作 D_k。然后计算 Exponential(A_i - D_k)。接着,我们将 Exponential(A_i - D_k) 在第 k 个块中所有数值的总和计算出来。
之前存储在 S 中的值 S_{k-1} 需要进行一个变化,才能与新的 D_k 配合。所以 S_{k-1} 需要乘以 Exponential(D_{k-1} - D_k),才能匹配新的 Exponential(A_i - D_k) 这些项,然后得到 S_k。这个操作反复进行,直到你到达最后一个块(大 B 个块)。接下来的操作都与之前讲过的类似:读取一串 A_i 进来,比较大小。最大的那个就是 D_B,而这个 D_B 实际上就是 Amax。在我们遍历完所有块之后,最终会放在这个区域的数值。我们每次都会用这个区域的数值去与当前所有块中读取到的数值进行比较。当所有块都比较过后,这个蓝色区域中存储的便是 Amax。
我们现在已经有了 Amax,接下来计算 Exponential(A_i - D_B)。由于 D_B 已经等于 Amax,所以这样计算出的数值是正确的。然后,我们将 L - N 加 1 到 L(即最后一个块的数值)的所有 Exponential(A_i - D_B) 取出,并对前面留下的 S_{B-1} 乘以 Exponential(D_{B-1} - D_B) 进行变化和调整,再加上后面这一项,得到 S_B。这个 S_B 就是我们所需的分母项,因为现在 D_B 已经是 Amax,所以这个计算会得到 summation over i 等于 1 到 N 的 Exponential(A_i - Amax)。
经过这一系列操作,S 中存储的数值已经就是分母项了。所以我们现在有了 Amax,有了分母项。接下来计算 A hat 似乎是举手之劳。现在工作台上有 D 和 S,D 就是 Amax,S 就是分母项。你要做的事情就是将 A_i 的值读取进来,然后取 Exponential(A_i - D)(D 就是 Amax),再除以分母项 S(即 summation over i 等于 1 到大 L)。你对每个块的数值都做反复同样的操作,就能得到 A hat。所以,经过上述操作后,我们已经将从 A_i 到 A hat 这一连串本来需要多次读写的任务,变为只需两次读写。刚才只需从仓库中搬运两次 A_i 的数值进来,就能将 A_i 转换为 A hat。因此,我们可以将多次读写减少到两次读写。
你可能会觉得这样已经很不错了。前面先将 Q 和 K 进行一次点积得到 A_i,然后两次读写得到 A hat,再进行加权和。这比原始未经优化的方法效率更高。但这并非Flash Attention的全部。Flash Attention最神奇的地方在于,它提出了一个关键的灵魂拷问:一定要算出注意力权重才能计算最终的加权和 O 吗?它跳过了真正计算出 A hat 的步骤,直接得到了 O。
因此,如果你使用Flash Attention,一个神奇之处在于你无法直接读取到注意力权重。许多人在分析Transformer时,会想绘制注意力矩阵来查看注意力权重的数值。但如果你选择使用Flash Attention计算注意力,例如使用Hugging Face的模块,它会报错,告诉你没有注意力权重可以读取。这是因为Flash Attention这个方法可以跳过计算真正的注意力权重 A hat,直接给出 O。而且它得到的 O 的结果,与先计算 A hat 再得到 O 的结果,在理论上是完全一致的。
Flash Attention的实现与性能优势
Flash Attention是如何做到这一点的呢?它的做法是:它希望一步到位。将 Q、K、V 全部放入工作台后,进行一次运算就得到最终结果。我们来看看这一连串的运算最终是如何实现的。
前面的操作与刚才大致相同:Q 需要与一个块中的 K 都进行点积计算,然后找出点积中最大的数值,记作 D。我们这里将 D1 存入 D 中,它是这些数值中最大的那个。然后计算 Exponential(A_i - D1) 作为 S。S1 就是 summation over i 等于 1 到 N 的 Exponential(A_i - D1)。这部分操作我们刚才已经见到过。
最后,它要做的事情是,在这个时间点直接将 V 读取进来,计算加权和。我们直接将 V_i 乘以 Exponential(A_i - D1) 除以 S1(也就是 S1 中存储的数值),然后 summation over i 等于 1 到 N,得到 O1。你可能会想,这个加权和对吗?这个加权和的权重根本不是我们想要的注意力,它根本不是 A hat。没关系,稍后会想办法修复这个问题。所以这里是“将错就错”,继续计算下去。要记住,这里这项注意力权重是错误的,它不是真正的 A hat。例如,它的分子是错误的,因为 D1 不是真正的 Amax。它的分母也是错误的,因为 S1 只对 i 等于 1 到 N 进行了总和,它的 Amax 也是错误的。它只是对这个区间进行了总和,还没有对整个序列进行总和,但它已经被用作分母了。
接下来我们来看看Flash Attention如何修复之前的错误。它的做法是这样的:读取第二个块,同样计算出 D2,它是到目前为止看到的点积中数值最大的那个。然后,我们需要计算 S2。这与刚才在讲解Softmax简化版时所说的一样:原有的 S1 需要进行处理和修复,然后才能加上这个 Exponential(A_i - D2) 得到 S2。刚才我们也提到,如果你一直进行这种修复,最终 S 中存储的就会是正确的分母项,即正确 A hat 的分母项。但实际上,对于真正的Flash Attention来说,这件事已经不重要了。它真正关心的是最终能否正确计算出 O。
那么,如何正确计算 O 呢?刚才我们已经计算出了 O1,它是前一个块中 V 对一个错误的注意力权重的加权和。现在我们计算出了 S2。这个 S2 即使它可能仍然不是一个正确的分母,但它至少比 S1 更正确一点。所以,我们现在新的注意力权重可以写成 Exponential(A_i - D2) 除以 S2。尽管它不是很正确,但也许比这一项更正确一点。所以我们用第二个块中的 V 来对 Exponential(A_i - D2) 除以 S2 进行加权和。
那前面的 O1 怎么办呢?前面的 O1 已经算错了,所以我们需要抹去 S1 和 D1 存在的痕迹。我们要将 S1 想办法换成 S2,将 D1 想办法换成 D2。我们要做的事情就是将 S1 分之 Exponential(A_i - D1) 乘以 V_i,换成 S2 分之 Exponential(A_i - D2) 乘以 V_i。
那么,这个置换如何实现呢?你看它们的分母不同,一个是 S1,一个是 S2。我就将 O1 乘以 S1 再除以 S2。乘以 S1 可以抵消这一项,除以 S2 可以将这一项加回来。所以我们就弥补了 S1 和 S2 之间的差异。然后指数项(Exponential Term)不同,如何弥补这个差异呢?其实和前面讲的一样,就是乘以 Exponential(D1 - D2)。乘以 Exponential(D1 - D2) 就可以弥补指数项的差异。所以,将 O1 乘以 S1 除以 S2,再乘以 Exponential(D1 - D2) 后,O1 与这一项计算的差异,即 D1、D2、S1、S2 造成的差异,就被抹平了。你得到的结果是 summation over i 等于 1 到 2N 的 S2 分之 Exponential(A_i - D2) 乘以 V_i。这就是Flash Attention的精神。我们真正关心的是最终得到的 O。
这个过程反复进行。在第 k 个块时,操作与刚才一样:计算注意力权重。这里我们没有画出查询向量,你可以想象有一个查询向量进来,与当前块中的 K 进行点积计算,得到一堆点积结果后,找出最大的就是 D_k。然后我们需要计算这个总和。需要对原有的 S_{k-1} 进行修正,然后再加上新的 Exponential(A_i - D_k) 得到 S_k。接着我们计算 O。O_k 是原有的 O_{k-1} 经过修正项(这个是修正项)再加上我们新得到的 V_i 的加权和。所以我们的过程就是不断修正前面的错误,再加入新的东西。前面有错没关系,因为接下来总能修正回来。
所以在最后第 B 个块的操作是:读取一个块的数据进来,计算加权和,找出最大的就是 D_B。然后计算 S_B。S_{B-1} 需要进行修正,再加上新的数值。O 也一样,最终计算出的 O_B 就是前面的 O_{B-1} 经过一系列修正后再加入新的数值。最终 D 中存储的是 Amax,S 中存储的是 A hat 的分母项。但你根本不在乎这两项,你想要的是 O。这个 O 是 summation over i 等于 1 到大 L(即整个序列)的 V_i 乘以正确的注意力权重的加权和。这一连串的修正会让我们得到正确的注意力与 V_i 的加权和结果,尽管在整个计算过程中注意力权重从未被真正计算出来。这就是Flash Attention。
Flash Attention实战与性能评估
以上就是Flash Attention算法的简单讲解,稍后在助教的讲解和作业中会有更详尽的说明。这里提供了一个示例程序,让大家感受一下Flash Attention实际使用时的效果。我们将比较使用Flash Attention和不使用Flash Attention时的速度差异。
首先,这是Colab的典型使用方式:先导入一些必要的库,然后确认GPU环境。这里我们使用的是A100,其内存为80GB。请注意,这80GB的内存指的是仓库容量,而工作台(SRAM)通常只有十几MB。
接着,我们实现两种不同的注意力机制:一种使用Flash Attention,另一种是朴素Attention(Naive Attention),也就是你直觉上的算法。这两种实现都调用了PyTorch的scaled_dot_product_attention模块。scaled_dot_product_attention允许你设置使用哪种注意力算法,默认情况下,如果没有特殊设置,通常会直接使用Flash Attention。所以,现在你在运行Transformer时,如果没有做特殊处理,很可能已经默认使用了Flash Attention。
这两个函数都接收QKV作为输入,并且causal参数表示我们正在处理的是一个语言模型,它只关注前面的信息,而不是双向注意力。QKV读取进来后,唯一的不同在于第二行。第三行完全一样,都是执行PyTorch的函数来计算注意力。run_Naive_Attention中设置attn_mask_config={'is_causal': True, 'attention_implementation': 'eager'},表示使用朴素注意力;而run_Flash_Attention中设置attn_mask_config={'is_causal': True, 'attention_implementation': 'flash_attention'},表示使用Flash Attention。
接下来是数值验证,目的是证明Flash Attention和朴素Attention计算出的数值几乎是一样的。我们初始化QKV,设置几个参数:B是Batch Size(批处理大小: 一次处理的样本数量),H是注意力头(Attention Head: 自注意力机制中的并行计算单元)的数量,L是序列长度,D是QKV向量的维度。这里分别设置为4、8、256和64。256的序列长度相对较小,实际使用中通常会更长。
然后,随机初始化QKV这三个对象。这里没有绑定任何语言模型,QKV的数值是随机设置的,因为我们现在只关注计算速度,而非真正运行一个语言模型。在衡量计算速度之前,我们分别执行朴素Attention和Flash Attention,看看它们计算出的数值是否一致。一个结果存储在out_naive中,另一个存储在out_flash中。我们检查这两个输出的最大差异。执行后会发现最大差异非常小,约为10的负7次方。在进行这种不同运算时,由于底层运算机制的差异,总会存在一点点差异,但这个差异非常小,始终在10的负7次方左右。
实际案例分析:长序列加速效果
接下来,我们真正比较使用Flash Attention与不使用Flash Attention的时间差异。我们编写了一个名为benchmark_attention的函数,用于计算不同注意力机制所花费的时间。它接收QKV作为输入,还接收一个mode参数,表示当前要运行朴素Attention还是Flash Attention。另外两个参数warm_up和repeat,你可以自行调整。warm_up表示在真正计算之前先试运行一次。因为你可能不清楚Colab背后GPU的调度和使用情况,所以先运行一次,让它加载所需的库。至于需要预热多少次,你可以自行尝试。为了加速,这里只运行一次。
然后,我们开始计时,从这一行代码开始计时。我们将同一种注意力机制运行三次(你也可以运行更多次,计算结果可能更准确)。然后结束计时,计算每次注意力需要多长时间,以毫秒为单位。我们使用不同的序列长度:64、128,直到4096,来观察在不同长度下,Flash Attention能加速多少。
运行后,你会发现它运行得很快。对于不同的序列长度,例如64、128、256、512等,都会显示朴素Attention和Flash Attention分别花费的毫秒数,并将两者相除,告诉你Flash Attention相对于没有使用它加速了多少。你会发现,当序列长度设置为4096时,它最多可以加速到9倍左右。Flash Attention仅仅通过减少数据搬运次数,就可以实现8到9倍的加速。所以这是一项非常有用的技术。它的代价很小,只是算法复杂一些,并且算法中为了修正进行了一些额外的运算。但瑕不掩瑜,减少数据搬运次数确实非常划算,可以大幅提升速度。
到目前为止,我们只是做了一个玩具示例,使用假的QKV进行计算。接下来,我们将使用一个真正的语言模型来比较安装Flash Attention与否的差异。这里我们使用Yi-34B模型(其他模型的结果也类似)。首先,我们从Hugging Face下载这个模型,这需要一些时间。
在模型下载过程中,我们来看看接下来要做什么。我们会先存储一个非常长的字符串,这将在稍后作为语言模型的Prompt(提示: 输入给语言模型,引导其生成内容的文本)。我们称之为long_text,因为它确实非常长。我们只是将一个无聊的句子反复重复很多次。我们将这个句子重复100次。稍后,我们将使用Hugging Face的pipeline来调用语言模型。这个序列有7300个token。
接下来,我们开始进行测试。首先测试没有使用Flash Attention的情况。我们使用Hugging Face的pipeline函数来调用语言模型。这里有一个pipeline eager,其设置告诉我们注意力机制的实现要使用eager方法。eager方法就是不使用Flash Attention。eager虽然有“急切”的意思,但在这里它代表没有使用Flash Attention。如果你没有特殊设置,默认通常是使用Flash Attention。所以,如果你不想使用Flash Attention,你需要特别进行设置。
现在,我们使用没有Flash Attention的pipeline进行推论。我们只让它输出一个token。输入一个非常长的序列,只预测下一个token,看看花了多少时间。结果花了0.15秒。然后,我们换成使用Flash Attention,这里是pipe_sdpa,表示使用了Flash Attention。如果使用了Flash Attention,同样只输出一个token,你会发现速度并没有快多少。为什么呢?因为我们的序列太短了。
在一个语言模型中,除了注意力机制外,还有非常多不同的组成部分,例如大量的前馈网络(Feed Forward Network)。甚至在计算时,仅仅将文字转换为嵌入向量(Embedding: 将离散的符号(如单词)映射到连续的向量空间中)也需要时间,嵌入向量转回文字也需要时间。许多部分都需要时间。所以,当序列太短时,Flash Attention的效果并不明显。
我们将序列长度改为1000,给它7万个token看看。结果,没有使用Flash Attention时花了2秒,而使用Flash Attention时花了1.3秒。所以你看,在序列较长时,使用Flash Attention可以显著加快速度。
你可能会想,如果再长一点呢?我们刚才使用了1000,如果使用1万个token呢?那将有73万个token。运行一下,看看会发生什么。
哇,无法运行!出现了CUDA out of memory。这个CUDA out of memory并非工作台的内存不足,而是仓库的内存不足。所以我刚才说,仓库虽然很大,但它也有极限。当序列太长时,最终会撑爆仓库。至于序列太长为何会撑爆仓库,我们将在下一堂课讲解KV Cache时详细说明。
📌 文中提及的人物和组织
公司/组织: Hugging Face
产品/模型: Yi-34B