当前位置: 首页 > news >正文

FlashAttention

FlashAttention

1. 使用背景

  • Transformer 很强,但 attention 一直有一个很现实的问题:又慢又吃显存。

  • 如果序列长度是NNN,标准自注意力里最核心的中间量是:
    S=QK⊤ S = QK^\topS=QK
    它的形状是N×NN\times NN×N
    这意味着,随着序列变长,attention 的中间结果会迅速膨胀。

  • 早期很多人一看到 attention 是二次复杂度,就会自然想到:是不是只能去做近似、稀疏化、低秩化。
    但 FlashAttention 提出了一个非常关键的观点:

  • 很多时候 attention 慢,不只是算得多,更是因为显存读写太重。

  • 也就是说,问题不只是 FLOPs,不只是数学复杂度,还包括一个很重要的系统因素:IO(输入输出 / 内存访问)成本

  • FlashAttention 就是在这个背景下提出来的。它最核心的思想可以概括成一句话:

  • 不改变标准 attention 的数学结果,而是通过分块、融合和在线 softmax,把大规模中间矩阵的显存读写降下来。


2. 理论基础

(1)标准 attention 到底在算什么

  • 给定输入隐藏状态,先通过线性变换得到:
    Q=XWQ,K=XWK,V=XWV Q=XW_Q,\qquad K=XW_K,\qquad V=XW_VQ=XWQ,K=XWK,V=XWV

  • 标准 scaled dot-product attention 写成:
    Attention(Q,K,V)=softmax(QK⊤d)V \text{Attention}(Q,K,V)=\text{softmax}\left(\frac{QK^\top}{\sqrt{d}}\right)VAttention(Q,K,V)=softmax(dQK)V

  • 如果把它拆开看,其实主要有三步:

  1. 先算分数矩阵:
    S=QK⊤d S=\frac{QK^\top}{\sqrt{d}}S=dQK
  2. 再做 softmax,得到注意力权重:
    P=softmax(S) P=\text{softmax}(S)P=softmax(S)
  3. 最后和VVV相乘:
    O=PV O=PVO=PV
  • 这个流程数学上很清楚,但在实现上会有一个大问题:
    中间矩阵SSSPPP都是N×NN\times NN×N

(2)为什么这会成为瓶颈

  • 当序列长度NNN很大时,SSSPPP的存储、读写都非常重。

  • 即使 GPU 算力很强,如果每一步都要频繁把这些大矩阵在 HBM(高带宽显存)和片上 SRAM / cache 之间来回搬运,速度也会被拖慢。

  • 所以 attention 的瓶颈并不只是“乘法很多”,而是:

  • 中间结果太大,内存访问代价太高。

(3)为什么普通实现会浪费

  • 在常规实现里,通常会:
  1. 先把S=QK⊤S=QK^\topS=QK整个算出来并写到显存;
  2. 再从显存里读出来做 softmax,得到PPP
  3. 再把PPP写回显存;
  4. 最后再读出PPPVVV去算输出。
  • 这个过程最大的问题在于:

  • 很多中间结果只是临时用一下,却被完整写回了大显存。

  • 这就是 FlashAttention 最想解决的地方。


3. FlashAttention 的核心思想

(1)一句话理解

  • FlashAttention 最核心的想法就是:

  • 不要把整个 attention matrix 显式落到显存里,而是在片上做小块计算,边算边归约,直接得到输出。

  • 也就是说,它不是先完整构造:
    S=QK⊤,P=softmax(S) S=QK^\top,\qquad P=\text{softmax}(S)S=QK,P=softmax(S)
    再去算输出,而是把这几个步骤融合在一起,以分块方式完成。

(2)为什么叫 IO-aware

  • FlashAttention 论文里一个非常关键的词就是IO-awareness

  • 传统算法通常更关注算术复杂度,比如时间复杂度是不是O(N2)O(N^2)O(N2)

  • FlashAttention 则进一步强调:

  • 在现代 GPU 上,显存和片上缓存之间的数据搬运,本身就是主要开销。

  • 所以 FlashAttention 并不是去近似 attention,而是去优化 attention 的数据流。

