2.1 FlashAttention-2 的核心思想:把注意力算在 SRAM 里 读者读完本节应带走的一句话:FlashAttention-2 不改注意力的数学结果,它只是换了一种「不把巨大分数矩阵写进显存」的计算顺序,靠 Tiling + 在线归一化 + 更聪明的 GPU 分工,把带宽瓶颈打掉,从而做到无损加速。 第一章我们算过一笔账:标准注意力需要物化一个 的分数矩阵,单序列显存占用是 O(n²),而且这个矩阵必须落在 HBM(显存)上反复读写。当序列长度 n 从 2K 涨到 32K、128K,O(n²) 会直接把显存和带宽同时压垮。FlashAttention-2 的全部努力,都围绕一个目标:让注意力计算全程待在 GPU 的 SRAM(片上高速缓存)里,尽量不碰 HBM。
读者读完本节应带走的一句话:FlashAttention-2 不改注意力的数学结果,它只是换了一种「不把巨大分数矩阵写进显存」的计算顺序,靠 Tiling + 在线归一化 + 更聪明的 GPU 分工,把带宽瓶颈打掉,从而做到无损加速。
第一章我们算过一笔账:标准注意力需要物化一个 [n, n] 的分数矩阵,单序列显存占用是 O(n²),而且这个矩阵必须落在 HBM(显存)上反复读写。当序列长度 n 从 2K 涨到 32K、128K,O(n²) 会直接把显存和带宽同时压垮。FlashAttention-2 的全部努力,都围绕一个目标:让注意力计算全程待在 GPU 的 SRAM(片上高速缓存)里,尽量不碰 HBM。
本节我们把它拆解成三个关键设计,逐个讲透,并告诉你它到底「无损」在哪里、又在什么情形下收益没那么夸张。
把标准缩放点积注意力(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;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 和在线归一化配合。
Tiling 的思想很朴素:把 Q、K、V 沿序列维度切成很多小块(tile),比如每次只加载 Q 的一个行块(对应若干 query 位置)和 K、V 的一个列块(对应若干 key/value 位置)到 SRAM,在 SRAM 内完成这一小块分数的计算与部分累加,算完就丢弃,绝不把整张 [n, n] 留在 HBM。
关键约束是:任意一个小块的计算,必须能「局部」地推进全局正确结果。这要求归一化也能「局部推进」——这正是第二个设计要解决的。
Tiling 带来的直接好处是显存从 O(n²) 降到 O(n):我们只需要在 SRAM 里保留「当前正在处理的 Q 块」以及尚未完成的「输出累加器 O(每个 query 位置一份)」,外加运行中的归一化统计量,而不需要为所有 query-key 对保留中间分数。这意味着长序列下显存几乎随 n 线性增长,而不是平方增长——这是 FlashAttention 系列能支撑超长上下文推理的底层原因。
标准 softmax 之所以必须先有整行数据,是因为它依赖整行的最大值 m 和指数和 l:
softmax(x)_i = exp(x_i - m) / sum_j exp(x_j - m), 其中 m = max_j x_j
m 和 l 都要看到整行才能算对。FlashAttention 引入的「在线(online)」技巧,把这件事变成可递增的:每读入一个新 K/V 分块,就更新运行最大值 m、运行指数和 l,并对已经累积的输出 O 做修正。数学上可以证明,这种增量更新与一次性 softmax 完全等价。
把更新规则写出来(设上一步的统计量记为 m_old、l_old、O_old,新分块贡献记为 m_new、l_new、O_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,之前累积的 O 和 l 都要按 exp(m_old - m) 这个因子「回退一下」重新加权,因为 softmax 对最大值极其敏感。这个修正因子保证最终结果和把整行一次性喂给 softmax 完全一致。
这就解释了为什么 Tiling 和在线归一化必须成对出现:Tiling 让我们一次只看一块,而在线归一化让「只看一块」也能逐步逼近全局正确结果。两者合起来,注意力就可以被「流式」地处理,永远不需要物化完整分数矩阵。
FlashAttention-1 已经解决了分块和在线归一化,但作者发现它还有冗余:softmax 及其前后的规约、重缩放(rescaling)、掩码等「非矩阵乘」操作,在朴素实现里仍会在 SRAM 和寄存器之间、以及 warp 之间产生大量通信与同步。FA-2 的核心改进之一是重新划分 warp 间的工作负载:
形像地说,FA-2 把「能不能让 Tensor Core 一直在忙、少做杂活」做到了极致。这也是为什么同样的分块思想,FA-2 比 FA-1 还能再快一截——它不是改了算法数学,而是改了 GPU 上的「 choreography( choreography 指任务编排)」。
这一点必须强调:FlashAttention 不改变注意力的数值结果,只是换了一种数值稳定的计算顺序。输出 O 在数学上与传统 softmax 注意力完全等价,只是因为在线归一化用 m 做了数值稳定(类似 log-sum-exp 技巧),浮点误差通常更小、更稳定,而非更大。
这意味着两件事:
具体数字强烈依赖硬件、头维度 d、序列长度 n 和实现版本,以下只是「数量级」和「趋势」,不视为精确基准;落地时请以你所用硬件上的官方 benchmark 或自测为准:
注意:这些数字会随版本迭代变化,不要把它们当承诺。做性能对比时,务必固定同一输入形状、同一精度(fp16/bf16)、同一 batch 设置再测。
第一章我们引过 Roofline 模型:一个算子的性能取决于它的「算术强度」(每搬一个字节能做多少次浮点运算)。注意力在长序列时,瓶颈恰恰在「搬数据」而不是「算」——因为分数矩阵要反复读写 HBM。FlashAttention-2 的本质,是把算术强度抬高:它不再物化 [n, n] 分数矩阵,于是每读一份 K/V 数据能多算出很多输出,等价于用同样的数据搬运量完成更多有效计算,从而把原本 memory-bound 的注意力往 compute-bound 方向推。
但请记住一个边界:如果模型的耗时主要被 MLP 里的大矩阵乘(GEMM)占据,或者受采样逻辑、通信、Embedding 限制,那么即便注意力再快,端到端收益也有限。所以第一章的「先定位瓶颈再上优化」在这里依然成立——FA-2 是治「注意力带宽瓶颈」的特效药,不是包治百病的万金油。
光说「等价」不够直观,我们来手算一遍,确认在线归一化确实给出和标准 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,并对历史权重做重缩放:
[2,5]:新最大值 m=5。未归一化权重 u_A=[e^(2-5), e^(5-5)]=[0.0498, 1.0],S=1.0498,归一化得 [0.0474, 0.9526]。[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 里「回退重加权」因子的具体表现。这个例子说明:在线更新不需要看到整行,也能逐步逼近全局正确概率分布。
FlashAttention-1 已经做到了分块与在线归一化,但它在 GPU 上仍不够「省」:softmax 前后的规约、重缩放、掩码等非矩阵乘操作,会在 warp 之间引发大量共享内存通信与同步。FA-2 的关键改进在「warp 间工作划分」:
作者把这种改进称为对 GPU 上「任务编排(choreography)」的优化。效果上,它把 HBM 往返次数进一步压到更少(论文里强调把非 matmul 的 HBM 访问次数降到接近下限),从而在 A100 上把 GPU 利用率显著抬高。需要诚实地说:利用率离理论上限仍有距离,说明这条路线还能继续挖。
几条现实提醒,帮你避免误用:
下一节 2.2 我们转入另一个战场:当模型要同时服务很多请求、KV Cache 越积越多时,如何用「分页」把显存碎片几乎消灭。