FlashAttention 深度解析

FlashAttention 深度解析 FlashAttention 深度解析FlashAttention 是大语言模型LLM发展史上的里程碑式技术。它通过算法与底层硬件GPU的深度协同从根本上解决了标准 Transformer 自注意力机制的O(N2)O(N^2)O(N2)复杂度瓶颈是实现几十万甚至上百万上下文窗口Long-context的核心基石。本文档将从物理硬件痛点、核心算法设计到严格的数学推演全面剖析 FlashAttention 的工作原理。1. 痛点分析标准注意力的“内存墙”在标准的 Transformer 中自注意力Self-Attention的数学表达式为Attention(Q,K,V)softmax(QKTd)VAttention(Q, K, V) \text{softmax}(\frac{QK^T}{\sqrt{d}})VAttention(Q,K,V)softmax(d​QKT​)V随着输入序列长度NNN的增长计算QKTQK^TQKT会产生一个N×NN \times NN×N的注意力分数矩阵Attention Score Matrix。标准框架如早期 PyTorch的执行逻辑是“计算-存储-读取”计算SQKTS QK^TSQKT将N×NN \times NN×N矩阵写入 GPU 的全局内存HBM。从 HBM 重新读取SSS计算 Softmax 得到概率矩阵PPP再写回 HBM。从 HBM 读取PPP和矩阵VVV计算最终输出OPVO PVOPV再次写回 HBM。致命瓶颈Memory-bound当NNN达到 8K 或 32K 时这个N×NN \times NN×N矩阵不仅会瞬间撑爆显存OOM更严重的是GPU 的计算核心Tensor Cores在疯狂等待数据从缓慢的 HBM 中读写。算力被闲置整个过程被极低的内存带宽死死卡住这就是所谓的“内存墙”。2. 硬件思维GPU 的分级存储与 IO-AwareFlashAttention 的突破在于提出了IO-aware感知 IO的设计理念即算法必须适配硬件的物理特性。现代 GPU 的存储层级极度不平衡存储层级物理位置容量示例 (A100)带宽速度访问特性HBM (全局显存)GPU 芯片外侧40GB - 80GB~1.5 TB/s巨大但极慢读写代价高昂SRAM (共享内存)计算核心 (SM) 内部约 20MB / 192KB per SM~19 TB/s极小但极快计算零延迟核心解法Kernel FusionFlashAttention 的终极目标是将数据从 HBM 加载到极速的 SRAM 中后在 SRAM 内一口气算完QKTQK^TQKT、Softmax 和 乘VVV的全套流程最后只将N×dN \times dN×d的最终结果写回 HBM从而彻底消灭N×NN \times NN×N中间矩阵的读写。3. 核心技术一Tiled Attention (分块注意力)由于 SRAM 的容量MB级别远远装不下完整的Q,K,VQ, K, VQ,K,V矩阵FlashAttention 采用了分块Tiling策略矩阵切块将 HBM 中的Q,K,VQ, K, VQ,K,V和最终输出OOO切分成大小适中的块Blocks。块的大小需精确计算以确保几块加起来刚好能塞满 SRAM。双重循环加载外循环加载KKK和VVV的小块到 SRAM。内循环加载QQQ的小块到 SRAM。就地计算在 SRAM 中直接让QQQ的块与KKK的块相乘计算局部的 Softmax然后立刻与VVV的块相乘将结果累加到输出块OOO中。这种分块计算虽然有效但立刻遇到了一个极其严峻的数学挑战Softmax 的计算依赖全局信息如何分块4. 核心数学基石Online Softmax (在线 Softmax)标准的 Safe Softmax 需要遍历整行数据才能找到全局最大值mmm和全局指数和lll分母这在分块且“阅后即焚”的 SRAM 中是不可能做到的。FlashAttention 引入了Online Softmax利用代数技巧在局部计算时动态修正历史结果实现 100% 精确的无损计算。动态修正推演假设我们在遍历QQQ的某一行与KKK的各个分块相乘。我们维护两个全局标量状态当前的最大值moldm_{old}mold​和当前的指数和loldl_{old}lold​以及当前的未归一化输出O~old\tilde{O}_{old}O~old​。当加载一个新的分块并算出局部最大值mlocalm_{local}mlocal​、局部指数和llocall_{local}llocal​与局部输出O~local\tilde{O}_{local}O~local​时1. 更新全局最大值mnewmax⁡(mold,mlocal)m_{new} \max(m_{old}, m_{local})mnew​max(mold​,mlocal​)2. 核心修正公式计算全局指数和由于最大值变了之前算过的所有指数项ex−molde^{x - m_{old}}ex−mold​都偏大或偏小了。利用指数的性质只需乘上一个缩放因子emold−mnewe^{m_{old} - m_{new}}emold​−mnew​即可完美修正历史lnewemold−mnew⋅loldemlocal−mnew⋅llocall_{new} e^{m_{old} - m_{new}} \cdot l_{old} e^{m_{local} - m_{new}} \cdot l_{local}lnew​emold​−mnew​⋅lold​emlocal​−mnew​⋅llocal​3. 修正输出矩阵O~\tilde{O}O~同理针对VVV的加权求和也用完全相同的缩放因子修正旧的累加器O~newemold−mnew⋅O~oldemlocal−mnew⋅O~local\tilde{O}_{new} e^{m_{old} - m_{new}} \cdot \tilde{O}_{old} e^{m_{local} - m_{new}} \cdot \tilde{O}_{local}O~new​emold​−mnew​⋅O~old​emlocal​−mnew​⋅O~local​当所有块遍历完毕后只需进行最后一步除法OfinalO~newlnewO_{final} \frac{\tilde{O}_{new}}{l_{new}}Ofinal​lnew​O~new​​结果即可写回 HBM。结论依靠这段优雅的代数公式FlashAttention 完美绕过了N×NN \times NN×N矩阵的存储将空间复杂度从O(N2)O(N^2)O(N2)降到了O(N)O(N)O(N)。5. 核心技术二反向传播的“重计算” (Recomputation)在深度学习的后向传播Backward Pass中计算注意力层的梯度通常需要用到前向传播时生成的N×NN \times NN×NSoftmax 概率矩阵。由于 FlashAttention 在前向时为了省显存根本没有保存这个矩阵它采用了极其反直觉的重计算Recomputation策略前向传播时FlashAttention 会在 HBM 中保存每个序列行的最终归一化统计量全局mmm和lll。这个统计量的大小是O(N)O(N)O(N)的极小。反向传播时利用保存下来的mmm和lll以及重新分块加载的Q,K,VQ, K, VQ,K,V在 SRAM 中直接把前向的 Attention 矩阵重新算一遍算完梯度就丢弃。为什么算两遍反而更快这就是 IO-aware 的威力。在 GPU 上重新执行几千次乘加运算FLOPs的时间远远少于去 HBM 里读取一个巨大矩阵MAC的时间。FlashAttention 用极其廉价的算力换取了极其昂贵的内存带宽。6. 总结与工程意义FlashAttention 并不改变 Transformer 的数学本质它的输出与标准 PyTorch Attention 的误差在极小浮点精度内实际上因为避免了大型矩阵累加FlashAttention 通常精度更高、更不容易溢出。它的三大核心贡献彻底打破显存瓶颈显存占用与上下文长度NNN呈线性关系单卡轻松支持 32K、64K 甚至更长文本。极致的训练/推理加速消除 IO 瓶颈让 GPU Tensor Core 跑满通常带来 2-4 倍的端到端速度提升。全行业标配如今无论 OpenAI 的 GPT-4、Anthropic 的 Claude还是开源的 Llama、Qwen底层全部依赖 FlashAttention 及其演进版本如 FlashAttention-2、FlashAttention-3。