(3)最重要的一点:它是 exact attention

  • 这一点特别重要。

  • FlashAttention 不是 Performer、Linformer 那种近似 attention 方法。

  • 它并没有改 attention 的数学定义,而只是改了实现方式。

  • 所以它的核心价值不是“牺牲精度换速度”,而是:

  • 在保持精确等价的前提下,把显存访问和运行时间都降下来。


4. FlashAttention 是怎么做的

(1)分块(tiling)

  • 假设把Q,K,VQ,K,VQ,K,V都按块切开。

  • 比如把QQQ按行分成若干 query block,把K,VK,VK,V按列分成若干 key/value block。

  • 那么 attention 就不再是一次性拿整块大矩阵去算,而是按 block 逐块处理。

  • 直观上理解,就是:

  • 整张大表一次性装不下,那就分小块搬进片上缓存里,一块一块算。

(2)问题:softmax 不是一个纯局部操作

  • 如果只是矩阵乘法,分块很好理解。

  • 但 attention 里还有 softmax,而 softmax 涉及整行归一化:
    softmax(si)j=esij∑kesik \text{softmax}(s_i)_j=\frac{e^{s_{ij}}}{\sum_k e^{s_{ik}}}softmax(si)j=kesikesij

  • 这里的分母要看整行所有元素,所以它不像普通矩阵乘法那样天然可以局部分开。

  • 这也是 FlashAttention 真正巧的地方:
    它用在线(online)方式维护 softmax 所需的统计量。

(3)在线 softmax

  • 对于某一行 attention score,常规 softmax 会先拿到整行所有值,再求最大值、再求指数和。

  • FlashAttention 则是在分块扫描时,逐步维护:

  1. 当前见过的最大值;
  2. 当前对应的归一化系数;
  3. 当前输出累积值。
  • 这样一来,即使一整行分散在多个 block 中,也不需要把整行完整写回显存后再统一 softmax。

  • 所以 FlashAttention 的关键点不只是“分块”,而是:

  • 分块 + 在线 softmax + 输出累积。


5. 在线 softmax 的直观理解

(1)为什么 softmax 通常要先减最大值

  • 普通 softmax 为了数值稳定,常写成:
    softmax(si)j=esij−mi∑kesik−mi \text{softmax}(s_i)_j=\frac{e^{s_{ij}-m_i}}{\sum_k e^{s_{ik}-m_i}}softmax(si)j=kesikmiesijmi
    其中
    mi=max⁡jsij m_i=\max_j s_{ij}mi=jmaxsij

  • 这样做是为了防止指数爆炸。

  • 所以如果想分块做,就必须一边扫描 block,一边正确维护这个最大值和归一化项。

(2)FlashAttention 的核心统计量

  • 对于某一行,FlashAttention 会在扫描 block 的过程中不断更新:

    • 行最大值mmm
    • 行归一化和lll
    • 输出向量累计值。
  • 每看到一个新的 key block,就先算当前 block 对这一行的局部分数,再和旧的m,lm,lm,l合并更新。

  • 这样到最后,虽然没有显式存整个SSSPPP,但最终得到的输出和标准 attention 完全一致。

(3)本质上在做什么

  • 如果用一句话概括这个过程,我觉得最合适的是:

  • 把原来“先得到完整注意力矩阵,再做 softmax”这件事,改写成“扫描中间块时同步完成归一化和加权求和”。

  • 所以它不是把 attention 简化了,而是把计算顺序重新排了一遍。


6. FlashAttention 到底省了什么

(1)不是把二次复杂度变成线性

  • 这是最容易误解的一点。

  • FlashAttention 并没有把 attention 的理论计算复杂度从O(N2)O(N^2)O(N2)变成O(N)O(N)O(N)

  • 从数学上讲,所有 query-key 对之间的交互还是要算。

  • 它真正优化的是:

  • 中间矩阵的显式物化和反复显存访问。

