【AI工程师精讲】05:FlashAttention:它不是算得更快,是搬得更少

理论上,注意力机制的计算量并不夸张:它主要是两次矩阵乘法。按显卡的算力算,处理一个几万 token 的序列应该很快。

但真实跑起来,显存占用会突然暴涨,速度也远不如预期。

问题的根源不在"算了多少",而在"搬了多少"。

要理解这件事,需要先接受一个反直觉的事实:

在现代加速器上,真正决定性能的往往是内存访问,而不是运算次数。

一篇被广泛引用的研究(关于 FlashAttention 的原始论文)就指出,很多操作的瓶颈是访存,而其算术强度低得惊人------算力单元大部分时间在等数据。

FlashAttention 正是冲着这件事去的。而且它的方案有点特别:它不改变计算结果,只改变计算的过程。


第一部分:硬件基础------GPU 的两级内存

所有加速器都有一个共同的层级结构。理解它就理解了后面的一切。

1.1 两个层级

层级 典型名称 容量 速度 住着什么
大容量内存 HBM / 显存 几十 GB 慢 模型权重、KV Cache、中间结果
片上高速缓存 SRAM / 共享内存 几百 KB ~ 几 MB 极快 当前正在处理的那一小块数据

两者之间的带宽差距通常在一个数量级以上。

这个差距直接导出一条黄金法则:

能待在片上的数据,就别往大内存里搬。

1.2 为什么"搬得少"比"算得快"更重要

这里需要引入算术强度的概念------之前讲 KV Cache 时提过:

算术强度 = 运算次数 ÷ 搬运字节数

  • 算术强度高 → 搬一次数据能算很多次 → 计算密集,瓶颈在算力;
  • 算术强度低 → 搬一次数据只算一两次 → 访存密集,瓶颈在带宽。

现代 GPU 的算力增长速度,远快于显存带宽的增长速度。结果是:越来越多的操作从"计算密集"滑向了"访存密集"。

算力在变便宜,带宽在变贵。

这就是"IO 墙"的来源:计算单元利用率很低,因为它在等数据。

1.3 把这条法则套到注意力上

现在看标准注意力的三步流程:

  1. 用 Q 乘 K 的转置,得到一个 N×N 的注意力分数矩阵;
  2. 对每一行做 softmax 归一化;
  3. 用归一化后的矩阵去乘 V,得到输出。

问题出在第一步和第二步之间:

那个 N×N 的矩阵,需要被写入大内存(HBM),然后再读回来做 softmax。

一次写入 + 一次读取,全是慢速内存访问。而这个矩阵的体积非常夸张。


第二部分:标准注意力的问题------中间矩阵的体积爆炸

2.1 平方级增长

序列长度翻倍,这个矩阵的元素数是四倍(因为它是 N×N)。

具体算一下(单头、单批次、fp16 存储):

序列长度 N N² 元素数 矩阵体积
8K 6710 万 约 134 MB
32K 10.7 亿 约 2.1 GB
128K 171.8 亿 约 34 GB

注意这是"单头、单批次"的数字。 实际模型有几十个头,批量又可能不止 1------虽然真实实现会做分块,不会真的物化完整的 N×N 矩阵,但即便是分块后的中间结果,搬运量也相当可观。

2.2 结果:算力闲着,带宽爆满

于是出现一个很讽刺的局面:

  • 计算量本身不大 ------ 两次矩阵乘而已,显卡的算力单元处理起来绰绰有余;
  • 搬运量极大 ------ 那个巨大的中间矩阵要在慢速内存里来回走。

显卡的算力单元大量时间都在等数据,利用率很低。这就是典型的"IO 墙"。

2.3 长序列下问题被放大

因为体积是平方级的,序列越长,这个瓶颈越严重:

  • 短序列(1K 以下):中间矩阵很小,不是主要瓶颈;
  • 长序列(32K 以上):中间矩阵成为显存和带宽的双重杀手。

这解释了一个常见现象:为什么模型标称支持 128K,但真的喂进 128K 时,显存会突然爆掉、速度会断崖式下跌?

因为在这一档,那个平方级的中间矩阵开始主导一切。


第三部分:FlashAttention 的三个关键动作

它的思路是把整个计算重构,让中间矩阵永远不落到大内存里。具体三个动作。

3.1 动作一:分块(Tiling)

