中文

FlashMask:FlashAttention 的高效且丰富的掩码扩展

机器学习 2025-03-04 v2

摘要

vanilla 注意力机制的计算和内存需求随序列长度 N 的平方增长,这在处理 Transformer 模型中的长序列时提出了重大挑战。FlashAttention 通过消除 O(N^2) 的内存依赖并通过 IO 优化降低注意力延迟来缓解这些挑战。然而,其对某些注意力掩码类型的本地支持有限,无法内在地容纳更复杂的掩码需求。以往方法采用密集掩码,导致 O(N^2) 的内存复杂度,效率低下。本文提出 FlashMask,这是 FlashAttention 的一种扩展,引入注意力掩码的列向稀疏表示。该方法高效地表示各种掩码类型,促进了优化内核实现的开发。通过采用这种新颖的表示方法,FlashMask 实现了线性内存复杂度 O(N),适用于建模长上下文序列。此外,这种表示方法使我们能够通过利用注意力掩码中的稀疏性来消除不必要的计算,而不牺牲计算精度,从而实现更高的计算效率。我们在 SFT、LoRA、DPO 和 RM 等大型语言模型的微调和对齐训练中评估了 FlashMask 的性能。FlashMask 在端到端速度上相较于现有 FlashAttention 稠密方法实现了 1.65 倍至 3.22 倍的显著提升。此外,我们的内核级比较表明,FlashMask 在内核 TFLOPs/s 方面超过最新的对应方法 FlexAttention 12.1% 至 60.7%,在 A100 GPU 上实现了 37.8% 至 62.3% 的理论最大 FLOPs/s。代码已开源于 PaddlePaddle,并集成到 PaddleNLP 中,支持超过 1000 亿参数的模型,上下文长度可达 128K token。

关键词

引用

@article{arxiv.2410.01359,
  title  = {FlashMask: Efficient and Rich Mask Extension of FlashAttention},
  author = {Guoxia Wang and Jinle Zeng and Xiyuan Xiao and Siming Wu and Jiabin Yang and Lujing Zheng and Zeyu Chen and Jiang Bian and Dianhai Yu and Haifeng Wang},
  journal= {arXiv preprint arXiv:2410.01359},
  year   = {2025}
}

备注

Published as a conference paper at ICLR 2025