从 FlashAttention 看懂算子融合:为什么融合、少读写了什么、又会错在哪里

用 FlashAttention 官方论文和 CUDA 源码拆解 attention fusion:从 QK、softmax、PV 的普通三段式,到 online softmax、HBM 读写减少、边界条件和 kernel 正确性风险。

一句话结论

算子融合不是把几行代码塞进一个 kernel 这么简单。它真正改变的是数据的命运:哪些中间结果要写回 HBM,哪些可以留在 shared memory 或寄存器里,哪些状态必须用更稳定的数学形式边算边维护。

FlashAttention 是最适合入门这个问题的真实案例。普通 attention 会显式形成一个巨大的 QK^T 分数矩阵,再 softmax,再乘 V。FlashAttention 的关键进化是:按块读取 K/V,在片上完成 QK、mask、softmax 统计、乘 V 和输出累加,不把完整 attention matrix 写回显存。

这篇的目标不是教你写 CUDA,而是训练一个源码阅读问题:当你看到一个 kernel solution,要能问出三件事:

  1. 它融合了哪些原本分开的算子?
  2. 它省掉了哪些 HBM 中间读写?
  3. 它为了正确性新增了哪些风险点?

先看硬件地图:数据到底在哪里流动

如果只看 QK -> softmax -> PV,很容易把算子融合理解成“把几段代码放进一个 kernel”。更贴近真实推理系统的看法是:数据在一组有明显速度、容量和距离差异的硬件层级之间移动。

我补了一张 draw.io 动态图,把这个层级放到 NVIDIA Vera Rubin 级别的推理系统里看。图里的动态箭头表示运行时数据流:tokens 进入 rack,Q/K/V 与 KV cache 进入 GPU,局部 scores 和 softmax tile 在片上短暂存在,最后 O 写回 HBM。

Rubin 级 AI 推理硬件内部架构图

下载 draw.io 源文件

这张图不是 NVIDIA die floorplan,而是一个教学抽象:外层是 Vera Rubin NVL72 这样的 rack-scale 推理系统,中间放大到单 GPU 内部的 HBM -> L2 Cache -> SM -> Shared Memory -> Registers -> Tensor Cores,最后再映射到 prefill、decode 和 FlashAttention。

图里的硬件锚点来自 NVIDIA 公开资料:Vera Rubin NVL72 的 rack-scale 组织、NVLink 6 的 scale-up 带宽,以及 Rubin CPX 面向 massive-context inference 的 prefill/generation 拆分方向。可以对照 Vera Rubin 平台页NVLink 介绍Rubin CPX 技术博客 看。

看这张图时,重点盯住一句话:融合改变的是中间数据的停留位置。普通 attention 会让 S = QK^TP = softmax(S) 这种大中间矩阵在 HBM 里长时间存在;FlashAttention 则让局部 scores、softmax tile、row_maxrow_sumacc_o 尽量停在片上,最后只把输出 O 和少量 softmax_lse 写回。

这里有两个容易混的点。

第一,SMshared memory 不是一回事。SM 是 Streaming Multiprocessor,是 GPU 里的执行单元;一个 CUDA block/CTA 被调度到某个 SM 上执行。shared memory 是 SM 内部可由 kernel 显式读写的一块片上 SRAM,常用来暂存 K/V tile 这种会被同一个 block 反复访问的数据。更准确的层级是:

GPU
  -> SM
      -> registers
      -> shared memory / L1 SRAM
      -> Tensor Cores
      -> schedulers / load-store units

现代 NVIDIA GPU 里,L1 cache 和 shared memory 往往共享同一片物理 SRAM 或可配置分区,但编程语义不同:L1 cache 由硬件自动管理,shared memory 由 CUDA kernel 显式布局、读写和同步。所以图里把 SM / CTA tile 画成外框,把 Shared MemoryRegistersTensor Cores 放进去,是为了表达“这些是 SM 内部或紧贴 SM 的执行/存储资源”。

第二,图里 only final O + small softmax_lse write back 的 “write back” 不是把完整 SP 写回 HBM,而是把最终 attention 输出写回去。也就是片上累加器里的 acc_o 经过归一化和 dtype 转换后,变成:

O = softmax(QK^T) V

这个 O 是下一层、residual、MLP 等后续计算必须消费的结果,所以必须落回 global memory/HBM。FlashAttention 省掉的是完整 S = QK^T scores matrix 和完整 P = softmax(S) attention probability matrix 的 HBM 物化。softmax_lse 则是每行一个很小的 logsumexp 统计量,训练或需要 backward 的 forward 路径会保存它,方便反向传播重建必要的 softmax 信息;纯推理路径是否写它,取决于具体 kernel/API,但它的量级和 N x NS/P 完全不是一回事。

先有一个朴素模型:attention 在算什么

