Learn AI · 模型硬功底(由内而外)
第 14 课 / 经典论文 P13  ·  阶段三:规模与效率(Dao et al., 2022)

FlashAttention:
attention 慢,不是因为算得多
是因为那张大矩阵在慢显存里反复搬——长上下文与低延迟能落地的工程底

上一课 LoRA 省的是微调成本(改 什么——只训低秩补丁)。这一课 FlashAttention 省的是跑 attention 时的显存与时间(改 怎么算——同样的结果、更聪明的内存顺序)。 它戳破一个流行的误解:attention 慢,大家以为卡在 O(N²) 的「计算量」上,真相却是卡在「把那张 N×N 大矩阵在 GPU 慢速显存里反复读写」的搬运上。 这直接对应实时语音、长上下文应用天天纠结的事:为什么上下文一长,延迟和显存就爆?怎么在不牺牲效果的前提下把它压下去?
FlashAttention 的一句话:attention 是「内存瓶颈」(memory-bound),不是「算力瓶颈」。
标准做法会把那张 N×N 的注意力矩阵整个物化到慢速大显存 HBM 里、再反复读写——时间全耗在搬运上。FlashAttention 用分块(tiling)+ 在线 softmax,让这张矩阵根本不落到慢显存,在片上高速缓存 SRAM 里算完。 结果完全一样(精确,不是近似!),但快好几倍、显存从 O(N²) 降到 O(N)。这就是长上下文和低延迟能落地的工程底。

一、立靶:attention 到底慢在哪、贵在哪?

回忆 self-attention:算 softmax(QKᵀ)·V,中间那个 QKᵀ 是一张 N×N 的注意力矩阵(N=序列长度)。直觉上大家把矛头指向计算量——N² 个分数要算,序列翻倍、计算量四倍。

立靶:瓶颈真的是「算得多」吗?

FlashAttention 的作者实测发现:GPU 的算力其实富余,真正卡脖子的是内存搬运——那张 N×N 矩阵要写进慢速大显存(HBM),softmax 时读出来,再写回、再读…… attention 的大部分时间,GPU 不是在算,而是在等数据搬来搬去既然瓶颈是搬运,那提速的正确方向就不是「少算」,而是「少搬」。

二、抽框架:GPU 的两层内存,和被忽视的「搬运」

要懂 FlashAttention,先懂 GPU 的内存层级——这是全篇的物理基础:

内存容量速度类比
SRAM(片上)很小(~20MB)极快(~10–20×)你的书桌,伸手就拿,但放不下几本书
HBM(显存)大(40–80GB)慢(带宽受限)书架仓库,啥都放得下,但来回取很费时
关键:标准 attention 把大矩阵堆进了「仓库」,于是不停跑仓库

标准实现把整张 N×N 矩阵物化在 HBM(仓库):算 QKᵀ 写进去 → 读出来做 softmax → 写回 → 再读出来乘 V。 每一步都在书桌和仓库之间长途搬运,而矩阵是 N²大的——序列一长,搬运量爆炸。这就是慢和吃显存的真因,和「算力」关系不大。

三、核心洞察:让大矩阵「永不落仓库」+ 结果还分毫不差

FlashAttention 三招连用,核心目标只有一个:别把 N×N 矩阵整个写进 HBM

做什么
① 分块(tiling)把 Q、K、V 切成小块,一次只搬一小块进 SRAM(书桌),在书桌上把这块的注意力算完,永不物化整张大矩阵。
② 在线 softmaxsoftmax 本需要看完一整行才能归一化。FlashAttention 边流过小块边算,靠维护「running max + running sum」两个跑动统计量,逐块累加、动态校正——不必先存下整行
③ 反向重计算训练反向传播本需那张大矩阵。FlashAttention 不存它,需要时用统计量当场重算——多花一点点计算,换巨量显存节省。因为本就是内存瓶颈,这笔交易稳赚。
最关键的一点:结果精确相同,不是近似

FlashAttention 没有改任何数学——它算出的结果和标准 attention 逐位相同,只是把计算重排成对硬件友好的顺序。 这和你在 GPT-3 见的稀疏注意力(少算一些对、近似、有损)是两条路:稀疏改数学换速度;FlashAttention 数学不动、只改 IO 顺序,白赚速度和显存、零质量损失

