跳到正文
动手理解 · CONCEPT LAB

不落整张注意力矩阵,Softmax 还能算对吗?

FlashAttention 把 Q/K/V 分块送入寄存器、shared memory 等片上存储,并用 online softmax 维护 running max、running sum 与未归一化 value 累加器。关键收益是少写、少读 HBM。

已经发生 正在观察 接下来
01 / 03

状态传导图

Qᵢ Q Tile
KⱼVⱼ K/V Tile
Sᵢⱼ 局部分数
MAX 局部最大值
ACC 开始累计
NEXT 下一 K/V 块

每个 Q tile 会依次扫过所需 K/V tiles;具体 tile 放在寄存器、shared memory 或 TMEM 取决于硬件和 kernel。

观察点 01

一小块 Q 与一小块 K/V 进入片上存储

kernel 不先生成完整 n×n 分数矩阵,而是选择 Q tile 与 K/V tile,计算这一小块分数和局部最大值。

HBM 中间矩阵写入几乎不落盘
片上数据复用很高
此刻要记住

对 dense 路径,分块不会少看 mask 允许访问的 token,只改变计算与保存顺序;稀疏/local 路径可以另行跳块。

算法工程 · Attention Kernel

FlashAttention:不是少算,而是少搬

Transformer 的注意力看起来是几次矩阵乘法,真正拖慢长上下文的却常常是显存读写。FlashAttention 的核心不是近似、不是剪掉 token,而是把注意力按块搬进更快的片上存储(on-chip storage),用在线 Softmax(online softmax)边算边归并,避免把巨大的注意力矩阵写回 HBM。显存/带宽底层概念可先读 显存与带宽

FlashAttention 直觉总览:HBM 中的大矩阵切成小块,送入片上存储计算,再流式写出结果,避免落盘完整注意力矩阵
直觉图:不要在 HBM 里摊开整张 N×N 注意力表,而是把 Q/K/V 切成 tile,在片上流式归并;图中的 “SRAM” 是 FA1/FA2 的两级内存直觉,具体数据会落在寄存器、shared memory 或 Tensor Memory(TMEM,张量内存)等位置,取决于 GPU 与实现代际。查看原图 ↗
00 · 先抓住

FlashAttention 不是少看内容,而是少搬材料

生活类比
普通注意力像厨师每做一道菜都去仓库搬一大箱食材、切完又搬回去;FlashAttention 像把一小批食材放到案板上,一次处理完再换下一批。菜没少做,来回搬运少了。
小例子
32K 上下文时,注意力矩阵本身可能比 Q/K/V 大得多。FlashAttention 的意义是:不把这张巨表完整写进 HBM,再读回来;只在片上小内存里流式算完。
一句话稠密路径算的是同一个精确 attention 定义,优化的是“先算什么、存在哪里、什么时候写回”。“精确”不等于不同浮点 kernel 必须逐位相同。
别混淆稠密 FlashAttention 不是线性注意力,也不是稀疏注意力;原论文另有会跳过指定块的 block-sparse 扩展,那条路线才改变被计算的关系集合。
先看什么初学者先看厨师/案板类比和 32K 矩阵例子,再看 online softmax 公式。
01 · 问题在哪

标准注意力的瓶颈,不只是 FLOPs

注意力公式很短,但物化中间量的朴素/eager 实现会把 QK^T 的分数矩阵、softmax 后的概率矩阵写到显存。长上下文时,这张表是 N×N,可能比 Q/K/V 本身大得多;现代 fused backend 不一定采用这种执行路径。

O=softmax ⁣(QKd)VO=\mathrm{softmax}\!\left(\frac{QK^\top}{\sqrt{d}}\right)V
稠密 FlashAttention 计算的是同一个公式,在实数算术下不是近似;浮点运算顺序、低精度指令和 dropout 的随机掩码仍会造成正常数值差异。原论文的 block-sparse 扩展则是另一条稀疏/近似路径。
算力 compute

矩阵乘法本身可以上 Tensor Core,GPU 很擅长。问题是 softmax、mask、dropout、反向传播会制造大量中间结果。

显存 IO HBM traffic

HBM 容量大但离计算单元远。把 N×N 矩阵写出去、再读回来,常比多做一点计算更贵。显存/带宽详解 →

片上存储 on-chip storage

片上容量小但离计算单元近。FlashAttention 的策略是让一个 tile 在寄存器、shared memory 等片上资源中尽量多完成工作;FA4 还利用 Blackwell 的 TMEM。具体放置是 kernel 与 ISA 的实现细节。