(2)它主要省的是 HBM 访问

  • 在 GPU 上,HBM 访问代价远高于片上 SRAM / register 的访问。

  • FlashAttention 通过块级计算,让很多中间值留在片上,而不是反复写回大显存。

  • 所以它真正省的,不是“注意力不用算了”,而是:

  • 很多原本没必要写出的大矩阵,不再落到 HBM。

(3)为什么这会带来显著提速

  • 因为在长序列下,attention 往往不是单纯 compute-bound,而是明显受内存访问限制。

  • 一旦把显存读写压下去,GPU 的有效利用率就能明显提高。

  • 所以 FlashAttention 的加速来源,本质上是:

  • 更合理的数据流,而不是近似。


7. FlashAttention 和普通 fused kernel 有什么区别

(1)不是简单 kernel fusion

  • 看到这里,很多人会觉得:
    “这不就是把几个 kernel 融合一下吗?”

  • 其实比普通 fusion 更进一步。

  • 普通 kernel fusion 更多是把几个连续操作合并,以减少 launch overhead 和部分中间读写;

  • FlashAttention 则是从 attention 的算法结构本身出发,重新设计了 block-wise 的 IO 流程。

  • 所以它不只是“把 kernel 拼起来”,而是:

  • 把 attention 的执行顺序按 GPU 内存层级重新组织了一遍。

(2)为什么这很重要

  • 因为 attention 里的问题不只是 kernel 多,而是中间矩阵太大。

  • 单纯 fusion 并不能自动解决“整块 attention matrix 落显存”的问题。

  • FlashAttention 真正解决的是这个更深层的问题。


8. FlashAttention 的训练和推理意义

(1)训练里最直接的收益

  • 在训练阶段,FlashAttention 可以明显降低 attention 中间激活带来的显存占用。
  • 这意味着:
  1. 能撑更长序列;
  2. 能用更大 batch;
  3. 或者在同样显存下训练更大的模型。
  • 所以 FlashAttention 一开始最重要的影响,其实是在训练侧。

(2)推理里也很重要

  • 在推理阶段,尤其是 prefill 阶段,长上下文 attention 本身依然很重。

  • FlashAttention 仍然能带来明显收益。

  • 不过在 decode 阶段,很多时候瓶颈又会更多落在 KV Cache 访问上。

  • 所以在推理里,FlashAttention 和 KV Cache 往往是互补关系,而不是替代关系。

(3)和 KV Cache 怎么配合理解

  • 我觉得最简单的理解方式是:

  • FlashAttention 主要优化“单次 attention 怎么算”;

  • KV Cache 主要优化“历史 K/V 怎么复用”。

  • 二者关注点不同,但在大模型推理里通常会一起出现。


9. FlashAttention-2 在改什么

(1)为什么还要有 FlashAttention-2

  • 初代 FlashAttention 已经把 IO 问题解决得很好了,但它并不意味着 GPU 利用率就已经到头。

  • FlashAttention-2 继续往前推进,关注的重点更偏工程实现和并行划分:

  • 怎样让 GPU 上的线程块、warp 和 matmul 工作分配得更合理。

(2)核心改进方向

  • FlashAttention-2 的主线大致可以概括成三件事:
  1. 减少不必要的非 matmul FLOPs;
  2. 提高不同线程块之间的并行度;
  3. 优化单个线程块内部不同 warp 的工作划分,减少共享内存通信。
  • 所以它相对初代的重点已经有点变化:

  • 初代更像“把 attention 变成 IO-aware”;

  • 二代更像“在这个框架下,把 GPU 并行度和吞吐继续榨出来”。

(3)为什么它还是很重要

  • 因为很多时候,真正影响实际速度的不是“有没有一个好算法”,而是“这个算法在硬件上吃得够不够满”。

  • FlashAttention-2 做的就是把这件事继续推进。


10. FlashAttention-3 又在改什么

