DISTFLASHATTN: Distributed Memory-efficient Attention for Long-context LLMs Training
作者: Dacheng Li, Rulin Shao, Anze Xie, Eric P. Xing, Xuezhe Ma, Ion Stoica, Joseph E. Gonzalez, Hao Zhang
分类: cs.LG, cs.AI, cs.DC
发布日期: 2023-10-05 (更新: 2024-03-31)
🔗 代码/项目: GITHUB
💡 一句话要点
提出DISTFLASHATTN以解决长上下文LLMs训练中的内存效率问题
🎯 匹配领域: 支柱九:具身大模型 (Embodied Foundation Models)
关键词: 长上下文 大型语言模型 分布式计算 内存效率 注意力机制 模型训练 性能优化
📋 核心要点
- 现有方法在训练长上下文LLMs时存在内存使用效率低下的问题,尤其是在单GPU环境下。
- 论文提出的DISTFLASHATTN通过分布式内存高效注意力机制,优化了长上下文LLMs的训练过程。
- 实验结果显示,DISTFLASHATTN在序列长度和训练速度上均显著优于现有方法,提升幅度可达8倍和5.64倍。
📝 摘要(中文)
FlashAttention(Dao, 2023)有效地将基于变换器的大型语言模型(LLMs)训练中的二次峰值内存使用减少到线性。本文介绍了DISTFLASHATTN,一种针对长上下文LLMs训练的分布式内存高效注意力机制。我们提出了三项关键技术:基于token的工作负载平衡、重叠的键值通信和考虑重材料化的梯度检查点算法。我们在Llama-7B及其变体上评估了DISTFLASHATTN,序列长度从32K到512K不等。DISTFLASHATTN实现了8倍更长的序列,相较于Ring Self-Attention提升了4.45-5.64倍的速度,相较于Megatron-LM与FlashAttention提升了2-8倍的序列长度和1.24-2.01倍的速度。与近期的Ring Attention和DeepSpeed-Ulysses相比,分别实现了1.67倍和1.26-1.88倍的速度提升。代码可在https://github.com/RulinShao/LightSeq获取。
🔬 方法详解
问题定义:本文旨在解决在训练长上下文大型语言模型时,现有注意力机制在内存使用上的低效率问题,尤其是在单GPU环境下的二次峰值内存使用问题。
核心思路:DISTFLASHATTN通过引入分布式内存高效注意力机制,结合token级的工作负载平衡、重叠的键值通信和梯度检查点算法,旨在提升长上下文训练的效率和可扩展性。
技术框架:整体架构包括三个主要模块:1) token级工作负载平衡,确保计算资源的高效利用;2) 重叠的键值通信,减少通信延迟;3) 考虑重材料化的梯度检查点算法,降低内存占用。
关键创新:最重要的技术创新在于将分布式内存管理与高效注意力机制结合,显著降低了内存使用并提升了训练速度,与传统的Ring Self-Attention和Megatron-LM方法相比,具有本质的效率提升。
关键设计:在设计中,采用了动态的工作负载分配策略,优化了通信协议,并引入了重材料化策略以减少内存占用,确保了在长序列训练中的高效性。具体参数设置和损失函数设计在实验部分进行了详细描述。
📊 实验亮点
实验结果表明,DISTFLASHATTN在Llama-7B模型上实现了8倍更长的序列,较Ring Self-Attention提升了4.45-5.64倍的速度,相较于Megatron-LM与FlashAttention提升了2-8倍的序列长度和1.24-2.01倍的速度,显示出显著的性能优势。
🎯 应用场景
该研究的潜在应用领域包括自然语言处理、机器翻译和对话系统等,能够有效支持长上下文的理解和生成任务。随着大型语言模型的广泛应用,DISTFLASHATTN的技术创新将推动更高效的模型训练,降低计算资源需求,具有重要的实际价值和未来影响。
📄 摘要(原文)
FlashAttention (Dao, 2023) effectively reduces the quadratic peak memory usage to linear in training transformer-based large language models (LLMs) on a single GPU. In this paper, we introduce DISTFLASHATTN, a distributed memory-efficient attention mechanism optimized for long-context LLMs training. We propose three key techniques: token-level workload balancing, overlapping key-value communication, and a rematerialization-aware gradient checkpointing algorithm. We evaluate DISTFLASHATTN on Llama-7B and variants with sequence lengths from 32K to 512K. DISTFLASHATTN achieves 8x longer sequences, 4.45 - 5.64x speedup compared to Ring Self-Attention, 2 - 8x longer sequences, 1.24 - 2.01x speedup compared to Megatron-LM with FlashAttention. It achieves 1.67x and 1.26 - 1.88x speedup compared to recent Ring Attention and DeepSpeed-Ulysses. Code is available at https://github.com/RulinShao/LightSeq.