做法 :把 Q、K、V 切成能塞进片上缓存的小块。一次只处理一块 Q 与一块 K,算出的分块注意力就留在片上用掉。

分块大小的选择:根据可用的 SRAM 容量精调。太小 → 片上的优势用不上;太大 → 塞不进片上缓存,又要落回大内存。

这一步的价值 :它把"必须一次性处理整行"变成了"可以一块一块处理"。 但这里立刻遇到一个障碍------softmax 需要看完整行才能算。

3.2 动作二:在线 softmax(最巧的一步)

问题所在 :softmax 需要每行的最大值 (用于数值稳定)和归一化分母 (所有指数的和),而这些需要看完整行才能算准------但分块计算时你只看得到一部分。

解决办法是"边走边修正":

  1. 维护一个滚动最大值 m 和一个滚动分母 l;
  2. 每处理一个新块,先看这个块里有没有更大的值,如果有就更新 m;
  3. 更新 l 时,把之前累积的结果按比例缩放(因为最大值的基准变了);
  4. 继续累加输出。

关键结论:

最终结果与一次性计算完全一致,但不需要保存整行。

为什么它能做到等价? 因为它维护的是可增量修正的统计量。softmax 的数学形式允许你"分步累加 + 事后缩放",只是顺序换了。

3.3 动作三:重计算而非存储

问题所在:反向传播需要注意力矩阵。既然不能存,怎么办?

做法 :在反向时用分块再算一遍。

为什么这笔交易划算?

方案 代价
存储中间矩阵 大量显存 + 大量搬运
反向时重算 一点额外计算

在这个瓶颈下(访存密集),用计算换搬运几乎总是赚的。 因为计算资源相对充裕,而带宽是稀缺的。

3.4 三个动作合起来的效果

维度 改善
显存占用 从随序列长度平方增长 → 线性增长
速度 提升数倍,长序列下收益更明显
结果 精确等价,不是近似
决策成本 可以直接替换,无需评估质量损失

第四部分:为什么它是"精确"的------这件事很重要

这一节专门澄清一个常见的误解。

4.1 误解的内容

网上常见一个说法:FlashAttention 是"稀疏注意力"或"近似注意力",会掉一点精度。

不是。

它做的事情只是换了计算的顺序和数据的存放位置,数学上等价于标准注意力。

所以它不会带来质量损失,也不需要考虑精度调参。

4.2 这个区别在工程上意味着什么

技术类型 本质 上之前要做什么
近似方法(稀疏化、低秩、窗口注意力) 用质量换速度 必须评估效果损失、做 A/B 测试
FlashAttention 用更聪明的调度换速度 可以直接替换,无需评估质量

这就是它能成为默认选项的原因------升级它没有决策成本。

对比一下两者的推广阻力:

  • 近似方法:要评估质量损失 → 要跑评测 → 要权衡 → 决策慢、易被否决;
  • 精确优化 :无质量损失 → 直接开 → 推广阻力最小。

4.3 一条通用的判断准则

只要不损失质量,升级就没有决策成本。

这解释了为什么"精确的工程优化"往往比"近似的算法优化"更快被全行业接受。不是因为它效果更好,而是因为它更容易被决策。


第五部分:它与其他优化的分工

初学者容易把这一堆技术混成一团。其实它们解决的是完全不同层面的问题。

5.1 四种技术的分工

技术 解决什么 优化的是 主要代价
FlashAttention 计算过程怎么不搬数据 注意力的执行方式 实现复杂度
KV Cache 历史结果怎么复用 跨步骤的重复计算 必须把历史 K/V 存在显存里
GQA / MQA 缓存体积怎么变小 缓存本身的体积 轻微质量损失
PagedAttention 缓存怎么分配不浪费 显存管理 实现复杂度

5.2 关键点:它们可以叠加

因为解决的是不同层面的问题,所以它们不互斥,而是互补:

text 复制代码
GQA            → 把每 token 的缓存压小(改模型结构)
PagedAttention → 让它分配得更密(改内存管理)
KV Cache       → 让历史不复算(改计算流程)
FlashAttention → 让每步计算不搬多余数据(改算子实现)

这也是为什么新一点的推理框架,这些技术基本都是标配------它们各自贡献一部分收益,加起来才是完整的优化。

5.3 一个容易混淆的点

很多人以为 FlashAttention 和 KV Cache 是"二选一"或者"作用重叠"。

