2.1 FlashAttention-2 的核心思想:把注意力算在 SRAM 里


文档摘要

2.1 FlashAttention-2 的核心思想:把注意力算在 SRAM 里 读者读完本节应带走的一句话:FlashAttention-2 不改注意力的数学结果,它只是换了一种「不把巨大分数矩阵写进显存」的计算顺序,靠 Tiling + 在线归一化 + 更聪明的 GPU 分工,把带宽瓶颈打掉,从而做到无损加速。 第一章我们算过一笔账:标准注意力需要物化一个 的分数矩阵,单序列显存占用是 O(n²),而且这个矩阵必须落在 HBM(显存)上反复读写。当序列长度 n 从 2K 涨到 32K、128K,O(n²) 会直接把显存和带宽同时压垮。FlashAttention-2 的全部努力,都围绕一个目标:让注意力计算全程待在 GPU 的 SRAM(片上高速缓存)里,尽量不碰 HBM。

2.1 FlashAttention-2 的核心思想:把注意力算在 SRAM 里

读者读完本节应带走的一句话:FlashAttention-2 不改注意力的数学结果,它只是换了一种「不把巨大分数矩阵写进显存」的计算顺序,靠 Tiling + 在线归一化 + 更聪明的 GPU 分工,把带宽瓶颈打掉,从而做到无损加速。

第一章我们算过一笔账:标准注意力需要物化一个 [n, n] 的分数矩阵,单序列显存占用是 O(n²),而且这个矩阵必须落在 HBM(显存)上反复读写。当序列长度 n 从 2K 涨到 32K、128K,O(n²) 会直接把显存和带宽同时压垮。FlashAttention-2 的全部努力,都围绕一个目标:让注意力计算全程待在 GPU 的 SRAM(片上高速缓存)里,尽量不碰 HBM

本节我们把它拆解成三个关键设计,逐个讲透,并告诉你它到底「无损」在哪里、又在什么情形下收益没那么夸张。

2.1.1 先确认敌人:标准实现到底慢在哪

把标准缩放点积注意力(scaled dot-product attention)朴素地写成 PyTorch,本质是三步:

# Q, K, V 形状均为 [batch, heads, n, d] scores = Q @ K.transpose(-2, -1) / math.sqrt(d) # 得到 [n, n] attn = torch.softmax(scores, dim=-1) # 沿最后一维归一化 out = attn @ V # 得到 [n, d]

问题不在矩阵乘本身——现代 GPU 的 Tensor Core 算 Q@K 非常快。真正的代价藏在「数据搬移」上:

  • scores[n, n],在算完之后要写回 HBM
  • softmax 又要把它从 HBM 读回来做归约(求最大值、求和);
  • 归一化后再写回,最后和 V 相乘时再读。

HBM 的带宽(A100 约 1.5–2 TB/s)虽然已经很大,但相比 SRAM(A100 约 19–20 TB/s)差了一个数量级,而且容量极小(A100 的 SRAM 每个 SM 约 220 KB,整卡约 20 MB 量级)。注意力的「算术强度」其实不低,但被这些反复的 HBM 往返吃掉了绝大部分收益。更致命的是 O(n²) 的显存:n=8192 时单序列分数矩阵就要约 268 MB(bf16),多卡多请求的上下文里这根本放不起。

所以 FlashAttention-2 的第一性原理是:能不能不算出完整的 [n, n],也把最终输出 O 算对? 答案是能,靠 Tiling 和在线归一化配合。

FlashAttention-2 的 Tiling 分块加载

2.1.2 设计一:Tiling(分块)—— 一次只看一小块

Tiling 的思想很朴素:把 Q、K、V 沿序列维度切成很多小块(tile),比如每次只加载 Q 的一个行块(对应若干 query 位置)和 KV 的一个列块(对应若干 key/value 位置)到 SRAM,在 SRAM 内完成这一小块分数的计算与部分累加,算完就丢弃,绝不把整张 [n, n] 留在 HBM。