先不要看 CUDA。把 attention 想成一个检索动作:

  • Q 是当前 token 提的问题。
  • K 是历史 token 的索引。
  • V 是历史 token 携带的内容。

一个 query 会先和每个 key 算相似度,得到分数;分数经过 softmax 变成权重;最后用这些权重加权求和 value。

最小例子:

query q0 对 3 个 key 的分数是 [2, 1, 0]
softmax 后大概是 [0.665, 0.245, 0.090]

如果 3 个 value 是 [10, 20, 30]
输出就是 0.665*10 + 0.245*20 + 0.090*30 = 14.25

真实模型里,value 不是一个数,而是一段向量;query 也不止一个。于是 attention 变成三步:

S = QK^T
P = softmax(S)
O = PV

这里最危险的不是乘法本身,而是 SP 的形状。假设序列长度是 N,每个 head 都会有一个 N x N 的 attention matrix。N 一大,中间矩阵就会快速膨胀。

这就是 FlashAttention 论文里说的核心问题:attention 不只是计算量大,更是读写 HBM 的成本大。GPU 的片上 SRAM/shared memory 很快但小,HBM 大但慢很多。普通实现的问题是把巨大的中间结果写到 HBM,再读回来继续算。

普通 attention 的数据流:中间矩阵被物化

可以把传统路径想成三段流水线:

Q, K 读入
  -> 计算 S = QK^T
  -> 把 S 写回 HBM

S 再读入
  -> 计算 P = softmax(S)
  -> 把 P 写回 HBM

P, V 再读入
  -> 计算 O = PV
  -> 把 O 写回 HBM

从框架视角看,这很自然:matmul、softmax、matmul 是三个清晰算子。每个算子都可以单独优化、单独测试、单独调度。

但从 GPU 内存视角看,这很浪费:SP 是中间产物,它们的最终使命只是帮助生成 O。如果能让 S/P 短暂存在于寄存器或 shared memory,然后马上参与下一步,就不需要把完整 N x N attention matrix 写回 HBM。

这就是融合的第一层含义:不是少调用两个函数,而是让中间结果不落地。

FlashAttention 的核心变化:不保存完整 P

FlashAttention 的做法可以概括成:

固定一块 Q
for 每一块 K/V:
  算当前块的 scores = Q_block @ K_block^T
  对 scores 做 mask
  更新这一行 softmax 的 max 和 sum
  用当前块的 softmax 权重乘 V_block
  累加到输出 O_block

最后把 O_block 和 logsumexp 写回 HBM

这里有两个关键点。

第一,QK^T 不是一次生成完整矩阵,而是按 block 生成局部 scores。局部分数用完就丢,不需要物化全局 S

第二,softmax 不能天真地分块。因为 softmax 的分母要看整行所有 key:

softmax(x_i) = exp(x_i) / sum_j exp(x_j)

如果你只看一块 key,就不知道全行最大值和全行分母。FlashAttention 因此维护每行的两个状态:

  • 当前看到的最大值 row_max
  • 当前看到的指数和 row_sum

当新 block 进来,如果发现更大的分数,就要把旧的 row_sum 和已经累加的输出 acc_o 重新缩放。这就是 online softmax 的核心。

用一个小数字感受一下。假设第一块分数是 [2, 1]

max = 2
sum = exp(2-2) + exp(1-2) = 1 + 0.367 = 1.367

第二块来了一个更大的分数 [4]。新的 max 变成 4。旧分母不能直接加,因为它原来是按 max=2 缩放的,现在要改到 max=4:

旧 sum 需要乘 exp(2-4) = 0.135
新 sum = 1.367*0.135 + exp(4-4) = 1.184

输出累加 acc_o 也一样要乘这个缩放因子。否则前面 value 的贡献会被放大,结果就错了。

这就是 FlashAttention 里“融合导致的正确性压力”:为了不保存完整 softmax 矩阵,kernel 必须自己维护数值稳定的 online softmax 状态。

源码地图:真实融合发生在哪里

下面用 FlashAttention 官方 repo 当前主分支源码做锚点。为避免把源码细节淹没在模板里,我只看 forward 主路径。

1. Python API 暴露的还是普通 attention

用户调用的是 flash_attn_func(q, k, v, ...)。接口文档仍然描述为 QK^T 先 scale 再 softmax,并返回 out。这很重要:对上层模型来说,语义没有变,变的是底层执行策略。

源码入口:

这个接口还暴露了几个容易出错的语义:causal mask、sliding window、ALiBi、dropout、GQA/MQA、是否返回 attention probabilities。也就是说,真实 kernel 不是只算一个理想公式,它要承接模型里各种变体。

2. C++ dispatch 选择具体 kernel

C++ 层会根据 head dimension、是否 causal、是否 dropout、shape 是否整齐等条件选择模板特化,然后 launch CUDA kernel。

