- 阶段一 · 地基:P01《Attention Is All You Need》 · P02 Bahdanau Attention · P03 seq2seq · P04 word2vec
- 阶段二 · 预训练范式:P05 BERT · P06 GPT-1 · P07 GPT-2 · P08 GPT-3 · P09 T5
- 阶段三 · 规模与效率:P10 Scaling Laws · P11 Chinchilla · P12 LoRA · P13 FlashAttention · 本篇 · P14 MoE
- 阶段四 · 对齐与推理:P15 InstructGPT(RLHF)· P16 思维链 CoT · P17 DPO · P18 Constitutional AI
- 阶段五 · 检索与 Agent:P19 RAG · P20 ReAct · P21 Toolformer
太长不读2022 年初,大模型的上下文窗口普遍卡在 2K token 上下——GPT-3 是 2048,BERT 只有 512。摁住它们的是注意力机制与生俱来的账单:序列里每个词都要对每个词打分,长度 N 的序列就是一张 N×N 的关注度大表,长度翻倍、开销翻四倍。一大批论文用「近似」砍这笔账——少看几个词、把表压扁——拿精度换速度,结果却很尴尬:纸面运算量(FLOPs)确实降了,真实训练时间常常没快,业界几乎没人真用。斯坦福的 Tri Dao 团队看出了病根:大家优化错了对象。现代 GPU 上算力的增长远快于显存带宽,注意力的时间大头根本不是「算」,而是把那张大表在 HBM(显存,大而慢)和片上 SRAM(缓存,小而快)之间来回搬。FlashAttention 数学一字不改——输出与标准注意力逐位相同,不是近似——只是把计算切成能塞进片上缓存的小块,用滚动统计量边搬边算,让 N×N 大表从头到尾不落地。代价是多算 13% 的浮点运算,回报是 HBM 读写少 9 倍:注意力模块最高快 7.6 倍,BERT 训练把 MLPerf 世界纪录再提 15%,GPT-2 训练快 3.5 倍,显存从平方降到线性——Transformer 第一次解出了 16K 长的 Path-X 谜题。今天它内置在 PyTorch 里,是几乎每个训练与推理框架的底层标配;你现在能用上 128K 甚至更长的上下文,这篇论文是最大的功臣之一。
上篇回顾 · 先答一题上一篇 LoRA 讲完,「适配」的账被砍平了:冻住底座,只训一块低秩补丁。先别往下翻,试着答:LoRA 凭什么敢只训万分之一的参数,效果还能追平全量微调?……答案:因为微调的「改动量」本身是低秩的——它不学新本领,只是把预训练权重里已有但未被强调的方向放大(论文实测约 21.5 倍),要改的信息本来就不多,薄补丁装得下。但 LoRA 省的是「让模型专精你的任务」的钱;模型本体的训练和推理里,还蹲着一头从 P01 就埋下的吞金兽——注意力对序列长度的平方级开销。今天这篇,轮到它了。
012022 年的死结:上下文为什么卡在 2K
先把这头吞金兽拉出来看清楚。P01 讲过,2017 年 Transformer 用注意力换掉循环网络,天才之处是让每个词直接看见所有词——但代价当时就写在纸上:每个词对每个词打一个关注度分数,长度 N 的序列就是 N² 个分数。512 个词,26 万格;2048 个词,420 万格;长度翻一倍,表大四倍。训练时反向传播还要求这张表留在显存里等着算梯度,于是时间和显存都随长度平方暴涨。这就是 2022 年所有大模型「书读不长」的直接原因:BERT 停在 512,GPT-2 停在 1024,GPT-3 财大气粗,也只敢开到 2048。
学界当然扑上去了,而且路数和上一篇的「低秩」一脉相承:近似。稀疏注意力让每个词只看一部分词,低秩注意力把大表压扁再算,还有两者的组合拳——Reformer、Linformer、Performer,2020 年前后热闹非凡,P01 结尾我还专门提过这批「让注意力变便宜」的尝试。它们确实把纸面运算量从平方压到线性或近线性。但这篇论文开篇就给整个流派定了性:其中许多在真实钟表上并不比标准注意力快,而且几乎没有被业界广泛采用。
降了运算量却没降时间——这话听着就蹊跷。论文一句话点破:这批工作都在优化 FLOPs(浮点运算次数),而 FLOPs 和真实耗时未必相关;它们集体忽略了另一本账:内存访问(IO)。这是全文的题眼,也是一条老工程教训的重演:动手优化之前,先确认瓶颈到底在哪。
于是问题被钉死:能不能不做任何近似——精确注意力,结果一位不差——却在真实时间上显著变快?这听起来像白日梦:活一点没少干,凭什么快?答案要从 GPU 的「地形图」里找。
02GPU 的地形:算得快,搬得慢
把一块 A100 掀开,里面不是一台匀质的计算机,而是一套分层的仓储系统(上图左的金字塔)。塔尖是片上 SRAM:带宽约 19TB/s,快得惊人,但总共只有约 20MB——每个计算单元分到 192KB,全卡 108 个单元凑起来还不够存一首无损音乐。塔腰是 HBM,也就是平时说的「显存」:40 到 80GB,带宽 1.5 到 2.0TB/s。塔底是 CPU 那边的主内存:容量上 TB,带宽只剩 12.8GB/s。规律赤裸裸:越快的越小,越大的越慢。打个比方:SRAM 是灶台边的料理台,快一个数量级但只放得下一盘菜;HBM 是后院冷库,什么都装得下,可来回一趟要时间。
再叠一个背景:近十几年 GPU 的算力增速一直远超显存带宽的增速——厨师的刀工越来越快,跑冷库的腿却没怎么变快。后果是深度学习里大量操作的耗时根本不由「算多少」决定,而由「搬多少」决定。论文按「算术强度」(每搬一字节数据做几次运算)把操作分成两类:compute-bound(计算受限,如大矩阵乘法——GPU 的甜区,芯片厂商为它堆满了专用电路)和 memory-bound(内存受限:激活函数、dropout、求和、softmax、层归一化……这些操作每搬一字节只算一两下,时间几乎全花在搬运上)。
现在解剖标准注意力的工程实现(论文的 Algorithm 0)。数学上就三步:Q 乘 K 转置得到分数表 S,softmax 归一化成权重表 P,P 乘 V 得到输出。但在 GPU 上跑起来是这样的:算完 S——一张 N×N 的大表——写回 HBM;再从 HBM 读回来做 softmax,得到 P,又写回 HBM;再读回来乘 V。N=1024 时这张表一百万格,16 个注意力头、batch 64,一步训练里这样的冷库往返成千上万次。而且中间两步恰恰是 memory-bound 的重灾区。上图右的耗时解剖直白到残忍:PyTorch 版注意力里,GPU 真正擅长的矩阵乘法只占一小截,softmax、dropout、mask 这些搬运工序吃掉了绝大部分时间。
病根至此确诊:注意力慢,不是 GPU 算不动,是它一直在等快递。药方也随之从玄学变成一个具体的工程问题——怎么让那张 N×N 的大表,从头到尾不去冷库?
03三板斧:切块、滚动 softmax、事后重算
板斧一:分块(tiling)。把 Q、K、V 切成小块,轮流搬进 20MB 的料理台:外循环取一块 K 和 V,内循环取一块 Q,在片上算出这一小块的贡献、直接累进输出,算完即走。整张 N×N 大表自始至终只以小块的形式在片上闪现,从不完整落地 HBM——上图(论文 Figure 1 中)虚线框画的就是这张「从未出生」的表。
但这里横着一道数学关卡:softmax。归一化要用整行的信息——减去整行最大值(保数值稳定),再除以整行的指数和——行没看完,怎么归一?FlashAttention 借来一手「在线 softmax」:随块维护两个很小的统计量——目前为止的行最大值 m,和按它缩放过的指数部分和 ℓ。每来一块新数据,就用新旧最大值的差把之前所有结果整体修正一遍。论文给出恒等式证明:看完最后一块,输出与一次看整行逐位相同。不是近似,是同一道数学换了一种加法顺序——这就是「数学一字不改」的底气所在。
板斧二:重算(recomputation)。训练还有反向传播:标准做法要把 S 和 P 两张 N×N 大表存到 HBM,等着算梯度。FlashAttention 干脆不存——只留输出 O 和那两个统计量,反向传播时在料理台上把需要的块现算出来。听着浪费(浮点运算确实变多了),但省下的 HBM 往返远远值回票价。这一招是「梯度检查点」思想的极致用法,而且破了它的老规矩:以往的检查点技术都是拿速度换显存,这里因为砍掉的是搬运,显存和速度居然双赢。
板斧三:熔成一个 kernel(kernel fusion)。有了分块,整个注意力——两次矩阵乘、softmax、mask、dropout——可以熔成一个 GPU kernel:从 HBM 读一次输入,全部工序在片上一气呵成,写一次输出。为了做到这一层,团队绕开 PyTorch 手写 CUDA(GPU 的底层编程语言)——这也给后文的「局限」埋了伏笔。
三板斧合起来,账面变化就是上图左表那三行全文最反直觉的数(GPT-2 medium,序列 1024):浮点运算量 66.6 → 75.2 GFLOPs,多算 13%;HBM 读写 40.3GB → 4.4GB,少 9 倍;耗时 41.7ms → 7.3ms,快 5.7 倍。多干活的反而先下班——因为账根本不记在「干活」上,记在「跑腿」上。
论文没停在实测,把账算成了定理。Theorem 2:标准注意力要 Θ(Nd+N²) 次 HBM 访问,FlashAttention 只要 Θ(N²d²/M)——d 是每个头的维度(64 到 128),M 是缓存大小,按典型值代入就是少一个数量级。更狠的是 Proposition 3 给出了下界:不存在任何精确注意力算法,能对所有缓存尺寸做到渐近更少的 HBM 访问。也就是说 FlashAttention 不是「暂时没被超越」,是在这个意义上数学上无法被超越。工程优化做到附赠最优性证明的,21 篇里这是独一份。顺带一提 Theorem 1:除输入输出外,它的额外显存只有 O(N)——平方降线性,这行小字是下一节所有新能力的地基。
04放榜一:破世界纪录,且答案一个小数点都不差
注意力模块自己快 7.6 倍,不等于整个模型快 7.6 倍——模型里还有别的层摊薄收益。所以放榜要看端到端。第一场对决论文挑了个狠角色:BERT-large 训练,对手是 Nvidia 创下 MLPerf 1.1 世界纪录的那套实现。同样的初始化、同样的达标线,8 张 A100 跑 10 次取平均:对手 20.0 分钟,FlashAttention 17.4 分钟,快 15%。别嫌 15% 小——这不是打学术 baseline,是把工业界卷到极致的世界纪录再往前推了一截。
第二场 GPT-2:对 HuggingFace 实现,small 版训练从 9.5 天压到 2.7 天,快 3.5 倍;对久经优化的 Megatron-LM 也快 1.7 倍;medium 版同趋势,快 3 倍。第三场长序列基准 LRA(序列 1024 到 4096):平均提速 2.4 倍。顺手把上一篇结尾我预告的「提速 2 到 4 倍」圆上——拆开就是这批数:端到端 1.7 到 3.5 倍,长序列基准 2.4 到 2.8 倍,注意力模块单看最高 7.6 倍。
但全表最关键的不是加速比,是困惑度那一列:18.2 对 18.2,14.2 对 14.3——训练随机性以内,分毫不差;附录里两条训练曲线干脆重合。数学没变,答案就不会变。这一列,近似派任何一家都给不出来;而「给不出这一列」,正是它们始终没被大规模采用的原因——没人敢拿大几百万美元的训练赌「近似应该没事」。
意义句只有一行:这不是新模型打败旧模型,是同一个模型换了个搬法。不改一行模型定义、不动任何超参,速度白拿。「免费的午餐」在深度学习里几乎不存在,所以它一出现,就注定铺满全行业。
05放榜二:不是模型变聪明了,是它终于能看完卷子
速度只是上半场。下半场的主角是那行小字——显存从平方降到线性(上图 Figure 3 右):实测比精确实现最多省 20 倍;序列拉到 64K 时,其余实现几乎全数爆显存出局,幸存的近似法 Linformer 还比它多吃 2 倍。显存账清了,一个被平方开销锁死的旋钮突然拧得动了:序列长度。
快省下来的预算,可以直接换「长」。论文给了个漂亮的换算(Table 4):GPT-2 把上下文从 1K 扩到 4K,训练时间还比 Megatron 跑 1K 短 30%,困惑度反而好 0.7。同一笔钱,书读四倍长,成绩还更好——「更长的上下文」第一次从奢侈品变成了划算的买卖。
长了到底有什么用?两个真实数据集见真章(Table 5):MIMIC-III 是 ICU 出院小结,平均 2395 个 token,最长 14562;ECtHR 是欧洲人权法院判例,最长近 5 万 token。把窗口从 512 拉到 16K,前者涨 4.3 分;拉到 8K,后者涨 8.5 分。此前不是模型读不懂病历和判例,是 512 的窗口只让它看见卷子第一页。
然后是名场面 Path-X(Table 6)。这是 LRA 基准里的珠穆朗玛:把一张 128×128 的黑白图拆成 16384 个像素逐个喂给模型,问图上两个点之间有没有一条路连通——考的就是超长程依赖。在此之前,所有 Transformer 在这道题上不是爆显存,就是成绩等于瞎猜(50%),以至于学界一度怀疑该换架构了。FlashAttention:61.4%,第一个越过随机线的 Transformer。块稀疏版更进一步,把序列推到 64K,在 Path-256 上拿 63.1%——又一个第一次。
这两个「第一次」最耐人寻味的地方在于:没有新架构、没有新损失函数、没有新数据,只是让同一个 Transformer 装得下更长的序列。换句话说:长度本身就是能力。这句话在后来几年被整个行业反复验证——而把长度价格打下来的,正是这篇论文。
06后日谈:一页许愿单,三年内全部兑现
按系列惯例讲讲它后来的命运——这可能是 21 篇里「从论文到基础设施」走得最快的一篇。有趣的是,路线图论文自己就写好了:第 5 节作者罕见地用一整节抱怨自己的方法(上图)——每设计一种注意力变体都得手写一遍 CUDA kernel,比 PyTorch 低太多层,费时费力,还可能换一代 GPU 就失效。然后许了三个愿:希望能用高层语言写注意力、自动编译成 IO 感知的实现;希望 IO-aware 思想走出注意力、覆盖网络每一层;希望 IO 分析推广到多 GPU。三年之内,三条全部兑现:OpenAI 的 Triton 语言和 PyTorch 的编译器让「用 Python 写 kernel」成了日常;IO 感知成了推理优化的默认思维方式;多卡拼接注意力的方案把上下文推向了百万级。
自家这条线也没停:FlashAttention-2(2023)重排了并行方式和任务划分,在一代基础上又快约一倍;FlashAttention-3(2024)贴着新一代 H100 的硬件特性再榨一轮。名字干脆长成了品类。铺开的程度更能说明问题:PyTorch 2.0 起,你调用一次内置的注意力函数,底层就自动走 FlashAttention 一族的 kernel——一篇论文,变成了一行你不需要知道它存在的代码。今天几乎每个主流训练框架和推理引擎都内置它;你在 2026 年用到的超长上下文产品,追根溯源都有它一份地基。
还有个值得回味的彩蛋:把注意力的搬运账一路算到理论下界的 Tri Dao,2023 年转身与 Albert Gu 提出了 Mamba——一个干脆不用注意力的架构,试图从根上绕开那张 N×N 的表。把一件事优化到数学极限的人,往往也最清楚这件事本身的天花板在哪。这条「注意力之外」的线,值得单独留意。
放进系列时间线收个尾:Scaling Laws 说「大就是好」,Chinchilla 说「大要喂饱」,LoRA 把「用起来」砍到平民价,FlashAttention 把「跑起来」的搬运账算清——纸面运算量与真实钟表之间的裂缝,被 IO 这个词填平了。但省钱这件事还有一条更激进的路线在生长:模型如果有一万亿参数,每个 token 真的需要全部一万亿都参与计算吗?能不能每次只叫醒其中一小撮?
✎读完自测:先别往下翻,试着答
主动回忆一遍才记得住。三问,想好再看答案:
- 1. 近似注意力把运算量降到了线性,为什么常常并不真的变快?
- 2. FlashAttention 反而多算 13% 的浮点运算,凭什么快 5.7 倍?
- 3. Transformer 第一次解出 Path-X,模型结构一个字没改——凭什么?
答案1. 因为瓶颈不在算,在搬:注意力是 memory-bound(内存受限)操作,耗时由 HBM 读写次数决定;近似派省的是算术、没省搬运,FLOPs 与真实钟表脱节。2. 它把 HBM 读写砍了 9 倍(40.3GB→4.4GB):多出来的计算发生在快一个数量级的片上 SRAM 里,省掉的搬运才是时间的大头——账记在跑腿上,不记在干活上。3. 因为显存从平方降到线性,16K 的序列第一次装得进显存、训得动;长度本身带来能力——此前模型不是不聪明,是看不完卷子。
→一句话带走,和下一步
只记一句:FlashAttention 数学一字不改,只按 GPU 的内存分层重排了注意力的计算顺序——切小块装进片上缓存、滚动 softmax、反向传播时现算、整个注意力熔成一个 kernel——HBM 读写少 9 倍,注意力最高快 7.6 倍,BERT 破纪录、GPT-2 快 3.5 倍,显存平方降线性,Transformer 第一次读完 16K 的卷子;它给整个行业上的一课是:优化运算量之前,先看看瓶颈是不是在搬运。
启后:FlashAttention 省的是每一层里「搬」的钱,但每个 token 仍要惊动模型的全部参数。2017 年起就有另一批人在琢磨一条更激进的路:把模型拆成一群「专家」,每个 token 进门只找其中一两位——万亿参数,只算零头。下一篇 P14《MoE 混合专家》,讲这套「稀疏激活」的省钱术如何从边缘技巧长成 GPT-4 时代的主流传闻与开源标配。
想再深入一点荐读一个最优源:Horace He 的《Making Deep Learning Go Brrrr From First Principles》(horace.io/brrr_intro.html,免费公开)——把 compute-bound、memory-bound、overhead 三种瓶颈讲得最清楚,是理解本篇「搬运账」的最佳前置读物。原论文:arXiv 2205.14135(官方代码 github.com/Dao-AILab/flash-attention);关键前作:在线 softmax(Milakov & Gimelshein 2018,arXiv 1805.02867)、Rabe & Staats 2021《Self-attention Does Not Need O(n²) Memory》(arXiv 2112.05682);延伸:FlashAttention-2(arXiv 2307.08691)。/ 互动钩子:你把最长多少的文档整个塞给过模型?有没有遇到「上下文一长就变慢变贵」的时刻?评论区聊聊。