观文听傑

返回

上一篇用 micro-batch 和 1F1B 填补流水线空洞,但 GPU 忙起来不等于算子已经高效。标准 Attention(注意力)会先写出完整分数矩阵,再读回来做 Softmax,最后又读一次乘 VV。长序列下,真正拖慢它的常常不是浮点乘法,而是 High Bandwidth Memory(高带宽显存,HBM)与片上 SRAM 之间的数据搬运。

FlashAttention(闪存注意力)没有改变注意力函数,也没有删掉任意一条注意力边。它通过 tiling(分块)、kernel fusion(算子融合)和 online softmax(在线 Softmax),让中间的 L×LL\times L 矩阵不必写回 HBM。

01 标准实现究竟把什么搬来搬去?#

单个 batch、单个 head 的缩放点积注意力为

S=QK/d,P=softmax(S),O=PV,S=QK^\top/\sqrt d,\qquad P=\operatorname{softmax}(S),\qquad O=PV,

其中 Q,K,VRL×dQ,K,V\in\mathbb R^{L\times d}S,PRL×LS,P\in\mathbb R^{L\times L}ORL×dvO\in\mathbb R^{L\times d_v}。朴素 GPU 流水线常把三个步骤拆成多个 kernel:

flowchart LR
  QK[从 HBM 读 Q,K] --> S[写回 S: L×L]
  S --> SM[再读 S 做 Softmax]
  SM --> P[写回 P: L×L]
  P --> PV[再读 P,V 做矩阵乘]
  PV --> O[写回 O: L×dᵥ]
mermaid

计算量仍为 Θ(L2d)\Theta(L^2d);但仅 SSPP 就有 2L22L^2 个元素。若 L=8192L=8192、FP16,每个矩阵约 128 MiB,每层每头组合后的中间读写很快压过 Q,K,V,OQ,K,V,O 的线性存储。

02 FlashAttention 改的是执行顺序,不是数学目标#

QQ 沿行切成块 QiRBr×dQ_i\in\mathbb R^{B_r\times d},把 K,VK,V 切成 Kj,VjRBc×dK_j,V_j\in\mathbb R^{B_c\times d}。一个 QiQ_i 留在片上,依次扫描各个 Kj,VjK_j,V_j

HBM:  Q 块   K₀,V₀   K₁,V₁   K₂,V₂ ...
        |       |       |       |
        v       v       v       v
SRAM: [Qᵢ] -> [局部分数] -> 更新 m,l,O -> 丢弃局部分数
                                      |
                                      v
HBM:                              只写最终 Oᵢ
text

局部分数 Sij=QiKj/dS_{ij}=Q_iK_j^\top/\sqrt d 只在片上短暂存在。难点是 Softmax 的分母依赖一整行全部 key;若按块各做一次 Softmax 再相加,结果一定错误。在线 Softmax 提供了可合并的行状态。

03 在线 Softmax 只需保留三个状态#

对一行分数,处理到当前块时保留:

  • mm:目前见过的最大分数,shape 为 [B_r, 1]
  • =testm\ell=\sum_t e^{s_t-m}:以 mm 为基准的指数和,shape 为 [B_r, 1]
  • a=testmvta=\sum_t e^{s_t-m}v_t:尚未除分母的加权值,shape 为 [B_r,d_v]

新块的行最大值是 mbm_b,指数和与加权值是 b,ab\ell_b,a_b。合并时令