关键约束是:任意一个小块的计算,必须能「局部」地推进全局正确结果。这要求归一化也能「局部推进」——这正是第二个设计要解决的。

Tiling 带来的直接好处是显存从 O(n²) 降到 O(n):我们只需要在 SRAM 里保留「当前正在处理的 Q 块」以及尚未完成的「输出累加器 O(每个 query 位置一份)」,外加运行中的归一化统计量,而不需要为所有 query-key 对保留中间分数。这意味着长序列下显存几乎随 n 线性增长,而不是平方增长——这是 FlashAttention 系列能支撑超长上下文推理的底层原因。

2.1.3 设计二:Online Softmax(在线归一化)—— 边读边归一化

标准 softmax 之所以必须先有整行数据,是因为它依赖整行的最大值 m 和指数和 l

softmax(x)_i = exp(x_i - m) / sum_j exp(x_j - m), 其中 m = max_j x_j

ml 都要看到整行才能算对。FlashAttention 引入的「在线(online)」技巧,把这件事变成可递增的:每读入一个新 K/V 分块,就更新运行最大值 m、运行指数和 l,并对已经累积的输出 O 做修正。数学上可以证明,这种增量更新与一次性 softmax 完全等价。

把更新规则写出来(设上一步的统计量记为 m_oldl_oldO_old,新分块贡献记为 m_newl_newO_new):

m = max(m_old, m_new) # 新的全局最大值 l = exp(m_old - m) * l_old + exp(m_new - m) * l_new O = (exp(m_old - m) * l_old / l) * O_old + (exp(m_new - m) * l_new / l) * O_new

直觉是:每当发现更大的 m,之前累积的 Ol 都要按 exp(m_old - m) 这个因子「回退一下」重新加权,因为 softmax 对最大值极其敏感。这个修正因子保证最终结果和把整行一次性喂给 softmax 完全一致。

Online Softmax 的增量更新状态机

这就解释了为什么 Tiling 和在线归一化必须成对出现:Tiling 让我们一次只看一块,而在线归一化让「只看一块」也能逐步逼近全局正确结果。两者合起来,注意力就可以被「流式」地处理,永远不需要物化完整分数矩阵。

2.1.4 设计三:把 GPU 工作重新分工,压到最少的非矩阵乘操作

FlashAttention-1 已经解决了分块和在线归一化,但作者发现它还有冗余:softmax 及其前后的规约、重缩放(rescaling)、掩码等「非矩阵乘」操作,在朴素实现里仍会在 SRAM 和寄存器之间、以及 warp 之间产生大量通信与同步。FA-2 的核心改进之一是重新划分 warp 间的工作负载

  • 在 FA-1 里,每个 warp 负责 Q 的一个分块,并通过共享内存彼此交换 K、V 分块,带来不少同步开销;
  • FA-2 改为让一个 warp 同时持有 Q 和 K/V 分块,并沿「外层循环 K/V、内层循环 Q」的组织方式,把矩阵乘交给 Tensor Core,把 softmax 等非 matmul 工作降到最少,同时减少共享内存的读写往返。

形像地说,FA-2 把「能不能让 Tensor Core 一直在忙、少做杂活」做到了极致。这也是为什么同样的分块思想,FA-2 比 FA-1 还能再快一截——它不是改了算法数学,而是改了 GPU 上的「 choreography( choreography 指任务编排)」。

2.1.5 为什么说它是「无损加速」

这一点必须强调:FlashAttention 不改变注意力的数值结果,只是换了一种数值稳定的计算顺序。输出 O 在数学上与传统 softmax 注意力完全等价,只是因为在线归一化用 m 做了数值稳定(类似 log-sum-exp 技巧),浮点误差通常更小、更稳定,而非更大。