四、动手玩:序列拉长,看「显存」怎么爆 vs 怎么不爆

点不同的序列长度 N。看标准 attention 要物化的注意力矩阵显存(∝ N²)怎么指数爆炸,而 FlashAttention 的额外开销只是 ∝ N

选序列长度 N: 
标准 attention(物化 N×N 矩阵)
FlashAttention(O(N) 额外开销)

看出门道了吗:序列越长,标准做法的矩阵显存按 N² 爆炸,FlashAttention 按 N 线性走。这就是为什么「几万、几十万 token 上下文」从不可行变可行——不是 GPU 变大了,是别再把那张大矩阵堆进仓库了。

五、对照表(一):标准 attention vs FlashAttention

维度标准 attentionFlashAttention
瓶颈认知以为是算力(O(N²) FLOPs)看穿是内存搬运(IO)
N×N 矩阵整个物化到 HBM分块进 SRAM,永不物化
显存O(N²)O(N)(线性)
速度基准快数倍(长序列更明显)
结果逐位精确相同(无损)

六、对照表(二):三条「让 attention 更快」的路

路线怎么做代价
稀疏注意力每个 token 只看一部分,少算近似、有损(改了数学)
低秩 / 线性近似用低秩结构逼近注意力近似、有损
FlashAttention不少算,只改内存搬运顺序无损、精确,纯白赚

🔮 FlashAttention 是「系统/硬件视角」的胜利:同样的数学,懂 GPU 内存层级就能快好几倍。它现在是几乎所有 LLM 训练/推理栈的默认件,后续还有 FlashAttention-2 / -3 继续榨硬件。它和 LoRA 合起来,是阶段三「效率」的两面:LoRA 改「算什么」、FlashAttention 改「怎么算」。

七、映射到你的项目

你工作里的东西其实就是这篇论文的什么
你做实时语音 agent 等低延迟应用,「上下文一长,延迟就飙」正是 attention 的内存搬运瓶颈——FlashAttention 直接压低这块延迟,是长上下文低延迟的工程底
你算「这个长上下文请求要吃多少显存 / 能不能放下」记住:标准是 O(N²)、Flash 是 O(N)——长上下文能不能落地,差别全在这
纠结「为了省钱要不要上稀疏/近似注意力,会不会掉效果」先确认有没有用上 FlashAttention:它是无损白赚的,该先吃干净;近似方案是有损的,最后再考虑
工程师说「换上 FlashAttention 快了 3 倍」,你想判断真假合理——它不改结果、只改 IO,长序列上数倍提速是常态,且不该有质量变化(若效果变了,那不是 FlashAttention 的锅)
更通用的判断力「慢」未必是「算得多」——先分清 compute-bound 还是 memory-bound,再决定优化方向。拍脑袋优化常优化错地方

八、检索练习(关掉上文,凭记忆答)

1. FlashAttention 看穿的真正瓶颈是?
2. FlashAttention 怎么避开瓶颈?
3. FlashAttention 和稀疏注意力的根本区别?

九、常见问题 FAQ

把这一课接到你真实工作上的几个关键问。点开看答。

FlashAttention 把显存降到 O(N),那「长上下文很贵」是不是就解决了?

解决了一大块,但没全解决。FlashAttention 让注意力矩阵的显存从 O(N²) 降到 O(N),这是长上下文的最大拦路虎之一,被它搬走了。但还有两笔账它管不到:① 计算量仍是 O(N²)(少搬不等于少算,长序列的 FLOPs 还在涨,所以更长还是更慢、更贵,只是不再爆显存);② 推理时的 KV cache 随上下文线性增长,仍占大量显存。

所以长上下文是「多线作战」:FlashAttention 砍注意力矩阵显存,KV cache 压缩 / 量化是另一条线,稀疏 / 线性注意力又是一条(用质量换 FLOPs)。FlashAttention 是其中最该先吃的那口——因为它无损白赚。

「在线 softmax」凭什么不看完整行就能算对?

softmax 要除以「一整行 exp 之和」,看着必须先有整行。诀窍是维护两个跑动统计量、边走边校正:当前见过的最大值和当前的exp 之和。每来一个新块,若出现更大的值,就把之前累计的结果按比例缩放一下(用新旧最大值之差校正),再把新块累加进去。