一张注意力矩阵有多吓人

Quadratic intermediate
若把 32K 取作 32,768,单个 attention head 的完整稠密表有 32,768²=1,073,741,824 个数。FP16/BF16 每项 2 bytes,因此是 2GiB/head;batch=1、32 个 head、单层的一张 SP 逻辑张量便是 64GiB。同时物化两张时,逻辑规模可到两倍,但实际峰值还取决于生命周期、复用和 backend。因果 attention 虽只有下三角有效,朴素稠密物化仍可能分配方阵。
Attention matrix memoryB×H×N2×bytes\text{Attention matrix memory} \approx B \times H \times N^2 \times \text{bytes}

公式中的 B 是 batch、H 是 head 数;上面的 64GiB 没有再乘层数,也不代表现代框架一定真的分配它。这个反事实数量级说明了核心矛盾:长上下文不仅怕计算量,也怕物化中间矩阵带来的容量与 IO。FlashAttention 正是从这个角度切入。

02 · 核心招式

Online softmax:不看完整行,也能算同一个 softmax

普通 softmax 似乎必须先拿到一整行分数:先找最大值,再求指数和,再归一化。FlashAttention 的关键是把这一行拆成多个块,边看边更新最大值和归一化因子。

公式

一块一块更新 running max / sum / output