m=max(m,mb),m'=\max(m,m_b), =emm+embmb,\ell'=e^{m-m'}\ell+e^{m_b-m'}\ell_b, a=emma+embmab.a'=e^{m-m'}a+e^{m_b-m'}a_b.

最后输出 o=a/o=a/\ell。当出现更大的最大值时,旧块的累计量会按 emme^{m-m'} 重新缩放,因此数值稳定且不需要保存旧分数。

04 用三个分数手算两次合并#

令一行分数为 [1, 2, 3],对应一维 value 为 [10, 20, 40]。先处理第一块 [1,2]

m=2,=e1+11.3679,m=2,\quad \ell=e^{-1}+1\approx1.3679, a=e110+12023.6788.a=e^{-1}\cdot10+1\cdot20\approx23.6788.

第二块只有分数 3,所以 mb=3,b=1,ab=40m_b=3,\ell_b=1,a_b=40。合并:

m=3,m'=3, =e11.3679+11.5032,\ell'=e^{-1}\cdot1.3679+1\approx1.5032, a=e123.6788+4048.7109.a'=e^{-1}\cdot23.6788+40\approx48.7109.

因此 o=a/32.405o=a'/\ell'\approx32.405。直接对 [1,2,3] 做 Softmax 后乘 [10,20,40] 也是同一结果;分块只改变求值顺序。

05 一个透明的教学版分块前向#

下面代码故意用 PyTorch 普通算子表达算法,不会自动获得定制 CUDA kernel 的速度,但适合与稠密基线逐元素对照。

真实 kernel 还会沿 query 维分块、控制寄存器与 shared memory 占用、融合 mask/dropout,并为反向传播设计重算。这里把累计状态保留为 FP32,是为了避免长行上的指数和在低精度下丢失有效位。

06 Causal Mask 怎样进入分块?#

自回归 attention 只允许 query 位置 ii 看 key 位置 jij\le i。分块不能用“块号相同就全部可见”的粗略判断,因为对角块内部仍是三角形。

key block ->   K0      K1      K2
Q0             三角    跳过    跳过
Q1             全部    三角    跳过
Q2             全部    全部    三角
text

完全位于对角线右侧的块可直接跳过;左侧块全算;对角块逐元素施加 causal mask。Padding、局部窗口和 attention bias 也必须在局部分数进入最大值与指数和之前处理,否则被屏蔽位置会污染归一化。

07 反向传播为何还能保持线性额外存储?#

朴素 autograd 会保存 PRL×LP\in\mathbb R^{L\times L}。FlashAttention 前向保存输出 OO 与每行 log-sum-exp 等线性大小统计量;反向时重新分块计算所需局部分数和概率,再累积 dQ,dK,dVdQ,dK,dV

这是一种 recomputation(重算):用额外 FLOPs 换掉二次方中间激活。它与上一篇之前讲过的 Activation Checkpointing 思想相似,但边界在专用 attention kernel 内,且重算公式针对 Softmax 导数高度融合。

08 用 PyTorch 2.14 当前 SDPA 选择后端#

当前官方入口是 torch.nn.functional.scaled_dot_product_attention。它会按 device、dtype、shape、mask 和 dropout 等条件选择后端;调试时可用 sdpa_kernel 强制 Flash backend,让“不支持而回退”变成明确错误或警告。

dropout_p 无论模块是否处于 eval 模式都会按传入值执行;推理时应显式传 0.0attn_maskis_causal 的组合限制、支持的 head dimension、dtype 和硬件能力都可能影响后端资格,不要把“调用了 SDPA”当作“运行了 FlashAttention”。

09 怎样证明后端、结果与梯度都正确?#

建议固定一个很小的 FP32 数学基线,再测试生产 dtype:

from torch.nn.attention import SDPBackend, sdpa_kernel

def run(backend, q, k, v):
    with sdpa_kernel(backend):
        return F.scaled_dot_product_attention(q, k, v, is_causal=True)

q0 = torch.randn(1, 2, 17, 32, device="cuda", dtype=torch.float32)
k0 = torch.randn_like(q0)
v0 = torch.randn_like(q0)
ref = run(SDPBackend.MATH, q0, k0, v0)

q1, k1, v1 = (x.to(torch.bfloat16) for x in (q0, k0, v0))
got = run(SDPBackend.FLASH_ATTENTION, q1, k1, v1).float()
torch.testing.assert_close(got, ref, rtol=2e-2, atol=2e-2)
python

梯度测试需为两条路径分别 clone requires_grad_(),对相同标量 loss 调 backward(),再比较 dQ,dK,dVdQ,dK,dV。包含全 mask 行、非整块长度、padding、causal、dropout 和真实 head dimension 的 case 才能覆盖边界。

10 性能验证不能只计一次 Python 时钟#

CUDA 是异步执行的。应先 warmup,再用 CUDA Event 或 torch.utils.benchmark,并在计时边界同步。训练要测 forward+backward,且记录:

  • tokens/s 与 step time,而不只是单 kernel 微秒数;
  • torch.cuda.max_memory_allocated() 的峰值;
  • profiler 中实际 SDPA kernel 名称、HBM 流量和 kernel gaps;
  • 多个 LL、head dimension、dtype 与 mask 组合。

短序列、小 batch 或不受支持的 shape 上,调度开销可能抵消收益。FlashAttention 省的是中间矩阵 IO,不会把 Θ(L2d)\Theta(L^2d) 的点积计算变成线性。

11 常见错误与最短调试路径#

症状常见原因最短检查
强制 Flash 后报不支持dtype、设备、shape 或 mask 不满足后端约束缩成已知支持的 BF16 CUDA case,再逐项加回
输出整行 NaN某 query 的所有 key 都被 mask检查每行至少一个有效位置及 mask 语义
推理结果每次变化dropout_p 仍非 0在 eval 路径显式传 0.0
内存仍呈 L2L^2 增长代码在 SDPA 前保存了完整 attention weights/biasprofiler 与 memory snapshot 找出 L×L 分配
与基线不逐 bit 相同浮点归约顺序不同改用 dtype 对应的 rtol/atol,比较统计误差
kernel 快但整步不快QKV 投影、通信或数据加载成为瓶颈看端到端 profiler,不只 microbenchmark

12 与稀疏注意力、Checkpointing 有何区别?#

FlashAttention 对稠密 attention 是 exact implementation(精确实现):边数和数学目标不变。滑动窗口、块稀疏会删除连接,计算图本身变了。Activation Checkpointing 可包住任意子图;FlashAttention 的重算专门利用 attention 的分块与在线归一化。

PagedAttention 则解决另一层问题:自回归服务中,历史 KV cache 怎样动态分配和寻址。它可以使用分块 attention kernel,但核心目标是减少并发请求的 KV 内存碎片;下一篇会把这条边界讲清楚。

13 今天真正需要记住什么?#

  1. 标准 attention 的瓶颈常是 L×LL\times L 中间矩阵反复进出 HBM,而非公式里的乘法数量本身。
  2. 分块在线 Softmax 用每行的最大值、指数和与加权值就能合并任意 key blocks,保持同一数学结果。
  3. FlashAttention 不物化完整分数/概率矩阵,反向重算局部量,以更多片上工作换更少 HBM IO。
  4. 调用 SDPA 不保证选中 Flash backend;必须强制后端做兼容测试,并用端到端指标验证收益。

14 思考题与小练习#

  1. 把分数 [0, 2, 1, 3] 分成 [0,2][1,3] 两块,手算 m,m,\ell 的两次状态并验证最终 Softmax 分母。
  2. L=4096L=4096、16 heads、FP16,估算显式保存一个 [H,L,L] 概率张量的 MiB;再与 [H,L,64] 输出比较。
  3. 为教学版 tiled_attention 增加 padding mask,并设计一个“整块被屏蔽但整行仍有有效 key”的测试。

相关工作#

  1. Dao et al., FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness,提出 IO-aware 的精确分块 attention。
  2. Dao, FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning,改进线程块与 warp 间工作划分。
  3. Milakov & Gimelshein, Online Normalizer Calculation for Softmax,给出可流式更新的稳定 Softmax 归一化。
  4. Rabe & Staats, Self-attention Does Not Need O(n2)O(n^2) Memory,讨论通过重算降低 attention 内存。
  5. PyTorch, Scaled Dot Product Attention 官方文档,说明当前 API、shape 与 backend 选择。

15 下一篇预告#

FlashAttention 解决了一次 attention 算子怎样少搬数据;但在线服务同时生成许多不同长度请求时,KV Cache 还会因预留与碎片让显存提前耗尽。下一篇将用 block table 手算 PagedAttention 如何让逻辑连续的 token 映射到非连续物理块。

注意力公式没变,为何还能快几倍?FlashAttention 的分块、在线 Softmax 与 IO
https://zwjcode.cn/blog/flashattention-online-softmax-tiling-io
作者
发布于 2026年9月17日
版权协议 CC BY-NC-SA 4.0
评论加载似乎遇到了问题,请尝试刷新页面。