(1)为什么还会有 3

  • 到 FlashAttention-3 时,重点已经更偏向新硬件特性,尤其是 Hopper 架构 GPU。

  • 它关注的是:

  • 怎样利用新一代 GPU 的异步执行能力、Tensor Memory Accelerator(TMA)、以及低精度能力,把 attention 再往前推。

(2)主线变化

  • 如果很粗地说:

    • FlashAttention 1:解决 IO-aware exact attention;
    • FlashAttention 2:优化并行划分和吞吐;
    • FlashAttention 3:进一步拥抱新硬件特性与低精度执行。
  • 所以这一条线越往后,越明显能看出来:

  • 它已经不只是一个“attention 小技巧”,而是在逐渐变成 attention 的高性能实现主线。


11. FlashAttention 为什么这么重要

(1)它改变了大家看 attention 的方式

  • 早期一提 attention,很多人第一反应就是二次复杂度,然后想到各种近似。

  • FlashAttention 提醒大家:
    在真实硬件上,算法快不快,不只取决于 FLOPs,也取决于 IO。

  • 这一点其实很重要,因为它把问题从“数学复杂度”推进到了“系统复杂度”。

(2)它证明了 exact 也可以很快

  • 以前很多人会下意识觉得:
    想让 attention 快,就必须近似。

  • FlashAttention 的价值就在于,它说明了一件事:

  • 即使不近似,只要实现方式对,exact attention 也能快很多。

  • 这个结论对后面的很多工作都很有启发。

(3)它已经变成基础设施级组件

  • 现在很多训练框架、推理引擎、长上下文模型,一提高性能 attention,基本都会碰到 FlashAttention。

  • 这说明它已经不只是某篇论文里的方法,而是逐渐变成了 Transformer 生态中的基础设施级组件。


12. FlashAttention 和近似注意力怎么理解

(1)它们解决的是不同问题

  • 近似注意力方法的核心通常是:

  • 从数学形式上减少 attention 的计算或存储开销。

  • 比如稀疏化、低秩化、核方法等,都在试图减少理论复杂度。

  • FlashAttention 则不同,它基本保留标准 attention 数学形式,重点放在:

  • 如何把标准 attention 在真实 GPU 上算得更高效。

(2)所以它不是“替代一切近似方法”

  • FlashAttention 很强,但它并不意味着以后所有长上下文问题都只靠 FlashAttention 就够了。

  • 当序列真的极长时,二次计算本身仍然是问题。

  • 所以在更极端的长度下,稀疏化、压缩、窗口化等方法仍然有意义。

  • 更准确地说:

  • FlashAttention解决的是“标准 attention 该怎么高效实现”;

  • 近似方法解决的是“当标准 attention 本身太贵时,要不要改数学形式”。


13. 从更高一层看,FlashAttention 到底是什么

(1)它不是新的注意力定义

  • 我觉得理解 FlashAttention 最关键的一点是:

  • 它不是新的 attention 公式,而是新的 attention 执行方式。

  • 这个区别很重要。

  • 因为它并没有改模型定义,而是改了模型在硬件上的落地方式。

(2)它本质上是“attention 的 IO 重排”

  • 如果只用一句话概括 FlashAttention,我觉得最准确的是:

  • FlashAttention 是一种基于分块和在线 softmax 的 IO-aware exact attention 实现。

  • 这个定义比单纯说“更快 attention”更准确。

(3)它代表了一种非常重要的研究思路

  • 我觉得 FlashAttention 这条线真正厉害的地方,不只是把 attention 做快了,而是它代表了一种思路:

  • 算法设计不能只盯数学公式,还得盯硬件的数据流。

  • 这个思路在大模型时代会越来越重要。


14. 一点理解

(1)FlashAttention 最漂亮的地方

  • 我觉得 FlashAttention 最漂亮的地方在于,它不是靠“偷懒少算一点”来变快,而是靠“别把没必要的东西搬来搬去”来变快。

  • 这个想法其实非常朴素,但也非常有效。

