封面

2022年,斯坦福大学等机构的研究团队(主要由博士生Tri Dao主导)发了篇论文,标题很朴实——FlashAttention。

听起来不像什么革命性的东西,对吧?但这篇论文的影响力是爆炸级的。它不是提出了一种新的模型架构,也不是发现了一种新的训练范式。它做的事情听起来特别无聊:让注意力机制算得更快、占内存更少

就这?

但结果是:几乎所有主流大模型都采用了它。GPT-4用了,Gemini用了,Llama用了,Claude也用了。在不到一年时间里,FlashAttention从一个学术项目变成了行业标配。

为什么?因为它解决了一个被所有人忽视、但卡住了所有人脖子的问题。

注意力机制不是"算得慢",而是"读得慢"

在讲FlashAttention之前,我们先要理解一个反直觉的事实。

大家总觉得Transformer慢是因为"计算量大"。毕竟注意力机制要对序列中每一对token计算关联度,复杂度是O(n²)嘛。

但FlashAttention的作者Tri Dao发现了一个被大多数人忽视的真相:在现代GPU上,注意力机制的瓶颈根本不是计算,而是内存读写。

什么意思?打个比方。想象你是一个厨师(GPU),要做1000道菜(矩阵运算)。你的手速极快(计算能力强),但你每次做菜都需要从仓库(GPU显存/HBM)跑到厨房操作台(SRAM缓存)拿食材。

问题来了:仓库很大但很远,操作台很小但就在手边。你做菜的速度很快,但大部分时间都花在来回跑仓库拿食材上了。等你跑回来,菜都快凉了。

这就是GPU面临的现实——内存墙(Memory Wall)

Image
FlashAttention的分块(tiling)策略示意图,以及GPU存储层级:SRAM(19TB/s,20MB)vs HBM(1.5TB/s,40GB)(来源:原论文Figure 1)

具体来说,GPU的SRAM(片上缓存)速度极快,但只有约20MB。而HBM(高带宽显存)有约40GB,但速度慢得多。注意力机制需要频繁在HBM和SRAM之间搬运大量中间结果——那个巨大的n×n注意力矩阵。这导致GPU大部分时间都在等数据搬运,而不是在计算。

FlashAttention的核心洞见就是一句话:别再搬那么多数据了!

分块计算:厨房虽小,但够用

FlashAttention的方案说起来其实很简单,这也是它最优雅的地方。

原来的注意力机制是这样干的:先把整个巨大的Q×K矩阵算出来,存到HBM里;再读回来算softmax;再读回来乘V。中间需要把整个n×n矩阵反反复复写回HBM又读回来,来来回回搬了好几趟。

FlashAttention说:别这么干。我们每次只取一小块Q和K到SRAM里,在SRAM里算好softmax,直接得到一小块输出,写回HBM。这样就不需要把那个巨大的中间矩阵存到HBM里了。

回到厨师类比:不要一次把所有食材都搬到操作台上,而是分批处理。每批拿够用的量,做完这一批,再拿下一批。操作台虽然小,但效率反而更高——因为你不用来回跑了。

Image
标准注意力与FlashAttention的对比:HBM读写从40.3GB降至4.4GB,运行时间从41.7ms降至7.3ms(来源:原论文Figure 2)

在线Softmax:分块了还能算对吗?

这里有个技术难点:softmax需要对所有数据做归一化,它得看到"全局"才能工作。你现在分块了,每个块只看到一小部分数据,怎么保证softmax的结果是对的?

FlashAttention用了一个叫"在线softmax"(Online Softmax)的技巧。简单说,它维护两个中间变量——当前的最大值和累加和。每处理一个新的块,就用数学技巧更新这两个值。等所有块都处理完了,最终结果和一次性算完全一样。

这就是FlashAttention最关键的一点:它不是近似算法,它是精确的。结果和原始注意力机制数学上完全一致,没有损失任何精度。

之前有不少"近似注意力"的工作,比如Linformer、Performer,用各种数学近似来加速,但都牺牲了精度。FlashAttention证明:你不需要牺牲精度,只要聪明地管理数据流动就够了。

重计算:扔掉比存着更快

在反向传播时,正常做法需要保存中间结果(比如softmax的输出矩阵),以供求导使用。但这些中间结果太大了——一个n×n的矩阵,序列一长就爆炸——会吃掉大量显存。

FlashAttention的策略听起来很"暴力":不存,用的时候重新算

听起来浪费?但实际上,在SRAM里重新算一遍,比去HBM里读一遍还要快。因为SRAM的速度比HBM快了一个数量级。这就像你背课文:与其把课本放远处每次跑过去翻(从HBM读取),不如直接在脑子里重新想一遍(在SRAM重计算)——因为"脑子"(SRAM)实在太快了。

这个"重计算"思路对工程师来说应该不陌生。我们在日常开发中也经常做类似的事情:为了减少数据库查询(慢IO),我们会在内存里用缓存或者直接重新算一遍。本质是一样的——计算便宜,IO昂贵

效果有多好?

论文给出的数据非常亮眼。

**训练速度:**在GPT-2上,训练速度提升了约3倍。在BERT-large上,训练速度提升了约15%。这可是不改模型、不改精度、纯靠优化内存访问得到的提升。

**内存占用:**显存占用随序列长度呈线性增长(而非原来的二次方增长),从而大幅节省了显存。这意味着同样的显存可以训练更长的序列,或者用更大的batch size。

**长序列能力:**原来注意力机制处理到约2K序列长度就会因为显存不够而OOM,FlashAttention可以轻松处理到64K甚至更长。这直接为后来的长上下文模型打下了基础。

Image
不同序列长度下的注意力运行时间(左)和显存占用(右)对比,FlashAttention在长序列场景优势显著(来源:原论文Figure 3)

为什么这个工作如此重要?

FlashAttention之所以成为行业标配,不仅仅因为它"更快"。它改变了一个思维方式。

**过去大家以为:**算法改进只能从数学层面入手——改模型结构、改训练方法、改损失函数。

**FlashAttention告诉我们:**理解硬件特性,从"数据怎么流动"的角度优化,同样可以获得巨大的收益,而且这种优化可以和算法改进叠加。

作为工程师,我从这篇论文里得到的一个重要直觉是:计算便宜,IO昂贵。这个道理在我们日常做系统优化时无处不在——数据库查询优化、缓存设计、网络请求合并——本质上都是在减少"慢路径"上的开销。FlashAttention把这个原则应用到了GPU计算的极致。

更重要的是,FlashAttention打通了"长上下文"这条路。没有FlashAttention,今天的大模型可能只能处理几千个token的上下文,而不是几万、几十万。想想看,如果ChatGPT每次只能记住你说的前几句话,那对话体验会有多差。长上下文是AI助手实用性的关键,而FlashAttention是长上下文的基石之一。

FlashAttention的数学和原始注意力一模一样,但通过对硬件的理解,做到了数量级的提升。它不是靠更聪明的算法赢了,而是靠更聪明地"搬数据"赢了。有时候,最优雅的优化不是让计算变少,而是让数据不动。

论文链接:https://arxiv.org/abs/2205.14135

kk的大模型论文学习笔记 · 第8篇 · FlashAttention

Logo

中国智能体开发者社区,聚焦智能体与大模型开发,提供前沿资讯、实用工具链、开源项目及行业案例。通过技术沙龙、开发者大赛等活动,促进经验交流与协作,助力开发者快速构建创新智能应用。

更多推荐