这意味着两件事:

  • 你用它替换朴素注意力,不需要重新训练模型,权重直接复用;
  • 它没有任何近似、不会牺牲质量,是「同样的结果、更少的显存和更高的速度」。在推理落地里,这是最舒服的一类优化——没有 accuracy 回退的代价。

2.1.6 实际收益大概是什么量级(带前提)

具体数字强烈依赖硬件、头维度 d、序列长度 n 和实现版本,以下只是「数量级」和「趋势」,不视为精确基准;落地时请以你所用硬件上的官方 benchmark 或自测为准:

  • 显存:从 O(n²) 降到 O(n),长序列下差异是数量级的;这也是它能跑长上下文的前提。
  • 速度:在长序列(如 n 从几千到几万)上,相比 PyTorch 朴素实现,FA-2 通常能拿到数倍到更高的加速,且序列越长、优势越大;短序列时 HBM 往返占比小,收益会更有限。
  • 与 FA-1 比:FA-2 在 A100 上进一步把 GPU 利用率(model FLOPs utilization)显著抬高,官方报告里可到 50%–70% 量级区间,但仍远低于理论上限,说明还有继续优化的空间。

注意:这些数字会随版本迭代变化,不要把它们当承诺。做性能对比时,务必固定同一输入形状、同一精度(fp16/bf16)、同一 batch 设置再测。

2.1.7 常见误解与踩坑

  • 误解一:「FlashAttention 是某种新注意力机制」。错。它是同一套数学的高效实现,权重可无缝替换。
  • 误解二:「只要开了 FA-2,推理就一定快好几倍」。不一定。短序列、或瓶颈本来就不在注意力(比如被 Embedding、LayerNorm、采样逻辑、通信带宽卡住)时,收益有限。先用第一章的 Roofline 思路定位瓶颈,再决定要不要上。
  • 误解三:「FA-2 能省 KV Cache 显存」。部分混淆了概念。FA-2 主要省的是「注意力分数矩阵」的 O(n²) 临时显存,并降低带宽压力;它不解决 KV Cache 在长对话、多请求下的碎片与预留浪费,那是 PagedAttention(本章 2.2)解决的问题。两者职责不重叠。
  • 踩坑:精度与 shape 支持。早期实现对某些头维度(如非 64/128 对齐)、变长序列、特定数据类型支持不完整,可能静默回退到朴素路径从而「没加速」。上线前用 Profiler 确认内核确实被调用。
朴素实现与 FlashAttention-2 的 HBM 访问对比

2.1.8 它是带宽受限还是计算受限:用 Roofline 再定位一次

第一章我们引过 Roofline 模型:一个算子的性能取决于它的「算术强度」(每搬一个字节能做多少次浮点运算)。注意力在长序列时,瓶颈恰恰在「搬数据」而不是「算」——因为分数矩阵要反复读写 HBM。FlashAttention-2 的本质,是把算术强度抬高:它不再物化 [n, n] 分数矩阵,于是每读一份 K/V 数据能多算出很多输出,等价于用同样的数据搬运量完成更多有效计算,从而把原本 memory-bound 的注意力往 compute-bound 方向推。

但请记住一个边界:如果模型的耗时主要被 MLP 里的大矩阵乘(GEMM)占据,或者受采样逻辑、通信、Embedding 限制,那么即便注意力再快,端到端收益也有限。所以第一章的「先定位瓶颈再上优化」在这里依然成立——FA-2 是治「注意力带宽瓶颈」的特效药,不是包治百病的万金油。

2.1.9 一个可以手算的在线归一化例子

光说「等价」不够直观,我们来手算一遍,确认在线归一化确实给出和标准 softmax 一样的结果。设某个 query 位置对四个 key 的原始分数(未缩放)为 [2, 5, 1, 4],分两块到达:块 A=[2,5],块 B=[1,4]

标准 softmax 作为对照:最大值 m=5,指数 [e^-3, 1, e^-4, e^-1] = [0.0498, 1, 0.0183, 0.3679],和 S=1.4359,最终概率 [0.0347, 0.6964, 0.0128, 0.2562]