直觉:不是「先收齐再算」,而是「来一块算一块,发现算多了/少了就回头按比例修正」。流过所有块后,结果和「一次看整行」分毫不差——这就是它能既分块、又精确的数学保证。(减最大值那步同时也是数值稳定的标准技巧。)

这是纯训练优化,还是推理时我也享得到?我不训模型有关系吗?

训练和推理都享得到(重计算那招只在训练反向用,前向的分块+在线 softmax 推理同样生效)。你不训模型也直接受益:你调的几乎所有现代推理引擎(vLLM、TGI 等)底层默认就用 FlashAttention(或其变体)——你的每一次长上下文请求,延迟和显存都被它压过了。

对你做产品的实际意义:选推理框架/服务时,确认它用了 FlashAttention-2/3 是个基本盘。这也是为什么同样的模型,不同推理栈延迟能差很多——底层 attention 内核的工程水平是关键变量之一。

怎么判断一个算子是 compute-bound 还是 memory-bound?

算术强度(arithmetic intensity)= 计算量 ÷ 搬运字节数(FLOP/byte),和这台 GPU 的「平衡点」(峰值算力÷峰值带宽,roofline 脊点,如 A100 ≈150 FLOP/byte)比:高于它 → compute-bound;低于它 → memory-bound

秒判法则——「每个加载进来的数被反复用很多次,还是碰一次就扔?」反复复用 → compute-bound(大稠密矩阵乘 GEMM);碰一次就过 → memory-bound(elementwise、LayerNorm、softmax、reshape、attention 矩阵的读写)。attention 落在后者,因为矩阵「瘦」(head_dim 小)、耗时在 softmax+搬运。

实测:profiler 看「SM 计算利用率 vs DRAM 带宽利用率」哪个贴近 100%。快速验证:降精度没变快、但算子融合减少搬运就变快 → memory-bound

为什么 GPU 算力会富余、带宽反而是瓶颈?

一个词:内存墙(memory wall)——几代 GPU 下来,算力增长远快于带宽。V100→A100→H100 峰值 FLOPs 靠 tensor core + 低精度(fp16/bf16/fp8)翻了几十倍,HBM 带宽只涨几倍。于是「FLOP/byte 平衡点」被越推越高,越来越多算子掉到脊点以下、变 memory-bound

为什么涨得不一样快:堆算力(加 ALU)便宜可扩展;提带宽受物理限制(引脚、信号、功耗、内存接口),追不上。所以「算力富余」不是算力无限,而是芯片每秒能做的数学远超它每秒能供上的数据——算术强度不够,ALU 就空等 HBM。这正是近年 LLM 系统工作(融合、FlashAttention、量化)几乎都在优化「搬运」而非「计算量」的原因。

FlashAttention 和 PagedAttention / KV cache 是一回事吗?

不是,三件不同的事,叠着用,各打长上下文成本的一部分:

管的是解决什么阶段
FlashAttention怎么算 attention注意力矩阵的搬运训练+推理
KV cache存什么避免重算decode 每步重算历史 K/V推理(decode)
PagedAttention怎么管 KV cache 显存显存碎片/浪费、提并发推理服务

KV cache:缓存历史 token 的 K/V 免重算,但随上下文线性增长、是显存大户——FlashAttention 不缩它。PagedAttention(vLLM):借操作系统分页思路,把 KV cache 切成固定页非连续存放,消碎片、相同前缀跨请求共享页、塞更多并发。一句话:FlashAttention 管「算」、KV cache 管「存」、PagedAttention 管「这块显存怎么排布」。

本课主源

读原典:Dao et al. — FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness(2022)

想看后续榨硬件:搜 FlashAttention-2(2023)/ FlashAttention-3(2024,针对新一代 GPU)。

阶段三「规模与效率」收尾在即:LoRA 省微调 · FlashAttention 省 attention 显存/延迟。下一篇 P14 MoE / Switch Transformers:为什么前沿模型能又大又(相对)便宜——靠「每次只激活一部分专家」。

我是你的老师,随时提问。 比如:「compute-bound 和 memory-bound 我怎么判断一个算子是哪种?」「为什么 GPU 算力会富余、带宽反而是瓶颈?」 「FlashAttention 和 PagedAttention / KV cache 优化是一回事吗?」——别带着模糊往下走。

参考:
Dao et al., FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness, 2022.