源码位置:

这一步是算子进化里经常被忽略的部分:高性能 kernel 不是一个万能函数,而是一组针对 shape、head dim、硬件和语义开关的特化路径。所谓“一个 attention 算子”,在源码里会展开成很多具体实现。

3. forward kernel 里融合了 QK、mask、softmax、PV

真正的融合主线在 compute_attn_1rowblock

源码位置:

你可以按变量名建立直觉:

  • gQ/gK/gV:global memory 里的 Q/K/V tile。
  • sQ/sK/sV:shared memory 里的 Q/K/V tile。
  • acc_s:当前 Q block 和 K block 算出来的分数块,也就是局部 S
  • softmax:维护每行的 row_maxrow_sum
  • acc_o:当前输出块的 fp32 累加器。
  • rP:当前 block 的 softmax 权重,马上用于乘 V

主循环里可以看到这条路径:

读 K/V tile 到 shared memory
gemm 得到 acc_s = Q @ K^T
mask.apply_mask(acc_s)
softmax.softmax_rescale_o(acc_s, acc_o)
rP = convert(acc_s)
gemm_rs(acc_o, rP, V)

对应源码:

这几行就是“算子融合”的现场。它不是三个 kernel:

matmul(Q, K) -> softmax -> matmul(P, V)

而是一个 block 内的循环:

局部 QK -> 局部 softmax 状态更新 -> 局部乘 V -> 输出累加

完整 P 没有成为主路径上的 HBM 中间结果。

4. softmax 的正确性被挪进 kernel 内部

softmax.h 是理解 FlashAttention 的关键。Softmax 结构维护 row_maxrow_sum。第一次看到一块 scores 时,它做 reduce max、exp、reduce sum。后续 block 进来时,它会:

  1. 保存旧的 row_max
  2. 用当前 scores 更新新的 row_max
  3. 用新旧 max 的差值缩放旧的 row_sum
  4. 同样缩放已经累加的 acc_o
  5. 再把当前 block 的 exp 分数加进来。

源码位置:

这就是为什么读 fused kernel 时不能只问“少了几个 kernel launch”。更深的问题是:原本由独立 softmax 算子保证的数值稳定性,现在谁来保证?在 FlashAttention 里,是 kernel 内部的 row_max/row_sum/acc_o 重标定逻辑在保证。

5. 输出只在 epilogue 写回

主循环里 acc_o 一直在片上累加。最后 epilogue 才归一化、转换 dtype、写出 Osoftmax_lse

源码位置:

这里 softmax_lse 很重要。它不是为了 forward 输出给用户看的主要结果,而是为了 backward 或调试保存每行 softmax 的 logsumexp。因为 forward 没保存完整 P,backward 需要一些足够小的统计量来重建必要信息。

这也是算子融合的另一个代价:你省掉了大中间矩阵,但要认真选择保留哪些小状态,才能让后续计算仍然可做。

它到底少读写了什么

用最粗略的方式看,普通 attention 会把 SP 这两个 N x N 矩阵写到 HBM 或至少在全局层面物化。FlashAttention 主路径避免了这件事。

如果 N=4096,单个 head 的 N x N 有 16,777,216 个元素。即使用 fp16,也大约是 32MB。S 一份、P 一份,再加上读回,HBM 流量会很快变成瓶颈。多 batch、多 head 后这个量继续放大。

FlashAttention 不改变 attention 的数学结果,它改变的是中间数据的生命周期:

普通路径:
scores 活到 HBM
probabilities 活到 HBM
output 写回 HBM

FlashAttention:
scores 活在当前 tile
probabilities 活在当前 tile
row_max / row_sum / acc_o 跨 tile 存活
output 最后写回 HBM

这就是“IO-aware”的含义。不是只数 FLOPs,而是把 HBM 访问当成一等公民。

为什么可能错

读一个 kernel solution,最有价值的不是只看它快在哪里,还要看它把哪些风险搬到了自己身上。FlashAttention 这个例子至少有六类风险。

1. softmax 重标定错

如果新 block 出现更大的 max,旧的 row_sumacc_o 必须缩放。漏掉输出缩放,或者缩放用错底数,结果会在数值上偏,但不一定立刻 NaN。

这类错误很危险,因为小输入可能过测试,长序列或极端 logits 才暴露。

2. mask 边界错

causal mask、local window、变长 batch、seqlen_q != seqlen_k 都会改变哪些 key 可见。FlashAttention 的 Python 文档甚至专门说明 causal mask 在 seqlen_qseqlen_k 不同的时候如何对齐。

mask 错一格,模型仍然能输出,但语义已经泄漏未来 token 或漏看历史 token。

3. OOB 读写错