实际上它们优化的是完全不同的阶段:

  • KV Cache 解决跨步骤的问题:第 N 步不用重算前 N-1 步;
  • FlashAttention 解决单步内部的问题:这一步的计算过程别搬多余数据。

一个是"时间维度上的复用",一个是"空间维度上的调度"。


第六部分:什么时候收益最大,什么时候感觉不到

收益大小取决于"你原来被 IO 卡得有多狠"。

6.1 收益最大的情况

① 长序列。 N 越大,标准实现的中间矩阵越夸张,而分块方法的优势越大。

长上下文场景是它最大的舞台。

② 训练的注意力层。 训练时需要保存激活以供反向,标准实现的显存压力更大。而且训练通常用更长的序列,收益叠加。

③ 批量小但序列长。 这类负载本来就 IO 受限,改善立竿见影。

6.2 收益不明显的情况

① 序列很短。 中间矩阵本来就很小,分块带来的调度开销可能反而吃掉收益。

② 瓶颈不在注意力上。 这一点最重要。如果整个推理的耗时主要花在:

  • 别的层(比如 FFN);
  • 访存密集的解码阶段;
  • 外部工具调用、网络往返;

那优化注意力只能改善其中一小部分。

6.3 正确的做法:先分段计时

别盲目上技术,先看分段耗时。

具体做法是把一次请求的耗时拆成几段:

分段 典型占比(长文本场景)
预填充(含注意力) 高
解码(逐 token) 中 ~ 高
工具调用 / 网络 视架构而定
后处理 低

如果注意力只占 15%,那把它优化到极致也只能拿到 15% 的改善。

一句话:它可能是长序列的强心针,但不是万能加速器。

6.4 不同问题的对应技术

你遇到的问题 FlashAttention 能帮吗 该看的技术
长序列时显存爆掉 能,减少中间矩阵占用 FlashAttention + KV Cache 量化
首字延迟高 部分能(预填充受益) 前缀缓存、并行化
并发上不去 间接能 GQA、PagedAttention、限上下文
长文本答不准 不能 检索、分层摘要、位置组织
训练时显存不够 能,收益明显 FlashAttention + 梯度检查点
短序列推理慢 基本不能 查分段耗时,先定位真瓶颈

注意第四行 :长文本答不准是能力问题 ,不是效率问题。这一类靠工程效率优化是解决不了的,必须回到检索和上下文组织。


第七部分:四条通用工程思维

这个案例的价值不只在它本身,还在于它示范了一种思路。这些思路可以迁移到其他优化问题上。

思维一:先找瓶颈类型,再选优化手段

"计算密集"和"访存密集"需要的解法完全不同。

往访存瓶颈上堆算力,就像给堵车的路加宽收费站。

具体到实践 :在你决定"要不要换更好的卡"之前,先确认瓶颈是在算力还是在带宽。如果是带宽,换卡可能几乎没有效果。

思维二:中间结果能不上大内存就不上

融合算子、算子合并、分块计算------背后都是同一条原则。

这个原则的应用很广:

  • 多个小算子合成一个大算子(减少中间张量的读写);
  • 归一化与激活函数融合(省掉一次写回);
  • 分块处理(把工作集限制在片上容量内)。

判断标准很简单:这个中间结果,能不能不落地?

思维三:用计算换搬运,在访存瓶颈下往往是赚的

反向重计算就是典型例子。类似的取舍在工程中很常见:

  • 用 CPU 卸载换显存(牺牲速度换容量);
  • 用重新计算换存储(牺牲算力换带宽);
  • 用压缩换显存(牺牲精度换容量)。

关键在于判断"你缺的是什么"。缺带宽就用计算换,缺容量就用速度换。

思维四:精确优化优先于近似优化

只要不损失质量,升级就没有决策成本,推广阻力最小。

这也是它比稀疏注意力更快被全行业接受的原因。

落到实践上有一条很实用的排序原则:

在同等收益下,优先选那些"不需要做质量评估"的优化。

因为它们能立刻上、不会引发争论、不需要等评测周期。而需要质量评估的优化,往往在评审会上就卡住了。


常见误区

误区一:"FlashAttention 是近似算法,会掉一点精度。"

它是精确等价的。

掉精度的是稀疏化、低秩、窗口注意力这类方法。这个区别在选型时很重要------前者可以直接开,后者必须做评测。