Streaming recurrence
对某个 query 行,假设已经处理过前面的 key/value 块,手里有当前最大值 m、指数和 、未归一化输出累加器 a。新来一块分数 s_j 和 value v_j,就这样更新:
m=max(m,maxjsj),=emm+jesjm,a=emma+jesjmvj,o=a/m'=\max(m,\max_j s_j),\quad \ell'=e^{m-m'}\ell+\sum_j e^{s_j-m'},\quad a'=e^{m-m'}a+\sum_j e^{s_j-m'}v_j,\quad o'=a'/\ell'

这里的 a 是未归一化 value 累加器,不是已归一化输出 O;如果新块里出现更大的分数,就把旧的 a 同时按 exp(m-m') 重标定,再加入新块,扫完后才做 o=a/ℓ

伪代码直觉 · FA2+ 常见 Q-block 外层顺序

for each Q_block:
  m = -inf
  l = 0
  acc = 0
  for each K_block, V_block:
    scores = Q_block @ K_block.T / sqrt(d)
    scores += mask_or_bias
    m_new = max(m, rowmax(scores))
    p = exp(scores - m_new)
    acc = exp(m - m_new) * acc + p @ V_block
    l = exp(m - m_new) * l + rowsum(p)
    m = m_new
  O_block = acc / l

这段伪代码让一个 Q block 扫完所需 K/V blocks,再写出完成的 O block,是 FA2 及后续常见的讲法。FA1 论文原始 Algorithm 1 采用 K/V block 外层循环,会在每轮从 HBM 读取 Q,并读写部分 O 与 m/ℓ;它仍然不物化完整 N×N 的 S/P。

手算:两块与整行一致
设分数按两块到达:先看 [1,2],对应标量 value 为 [10,20];再看 [3],value 为 [30]。第一块得到 m=2ℓ=e⁻¹+1≈1.3679a=e⁻¹×10+20≈23.6788。第二块把基准更新为 3:ℓ'=e⁻¹ℓ+1≈1.5032a'=e⁻¹a+30≈38.710,所以 o≈25.75;直接对 [1,2,3] 做 softmax 加权也得到同一结果(浮点舍入范围内)。
类比
你要算全班最高分和平均贡献,不必把所有试卷摊满操场;可以一摞一摞看,每看一摞就更新“当前最高分”和“累计贡献”。FlashAttention 就是把“操场”换成 HBM,把“桌面”换成 SRAM。
专业延伸
这不是“近似 softmax”。在实数算术下,online softmax 与整行 softmax 等价;工程上不承诺逐位一致。训练若启用 dropout,反向重算还必须复现同一随机掩码,因此实现会保存或可重建相应随机数状态。
03 · 机制图

物化式注意力 vs FlashAttention:差在中间矩阵落不落盘

下面左图是教学用的物化式执行路径:把 S=QK^TP=softmax(S) 写进 HBM;右图采用 FA2 及后续常见的 Q-block-owned 前向顺序,让一个 Q tile 扫完 K/V tiles 后写出完成的 O tile。FA1 原始 K/V 外层循环会重复读写部分 O 与归一化统计,但两种顺序都不在 HBM 物化完整 S/P;左图也不代表所有“非 FlashAttention”backend 都必然物化两张矩阵。

← 左右滑动查看完整链路 · 打开原图

左侧把完整 S 与归一化后的 P 写入 HBM;右侧是 FA2+ 常见顺序,P̃_tile=exp(scores-m) 只是当前 tile 的临时未归一化指数权重,和 V_tile 相乘后进入累加器,扫完才除以 。FA1 原始循环会反复读写部分 O/m/ℓ,但同样不物化完整 S/P。

FlashAttention 省了什么

  • 省掉注意力分数矩阵和概率矩阵的 HBM 往返读写。
  • 前向保存输出及 log-sum-exp 等少量统计量而非完整概率矩阵;反向时重算局部 score/softmax。启用 dropout 时还要复现对应随机掩码。
  • 原论文把标准注意力的额外显存随序列长度从二次降为线性量级;端到端峰值仍包含 Q/K/V、输出及模型其他激活。

它没省什么

  • 没有把全注意力的两两 token 关系改成线性复杂度。
  • 没有消除 KV Cache 随上下文线性增长的容量账;但官方实现可在 decode kernel 中读取/更新连续或 paged KV Cache,两者是可组合的不同层次。
  • 不是所有 mask、head_dim、dropout、硬件组合都能走最快 kernel。
IO

原论文算的是“搬多少个 word”,不是只数 FLOPs

Two-level I/O model
在 FA1 的两级内存模型中,N 是序列长度,d 是单头维度,M 是可用片上快存储能容纳的标量数。对论文定义的物化式基线与 d ≤ M ≤ Nd 范围,前向 HBM 访问量为:
Θ(Nd+N2)materialized attentionvs.Θ ⁣(N2d2M)FlashAttention-1HBM word accesses\underbrace{\Theta(Nd+N^2)}_{\text{materialized attention}}\quad\text{vs.}\quad\underbrace{\Theta\!\left(\frac{N^2d^2}{M}\right)}_{\text{FlashAttention-1}}\quad\text{HBM word accesses}

两边的 dense 计算量主项仍是 Θ(N²d);FlashAttention 是用分块与局部重算换掉 中间张量的 HBM 往返。原论文只证明“没有一个算法能在所有 M 上同时渐近更好”;2024 年后续理论工作才给出更细的点态结论:当 M ≥ d² 时该上界在常数因子内匹配下界,而 M < d² 另有更优算法。真实 kernel 还受 tile 对齐、occupancy、寄存器压力和硬件流水线影响,公式不是运行时预测器。

04 · 演进

FA1 到 FA4:从 IO-aware 到硬件流水线共设计

FlashAttention 每一代都在回答同一个问题:当 GPU 的瓶颈变了,attention kernel 该怎么重新排流水线。

版本年份核心变化关键数字 / 适用硬件
FlashAttention2022tiling + online softmax,减少 HBM 与片上存储间读写;dense 路径保持精确注意力定义。论文在其特定训练基线中报告:GPT-2、1K 序列约 ,Long Range Arena、1K-4K 约 2.4×
FlashAttention-22023减少非矩阵乘 FLOPs,改进 thread block / warp 的工作划分,提高并行度。论文在 A100 kernel benchmark 中报告约为 FA1 的 ,达到理论峰值 FLOP/s 的 50%-73%;不是任意模型的固定加速比。
FlashAttention-32024面向 Hopper:TMA + Tensor Core 异步重叠、warp specialization、FP8 路径。最终论文在 H100 上报告 BF16 最高约 840 TFLOP/s(85%)、约为 FA2 的 1.5-2×;FP8 最高约 1.3 PFLOP/s
FlashAttention-42026面向 Hopper/Blackwell:算法与 kernel pipelining 共设计,CuTe-DSL 实现;Blackwell 路径使用 TMEM、2-CTA MMA 等能力。论文在 B200 BF16 benchmark 中报告最高 1613 TFLOP/s(71%)、最高为 cuDNN 9.13 的 1.3×;截至 2026-07-14,PyPI 最新为 4.0.0b21 预发布版。
一句话趋势
FA1 解决“别把大矩阵搬来搬去”;FA2 解决“GPU 线程怎么分工更满”;FA3 解决“Hopper 上计算和搬运怎么异步重叠”;FA4 解决“Blackwell 上 matmul 变得更快后,softmax / shared memory / pipeline 又成了新瓶颈”。
05 · 落地

在训练和推理里,它通常藏在框架后面

你不一定直接调用 flash_attn_func。很多时候,PyTorch 或上层训练/推理框架会根据输入形状、硬件与已安装版本选择 kernel;“配置里写了 flash”不等于每个 batch 都走同一条实现。

API

PyTorch SDPA:自动在多种实现里选

scaled_dot_product_attention
PyTorch 的 torch.nn.functional.scaled_dot_product_attention 会按设备和输入在多个 backend 间选择,包括 FlashAttention、Memory-Efficient、C++ math 与可用时的 cuDNN 路径。只有强制指定 fused backend 而它不适配时,系统才会用 warning 解释原因;默认自动选择不等于一定走 FlashAttention。
import torch.nn.functional as F

out = F.scaled_dot_product_attention(
    query, key, value,
    attn_mask=None,
    dropout_p=0.0,
    is_causal=True,
    enable_gqa=True,
)
容易踩坑
SDPA 会始终按传入的 dropout_p 应用 dropout,不会自行读取模块的 train/eval 状态;模块应显式写成 dropout_p=(self.p if self.training else 0.0)。布尔 attn_maskTrue 表示“参与注意力”,且当前 API 不允许同时传 attn_maskis_causal=Trueenable_gqa 仍标为实验性:CUDA 上只支持 Flash 与 math backend,要求 Q 头数能整除 KV 头数、K/V 头数相同,且当前不支持 Nested Tensor(嵌套张量);具体限制应以所用 PyTorch 版本为准。
训练时 training

收益常体现在更长序列、更大 batch、更低激活显存。打包短样本时,优先使用实现原生支持的 varlen/unpadding 边界;显式构造 dense block-diagonal mask 可能自己产生 成本,或让输入退出最快路径。

推理 prefill / decode

长 prompt 的 prefill 通常更能吃满 attention kernel;decode 每步 query 很短,常更受权重/KV 读取与调度影响,但并非“FlashAttention 无用”。官方实现自 2.2 起有小 query 的 split-KV decode 优化,flash_attn_with_kvcache 还能融合旋转位置编码、cache 更新与 paged KV 读取。

你看到的现象可能原因怎么判断
开了 FlashAttention 但没变快序列太短、head_dim/mask/dropout 不适配、batch 太小或瓶颈不在 attention。看 profiler 和实际 kernel 名;需要诊断时用 sdpa_kernel 强制 backend,让 PyTorch 报告不适配原因。
训练显存明显下降没有保存完整 attention probabilities,反向局部重算。对比峰值显存;长序列越明显。
长 prompt 首字更快prefill 阶段 attention kernel 更高效。分别看 TTFT 和 decode TPOT;不要只看总耗时。
输出有细微数值差异fused kernel 改变浮点运算顺序,低精度路径误差不同。对确定性/高精度敏感任务,固定 backend 或使用 math 实现复核。
06 · 常见误区

别把 FlashAttention 和所有“高效注意力”混成一类

误区 1:复杂度变线性

错。FlashAttention 仍是全注意力,计算量仍随 增长;线性注意力、SSM、稀疏注意力才是在关系图上少算。

误区 2:它替代或不能结合 KV Cache

都不对。FlashAttention 不消除历史 K/V 的容量,但 decode kernel 可以直接消费、更新甚至分页读取 KV Cache;算子 IO 重排与缓存生命周期/分页调度是不同层次,可协同实现。

误区 3:永远无脑更快

不一定。短序列、特殊 mask、老 GPU、CPU/MPS、需要 float64 或严格确定性时,可能回退或收益很小。

速查

一句话总结

Cheat sheet
FlashAttention 是什么IO-aware exact attention kernel:用 tiling + online softmax 减少 HBM 读写。
为什么快不落盘完整 N×N 注意力矩阵,让局部分数、指数权重与归一化统计在片上 tile 循环中短暂存在。
最适合哪里训练长序列、推理 prefill,以及长 KV 的特定 decode 形状;收益必须按目标 workload profile。
相邻但不同KV Cache / PagedAttention 管历史状态,continuous batching 管请求调度,稀疏/线性注意力改变关系或算术;它们不能被 FlashAttention 替代,但可以组合。
截至 2026-07-14官方 repo 提供 FA1/2、独立的 FA3 beta,以及可 pip install flash-attn-4 的 CuTe-DSL 版 FA4;PyPI 最新 FA4 为 4.0.0b21 预发布版,框架实际 dispatch 仍取决于版本、硬件和输入。
资料来源

主要参考

核查日期:2026-07-14。论文中的 TFLOP/s、利用率与加速比都是指定 GPU、形状、精度和对照实现下的 benchmark,不应直接外推为端到端模型的固定收益。