2307.08691-flashattention-2-parallelism-work-partitioning
FlashAttention 2: Faster Attention with Better Parallelism and Work Partitioning
FlashAttention-2 的核心贡献是把 FlashAttention v1 解决 HBM traffic 后剩下的性能瓶颈继续拆开:减少 expensive non-matmul FLOPs,把 attention 计算沿 sequence dimension 分给更多 thread blocks 提高 SM occupancy,并把 warp 内 work partition 从 sliced-K 调整为 sliced-Q 以减少 shared memory 通信,从而在保持 exact softmax attention 语义的同时,把 A100 上的 attention 吞吐提升到接近 GEMM 的一半以上效率。
Source
- Title: FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning
- arXiv: https://arxiv.org/abs/2307.08691
- PDF: https://arxiv.org/pdf/2307.08691
- Code/Project: https://github.com/Dao-AILab/flash-attention
- Blog: https://hazyresearch.stanford.edu/blog/2023-07-17-flash2
- Authors: Tri Dao
- Submitted: 2023-07-17
- Current version read: v1, submitted 2023-07-17
- Subjects: Machine Learning (cs.LG)
作者与关系
- Tri Dao: Stanford University.
阅读目标与判断边界
本笔记关注:
- FlashAttention v1 已经减少 HBM traffic 后,为什么仍然只有 25-40% 左右的理论 FLOPs/s。
- FlashAttention-2 如何通过 algorithm tweak、block-level parallelism 和 warp-level work partitioning 提升 GPU 利用率。
- 论文实验如何支撑 2x 左右 kernel speedup 与 1.3x 左右端到端训练 speedup。
- 它在 FlashAttention、Lightning Attention、long-context serving 和 RL rollout/trainer consistency 图谱中的位置。
判断边界:
- FlashAttention-2 仍是 exact softmax attention kernel;它改变执行路径和并行分解方式,attention 数学语义保持不变。
- 论文主实验基于 A100 80GB SXM4,另给出 H100 上无需特殊 Hopper 指令的初步结果;后续 FlashAttention-3/4 已继续改变 baseline。
- 它优化训练、finetuning 和 inference 中的 attention primitive;对于 1M context 级别的总成本,quadratic arithmetic 仍然存在。
- 论文目标是性能与内存效率;determinism、batch-invariant behavior、trainer/rollout logprob 一致性需要结合具体 kernel 和 serving engine 单独审计。
论文脉络
1. 研究问题、背景和价值
标准 attention 计算为:
其中
2205.14135 的 FlashAttention v1 已经把这个问题从“降低 FLOPs”推进到“降低 HBM traffic”:它用 tiling、online softmax 和 backward recomputation 避免将完整
FlashAttention-2 处理的是 v1 之后的下一层问题:当 HBM traffic 已经显著降低,attention kernel 仍然明显慢于高度优化的 GEMM。论文给出的判断是,FlashAttention v1 的 forward pass 约达到理论峰值的 30-50%,backward pass 约达到 25-35%;优化 GEMM 可以达到 80-90%。这说明 kernel 仍有结构性效率损失。
这个问题很重要,因为长上下文模型让 attention kernel 成为训练和推理共同瓶颈。2023 年时 GPT-4 32k、MPT 65k、Claude 100k 等长上下文需求已经出现;如果 exact attention kernel 不能更接近 GEMM 效率,工程上会更早转向 sparse、linear 或 hybrid attention。FlashAttention-2 的价值在于,它继续延长了 standard softmax attention 的工程可用边界。
2. 已有解决方案与不足
已有路线大致分为三类。
第一类是 approximate attention,例如 sparse、low-rank、kernelized attention。这些方法降低了 theoretical complexity,但可能改变模型质量、需要结构重训,并且实际 GPU kernel 不一定快。
第二类是 fused attention / xFormers / Triton 实现。它们在 mask、softmax、dropout、matmul 融合上有效,但性能取决于具体 work partitioning、head dimension、causal pattern 和 GPU 占用率。
第三类是 FlashAttention v1。它保留 exact attention 语义,并通过 IO-aware tiling 显著减少 HBM 读写。剩余不足来自三处:
- non-matmul FLOPs 在 A100 上相对昂贵。论文用 A100 说明:FP16/BF16 matmul 理论峰值约 312 TFLOPs/s,FP32 non-matmul 约 19.5 TFLOPs/s,单个 non-matmul FLOP 的机会成本约是 matmul FLOP 的 16 倍。
- v1 主要并行在 batch 和 head 维度。长序列训练经常 batch size 较小,或者 head 数较少,线程块数不足以填满 A100 的 108 个 SM。
- v1 的 warp partition 使用 sliced-K:不同 warps 切分
,中间结果需要写入 shared memory、同步、再相加,导致 shared memory 读写和 warp 间通信开销。
3. 作者可能的思考路径
可以把作者思路重建为一个逐层 profiling 的过程。
首先,FlashAttention v1 已经证明 exact attention 的主要内存问题可以通过 tiling 和 online softmax 解决。此时继续把问题归因于
其次,attention 与 GEMM 的差距提供了明确参照物。两者都做大量矩阵乘,但 attention 额外包含 softmax、rescale、mask、dropout、row reduction、online statistics update。GPU 对 matmul 有专门 Tensor Cores,而这些 elementwise / reduction 更难达到高吞吐。因此一个自然 intuition 是:保留 FA1 的 IO-aware block structure,同时把所有可避免的 non-matmul 操作推迟、合并或删除。
再次,长序列导致 batch size 下降。FA1 只按 batch 和 head 生成 thread blocks,在 batch/head 数较小时会产生低 occupancy。既然 attention matrix 在 row block 上可以独立计算输出,就可以把 sequence row blocks 也作为 parallelism 维度,让单个 head 内部也能拆给多个 thread blocks。
最后,warp-level partitioning 暴露出更底层的通信问题。FA1 sliced-K 会让多个 warps 分别得到 partial output,再通过 shared memory 合并。若改成 sliced-Q,每个 warp 负责不同 query rows,读共享的
4. 核心假设或切入点
FlashAttention-2 的核心假设是:在 FA1 已经降低 HBM traffic 后,attention kernel 的主要可优化空间转向“把执行路径变得更像 GEMM”。具体包括:
- 尽量减少 non-matmul FLOPs。
- 增加 thread block 数量,让 SM 有足够工作。
- 减少 shared memory 上的中间结果交换。
- 保持 exact online softmax,不牺牲模型语义。
5. 方法 / 系统 / 理论框架
FA2 的方法容易被混在一起理解,因为它同时改了算法状态、block tiling、thread block 调度和 warp-level 分工。更清晰的读法是按 GPU 执行层级拆开:
| 层级 | 操作对象 | FA1 已有内容 | FA2 的主要变化 | 解决的瓶颈 |
|---|---|---|---|---|
| 算法语义层 | online softmax / output accumulator | 分块 exact softmax,保存 |
维护 unscaled |
减少 non-matmul FLOPs 和状态写回 |
| Block tiling 层 | 沿 sequence 切 block,避免物化 |
保留同一类 tiling,但调整 loop order 让 |
保持 IO-aware exact attention | |
| Thread-block / SM 调度层 | GPU thread blocks / SM occupancy | 主要按 batch 和 head 并行 | 额外沿 sequence row blocks 并行 | 长序列小 batch 时提高 occupancy |
| Warp-level 分工层 | 一个 thread block 内的 4/8 个 warps | sliced-K,warps 切 |
sliced-Q,warps 切 |
减少 shared memory 通信和同步 |
| 特殊模式层 | causal mask、MQA/GQA、head dim | 基础支持有限 | block-level causal skip、MQA/GQA index mapping、head dim up to 256 | 覆盖真实 LLM 架构与 inference 需求 |
下面的 5.1-5.6 按这几个层级展开。
5.1 Block Tiling 层:一次 forward kernel 如何跑
先把一个 attention head 看成矩阵:
FA1 和 FA2 都按 block 处理 sequence,这是 block tiling 层的共同基础:
其中
在 FA2 forward 中,每个 thread block 负责一个
对每个
这一步是主 matmul,目标是尽量让 Tensor Cores 持续工作。随后对
一个
- 初始化
、 、 。 - 对每个
block 计算 。 - 用 online softmax statistics 合并当前 block。
- 循环结束后把
统一除以 ,得到最终 。 - 只把
和 row-wise logsumexp 写回 HBM。
因此 FA2 forward 的 HBM 写回对象很少:
其中
5.2 算法语义层:减少 non-matmul FLOPs
FlashAttention v1 的 online softmax 已经能 exact 地分块计算 softmax,但循环内部还有较多 rescale、除法、bound check 和 mask 处理。FA2 的核心调整是把循环内部的非矩阵乘操作降到更少,因为 A100 上 FP16/BF16 matmul 理论吞吐约 312 TFLOPs/s,FP32 non-matmul 理论吞吐约 19.5 TFLOPs/s。一次 elementwise / reduction 操作在性能上会占用更高机会成本。
FA2 的第一处改动是维护 unscaled output
先更新当前已经看过的 key 范围内的 row max:
再用新的
row sum 的更新为:
output accumulator 的更新为:
循环内不再每个 block 都把 output 变成 normalized
这样做的收益来自两点:
- 每次 block update 少做若干除法和 rescale。
- output accumulation 更接近 matmul + simple elementwise 的形式,让时间更多花在 Tensor Core matmul 上。
第二处改动是 forward 只保存 logsumexp:
在 backward 中,softmax probability 可由
这比同时保存
5.3 Backward 数据流层:用重算换 HBM traffic
FA2 的 backward 延续 FA1 的思想:forward 不保存完整
其中
这个量可先按 row 计算并写入 HBM,之后每个 block 读取对应
按 block 看,backward 的执行过程是:
- 选择一个 column block
,load 到 SRAM。 - 在 SRAM 中初始化
、 。 - 遍历所有 query row blocks
。 - 对每个
重算 。 - 用 forward 保存的
重建:
- 累加 value gradient:
- 计算 probability gradient:
- 计算 score gradient:
- 更新 query gradient:
- 累加 key gradient:
最后把
Backward 的计算量更大,因为它包含更多 matmuls,并且需要重算
5.4 Thread-block / SM 调度层:沿 sequence dimension 增加 parallelism
FA1 主要按 batch 和 head 维度生成 thread blocks:
A100 有 108 个 SM。短序列、大 batch 时,这个并行度通常足够;长序列训练会提高单样本显存,batch size 往往下降,thread blocks 数可能不足以填满 SM,GPU 出现 low occupancy。
FA2 将 forward 外层循环改成 query row blocks,让每个
其中:
这对长序列特别有效。举例来说,若 batch size 很小、head 数也不大,FA1 可能只能发出几十个 thread blocks;FA2 可以把每个 head 内部按 sequence rows 拆成更多 blocks,让 108 个 SM 都有工作。
Forward 中不同
Backward 的并行更复杂。FA2 按 column blocks 分配 worker,因为
5.5 Warp-level 分工层:从 sliced-K 到 sliced-Q
Thread block 内部通常有 4 或 8 个 warps。这里讨论的是一个 thread block 内部如何把当前 tile 的计算分给多个 warps,层级低于
FA1 forward 使用接近 sliced-K 的方式:多个 warps 切分
这个流程的瓶颈是 shared memory communication:
FA2 改成 sliced-Q:多个 warps 切分
由于 query rows 彼此独立,每个 warp 得到的 output slice 天然属于最终输出的不同 rows。warps 之间无需合并 output,shared memory 上的中间结果交换显著减少。
可以把两种分工理解为:
| 方案 | 切分对象 | 每个 warp 产物 | 需要合并吗 | 主要成本 |
|---|---|---|---|---|
| FA1 sliced-K | 同一 |
需要 | shared memory 写读、同步、reduction | |
| FA2 sliced-Q | 不同 |
基本不需要 | 更多 |
FA2 的选择更适合 forward,因为 forward 的输出天然按 query rows 分块。Backward 涉及
5.6 特殊模式与 tuning 层:causal mask、MQA/GQA、block size 与 head dimension
Causal attention 中,若某个 block 的 key column indices 全部位于 query row indices 之后,这个 block 对应未来 token,可以直接跳过。对于足够大的
需要真正应用 mask 的主要是对角线附近的 block。对角线以下的 blocks 全部合法;对角线以上的 blocks 全部跳过;只有跨过对角线的 block 需要在 block 内按元素设为
FA2 也支持 MQA/GQA。设 query heads 数为
其中
block size 方面,论文通常在
Hazy Research blog 补充说明 FA2 支持 head dimension up to 256,使 GPT-J、CodeGen、CodeGen2、Stable Diffusion 1.x 等模型可以使用这一实现。head dimension 越大,单个 tile 的 SRAM / register 需求越高,block size 和 warp 数的选择也更敏感。
关键实验/定理
结果 1:attention kernel microbenchmark
- 设置:A100 80GB SXM4;sequence length 从 512 到 16k;总 tokens 固定为 16k;hidden dimension 为 2048;head dimension 为 64 或 128;比较无 causal mask / 有 causal mask。
- 指标:forward、backward、forward+backward 的 TFLOPs/s。
- FLOPs 计算:forward 使用
causal mask 下约减半;backward FLOPs 约为 forward 的 2.5 倍,因为 forward 有 2 个 matmuls,backward 加 recomputation 后有 5 个 matmuls。
- 结果:FA2 相比 FlashAttention v1 快 1.7-3.0x;相比 Triton 版本 FA1 快 1.3-2.5x;相比 PyTorch standard attention 快 3-10x。A100 上最高约 230 TFLOPs/s,达到理论峰值 73%。
- 解读:这直接支撑论文的主要系统结论:FA1 之后仍有大量效率损失来自 work partitioning 和 non-matmul overhead。
结果 2:H100 初步 benchmark
- 设置:相同实现直接运行在 H100 SXM5,未使用 TMA、4th-gen Tensor Cores、FP8 等 Hopper-specific features。
- 指标:forward+backward TFLOPs/s。
- 结果:最高达到 335 TFLOPs/s。
- 解读:FA2 的基本 work partitioning 能从 A100 迁移到 H100,但论文也明确将 Hopper-specific optimization 留给后续工作。这个结果说明方向可迁移,尚未代表 H100 上的上限。
结果 3:GPT-style model end-to-end training
- 设置:8 x A100 80GB SXM;GPT-style 1.3B / 2.7B models;context length 为 2k / 8k。
- 指标:TFLOPs/s/GPU。论文使用 Megatron-LM 等常见口径:
- 结果:
| Model | Without FlashAttention | FlashAttention v1 | FlashAttention-2 |
|---|---|---|---|
| GPT3-1.3B, 2k context | 142 TFLOPs/s | 189 TFLOPs/s | 196 TFLOPs/s |
| GPT3-1.3B, 8k context | 72 TFLOPs/s | 170 TFLOPs/s | 220 TFLOPs/s |
| GPT3-2.7B, 2k context | 149 TFLOPs/s | 189 TFLOPs/s | 205 TFLOPs/s |
| GPT3-2.7B, 8k context | 80 TFLOPs/s | 175 TFLOPs/s | 225 TFLOPs/s |
- 解读:FA2 在端到端训练中最高达到 225 TFLOPs/s/GPU,即 72% model FLOPs utilization。相对没有 FlashAttention 的 baseline,8k context 下提升最大;相对 FA1,提升约 1.15-1.29x。kernel microbenchmark 的 2x 收益会被 MLP、communication、optimizer、data pipeline 等非 attention 部分摊薄,这是合理现象。
证据链强度评估
强证据
- A100 attention microbenchmark 覆盖 causal / non-causal、head dim 64 / 128、sequence length 512-16k,直接对应论文声称的 kernel-level speedup。
- 端到端 GPT-style 训练表说明 kernel 改动确实能传导到 training throughput,尤其在 8k context 下收益明显。
- 方法分析与 profiling 目标一致:减少 non-matmul、增加 sequence parallelism、减少 shared memory communication,都直接作用于 GPU roofline 中的具体瓶颈。
中等强度证据
- H100 结果证明可迁移趋势,但论文未使用 Hopper-specific features,因此它是过渡性 baseline。
- MQA/GQA 和 head dimension 256 支持来自实现与 blog 描述,论文正文给出机制说明,但实验表主要围绕 MHA、head dim 64/128。
- “可以用同等价格训练 16k context,接近此前 8k context 成本”是工程直觉,实际取决于模型尺寸、batching、通信和整体训练 stack。
需要谨慎的推论
- FA2 仍保留
arithmetic;超长上下文成本会继续随 增长。 - 论文没有系统比较后续 PyTorch SDPA、FlashAttention-3/4、FlashInfer、PagedAttention 等更新实现。
- Performance kernel 的正确性与 determinism 是两个问题。对于 RL rollout/trainer logprob matching,仍需测试 batch invariance、atomic add 顺序、precision path 和 backend 配置。
OpenReview / 审稿意见吸收
- Venue status: 当前档案未记录公开 peer-review 状态。
- Public reviews: 当前档案未记录可可靠匹配的 OpenReview / ARR / 会议 reviewer comments。
- Ratings / confidence: 无公开评分可用于校准。
- Reviewer consensus: 暂无。
- Main criticisms: 暂无公开 reviewer 质疑可引用;可信度主要由论文、技术报告、项目证据和本地一致性检查决定。
- Author response: 暂无公开 rebuttal 记录。
- 对本文可信度的影响: 按未完成公开审稿吸收处理,结论需要依赖实验设置、baseline 强度、复现证据和跨论文一致性校准。
本地讨论补充
1. 讨论收敛点
- 当前初版归档将 FA2 定位为 FlashAttention family 的第二个关键节点:FA1 解决
intermediate materialization 和 HBM traffic;FA2 解决 FA1 kernel 内部的 parallelism / work partitioning / shared memory communication。 - 对后续 long-context architecture 来说,FA2 是 exact softmax attention 能力边界的一次右移。Lightning Attention、DLA、DeepSeek-V4 CSA/HCA 等路线要放在这个边界之后理解。
2. 修正后的理解
- FA2 的“更快”主要来自执行路径优化,并未改变 attention family。它与 linear attention 的关系是系统思想相通,数学结构不同。
- FA2 kernel microbenchmark 中 2x 左右的提升进入端到端训练后会变成约 1.3x 以内,因为训练 step 还包含 MLP、normalization、optimizer、communication 和 data pipeline。
- FA2 的 sequence parallelism 解决的是单个 attention head 内 thread block 数不足的问题,特别适合长序列小 batch regime。
block tiling 和 sliced-K / sliced-Q 是两层概念。FA1 和 FA2 都会把整段 沿 sequence 维度切成 blocks,这是为了让当前 tile 能进入片上存储并配合 online softmax。sliced-K / sliced-Q 指当前 thread block 内多个 warps 如何分工。FA1 在 warp-level 上切 ,可理解为把当前 tile 再分给不同 warps,每个 warp 产生同一批 query rows 的 partial output,随后需要 shared memory 合并。FA2 在 warp-level 上切 rows,让每个 warp 产出不同 query rows 的 output slice,当前 tile 对这些 warps 共同可读。 - sliced-Q 中“共享
”的作用域是当前 thread block / warp group 内的 tile 共享:kernel 会把当前 tile 从 HBM 读入片上 SRAM / shared memory 和 registers,让同一个 thread block 里的多个 warps 共同读取。这个共享范围限于正在执行的 block;不同 row blocks 所在的 thread blocks 仍可能重新读取同一个 tile,L2 cache 可能提供命中,但 SRAM 是每个 SM 的片上资源,容量和生命周期都局限于当前 block。FA2 的收益主要来自避免 sliced-K 下的 partial output 合并和 shared memory reduction,整段 KV 长期缓存属于另一类问题。
3. 后续复验指标
- 在当前 PyTorch SDPA / FlashAttention-3 / FlashInfer baseline 下复测 A100/H100/H200/B200。
- 对 RL training stack 复测 rollout engine 与 trainer engine 的 logprob consistency,包括 causal mask、GQA、different precision、batch shape、padding layout。
- 对 deterministic inference 复测 atomic adds 与 block scheduling 对 backward / training reproducibility 的影响。
主要启发
- Long-context systems 的核心路径包含多层优化,单一“降低复杂度”不足以解释实际边界。即使 exact softmax attention 保持
arithmetic,HBM traffic、SM occupancy、warp communication 和 non-matmul FLOPs 也能决定可训练 context 的实际边界。 - GPU kernel 优化需要分层看:FA1 先改变 IO pattern,FA2 再改变 thread block / warp partitioning。许多系统论文也需要这样分解瓶颈,避免只用 FLOPs 或显存解释性能。
- 对 RL / agent 系统而言,attention kernel 是吞吐基础层;在追求 rollout speed 时,还要同时审计数值一致性、determinism 和 trainer/rollout backend mismatch。
局限
- 实验模型规模为 1.3B / 2.7B GPT-style models,不能直接代表更大 MoE、长输出 RL 或复杂 serving workload。
- 主实验基于 A100;H100 结果未使用 Hopper-specific features,后续硬件代际会改变性能排序。
- FlashAttention-2 保持 exact softmax attention,因此 arithmetic complexity 仍是 quadratic。
- block size 需要按 head dimension 和 device shared memory 手工 tuning;论文把 auto-tuning 留作 future work。
- 论文没有围绕 deterministic / batch-invariant kernel behavior 展开,这对 RL training-inference consistency 是额外工程问题。
跨论文关系
- 与 2205.14135:FA1 是 IO-aware exact attention 的基础;FA2 是同一作者线的直接延续,继续保留 exact softmax 语义,并把优化焦点转向 non-matmul FLOPs、sequence-level parallelism 和 warp-level work partitioning。
- 与 2405.17381:Lightning Attention 借鉴 IO-aware / block-wise 视角,但服务 causal linear attention。它的实验也把 FlashAttention-2 作为重要 softmax attention baseline。
- 与 2506.13585:MiniMax-M1 选择 Lightning Attention 支撑 long-output RL,说明在 1M input / 80K output 这类尺度上,FA2 这类 exact kernel 优化仍需与 architecture-level change 配合。
- 与 2606.10650:DLA 解决 linear attention memory state 的信息分配;FA2 解决 exact softmax kernel 的执行效率,两者共同构成长上下文 attention 系统基础。
- 与 2026-04-24:DeepSeek-V4 的 CSA/HCA、heterogeneous KV cache、deterministic kernels 延续了 FA/FA2 的硬件 data movement 视角,并把问题推到 million-token context。
- 与 2025-09-10 和 2605.14220:FA2 提供高吞吐 attention primitive;后续材料提醒,在 RL 闭环中还要检查 batch-invariant kernels、deterministic execution、precision path 和 rollout/trainer logprob consistency。
Reference Intake Brief
Target
- Intended target system: 新增论文笔记与现有 FlashAttention / long-context systems 图谱。
- Existing related assets:
content/utility/papers-index.md;2205.14135;2405.17381;2506.13585。 - Proposed form: 新建独立 Markdown 文档;更新当前收录和 FlashAttention 跨论文关系。
Reusable Elements
- FA1 -> FA2 的瓶颈演进:HBM traffic -> non-matmul FLOPs / occupancy / warp communication。
- sequence-dimension parallelism 解释:长序列小 batch 时,单个 head 内部也要拆成更多 thread blocks。
- sliced-K -> sliced-Q 的 warp partitioning 对比,可复用于讲解 GPU shared memory communication。
- 端到端 speedup 小于 kernel speedup 的解释,可复用于系统性能归因。
Risks
- Copyright/over-copying: 笔记只做结构化转述和少量公式重写,避免复制长段论文原文。
- Unsourced or unverifiable claims: 论文事实来自 arXiv v1、TeX source、Hazy Research blog 和 official GitHub README;后续硬件 baseline 需要另行复验。
- Tone/brand mismatch: 保持技术分析语气,避免营销化表达。
- Safety/compliance issues: 无直接安全或双用途操作风险;与 RL consistency 关系属于本地跨论文分析。
- Overlap with existing assets: 与
2205.14135-flashattention-io-aware-exact-attention.md有强重叠,本笔记专注 FA2 的 parallelism/work partitioning,不重复 FA1 的 IO lower bound 细节。
Skipped
| Material | Reason |
|---|---|
| FlashAttention-3 / FlashAttention-4 | 属于后续硬件代际优化,应独立归档 |
| 所有 benchmark 图的逐点数值 | TeX source 未提供原始数值表,当前记录论文给出的区间和端到端表 |
| 完整 CUDA/CUTLASS implementation walk-through | 当前任务为论文分析,kernel code 可在后续专题中单独读 |
Recommendation
Decision: merge
Why: FlashAttention-2 是已归档 FlashAttention 的直接延续,也是 Lightning Attention、MiniMax-M1、DLA、DeepSeek-V4、TIM/VeXact 等材料共同依赖的 attention kernel baseline;它补齐了“IO-aware exact attention 之后还要优化 GPU parallelism 和 work partitioning”的系统层链条。