误区二:"上了 FlashAttention,长上下文就没问题了。"

它只解决注意力的执行效率。

长上下文的另外两座山------KV Cache 显存占用 和模型有效上下文不足------它一个都管不了。

第三座山尤其重要 :就算显存够、速度也快,模型在很长的位置上也未必能准确取到信息。那是能力问题,不是效率问题。

误区三:"它是推理优化技术,训练用不上。"

恰恰相反。训练侧的收益常常更明显,因为:

  • 训练需要保存更多中间结果以供反向;
  • 训练的序列通常更长;
  • 显存压力更大(还要存优化器状态、梯度)。

误区四:"升级到最新版本一定更好。"

新版本通常针对新硬件特性做了优化(比如新的数据类型、新的指令集)。

在老硬件上,收益可能有限,需要实测而不是照搬结论。

这一点在采购决策上尤其重要:不要因为"论文里提升 3 倍"就认为自己的三年前的老卡也能提升 3 倍。

误区五:"注意力是瓶颈,所以优化注意力就对了。"

先确认它是不是瓶颈。 做一个分段计时,看注意力占多少。

在很多实际系统里,瓶颈可能根本不在模型计算上------工具调用、网络往返、检索延迟、后处理,这些加起来可能超过一半的时间。

误区六:"分块越大越好,减少分块次数。"

分块大小受片上缓存容量 限制。分得太大塞不进片上,就要落回大内存------那反而回到了原来的问题。

分块大小是一个需要针对具体硬件调优的参数,不是"越大越好"。


实战问答

Q1:面试被问"FlashAttention 解决了什么问题",怎么答能显出层次?

A:按"判断 → 机制 → 结论"三层走。

第一层,给出核心判断:

注意力的瓶颈是访存,不是算力。

第二层,说清机制 :标准实现会写出一个 N×N 的中间矩阵(体积随长度平方增长,要在大内存里写一遍读一遍)。FlashAttention 用分块 + 在线 softmax 把中间结果留在片上,并用重计算换掉存储。

第三层,给结论:

显存从平方级降到线性、速度提升数倍,且是精确算法,可以直接替换、无质量损失。

最后这一点往往是加分项------它展示的是"我知道这个优化能不能直接上"。


Q2:为什么我换了支持 FlashAttention 的框架,速度只提升了一点?

A :大概率是你的瓶颈不在注意力上。

第一步先做分段计时:预填充多久、注意力多久、解码多久、工具调用多久。

可能的情况:

观察 原因
解码阶段占大头 那是访存密集的逐 token 生成,注意力优化帮不上
外部接口耗时占大头 优化模型根本不影响总耗时
序列很短 中间矩阵本来就小,分块收益有限

"先测量,再优化"这条原则,在性能优化里永远成立。


Q3:长上下文场景,应该优先做哪几件事?

A:按投入产出比排序:

优先级 动作 成本 收益
① 前缀缓存 几乎零成本 首字延迟立降
② 换 GQA 模型 需要重新评测 缓存体积直接小几倍
③ FlashAttention 通常只是开关 长序列收益明显
④ KV 量化 需要评估质量 显存再压一档

四件事可以叠加,不用二选一。

顺序设计的逻辑是"成本从低到高":前两项几乎立刻见效,第四项需要评估质量损失,所以排在最后。


Q4:在线 softmax 为什么能做到和一次性计算完全一样?

A :因为它维护的是可增量修正的统计量。

具体是两个值:

  • 滚动最大值 m:用于数值稳定(防止指数溢出);
  • 滚动分母 l:所有指数值的和。

每处理一个新块:

  1. 看这个块里有没有更大的值,有就更新 m;
  2. 因为基准变了,把之前累积的输出按新分母重新缩放;
  3. 继续累加。

数学上等价于"先看完整行再算",只是顺序换了。

这个技巧的价值在于:它把"必须看到全部数据才能算"的操作,变成了"可以流式处理"的操作。 这个思路在其他流式计算场景里也很常见。


Q5:怎么判断我该上 FlashAttention 还是该上稀疏注意力?

A:看你能接受什么代价。

FlashAttention 稀疏 / 窗口注意力
质量 完全等价 有损失,需评估
显存 降到线性 可能更低
适用长度 受硬件限制 可做更极端的长文本
决策成本 低(直接开) 高(要评测)

