通俗解释:Flash Attention
FlashAttention 是一种快速且内存高效的精确注意力机制,通过 IO 感知设计利用 GPU 内存层次(HBM 和 SRAM)减少数据搬运。
ELI5: FlashAttention
这篇博客的目标是解释 FlashAttention,希望任何已经理解注意力机制的人在阅读后都会自问:
“为什么我之前没想到这个?”接着是“这太简单了。”
我们将从基本原理开始。首先理解标准/普通注意力是如何实现的,然后逐一解决低效问题——就好像我们自己独立发现了 FlashAttention 一样。
另外,我的一个次要目标是揭开编译器领域术语的神秘面纱:kernel、kernel fusion、materialization 等。
注意:我不会解释注意力机制本身,请参考 Jay Alammar 的精彩博客或我对原始 Transformer 论文的实现。

看完这篇博客,你应该就能完全理解这张图了。
话不多说,让我们从分解论文标题开始:
“FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness”
核心在于 FlashAttention 是:
- Fast——论文摘录:“我们在 MLPerf 1.1 中训练 BERT-large(序列长度 512)比之前的训练速度记录快 15%,GPT2(序列长度 1K)比 HuggingFace 和 Megatron-LM 的基线实现快 3 倍,长距离竞技场(序列长度 1K-4K)比基线快 2.4 倍。”
- Memory-efficient——相比于普通注意力(在序列长度上是二次的,$O(N^2)$),这种方法在 N 上是次二次/线性的($O(N)$)。
- Exact——意味着它不是注意力机制的近似(例如稀疏或低秩矩阵近似方法)——其输出与“普通”注意力机制相同。
- IO aware——相比于普通注意力,FlashAttention 是有感知的。

你刚刚被“Blake Lemoine 化”了。
开个玩笑 :) —— 这只是意味着它不把底层硬件当成黑盒。相反,它 利用底层硬件(例如 GPU,但其他 AI 加速器也应该适用,我将以 GPU 为例)的内存层次结构知识。
让我们进一步深入探讨 IO 感知部分。“IO”是更多 FLOPS 不一定转化为更长的挂钟时间的原因(可能有点反直觉,但如果你了解硬件工作原理,这很明显)。
论文的相关摘录:
“尽管这些 [近似] 方法将计算需求降低到序列长度的线性或接近线性,但其中许多方法在挂钟时间上并未显示出对标准注意力的加速,并且尚未被广泛采用。一个主要原因是它们专注于减少 FLOP(这可能与挂钟速度无关),而倾向于忽略内存访问(IO)的开销。”
诀窍是什么?
是硬件:

https://www.semianalysis.com/p/nvidiaopenaitritonpytorch#%C2%A7the-memory-wall
多年来,GPU 增加计算能力(FLOPS)的速度快于增加内存吞吐量(TB/s)的速度。
如果没有数据可处理,即使能以 exaFLOPS 速度计算也没用。 这两者需要紧密对齐,由于硬件失去了这种平衡,我们必须让软件对此进行补偿。
因此是“IO-aware”。
根据计算与内存访问之间的比例,操作可以分为:
- compute-bound(计算密集型)(例如:矩阵乘法)
- 或者 memory-bound(内存密集型)(例如:逐元素操作(激活、dropout、掩码),归约操作(softmax、层归一化、求和等)…)
关于术语的说明:这个比例通常通过算术强度来衡量,即每次内存访问的算术操作数。
注意 2:我强烈推荐阅读 Horace 的博客 https://horace.io/brrr_intro.html —— 它将有助于进一步澄清计算密集型/内存密集型/开销密集型之间的区别。
事实证明(在当前 AI 加速器上)注意力是内存密集型的。
为什么?
因为“它主要由逐元素操作组成”,或者更准确地说,注意力的算术密度不是很高。
让我们放大论文中的这个图:

