ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

【NeurIPS 2022】FlashAttention:IO 感知的快速内存高效精确注意力|从高效Transformer内核视角

【NeurIPS 2022】FlashAttention:IO 感知的快速内存高效精确注意力|从高效Transformer内核视角 摘要本文解读NeurIPS 2022杰出论文《FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness》。该论文提出FlashAttention——一个 IO 感知的精确注意力算法通过融合Tiling 分块计算、在线 softmax 增量聚合与反向重计算避免在 GPU 显存 HBM 上物化 $N\times N$ 的注意力矩阵把 HBM 访问从 $\Theta(NdN^2)$ 降到 $\Theta(N^2d^2/M)$典型配置下减少9 倍。实验表明GPT-2 端到端训练加速最高 3.5 倍、BERT-large 比 MLPerf 1.1 纪录快 15%且首次让 Transformer 在 16K/64K 超长序列上超越随机水平Path-X 61.4%为长上下文大模型训练提供了最重要的基础设施级借鉴。视频讲解点击观看 B 站视频摘要论文基本信息背景与动机研究主线从问题到结论基准/方法设计分类全景方法细节实验设计与结果结果对比总结关键发现局限性常见问题FAQFlashAttention 是近似注意力吗为什么减少 FLOPs 的近似方法反而不快FlashAttention 为什么增加 FLOPs 反而更快FlashAttention 如何解决 softmax 的数值稳定性FlashAttention 对模型质量有影响吗FlashAttention 现在的生态地位如何参考链接论文基本信息项目内容标题英文FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness标题中文FlashAttentionIO 感知的快速内存高效精确注意力作者Tri Dao, Daniel Y. Fu, Stefano Ermon, Atri Rudra, Christopher Ré机构Stanford University · University at Buffalo会议NeurIPS 2022Outstanding Paper 杰出论文奖arXivhttps://arxiv.org/abs/2205.14135项目网站https://github.com/HazyResearch/flash-attention背景与动机自注意力是 Transformer 的核心模块但它的时间与内存复杂度都随序列长度 $N$ 平方增长标准实现需要把 $\mathbf{S}\mathbf{Q}\mathbf{K}^{\top}$ 和 $\mathbf{P}\mathrm{softmax}(\mathbf{S})$ 两个 $N\times N$ 中间矩阵完整写回 GPU 显存HBM。序列越长这个矩阵越庞大这成为长上下文建模的根本瓶颈。此前的主流路线是近似注意力稀疏近似Reformer、Smyrf、Longformer、BigBird和低秩近似Linformer、Performer、Linear Attention把计算量降到近线性。但这些方法普遍只优化 FLOPs忽略了内存访问开销——现代 GPU 上计算速度远超内存速度A100 HBM 带宽约 1.5–2.0 TB/s片上 SRAM 带宽约 19 TB/s快一个数量级大部分 Transformer 算子其实是内存受限的。因此许多近似方法理论线性复杂度、实际墙钟时间却毫无优势这也是硬件彩票现象的根源。论文的核心论证是FLOPs 减少不等于墙钟加速注意力算法必须成为 IO 感知的——把 GPU 内存层次结构HBM vs SRAM放进算法设计的一等公民。IO 感知的思想在数据库连接、Halide 图像处理、数值线性代数中早有成熟应用但 PyTorch/TensorFlow 的高层接口无法表达细粒度的内存控制这正是 FlashAttention 用 CUDA 内核实现的原因。研究主线从问题到结论图 5研究主线Mermaid 流程图问题 → 动机 → 洞察 → 设计 → 方法 → 实验 → 结论基准/方法设计FlashAttention 的目标是不读取、不写入 $N\times N$ 注意力矩阵用两个成熟技术实现Tiling 分块计算把 $\mathbf{Q},\mathbf{K},\mathbf{V}$ 切成适配 SRAM 容量 $M$ 的块$B_c\lceil M/4d\rceil$$B_r\min(\lceil M/4d\rceil,d)$外层循环遍历 $\mathbf{K},\mathbf{V}$ 块内层循环遍历 $\mathbf{Q}$ 块在片上依次完成 QK 转置、softmax、PV 乘法。在线 softmax 增量聚合softmax 的行归一化需要看到整行论文维护行最大值 $m$ 与指数和 $\ell$ 两个统计量用 $m^{new}\max(m,\tilde m)$、$\ell^{new}e^{m-m^{new}}\elle^{\tilde m-m^{new}}\tilde\ell$ 逐块合并保证与全局 softmax严格一致且数值稳定。图 1FlashAttention 用 Tiling 避免在慢速 HBM 上物化 N×N 注意力矩阵左右图为对 PyTorch 注意力实现的 7.6 倍加速分类全景图 6高效注意力方法分类全景Mermaid 流程图方法细节反向重计算是第二个关键设计反向传播通常需要 $\mathbf{S},\mathbf{P}$ 两个中间矩阵求梯度。FlashAttention 只在前向保存输出 $\mathbf{O}$ 与归一化统计量 $(m,\ell)$反向时在片上用 $\mathbf{Q},\mathbf{K},\mathbf{V}$ 的块重算注意力矩阵——这是一种选择性梯度检查点但因为省去了海量 HBM 访问重算反而比存储更快。内核融合Tiling 使所有步骤矩阵乘、softmax、掩码、dropout、矩阵乘能在单个 CUDA 内核内完成输入只从 HBM 加载一次、输出只写回一次。图 2左标准注意力 40.3GB HBM 读写 vs FlashAttention 4.4GB运行 41.7ms→7.3ms中块大小与运行时间右稀疏扩展加速算法的正确性由定理保证返回结果与 $\mathrm{softmax}(\mathbf{Q}\mathbf{K}^{\top})\mathbf{V}$ 逐元素一致FLOPs 为 $O(N^2d)$额外内存仅 $O(N)$附录 A 完整伪代码。IO 复杂度理论附录 B标准注意力需要 $\Theta(NdN^2)$ 次 HBM 访问FlashAttention 只需 $\Theta(N^2d^2/M)$。论文更进一步证明了下界不存在精确注意力算法能对所有 SRAM 尺寸 $M\in[d,Nd]$ 渐近优于该复杂度——也就是说 FlashAttention 在 IO 意义下已经最优无可再省。Block-Sparse 扩展附录 D用 butterfly 模式固定稀疏掩码保持 SIMD 友好的稠密小块比 FlashAttention 再快2–4 倍序列长度可达 64KLRA 平均准确率几乎不掉点59.6 vs 59.8证明了 IO 感知框架能让近似方法真正兑现墙钟加速。实验设计与结果评测协议8×A100 上对比端到端训练墙钟时间与验证困惑度单卡 A100 40GB 上评测注意力前向反向运行时间与峰值内存。基线包括 PyTorch 标准注意力、HuggingFace、Megatron-LM、Linformer、Performer、Reformer 等。GPT-2 训练时间主表GPT-2 实现困惑度训练时间加速比small – HuggingFace18.29.5 天1.0×small – Megatron-LM18.24.7 天2.0×small –FlashAttention18.22.7 天3.5×medium – HuggingFace14.221.0 天1.0×medium – Megatron-LM14.311.5 天1.8×medium –FlashAttention14.36.9 天3.0×BERT-large 用 17.4 分钟达到 72.0% 目标准确率比 MLPerf 1.1 的 Nvidia 纪录20.0 分钟快 15%。加速的同时困惑度与基线完全一致——因为算法是精确的数值稳定性等价附录 E 训练曲线重合。LRA 基准附录 E 转写模型平均准确率加速比Transformer59.31.0×FlashAttention59.82.4×Block-Sparse FlashAttention59.62.8×Linformer54.92.5×Linear Attention59.62.3×Performer58.91.8×Reformer57.61.3×图 3注意力运行时间左与内存占用右FlashAttention 短序列最快、内存线性增长块稀疏版全面领先近似基线长序列新能力附录 F把 GPT-2 上下文从 1K 扩到 4K困惑度 18.2→17.5训练仍比 Megatron 的 1K 版本快 30%Path-X16K 序列上 FlashAttention 成为第一个超越随机水平的 Transformer61.4%Block-Sparse 版在 Path-25664K达 63.1%长文档分类在 MIMIC-III 提升 4.3 分、ECtHR 提升 8.5 分。图 4GPT-2 训练过程中 FlashAttention 与 HF/Megatron 基线的验证困惑度曲线几乎完全重合结果对比总结图 7结果对比总结Mermaid 流程图40.3GB/41.7ms → 4.4GB/7.3ms关键发现HBM 访问减少最多 9 倍GPT-2 medium 上 40.3GB → 4.4GB运行时间 41.7ms → 7.3ms。GPT-2 端到端训练加速 3.0–3.5 倍medium 从 21.0 天缩至 6.9 天small 从 9.5 天缩至 2.7 天困惑度不变。BERT-large 比 MLPerf 1.1 纪录快 15%17.4 分钟 vs 20.0 分钟8×A100。首个在 Path-X 超越随机水平的 Transformer16K 序列准确率 61.4%Block-Sparse 版在 Path-25664K达 63.1%。长上下文直接提升质量GPT-2 4K 上下文困惑度 18.2→17.5长文档分类 MIMIC 4.3 分、ECtHR 8.5 分。内存随序列长度线性增长比精确注意力基线最多省 20 倍显存短序列≤512快于所有已知注意力方法。局限性工程成本高每种新的注意力变体新掩码、dropout、稀疏模式都要手写 CUDA 内核开发成本高。可移植性差内核针对特定 GPU 架构优化跨架构迁移需要重写。单 GPU 最优多卡注意力还需额外的 GPU 间数据传输层分析。高层语言缺失PyTorch/TensorFlow 无法表达细粒度内存控制作者期望出现类似 Halide 的高层语言写注意力、自动编译为 IO 感知 CUDA的编译器。常见问题FAQFlashAttention 是近似注意力吗不是。它计算的是精确softmax 注意力输出与标准实现逐元素一致只是通过 Tiling 改变了计算顺序和内存访问模式没有任何精度损失。为什么减少 FLOPs 的近似方法反而不快因为现代 GPU 上注意力是内存受限算子运行时间由 HBM 读写决定而非计算量。近似方法降低了 FLOPs 但内存访问模式没有本质改善甚至引入了额外开销所以墙钟时间没有优势。FlashAttention 为什么增加 FLOPs 反而更快反向传播采用重计算FLOPs 增加约 13%但免去了读取 $O(N^2)$ 中间矩阵的 HBM 访问。HBM 访问才是瓶颈省下的时间远超多算的 FLOPs。FlashAttention 如何解决 softmax 的数值稳定性维护每块的行最大值 $m$ 与指数和 $\ell$增量合并时用 $e^{m-m^{new}}$ 重新缩放与标准 softmax 的 max-subtraction 技巧完全等价保证数值稳定。FlashAttention 对模型质量有影响吗没有负面影响反而因支持更长序列带来质量提升GPT-2 4K 上下文困惑度 18.2→17.5Path-X/Path-256 首次被 Transformer 解决。序列长度本身成为免费的模型改进维度。FlashAttention 现在的生态地位如何它已成为事实上的行业基础设施PyTorch 原生 SDPA、HuggingFace、vLLMPagedAttention、xFormers 均采用其内核后续 FlashAttention-2/3 与 Mamba 的硬件感知扫描延续了同一 IO 感知思想。参考链接FlashAttention 论文https://arxiv.org/abs/2205.14135官方开源代码https://github.com/HazyResearch/flash-attentionAttention Is All You NeedVaswani et al., 2017https://arxiv.org/abs/1706.03762ReformerKitaev et al., ICLR 2020https://arxiv.org/abs/2001.04451LinformerWang et al., 2020https://arxiv.org/abs/2006.04768The Input/Output Complexity of Sorting and Related ProblemsAggarwal Vitter, 1988https://dl.acm.org/doi/10.1145/48529.48535FlashAttention-2Dao, 2023https://arxiv.org/abs/2307.08691给大家推荐一款自用写文献综述、无虚构文献的 AI复旦大学 FudanNLP 团队自研 切问学术官网qiewenpaper.com覆盖3.6 亿篇可溯源真实中英文文献能自动整合文献观点生成规范综述还能挖掘研究创新点、复现实验配合视频教学新手快速上手文献综述写作后记博客的关键词集中在编程、算法、机器人、人工智能、数学等等持续高质量输出中。讨论QQ群白拾的小屋 (750365700)⭐B站账号白拾的物理AI组会活跃于知识区和动画区✨GitHub主页YhbCode000工程文件
返回列表