现在用在线方式逐块推进,维护运行最大值 m 与运行权重和 S,并对历史权重做重缩放:

  • 处理块 A [2,5]:新最大值 m=5。未归一化权重 u_A=[e^(2-5), e^(5-5)]=[0.0498, 1.0]S=1.0498,归一化得 [0.0474, 0.9526]
  • 处理块 B [1,4],旧 m=5、旧 S=1.0498:新最大仍是 5,m 未变,历史权重无需重缩放。未归一化 u_B=[e^(1-5), e^(4-5)]=[0.0183, 0.3679],新总和 S'=1.0498+0.3862=1.4359。全量归一化 = 旧权重×(旧S/新S) + u_B/新S = [0.0474, 0.9526]×0.7312 + [0.0128, 0.2563] = [0.0347, 0.6965, 0.0128, 0.2563],与标准 softmax 一致。

如果中间 m 发生变化(比如再来块 C=[10]),历史权重会整体乘 exp(旧m - 新m)=exp(5-10)=e^-5 再重新归一化——这正是 2.1.3 里「回退重加权」因子的具体表现。这个例子说明:在线更新不需要看到整行,也能逐步逼近全局正确概率分布。

2.1.10 从 FA-1 到 FA-2:不只是分块,更是「编排」

FlashAttention-1 已经做到了分块与在线归一化,但它在 GPU 上仍不够「省」:softmax 前后的规约、重缩放、掩码等非矩阵乘操作,会在 warp 之间引发大量共享内存通信与同步。FA-2 的关键改进在「warp 间工作划分」:

  • FA-1 让每个 warp 负责 Q 的一个分块,不同 warp 需通过共享内存彼此交换 K/V 分块,同步代价高;
  • FA-2 让一个 warp 同时持有 Q 和 K/V 分块,并调整内外层循环顺序,把矩阵乘尽量交给 Tensor Core,把非 matmul 的杂活降到最少。

作者把这种改进称为对 GPU 上「任务编排(choreography)」的优化。效果上,它把 HBM 往返次数进一步压到更少(论文里强调把非 matmul 的 HBM 访问次数降到接近下限),从而在 A100 上把 GPU 利用率显著抬高。需要诚实地说:利用率离理论上限仍有距离,说明这条路线还能继续挖。

2.1.11 何时不该指望 FA-2 单挑

几条现实提醒,帮你避免误用:

  • 短序列(如 n 只有几百)时,分数矩阵本就小,HBM 往返占比低,FA-2 收益有限;它的优势随序列变长而放大。
  • 若端到端耗时主要在 MLP 的 GEMM、采样、或网络/通信,FA-2 救不了整体延迟。先 profiling 再决定。
  • 某些头维度、变长序列、特定精度的早期实现可能静默回退到朴素路径(「没加速」),上线前务必用 Profiler 确认内核真的被调用。
  • 它与 PagedAttention 职责不重叠:FA-2 不解决 KV Cache 碎片,那是 2.2 的主场。

2.1.12 小结:读者现在应该能回答的三个问题

  1. FlashAttention-2 靠哪三个设计把 GPU 利用率拉满?——Tiling(分块避免物化大矩阵)、Online Softmax(边读边归一化)、以及重新划分的 warp 分工(压到最少非 matmul 操作)。
  2. 它为什么无损?——只是换了数值稳定的计算顺序,数学结果与标准 softmax 完全等价。
  3. 它治的是哪类瓶颈?——注意力计算中的 HBM 带宽瓶颈与 O(n²) 临时显存,而不是 KV Cache 的碎片问题。

下一节 2.2 我们转入另一个战场:当模型要同时服务很多请求、KV Cache 越积越多时,如何用「分页」把显存碎片几乎消灭。


发布者: 作者: 搬砖请按F5的小龙虾 转发
评论区 (0)
U