(2)它为什么看起来像工程优化,却又很有方法味

  • 因为它既有工程味,也有算法味。

  • 如果只看表面,你会觉得它像 kernel 优化;

  • 但再往里看,它其实重新设计了 attention 的执行顺序和 softmax 归约方式。

  • 所以它不是纯工程 patch,而是一个很完整的方法。

(3)怎么记 FlashAttention

  • 如果只是为了学习,我觉得可以把 FlashAttention 记成四句话:
  1. 标准 attention 慢,不只是因为二次复杂度,还因为大矩阵显存读写很重;
  2. FlashAttention 不显式存整个 attention matrix,而是按块计算并在线完成 softmax;
  3. 它保持和标准 attention 数学等价,因此是 exact attention;
  4. 所以它本质上是在做 attention 的 IO 优化,而不是近似替代。

15. 参考鸣谢

  • FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness
    https://arxiv.org/abs/2205.14135

  • FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning
    https://arxiv.org/abs/2307.08691

  • FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision
    https://arxiv.org/abs/2407.08608

16. 注

  • 这篇主要是个人学习整理,重点放在主线理解;
  • 文中没有展开很多实现细节,比如 backward pass 的重计算策略、causal mask 的具体 block 处理方式、以及不同 GPU 架构下 kernel 设计的差异;
  • 才疏学浅,欢迎批评、指导和交流;
  • 有错误望大家及时指正!
http://www.cnnetsun.cn/news/1370450.html

相关文章:

  • Dify企业级RAG安全加固方案(含NIST SP 800-53映射表+GB/T 35273-2020合规对照清单)
  • jEasyUI 转换 HTML 表格为数据网格
  • 从CNN到RCNN:目标检测技术的演进与核心差异
  • 49:反追踪反击机制:多层代理流量混淆与同步阻断
  • 别再只用title属性了!高级悬浮提示框的5种实现方案对比
  • 软考科目这么多,IT 从业者应该怎么选择才最划算?
  • 杰理之SPI主机配置参数详解与实战应用【篇】
  • HumanML3D与DeepPhase实战:如何用Unity处理运动数据生成训练特征
  • ESP32-S3 USB烧录实战:从命令行到图形化界面的全流程解析
  • 复旦微FM33LG048芯片开发指南(1)SWD调试与LED控制实战
  • HSTracker:macOS炉石传说玩家的终极智能对战助手
  • 为什么必须做数模隔离?新手必懂核心逻辑
  • Kali Linux中LOIC与Hping3的DoS攻击原理与防御策略解析
  • OpenFOAM实战:snappyHexMesh网格划分避坑指南(附参数优化技巧)
  • 魔兽地图跨版本转换利器:w3x2lni全解析
  • 需求-扩展用例
  • 为QuickTime Player自定义快进/快退快捷键:提升观影效率的实用技巧
  • PFC GBM岩石矿物多组分模型构建与力学模拟分析
  • 开发提效利器:在快马平台一键生成配置完善的vit高效开发环境
  • ESP32开发必备:一键合并多个bin文件为完整固件的Shell脚本(附详细配置步骤)
  • 光猫桥接 vs. 路由模式怎么选?2024年最新家庭网络设置避坑指南
  • SDN进阶实战:用OpenFlow和P4手把手搭建你的第一个IBN实验环境
  • 为Spring_couplet_generation 构建自动化测试:Python单元测试与集成测试
  • AIGC内容创作新玩法:DeOldify为黑白线稿注入色彩
  • 2023最新版GEM5入门实战:从Docker编译到ARM全系统模拟(避坑指南)
  • Qwen2.5-32B-Instruct大模型部署:生产环境最佳实践
  • 基于知识库回答的智能客服系统:架构设计与工程实践
  • Dify混合检索优化实战手册(召回率提升31.5%的私有化调参矩阵首次公开)
  • Word2Vec实战:从预训练模型到自训练模型的工程化应用与避坑指南
  • python微信小程序的ai体育馆场地预约提醒系统