建议的顺序:

  1. 先开 FlashAttention ------ 无成本,先拿到这部分收益;
  2. 测量是否还缺显存或速度;
  3. 如果还缺,再考虑近似方法,并配套做质量评测。

先摘低垂果实,再考虑需要权衡的方案。


Q6:训练时显存不够,FlashAttention 能帮多少?

A:能帮一部分,但通常需要和其他手段组合。

单独看 :它消除了注意力那个 N×N 中间矩阵的存储需求,在长序列训练上收益明显。

但训练显存还有另外几个大头:

占用项 对应手段
激活值(供反向使用) 梯度检查点(重计算换显存,同一个思路)
优化器状态(Adam 会占权重 2 倍) 优化器状态分片
梯度 梯度累积、分片
权重 混合精度、ZeRO 分片

所以正确的说法是:FlashAttention 是训练显存优化的一块拼图,不是万能药。

注意这里有个有趣的呼应 :梯度检查点用的也是"重计算换存储"这个思路------和 FlashAttention 的反向重计算是同一个哲学。


术语表

术语 含义
FlashAttention 通过分块与在线 softmax 避免写出 N×N 注意力矩阵的精确注意力实现
HBM / SRAM 大容量慢速内存 / 片上高速缓存,两者带宽差距是优化的起点
算术强度 每搬运 1 字节数据能完成多少次运算,用于判断计算密集还是访存密集
IO 墙 计算单元因等待数据搬运而利用率低下的现象
分块(Tiling) 把大计算切成能放进片上的小块,逐块处理
在线 softmax 滚动维护最大值与归一化分母,使分块计算与整体计算等价
重计算(Recomputation) 反向传播时重新计算中间结果,以省下存储与搬运
梯度检查点 训练中只保存部分激活、其余反向时重算,与重计算是同一思路

结语

把这篇压缩成三句话:

第一,现代加速器上,"搬得少"往往比"算得快"更重要。 算力在变便宜,带宽在变贵,越来越多的操作滑向了访存密集------注意力就是典型。

第二,FlashAttention 的一切动作,都围绕一个目标:让那个平方级的中间矩阵不落到大内存里。 分块解决"怎么切",在线 softmax 解决"切了之后怎么还能算准",重计算解决"反向怎么办"。

第三,它最容易被低估的优点,是"无质量损失"。

一个不损失质量的优化,可以直接上;一个需要权衡的优化,要在评审会上耗掉几周。

这就是精确优化在工程推广上的结构性优势。

而如果只允许带走一条思维,那就是:遇到性能问题,先问"这是算力问题还是访存问题" ------ 这一个问题,能过滤掉相当一部分"花了大钱却没用"的优化尝试。

顺着这个思路往下,还有一个经常被误解的架构值得一看:MoE(混合专家)------为什么一个"万亿参数"的模型,实际只激活一小部分,还能更强? 它的收益和代价,比宣传里说的都更具体。

相关推荐
C++ 老炮儿的技术栈1 小时前
char (*csConnectName)[256]; 和 char csConnectName [256] 有什么区别
java·c语言·开发语言·c++·人工智能·算法·c
Rocky Ding*1 小时前
VideoChat3 深度解析:视频理解的效率,来自时空压缩与主动感知的共同设计
论文阅读·人工智能·深度学习·机器学习·aigc·ai-native·视频理解
禹凕1 小时前
深度学习之激活函数(Deep Learning about Activation function)
人工智能·深度学习
loulanyue_1 小时前
脑机接口:意动·芯动·行动——读同济医院副院长廖家智2026云栖演讲
人工智能·深度学习·云栖大会
Dcr_stephen1 小时前
电商 Agent 实践:为什么运营分析与产品调研,需要两套相反的工作流?
人工智能
Omics Pro1 小时前
全新可重分析!代谢组质谱专用
数据库·人工智能·算法·机器学习·自然语言处理
一只废狗狗狗狗狗狗狗狗狗1 小时前
Network架构1——卷积神经网络CNN
人工智能·算法·cnn
鲲穹AI种草1 小时前
桌面与手机美化怎么做,多款 AI 壁纸生成工具使用记录
人工智能·壁纸工具
Rain的Java大神实战圈1 小时前
🔥8年Java老兵转型AI Agent:90%的人挂在同一个坑,根本不用学Python!(万字实战复盘,建议收藏)
ai编程