公司动态

大模型之FlashAttention技术详解:用 IO-aware 重写精确注意力,让长上下文不再困在显存里

📅 2026/9/2 7:47:12
大模型之FlashAttention技术详解:用 IO-aware 重写精确注意力,让长上下文不再困在显存里
写在前面【从零走向AGI】旨在深入了解通用人工智能AGI的发展路径从最基础的概念起逐步构建完整的知识体系。项目地址https://github.com/AI-mzq/From-Zero-to-AGI.git魔方AI空间猫先生从零走向AGI面试面经AIGC算法岗/开发岗面试面经交流社群涵盖AI Agent、AIGC图像创作、AI视频、LLM大模型、AI多模态、数字人、传统深度学习、具身智能等AIGC面试干货资源欢迎大家加入https://t.zsxq.com/YtJ09FlashAttention用 IO-aware 重写精确注意力让长上下文不再困在显存里导读FlashAttention 处理的不是“注意力公式算得还不够快”而是一个更贴近 GPU 现实的问题标准实现会把N × N N\times NN×N的分数矩阵和 softmax 后的概率矩阵反复写入、读出 HBM。序列变长时瓶颈往往不是乘法器而是显存带宽与中间张量占用。Tri Dao 等人把注意力改写为一个IO-aware 的精确算法。它仍计算同一个 softmax attention没有用低秩、随机特征或稀疏近似替换模型变化在执行顺序。Q、K、V 按块搬入片上 SRAM在一个 fused CUDA kernel 中完成矩阵乘、在线 softmax 和输出累积避免将巨大的S Q K ⊤ SQK^\topSQK⊤、P softmax ⁡ ( S ) P\operatorname{softmax}(S)Psoftmax(S)物化到 HBM。训练反向传播则保存少量归一化统计量再在 SRAM 中重算局部结果。原文的路线很清楚先用 GPU 内存层级解释为什么 FLOPs 不是充分指标再给出 tiling 与 recomputation 的算法随后证明 HBM 访问量更低、在一段 SRAM 范围内达到最优最后用训练、长序列任务和算子基准验证收益。论文的历史意义也在于它把“算法复杂度”向“数据搬运复杂度”推进了一步后来各代 FlashAttention 的优化都延续了这个问题意识。猫先生认为FlashAttention 最重要的贡献不是一段 CUDA 代码而是把注意力的优化目标从“少算一点”改成了“少把不该落盘的中间结果搬来搬去”这也是大模型算子走向硬件协同的典型转折。论文标题FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness论文地址https://arxiv.org/abs/2205.14135HTML 正文与原图https://arxiv.org/html/2205.14135v2GitHubhttps://github.com/Dao-AILab/flash-attention原文脉络第 1 节指出许多近似注意力虽然降低了理论 FLOPs却未必获得 wall-clock 加速因为它们没有处理访存开销。第 2 节把 GPU 视为 HBM 与 SRAM 组成的分层存储系统回顾标准 attention 如何生成并保存两个N × N N\times NN×N张量。第 3 节是核心分块、在线归一化、反向重算与 IO 复杂度分析随后把同一原语推广到 block-sparse attention。第 4 节分别测试训练速度、可用上下文长度带来的模型质量以及 attention kernel 的运行时间和显存占用第 5 节讨论多 GPU、其他核方法等尚未完成的方向。1-2. 从二次张量到内存墙标准 attention 卡在哪里设每个头的Q , K , V ∈ R N × d Q,K,V\in\mathbb{R}^{N\times d}Q,K,V∈RN×d标准计算为S Q K ⊤ , P softmax ⁡ ( S ) , O P V . SQK^\top,\qquad P\operatorname{softmax}(S),\qquad OPV.SQK⊤,Psoftmax(S),OPV.数学上最显眼的是S , P ∈ R N × N S,P\in\mathbb{R}^{N\times N}S,P∈RN×N的二次规模。工程上更糟的是通常的实现把S SS写到 HBM读回去做 softmax 后再写出P PP最后又读P PP与V VV计算输出。mask、dropout 等操作还会加剧往返。即使矩阵乘本身能高效使用 Tensor Coresoftmax 和逐元素操作仍常是 memory-bound。图 1论文 Figure 1。左侧把 SRAM、HBM、DRAM 的带宽与容量差异摆在一起中间显示 FlashAttention 按 K/V 块和 Q 块循环虚线框中的N × N N\times NN×N注意力矩阵不再落到 HBM右侧给出 GPT-2 attention kernel 相对 PyTorch 的 7.6× 加速。以 A100 为例原文给出的 HBM 带宽约为 1.5–2.0 TB/s而每个 SM 的片上 SRAM 约 192 KB、估算带宽约 19 TB/s。SRAM 快得多却装不下完整 attention matrixHBM 容量大却是数据移动的关口。因而“融合 kernel”还不够若中间矩阵为了反传仍要保存在 HBM融合只能消去一部分读写。近似 attention 的经验教训在此处尤为关键。把复杂度从二次写成线性不能自动保证更快不规则访问、额外投影和低算术强度可能把节省的乘法又花在 IO 上。FlashAttention 首先保留精确性然后重排计算以压缩 HBM traffic。3. FlashAttention分块、在线 softmax 与重算3.1 Tiling让计算留在 SRAM别让中间矩阵出门算法把 Q 切成行块把 K、V 切成列块。外循环将一块 K/V 从 HBM 读入 SRAM内循环依次载入 Q 块在 SRAM 中求局部Q i K j ⊤ Q_iK_j^\topQi​Kj⊤​、局部 softmax 所需统计量并更新输出块。完成后只把最终的O i O_iOi​及少量统计量写回 HBM而不保存全局S SS或P PP。困难在于 softmax 不能粗暴地“每块各算一次再拼接”。对一行分数维护当前最大值m mm与指数和ℓ \ellℓ新块到来时先以m ′ max ⁡ ( m , m new ) m\max(m,m_{\text{new}})m′max(m,mnew​)对旧输出和旧分母重新缩放再加入新块的exp ⁡ ( S − m ′ ) \exp(S-m)exp(S−m′)。这就是 online softmax每次只见到一小块最终的归一化却与一次性看见全行完全一致。图 2论文 Figure 2。原文以不同 SRAM 容量和序列长度比较 HBM 访问量并展示 FlashAttention 在相同精确计算下显著压缩显存与运行时间图中关键信息是 IO 降低而非 FLOPs 降低。这种循环顺序带来一个常被误解的取舍FlashAttention 可能做更多 FLOPs。分块造成重复加载/局部计算反传又要重算分数与概率但这些额外算术发生在高速片上存储附近换回了昂贵的 HBM 读写。对 memory-bound attention增加可控计算反而缩短端到端时间。猫先生认为online softmax 才是这项工作的“算法铰链”。没有它分块只是矩阵乘优化有了它精确的全行归一化才能在不保存全行分数的条件下成立。3.2 Recomputation反传不保存N 2 N^2N2保存能重建它的线索训练时标准反向传播通常需要S SS或P PP。FlashAttention 前向保存输出O OO和每一行的 softmax 归一化统计量( m , ℓ ) (m,\ell)(m,ℓ)而不是保存两个二次矩阵。反向阶段重新把 Q/K/V 分块调入 SRAM重建当前块的S , P S,PS,P立即累积d Q , d K , d V dQ,dK,dVdQ,dK,dV。这可视作选择性 gradient checkpointing但结论与常见“省显存必然慢”不同重算替代的是 HBM 中的大量读写在本论文的硬件条件下反而加速。代价也不能忽略kernel 实现更复杂数值稳定性、mask、dropout、不同 head dimension 与 GPU 架构都要细致处理它不是把三行 PyTorch 自动融合即可得到的效果。3.3 IO 复杂度把理论下界写到访存上论文设 SRAM 容量为M MM给出 FlashAttention 的 HBM 访问量为O ( N 2 d 2 M ) , O\left(\frac{N^2d^2}{M}\right),O(MN2d2​),而标准实现至少需要Ω ( N d N 2 ) \Omega(NdN^2)Ω(NdN2)的 HBM 访问。直觉上SRAM 越大每次可容纳的 Q/K/V 块越大越少回到 HBM但它仍不能容纳完整N 2 N^2N2矩阵。论文还给出下界说明在一段 SRAM 大小范围内没有另一个精确 attention 算法能在渐近意义上普遍少于它的 HBM 访问。这项分析值得单独记住同是O ( N 2 ) O(N^2)O(N2)算术执行图不同内存系统看到的工作量可以相差数量级。它也解释了为什么 FlashAttention 不是对 softmax attention 的近似替身而是可作为基础算子直接替换。3.4 Block-sparse 扩展先用 Flash 的 IO 纪律再引入稀疏性原文进一步选择块稀疏模式只对保留的 Q-K 块进行计算。稀疏带来更少的算术与更少的访问但规则以块为单位避免细粒度稀疏的索引和访存碎片。论文报告它相对 FlashAttention 本身可快 2–4×并把可处理长度推到 64K。需要区分两层价值FlashAttention 的主算法是 exact attentionblock-sparse FlashAttention 才是近似 attention。前者解决“同一模型怎样跑得更像硬件”后者才交换了模型连接模式与效率。把两者混称为“FlashAttention 近似注意力”会丢掉这篇论文最关键的定位。4. 实验速度、可用长度与质量的证据4.1 训练吞吐并非只有 kernel 跑分原文在 BERT-large长度 512上相对 MLPerf 1.1 训练速度纪录实现 15% 端到端加速GPT-2长度 1K相对 HuggingFace 与 Megatron-LM 基线约 3×Long Range Arena 的 1K–4K 设置约 2.4×。Figure 1 中单 attention 的 7.6× 不应直接外推到整模型端到端还包含 MLP、通信、优化器和数据流水线因此 15% 反而是更诚实的系统指标。4.2 更长上下文才是显存节省的用途论文不把内存曲线当作终点而是把释放出来的容量用于训练更长序列。GPT-2 的长上下文设置得到 0.7 的 perplexity 改善长文档分类得到 6.4 个点的提升。Path-X16K上FlashAttention 支持的 Transformer 首次超过随机猜测报告 61.4% accuracyblock-sparse 版本在 Path-25664K达到 63.1%。这些结果支持的是一条具体链路更省的 attention 内存 → 能放入更长输入 → 下游任务获得更多上下文信息。它们并不证明所有任务都仅靠加长窗口变好也不能替代对数据、位置编码和优化稳定性的评估。图 3论文 Figure 3。不同序列长度下的 runtime / memory benchmark。精确 FlashAttention 在常用短中长度上相对标准实现更快且更省长度继续增长后某些近似方法可能更快而 block-sparse FlashAttention 则把两类优势组合。4.3 基准的边界论文报告 FlashAttention 在长度 128–2K 的常用范围可比标准实现快至 3×并能扩展到 64K长度不超过 512 时它在文中比较的 attention 方法中同时更快、更省。超过 1K 后Linformer 一类近似方法可能开始更快这恰好说明“精确、速度、长度”并非永远同时最优。硬件版本也很重要。原始实验围绕当时的 GPU、CUDA 和框架基线后来的 FlashAttention-2/3 继续通过并行划分、warp 调度、低精度与异步流水改善常数项。读 v1 时应把这些数字看作 IO-aware 原理的首证而不是今天任意显卡、任意模型的承诺。5. 局限、复现与后续问题FlashAttention 没有改变全连接 attention 的O ( N 2 ) O(N^2)O(N2)算术本性当序列极长且 SRAM 相对有限IO 虽降为次二次形式计算量仍会成为问题。block-sparse 能更进一步但稀疏图案是额外归纳偏置是否适合任务不能由 kernel benchmark 决定。多 GPU attention、其他 kernel regression 类操作和更广泛硬件支持在原文中也仍属未来方向。复现时应避免只测forward。训练场景要同时核对前反向、dropout、causal mask、mixed precision、head dimension、batch size 和真实峰值显存服务推理还要区分 prefill 与 decode后者常受 KV cache 和小批量访存支配。现代框架中的scaled_dot_product_attention或flash-attn已可能自动选择内核但回退路径与实际 GPU 能力仍会决定是否得到预期收益。总结把注意力优化从算式带到数据移动FlashAttention 的核心思想可以收在三点第一精确 softmax attention 可以通过分块与在线归一化避免物化二次中间矩阵第二保存少量统计量、在反向重算局部块能以额外算术换掉更昂贵的 HBM 流量第三评估应同时看 IO、峰值显存和端到端时间而不是只比较 FLOPs。如果只带走一个判断那就是在现代 GPU 上决定一个深度学习算子快不快的不只是它计算了什么也是谁以什么顺序搬运了数据。FlashAttention 给出的不是一个孤立技巧而是一种把算法、编译与硬件层级一起纳入建模的方式。推荐阅读► 技术资讯 魔方 AI 新视界► 项目应用开源视界► 技术专栏 多模态大模型最新技术解读专栏 | AI 视频最新技术解读专栏 | 大模型基础入门系列专栏 | 视频内容理解技术专栏 | 从零走向 AGI 系列