Dropout、softmax、掩码——所有逐元素操作,所有都是内存密集型。我们可以看到它们主导了运行时间。
在左边的条形图中,可以看到掩码、softmax 和 dropout 是占用大部分时间的操作,而不是矩阵乘法(尽管大部分 FLOP 都在矩阵乘法中)。
但并非无计可施。内存不是单一的工件,它在本质上是层次化的,一般规则是:内存越快,成本越高,容量越小。
让我们放大图的这一部分:

内存在本质上是层次化的。我们可以通过进一步添加 SSD(更高容量但更慢)、HDD(硬盘)等来延续金字塔(AWS S3? :) )。你明白意思了。
“IO-aware” 在实践中归结为利用 SRAM 比 HBM(“高带宽内存”——不幸的名字)快得多的事实,确保减少两者之间的通信。
为了更具体一些,这里有一个具体例子:
A100 GPU 有 40–80GB 的高带宽内存(HBM,就是给你可爱 CUDA OOM 的东西),带宽为 1.5–2.0 TB/s,每个 108 个流式多处理器有 192KB 的片上 SRAM,带宽估计约为 19TB/s。
H100 和其他加速器仍然保持类似的比例。
现在,让我们看看标准注意力实现背后的计算:

符号说明:$Q$ — 查询,$K$ — 键,$V$ — 值,$S$ — 分数,$P$ — 概率,$O$ — 输出。
你可以看到标准实现如何表现出对硬件工作方式的极大不尊重。它基本上将 HBM 的加载/存储操作视为零成本(它不是“IO-aware”)。
现在让我们从基本原理思考如何使这个实现更高效(时间和内存方面)。
最低悬的果实是消除冗余的 HBM 读写。
为什么要把 $S$ 写回 HBM 只是为了(重新)加载它以计算 softmax?让我们把它保留在 SRAM 中,执行所有中间步骤,然后只将最终结果写回 HBM。
这就是编译器人员所说的**“kernel fusion”**,深度学习中最重要的一种底层优化:

不,不是那个,而是简单的这个:

https://horace.io/brrr_intro.html
Kernel 基本上是一种花哨的说法,意思是“GPU 操作”。
Fusion 意味着你将多个操作融合/组合在一起。
所以,你只从 HBM 加载一次,执行融合后的操作,然后只将结果写回。通过这样做,你减少了通信开销。
顺便说一句,我真心认为人们应该停止仅仅因为听起来酷而用单词命名概念。Kernel 是计算机科学世界中最过载的词(可能仅次于“模型”)。它可以指代任何东西:Linux 内核(Linux 操作系统的核心软件组件)、神经正切核、SVM 核、GPU 操作等。它是计算机科学世界的 HIV Aladeen。😂
你还会遇到的最后一个术语是 “materialization”(物化)。它指的是在上述标准注意力实现中,我们分配了完整的 $N \times N$ 矩阵($S$,$P$)。我们很快就会看到,这正是 FlashAttention 直接解决的瓶颈,将内存复杂度从 $O(N^2)$ 降低到 $O(N)$。
现在背景已完全设定,让我们深入探讨 FlashAttention 算法。
FlashAttention 基本上归结为两个主要思想:
-
Tiling(分块)(在前向和后向传播中都使用)——基本上将 $N \times N$ 的 softmax/分数矩阵分块成块。
-
Recomputation(重新计算)(仅在反向传播中使用——如果你熟悉激活/梯度检查点,这将很容易理解)
就是这样。
这里是算法:

FlashAttention 的要点(不包括 dropout/掩码)。
我的工作完成了。希望你喜欢这篇博客,订阅以获取更多未来的内容!🚀
开玩笑的。
让我们再理解几个使分块方法工作所需的概念,然后我会逐行解释算法。
FlashAttention 算法
使分块方法工作的主要障碍是 softmax。特别是,softmax 将所有分数列耦合在一起的事实。这是我们如何计算 softmax 的第 $i$ 个输出的方式。

