注意力公式没变,为何还能快几倍?FlashAttention 的分块、在线 Softmax 与 IO
从显存读写瓶颈出发,手算分块在线 Softmax,解释 FlashAttention 如何不物化 L×L 矩阵、保持精确结果,并用 PyTorch 2.14 SDPA 验证后端与性能。
上一篇用 micro-batch 和 1F1B 填补流水线空洞,但 GPU 忙起来不等于算子已经高效。标准 Attention(注意力)会先写出完整分数矩阵,再读回来做 Softmax,最后又读一次乘 。长序列下,真正拖慢它的常常不是浮点乘法,而是 High Bandwidth Memory(高带宽显存,HBM)与片上 SRAM 之间的数据搬运。
FlashAttention(闪存注意力)没有改变注意力函数,也没有删掉任意一条注意力边。它通过 tiling(分块)、kernel fusion(算子融合)和 online softmax(在线 Softmax),让中间的 矩阵不必写回 HBM。
01 标准实现究竟把什么搬来搬去?#
单个 batch、单个 head 的缩放点积注意力为
其中 ,,。朴素 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计算量仍为 ;但仅 与 就有 个元素。若 、FP16,每个矩阵约 128 MiB,每层每头组合后的中间读写很快压过 的线性存储。
02 FlashAttention 改的是执行顺序,不是数学目标#
把 沿行切成块 ,把 切成 。一个 留在片上,依次扫描各个 :
HBM: Q 块 K₀,V₀ K₁,V₁ K₂,V₂ ...
| | | |
v v v v
SRAM: [Qᵢ] -> [局部分数] -> 更新 m,l,O -> 丢弃局部分数
|
v
HBM: 只写最终 Oᵢtext局部分数 只在片上短暂存在。难点是 Softmax 的分母依赖一整行全部 key;若按块各做一次 Softmax 再相加,结果一定错误。在线 Softmax 提供了可合并的行状态。
03 在线 Softmax 只需保留三个状态#
对一行分数,处理到当前块时保留:
- :目前见过的最大分数,shape 为
[B_r, 1]; - :以 为基准的指数和,shape 为
[B_r, 1]; - :尚未除分母的加权值,shape 为
[B_r,d_v]。
新块的行最大值是 ,指数和与加权值是 。合并时令
最后输出 。当出现更大的最大值时,旧块的累计量会按 重新缩放,因此数值稳定且不需要保存旧分数。
04 用三个分数手算两次合并#
令一行分数为 [1, 2, 3],对应一维 value 为 [10, 20, 40]。先处理第一块 [1,2]:
第二块只有分数 3,所以 。合并:
因此 。直接对 [1,2,3] 做 Softmax 后乘 [10,20,40] 也是同一结果;分块只改变求值顺序。
05 一个透明的教学版分块前向#
下面代码故意用 PyTorch 普通算子表达算法,不会自动获得定制 CUDA kernel 的速度,但适合与稠密基线逐元素对照。
import math
import torch
def tiled_attention(q, k, v, block_k=64, causal=False):
# q:[N,H,L,D], k:[N,H,S,D], v:[N,H,S,Dv]
n, h, l, d = q.shape
s, dv = k.size(-2), v.size(-1)
m = torch.full((n, h, l, 1), -torch.inf, device=q.device)
ell = torch.zeros_like(m)
acc = torch.zeros((n, h, l, dv), device=q.device, dtype=torch.float32)
qf = q.float()
q_pos = torch.arange(l, device=q.device)[:, None]
for start in range(0, s, block_k):
end = min(start + block_k, s)
kb = k[..., start:end, :].float()
vb = v[..., start:end, :].float()
scores = qf @ kb.transpose(-2, -1) / math.sqrt(d)
if causal:
k_pos = torch.arange(start, end, device=q.device)[None, :]
scores = scores.masked_fill(k_pos > q_pos, -torch.inf)
mb = scores.amax(dim=-1, keepdim=True)
p = torch.exp(scores - mb)
lb = p.sum(dim=-1, keepdim=True)
ab = p @ vb
m_new = torch.maximum(m, mb)
old_scale = torch.exp(m - m_new)
new_scale = torch.exp(mb - m_new)
ell = old_scale * ell + new_scale * lb
acc = old_scale * acc + new_scale * ab
m = m_new
return (acc / ell).to(q.dtype)python真实 kernel 还会沿 query 维分块、控制寄存器与 shared memory 占用、融合 mask/dropout,并为反向传播设计重算。这里把累计状态保留为 FP32,是为了避免长行上的指数和在低精度下丢失有效位。
06 Causal Mask 怎样进入分块?#
自回归 attention 只允许 query 位置 看 key 位置 。分块不能用“块号相同就全部可见”的粗略判断,因为对角块内部仍是三角形。
key block -> K0 K1 K2
Q0 三角 跳过 跳过
Q1 全部 三角 跳过
Q2 全部 全部 三角text完全位于对角线右侧的块可直接跳过;左侧块全算;对角块逐元素施加 causal mask。Padding、局部窗口和 attention bias 也必须在局部分数进入最大值与指数和之前处理,否则被屏蔽位置会污染归一化。
07 反向传播为何还能保持线性额外存储?#
朴素 autograd 会保存 。FlashAttention 前向保存输出 与每行 log-sum-exp 等线性大小统计量;反向时重新分块计算所需局部分数和概率,再累积 。
这是一种 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,让“不支持而回退”变成明确错误或警告。
import torch
import torch.nn.functional as F
from torch.nn.attention import SDPBackend, sdpa_kernel
q = torch.randn(2, 16, 2048, 64, device="cuda", dtype=torch.bfloat16,
requires_grad=True)
k = torch.randn_like(q)
v = torch.randn_like(q)
with sdpa_kernel(SDPBackend.FLASH_ATTENTION):
out = F.scaled_dot_product_attention(
q, k, v,
attn_mask=None,
dropout_p=0.0,
is_causal=True,
)
assert out.shape == (2, 16, 2048, 64)
out.float().square().mean().backward()pythondropout_p 无论模块是否处于 eval 模式都会按传入值执行;推理时应显式传 0.0。attn_mask 与 is_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(),再比较 。包含全 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;
- 多个 、head dimension、dtype 与 mask 组合。
短序列、小 batch 或不受支持的 shape 上,调度开销可能抵消收益。FlashAttention 省的是中间矩阵 IO,不会把 的点积计算变成线性。
11 常见错误与最短调试路径#
| 症状 | 常见原因 | 最短检查 |
|---|---|---|
| 强制 Flash 后报不支持 | dtype、设备、shape 或 mask 不满足后端约束 | 缩成已知支持的 BF16 CUDA case,再逐项加回 |
| 输出整行 NaN | 某 query 的所有 key 都被 mask | 检查每行至少一个有效位置及 mask 语义 |
| 推理结果每次变化 | dropout_p 仍非 0 | 在 eval 路径显式传 0.0 |
| 内存仍呈 增长 | 代码在 SDPA 前保存了完整 attention weights/bias | profiler 与 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 今天真正需要记住什么?#
- 标准 attention 的瓶颈常是 中间矩阵反复进出 HBM,而非公式里的乘法数量本身。
- 分块在线 Softmax 用每行的最大值、指数和与加权值就能合并任意 key blocks,保持同一数学结果。
- FlashAttention 不物化完整分数/概率矩阵,反向重算局部量,以更多片上工作换更少 HBM IO。
- 调用 SDPA 不保证选中 Flash backend;必须强制后端做兼容测试,并用端到端指标验证收益。
14 思考题与小练习#
- 把分数
[0, 2, 1, 3]分成[0,2]与[1,3]两块,手算 的两次状态并验证最终 Softmax 分母。 - 对 、16 heads、FP16,估算显式保存一个
[H,L,L]概率张量的 MiB;再与[H,L,64]输出比较。 - 为教学版
tiled_attention增加 padding mask,并设计一个“整块被屏蔽但整行仍有有效 key”的测试。
相关工作#
- Dao et al., FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness ↗,提出 IO-aware 的精确分块 attention。
- Dao, FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning ↗,改进线程块与 warp 间工作划分。
- Milakov & Gimelshein, Online Normalizer Calculation for Softmax ↗,给出可流式更新的稳定 Softmax 归一化。
- Rabe & Staats, Self-attention Does Not Need Memory ↗,讨论通过重算降低 attention 内存。
- PyTorch, Scaled Dot Product Attention 官方文档 ↗,说明当前 API、shape 与 backend 选择。
15 下一篇预告#
FlashAttention 解决了一次 attention 算子怎样少搬数据;但在线服务同时生成许多不同长度请求时,KV Cache 还会因预留与碎片让显存提前耗尽。下一篇将用 block table 手算 PagedAttention 如何让逻辑连续的 token 映射到非连续物理块。