真实 head dim 不一定正好等于模板里的 rounded dim,序列长度也不一定正好是 block size 的倍数。源码里有大量 Is_even_KIs_even_MN、predicate、early exit 和 “不要把 OOB 写回 gmem” 这类逻辑。

这说明融合 kernel 的正确性不只在数学公式,还在边界条件。

4. 同步和 pipeline 错

kernel 里用 cp_async 把数据搬到 shared memory,并用 fence/wait/syncthreads 协调。源码注释里有一处直接说明:某个 cp_async_fence 必须放在 if block 里,否则同步不对,会出现 race condition。

这类 bug 不像 Python shape error 那样稳定复现。它可能和 GPU 架构、block 调度、输入大小有关。

5. dropout 和随机性错

dropout 不只是把概率乘 0。为了 backward 和可复现,kernel 要管理 RNG seed、offset、dropout mask 编码。源码里 forward 一开始就保存 RNG state,Python 接口也把 return_attn_probs 标成 testing-only,并提醒返回概率不保证是普通理解里的完整 attention probabilities。

这说明 fused kernel 有时为了性能暴露的调试输出并不是业务语义的一等结果。

6. split-KV 的并行收益和 HBM 代价冲突

为了提高 occupancy,推理时可能把 KV 分成多个 split 并行做。但 split 多了以后,需要额外的 softmax_lse_accumout_accum,再 combine。源码里的 heuristic 明确权衡:split 可以提高 SM 利用率,但太多 split 会增加 HBM 读写。

源码位置:

这就是性能工程的灰度:并行不是越多越好,融合也不是越大越好。真正的优化是在 occupancy、寄存器、shared memory、HBM 流量、数值稳定和编译复杂度之间找平衡。

算子是怎么进化的

我现在对 attention 算子的进化可以这样理解:

阶段 1:数学表达
O = softmax(QK^T)V

阶段 2:框架算子
matmul -> softmax -> matmul

阶段 3:手写 fused kernel
在一个 kernel 内按 tile 完成 QK、softmax、PV

阶段 4:IO-aware 算法
不保存 N x N 中间矩阵,只保存必要统计量

阶段 5:硬件协同
围绕 shared memory、寄存器、warp、cp_async、TMA、WGMMA、scheduler 继续重写

FlashAttention-2 论文说,FlashAttention 已经省了内存和带宽,但还没有接近 GEMM 的效率,于是 FlashAttention-2 继续改 work partitioning,减少非 matmul FLOPs,提高 occupancy,减少 shared memory 读写。官方 README 里又继续列出 FlashAttention-3 面向 Hopper,FlashAttention-4 面向 Hopper/Blackwell 和 CuTeDSL。

这条线给我的最大启发是:算子不是静态函数,而是数学、内存层级、硬件指令和模型需求不断互相挤压后的结果。

同一个公式:

softmax(QK^T)V

在不同阶段会长成完全不同的代码形态。越往下走,越不像“公式翻译”,越像“管理数据的生命”。

读 kernel solution 时的检查清单

以后读 Fable、Triton、CUDA 或 TorchInductor 生成的 fusion kernel,我会按这张清单看:

  1. 融合边界:它把哪些 op 合到一起?有没有跨越 reduction、softmax、dropout、mask 这种难融合边界?
  2. 中间结果:原来会写回 HBM 的张量是什么?现在是在 register、shared memory,还是根本不保存?
  3. 状态变量:为了不保存中间矩阵,它保留了哪些统计量?比如 max、sum、lse、partial output。
  4. 数值稳定:它有没有用 max-subtraction、logsumexp、重标定?极端值会不会 NaN 或 overflow?
  5. 边界条件:非整齐 shape、变长序列、causal/local mask、GQA/MQA、dropout、dtype 是否都覆盖?
  6. 同步语义:有没有异步拷贝、shared memory 复用、barrier、race condition 注释?
  7. 性能取舍:它是在省 HBM,增 occupancy,减 launch,还是减少 shared memory 读写?这些目标有没有互相冲突?
  8. 验证方式:它和 PyTorch reference 比什么?容忍误差是多少?长序列、奇怪 head dim、mask 边界有没有测?

能回答这些问题,就算暂时写不出 CUDA,也已经开始看懂算子融合了。

总结

FlashAttention 让我感知到的“算子进化”不是抽象的工程优化,而是一条很具体的路径:

先发现瓶颈不是公式,而是中间矩阵读写;
再把大矩阵拆成 tile;
再用 online softmax 保住数学等价;
再把 QK、softmax、PV 放进同一个 kernel;
最后围绕具体 GPU 的内存层级和并行模型继续重写。

所以,当我们说“读一个 kernel solution”,真正要读的是它如何重新安排数据的生命:什么时候出生,在哪里停留,什么时候被消费,哪些状态必须留下,哪些中间结果应该立刻消失。

这就是从“会用框架”走向“理解 AI infra”的分界线。

资料