$z_i$ 是第 $i$ 个分数(键-查询点积),输出是第 $i$ 个 token 的概率,我们稍后用它对值向量进行加权(再次,我假设你知道注意力如何工作)。
看到分母了吗?
这就是问题所在。
要计算输入序列中某个特定的第 $i$ 个 token 对其他 token 的注意力程度,你需要让所有这些分数(此处表示为 $z_j$)在 SRAM 中立即可用。
但让我提醒你:SRAM 的容量非常有限。你不能一次性加载整个东西。$N$(序列长度)可以是 1000,甚至10 万个 token。所以 $N^2$ 会迅速爆炸。
所以技巧是:我们实际上可以将 softmax 计算分解成更小的块,并且仍然得到完全相同的结果。
这里是主要公式:

公式 1. softmax 的部分计算(我们将迭代得到正确的 softmax 数值)。
我们可以只取前 $B$ 个分数($x_1$ 到 $x_B$)并计算它们的 softmax。
这些数字,至少目前是,不正确的。 但请耐心等待,通过迭代,我们将“收敛”到正确的结果。
注意:你可以忽略 $m(x)$ 部分,至少目前我们还在柏拉图的思想世界中。它的目的仅仅是为了避免数值不稳定。在未来的某种假设硬件上(例如,我们用更多比特表示数据),这可能不需要。$m(x)$ 不会以任何方式改变最终结果。
注意:另请参阅引入在线 softmax 的原始论文:https://arxiv.org/abs/1805.02867
现在诀窍是,我们可以以一种聪明的方式组合那些逐块的部分 softmax 数值,使得最终结果实际上是正确的。这是主要思想:

公式 2. 这是 softmax 分块的核心思想。通过在所有块上递归重复这个计算,我们最终得到正确的 softmax 输出。
所以基本上,为了计算属于前两个块(大小为 $B$)的分数的 softmax,你必须为每个块跟踪两个统计量:$m(x)$(最大分数)和 $l(x)$(指数分数的总和)。
然后你可以使用归一化系数将它们无缝融合在一起。
注意:如果你做一些非常基本的代数,你很容易说服自己系数是合理的。通过展开 $f(x)$ 和 $l(x)$ 项并用 $e^x$ 相乘,一些项会相互抵消,这就是基本的操作。
这个逻辑递归地一直持续到最后一个 $(N/B)$ 块,此时你拥有 $N$ 维的正确 softmax 输出!
好了,我们现在有了理解 FlashAttention 算法前向传播所需的所有要素。
注意:下面的算法假设我们有一个批次大小为 1(即单个序列)和一个注意力头,我们稍后可以轻松扩展(只需在 GPU 的流式多处理器上并行化——稍后详述)。此外,我们暂时忽略 dropout 和掩码,稍后添加它们很容易。

符号说明:$d$ — 注意力头维度,$B_r$(行块大小),$B_c$(列块大小)——其余符号之前已介绍。
现在让我们逐步分解!
步骤 0: HBM 的容量以 GB 计(例如,RTX 3090 有 24 GB 的 VRAM/HBM,A100 有 40–80 GB 等),所以分配 $Q$、$K$ 和 $V$ 不是问题。
步骤 1: 计算行/列块大小。为什么是 $\text{ceil}(M/4d)$?因为查询、键和值向量是 $d$ 维的,并且我们还需要将它们组合成输出的 $d$ 维向量。所以这个大小基本上允许我们用 $q$、$k$、$v$ 和 $o$ 向量最大化 SRAM 容量。
玩具示例:假设 $M = 1000$,$d = 5$。在这个例子中,块大小是 $(1000/4*5) = 50$。所以在这个例子中,我们一次加载 50 个 $q, k, v, o$ 向量的块,以确保减少 HBM/SRAM 之间的读写次数。
值得在脑海中保留这个图像(很快会更有意义):

至于 $B_r$,我不太确定他们为什么与 $d$ 执行 min 操作?如果有人知道,请随意评论!
步骤 2:

我们用全 0 初始化输出矩阵 $O$。它将作为一个累加器,所以是那个初始值。类似地对于 $l$(记住:它的目的是保存 softmax 的累积分母——指数分数的总和)。$m$(保存行方向最大分数)初始化为 $-\infty$,因为我们将在其上执行 max 操作,所以无论第一个块的最大值是多少,它肯定大于 $-\infty$——因此这是自然的初始值。
步骤 3:

我们使用步骤 1 中的块大小将 $Q$、$K$ 和 $V$ 分割成块。也请参见上面的图。
步骤 4:

类似地将 $O$、$l$、$m$ 分割成块(与 $Q$ 相同的块大小)。
步骤 5:

让我们开始循环遍历列,即键/值向量(上图中的外层循环)。
步骤 6:

让我们将 $K_j$ 和 $V_j$ 块从 HBM 加载到 SRAM。记住,由于我们构造块大小的方式,此时 SRAM 仍有 50% 未占用(专用于 $Q$ 和 $O$)。

大致且抽象地。GPU 内存布局显然会不同。
步骤 7:

开始内层循环遍历行,即查询向量(再次参见图)。
登录以更快地登录
步骤 8:

将 $Q_i$($B_r \times d$)和 $O_i$($B_r \times d$)块以及 $l_i$($B_r$)和 $m_i$($B_r$)加载到 SRAM。
当我们以只留有足够空间给 $K_j$、$V_j$、$Q_i$ 和 $O_i$ 的方式计算块大小时,$l_i$ 和 $m_i$(以及所有中间变量)如何适应 SRAM?我想答案是:寄存器(参见这个 CUDA 视频系列以了解 GPU 内存层次结构)。但我可能错了,有实际在 CUDA 中实现的人请纠正我。🙏 我相信仅通过分析伪算法,我会遗漏重要的实现细节。
步骤 9:

计算 $Q_i$($B_r \times d$)与 $K_j$ 转置($d \times B_c$)的点积,得到分数($B_r \times B_c$)。如你所见,我们没有“物化”整个 $N \times N$ 的 $S$(分数)矩阵。只有它的一小部分($S_{ij}$)!
玩具示例:假设外层循环索引是 $j$($j=3$),内层循环索引是 $i$($i=2$),$N$ 是 25,块大小是 5,这就是我们刚刚计算的(假设从 1 开始索引):

基本上,这是输入序列中 token 6–10 与 token 11–15 的注意力分数。但重要的是,这些是精确的分数,它们永不改变(与 softmax 结果不同,后者会逐渐改进)。
步骤 10:

使用上一步计算的分数计算 $\tilde{m}{ij}$、$\tilde{l}{ij}$ 和 $\tilde{P}_{ij}$。这很简单。
$\tilde{m}_{ij}$ 按行计算,找到上述每一行的最大元素。
我们通过逐元素操作得到 $\tilde{P}_{ij}$:
- 归一化——取行最大值并从行分数中减去它
- 指数运算
$\tilde{l}_{ij}$ 仅仅是矩阵 $P$ 的行求和。
步骤 11:

计算 $m_{\text{new},i}$ 和 $l_{\text{new},i}$。再次很简单,让我们重用上面的图:

$m_i$ 包含之前所有块($j=1$ 和 $j=2$,以绿色表示)的行方向最大值。$\tilde{m}{ij}$ 包含当前块(以黄色表示)的行方向最大值。要得到 $m{\text{new},i}$,我们只需对 $\tilde{m}{ij}$ 和 $m_i$ 应用最大值。类似地对于 $l{\text{new},i}$(它还需要按照之前我们在公式 2 中看到的乘以系数)。
步骤 12(最重要的步骤):

这是算法中最难的部分,但仍然不太复杂,特别是当你内化了部分 softmax 计算的公式 1 和公式 2 时。

让我们先分解 $\operatorname{diag}(l)$ 部分。
它基本上允许我们以矩阵形式进行行方向标量乘法。如果你有一个标量列表 $s$($N$)和一个矩阵 $A$($N \times N$),如果你做 $\operatorname{diag}(s) \cdot A$,你基本上是在用这些标量对 $A$ 的行进行逐元素乘法。
接下来注意步骤 12 与公式 1 的相似性(为了方便再次粘贴在这里):

