FlashAttention是一系列面向硬件特性设计的算法及其软件实现,用于计算Transformer架构中的注意力机制。其核心创新在于围绕图形处理器(GPU)的存储层次重新组织注意力计算,而非仅仅减少算术运算量。最初的算法由Tri Dao、Daniel Y. Fu、Stefano Ermon、Atri Rudra和Christopher Ré于2022年5月提出,能够计算精确的稠密注意力,同时避免将完整的注意力矩阵存储在GPU片外内存中。后续版本进一步改进了任务划分,并利用了新一代硬件的能力。(arxiv.org)
注意力与内存瓶颈
对于单个注意力头,缩放点积注意力以查询、键和值矩阵 、 和 为输入,计算:
其中, 是查询和键的维度,Softmax函数对分数矩阵的每一行进行运算。掩码可以限制允许的查询与键之间的交互。在长度为 的序列上计算自注意力时,分数矩阵包含 个元素。传统实现通常会在各个独立操作之间,将这个矩阵以及随后归一化得到的注意力权重写入设备内存。(arxiv.org)
FlashAttention针对的是由此产生的数据搬移问题。GPU高带宽内存的容量远大于片上存储,但访问成本也更高。通过将注意力计算划分为能容纳于片上存储的块,并融合多个操作,FlashAttention避免了反复传输大型中间矩阵。因此,其性能优势来自内存传输量的减少,而不是省略查询与键之间的比较。(arxiv.org)
分块计算与在线归一化
分块计算注意力的一个明显障碍是,softmax需要一个覆盖整行分数的归一化因子。如果在每个块内独立计算softmax,结果就会改变。FlashAttention改用在线归一化,为每个查询行维护一个动态更新的最大值、一个指数和,以及一个累积的加权输出。(tridao.me)
遇到新的分数块 时,这些量可以按如下方式更新:
最终输出为 。通过重新缩放此前的累积量,就能将新块纳入计算,而无需保留先前的分数。减去动态更新的最大值可以确保指数函数计算的数值稳定性。这种表达形式既保留了全局归一化,又允许逐块计算。(tridao.me)
“精确”一词用于将这种稠密算法与通过稀疏性或低秩表示来近似注意力的方法区分开来。它并不意味着与其他实现的结果逐位一致:改变浮点运算的顺序可能产生微小的数值差异。官方实现会将输出和梯度与参考实现进行比较,检验其差异是否在数值容差范围内。(arxiv.org)
训练与计算复杂度
在反向传播过程中,计算注意力的导数通常需要中间概率值。FlashAttention保留紧凑的归一化信息,并按需重新计算注意力块,而不是存储完整的概率矩阵。这是在注意力操作内部应用激活检查点的一种方式。尽管重新计算增加了算术运算量,但减少内存访问也可能使反向传播更快。(arxiv.org)
当注意力头的维度固定时,注意力计算所需辅助存储的空间复杂度相对于序列长度由二次降为线性。不过,稠密注意力算术运算的时间复杂度仍然是二次的:所有允许的查询与键配对仍需计算。因此,FlashAttention将存储效率与算术运算量的增长规律区分开来,这是计算复杂性中的一个重要区别。(arxiv.org)
各代版本的发展
FlashAttention-2于2023年7月推出,改进了并行计算和任务划分。它减少了矩阵乘法以外的操作,将单个注意力头的计算分配到多个线程块,并重新组织GPU线程束之间的任务,以减少通过共享内存进行的通信。论文报告称,在其基准测试中,该版本的速度约为原始实现的两倍,在A100 GPU上达到了理论峰值算术吞吐量的50%至73%。(arxiv.org)
FlashAttention-3于2024年7月推出,面向NVIDIA Hopper GPU。它让数据传输与计算重叠执行,将矩阵乘法与softmax交错安排,并利用硬件对低精度FP8计算的支持。这些技术不仅解决最初的内存传输问题,还进一步提升了硬件利用率。其FP8计算路径还引入了分块量化以及旨在减少量化误差的处理;这类低精度执行应与较高精度下的精确稠密注意力计算加以区分。(arxiv.org)
FlashAttention-4在2026年3月的一篇论文中得到介绍,针对的是硬件性能增长日益不均衡的问题:矩阵乘法吞吐量的增长速度已经超过了GPU某些其他资源的增长速度。它将算法改动与内核流水线化相结合,并使用嵌入Python的CuTe-DSL实现。论文报告称,在其基准测试配置下,B200 GPU上的BF16注意力计算吞吐量最高可达1,613 TFLOPs/s。这些是针对特定工作负载的内核测量结果,并非适用于所有应用的加速幅度。(arxiv.org)
软件集成与局限性
官方开源软件仓库提供了支持多头注意力、因果掩码和滑动窗口掩码、可变长度序列,以及使用缓存的键和值进行推理的接口。具体支持情况因实现路径、硬件代际和数值格式而异。这些接口以张量为操作对象,可将注意力计算集成到训练或神经网络推理工作流程中。(github.com)
PyTorch在2.2版本中,将FlashAttention-2集成到其缩放点积注意力操作中。当满足相应约束条件时,框架的分派机制可以选择优化后的注意力内核,否则则使用其他实现。实际收益取决于序列长度、注意力头维度、掩码、精度和硬件;注意力内核的加速并不会直接转化为整个大语言模型同等幅度的加速。(pytorch.org)
参考来源
- FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awarenessarxiv.org
- FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awarenessarxiv.org
- FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioningarxiv.org
- FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioningtridao.me
- FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precisionarxiv.org
- FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scalingarxiv.org
- flash-attention/README.md at main · Dao-AILab/flash-attention · GitHubgithub.com
- PyTorch 2.2: FlashAttention-v2 integration, AOTInductorpytorch.org
- Accelerated PyTorch 2 Transformerspytorch.org