所以步骤 12 的第一项(用绿色下划线标注)所做的是更新当前 softmax 估计,针对同一行块中当前块之前的块。如果 $j=1$(即该行中的第一个块),第一项将为 0,我们最终就只剩下第二项。
第一项乘以 $\operatorname{diag}(l_i)$ 是为了抵消上一次迭代中除以相同常数(该常数隐藏在 $O_i$ 内部)的操作。
表达式的第二项(用黄色下划线标注)不需要这种抵消,因为,如你所见,我们直接将 $\tilde{P}_{ij}$ 矩阵与 $V$ 向量的块($V_j$)相乘。
$e^x$ 项的存在是为了修改矩阵 $\tilde{P}{ij}$ 和 $O_i$,通过抵消上一次迭代的 $m$,并将其更新为包含到目前为止行方向最大值的最近估计($m{\text{new},i}$)。
说服自己这是有道理的最简单方法就是自己模拟几次迭代——如果你还没有完全理解的话。
这只需 5 分钟。这是我的逐步分析(希望有帮助!):

如你所见,主要点是外部的 $e$ 项与 $P/O$ 矩阵内部的 $e$ 项相互抵消,我们总是以最新的 $m_{\text{new},1}$ 估计结束!

第三次迭代类似地,我们以正确和最终结果结束!
记住:这只是最终 $O_i$ 的当前估计。只有当我们迭代完上图中所有红色块后,才会得到精确结果。就是这样!
步骤 13:

将最新的累积统计量($l_i$ 和 $m_i$)写回 HBM。注意这些是 $B_r$ 维度的。
步骤 14、15、16:

一旦嵌套 for 循环结束,$O$($N \times d$)将包含最终结果:每个输入 token 的注意力加权值向量!
就是这样。这就是 FlashAttention 的前向传播!
这个算法可以很容易地扩展到“块稀疏 FlashAttention”,一种稀疏注意力算法,甚至比 FlashAttention 快 2–4 倍,可扩展到 64k 的序列长度!思路是我们使用块形式的掩码矩阵,并从上述嵌套 for 循环中简单地跳过某些加载/存储操作,从而可以按稀疏系数成比例地节省时间。

块形式掩码矩阵的示例。假设我们的块长度为 3,在这个玩具示例中我们只需要 3/9 次迭代。因此快了 3 倍!
现在简要谈谈复杂度。
复杂度分析
空间: 我们在 HBM 中分配了 $Q$、$K$、$V$、$O$($N \times d$)、$l$ 和 $m$($N$)。那是 $4Nd + 2N$。去掉常数(大 $O$ 的东西),并且知道 $d$ 也是一个常数,通常比 $N$ 小得多(例如通常 $d={32, 64, 128}$,$N={1024, \dots, 100k}$),我们得到空间复杂度 $O(N)$。大胜利!这帮助我们“轻松”将 Transformers 扩展到 64k 序列长度(再加上一些其他“技巧”,如 ALiBi,我将在后续博客中覆盖)。
时间: 我们不会严格进行时间复杂度分析,而是使用一个好的代理:HBM 访问次数。
这是论文中的摘录:

他们如何得到那个数字?好吧,让我们分析嵌套的 for 循环:
- 我们的块大小是 $M/4d$。这意味着向量被分割成 $N/(M/4d)$ 个块。
- 将其提升到 $2$ 次方(因为我们在行/列块上循环),得到 $O(N^2 d^2 / M^2)$
- 现在我有 $M^2$,他们只有 $M$——我的假设是,在循环内部我们可能需要 $M/4d$ 次内存访问来实际获取所有向量,即我们不能一次获取整个块。这可以解释缺少的 $M$?(在大 $O$ 中 $d$ 可以忽略为常数)。
如果我们做一个大 $O$ 分析,可能会让我们认为这比标准注意力好不了多少,但对于典型数字,这会导致访问次数减少多达 9 倍(根据上面的摘录)。
就是这样,你现在(希望)已经理解了 FlashAttention!
让我们通过缩小与现实世界的差距来总结。到目前为止,我们只分析了针对单个注意力头且批次大小为 1 的伪算法。我们也忽略了反向传播。
batch_size > 1, num_heads > 1, 反向传播
让我们从低垂的果实开始。把前面看到的实现扩展到支持 $batch_size > 1$ 和 $num_heads > 1$ 实际上并不难。
到目前为止我们看到的算法基本上是由一个 thread block(CUDA 编程术语)处理的。这个线程块在一个 streaming multiprocessor(SM)上执行(例如,A100 上有 108 个)。为了并行化我们的计算,我们只需在不同的 SM 上并行运行 $batch_size \times num_heads$ 个线程块。这个数字越接近系统上可用的 SM 数量,利用率就越高(理想情况下是倍数,因为每个 SM 可以运行多个线程块)。
当这个数字大于可用 SM 数量时会发生什么?我不确定,但我假设有一个队列来跟踪等待的 kernel(更新:显然 CUDA 运行时负责处理,它使用某种队列来实现该逻辑)。
接下来我们简要讨论反向传播。
反向传播依赖于相同的一组概念 + recomputation(重新计算)。
为了演示重新计算的概念,我将使用“activation/gradient checkpointing”方法的例子。
我们知道在前向传播期间计算的激活值需要在反向传播期间立即可用,以便计算关于损失函数的梯度。
这里的技巧是不在前向传播期间存储它们(因为它们占用巨大的内存),而是在反向传播期间从头开始重新计算它们。这里有一个内置的 tradeoff:我们减慢了反向传播以降低内存占用。
注意:这个权衡是一个谱系,例如,你可以每 $n$ 层存储激活值,然后在计算第 $i$ 层的激活值时,你不必从输入开始,而是从最接近的存储激活值开始。
同样的重新计算概念在这里被重用——但有一个变化!幸运的是对于 FlashAttention,我们不必牺牲运行时间或内存!
通过存储输出 $O$($N \times d$)和 softmax 归一化统计量($N$),我们可以在反向传播中直接从 SRAM 中的 $Q$、$K$ 和 $V$ 块($N \times d$)重新计算注意力矩阵 $S$($N \times N$)和 $P$($N \times N$)!从而保持内存为 $O(N)$。如果你对细节感兴趣,我鼓励你阅读论文,但我向你保证,你现在已经具备了理解它所需的所有工具。
最后,让我们看看实现 FlashAttention 可能会遇到的一些问题。
现实世界是……混乱的
赋予 FlashAttention 力量的东西也是其问题的根源。让我们看看论文中的这个摘录:
“我们当前构建 IO-aware 注意力实现的方法需要为每种新的注意力实现编写一个新的 CUDA kernel。这需要在比 PyTorch 低得多的语言中编写注意力算法,并需要大量的工程工作。实现也可能无法跨 GPU 架构移植。这些限制表明需要一种方法,支持用高级语言(例如 PyTorch)编写注意力算法,并编译为 CUDA 中的 IO-aware 实现……”
因此,原始的 FlashAttention 只支持 GPU 的一个子集。例如,V100 不受支持。请参阅他们在 GitHub 上的这个问题:

https://github.com/HazyResearch/flash-attention/issues/148
注意:这里还有一个相关问题:https://github.com/HazyResearch/flash-attention/issues/190
为了让这个问题更加直观,这里是原始实现中的实际 CUDA 代码:

来自原始代码库的 CUDA 实现片段 https://github.com/HazyResearch/flash-attention/blob/main/csrc/flash_attn/src/fmha_fprop_kernel_1xN.h#L383
如你所见,编写 CUDA 是……混乱的。此外——这是一个研究代码库,这没有帮助,但即使不是,对于来自 ML 研究背景、只熟悉 Python 的人来说,这可能是一个(潜在的)交易破坏者。
这就是像 OpenAI 的 Triton 这样的项目可能成为游戏规则改变者的地方(参见他们的 FlashAttention 实现)。Triton 基本上是一个介于 CUDA 和其他 DSL(例如 TVM)之间抽象级别的 DSL(领域特定语言)。你可以编写经过超级优化的 Python 代码(编译后),而不是直接处理 CUDA。同样的 Python 代码可以部署在任意加速器上(这个责任落在 Triton 开发者和硬件制造商身上)。
Triton 最近已经与 PyTorch 2.0 集成,所以一定要关注这个项目!特别感谢 Philippe Tillet,他在博士期间开始构建 Triton,后来成为 OpenAI 的一员。
最后,值得一提的是,对于某些用例,你可能仍然更倾向于其他方法。例如,对于超过 1K 的序列长度,一些近似注意力方法(例如 Linformer)开始变得更快。但最终,据我所知,FlashAttention 的块稀疏实现优于所有其他方法。
结束语
你可能会问自己:为什么以前没有人发明 FlashAttention?考虑到这个计算模块对于所有现代 ML 工作负载的重要性。考虑到不同技术组织花费了多少工程小时,试图从这些系统中榨取最后一点性能。为什么是一个 斯坦福学生(特别感谢 Tri Dao 的出色工作!)想到了这个,而不是比如 NVIDIA 的工程师?
我能看到几种可能的解释:
- FlashAttention,据我所知,更容易/只有在最新的 GPU 上才能实现(因此原始代码库不支持 V100)。再加上一种现象,即“局外人”常常用初学者的眼光看待问题,并从基本原理解决,我们可能得到解释。
- 最后,我非常喜欢 Nat Friedman 关于世界效率的观点:

最后,让我用一些值得思考的东西来结束:
考虑到训练这些模型的成本(参见 MosaicML 的这篇博客以获得最乐观的估计),能够在 BERT-large 训练中缩减 15%,或将 GPT 训练加速 2/3 倍,考虑到当前和未来全球范围内的模型训练,具有如此巨大的经济影响。
想象一下,如果我们生活在一个 $Y$(捕获的价值)与 $X$(产生的价值)实际相关的世界里,作者将在未来几年内成为万亿富翁。:)
唉,研究人员很少捕获那种价值,无论好坏(想象一下如果毕达哥拉斯为他的定理申请了专利:P)。然而,那些常常只是重新包装东西的人才是捕获大部分价值的人。这里完全是题外话,下次见!;)
更多资源
- Flashier Attention 博客——https://www.adept.ai/blog/flashier-attention -> 他们展示了如何进一步优化 FlashAttention 以用于高度分布式的设置,其中批次大小变得非常小(流水线并行)且序列长度非常长。
- Tri Dao 的演讲:FlashAttention — Tri Dao | Stanford MLSys #67
- Tri Dao 的演讲:MedAI #54: FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness | Tri Dao
- 论文:FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness
致谢
感谢 Tri Dao、Horace He 和 Amrit Sahu 阅读本文的早期草稿并提供反馈!
BibTeX 引用
@article{gordiceli5flashattention,
author={Aleksa Gordic},
title={ELI5: FlashAttention},
year={2023},
url={https://gordicaleksa.medium.com/eli5-flash-attention-5c44017022ad},
}
与我联系
最后但同样重要的是,请随时给我留言或:
- 在 LinkedIn 和/或 Twitter 上连接并联系
- 订阅我的 YouTube 频道 获取更多 ML 内容
- 在 Medium 和 GitHub 上关注我
- 订阅我的 月度 AI 新闻通讯 并加入 Discord 社区!
- 📄 网站
如果你觉得我创建的内容有用,考虑成为 Patreon 支持者!
满满的爱 ❤️
- 原文链接: gordicaleksa.medium.com/...
- 登链社区 AI 助手,为大家转译优秀英文文章,如有翻译不通的地方,还请包涵~