Back to blog

Training Ultra Long Context Language Model with Fully Pipelined Distributed Transformer

提出全流水线分布式Transformer架构,支持超长上下文语言模型的高效训练。

Training Ultra Long Context Language Model with Fully Pipelined Distributed Transformer

1.1 基本信息

项目内容
论文标题Training Ultra Long Context Language Model with Fully Pipelined Distributed Transformer
arXiv ID2408.16978
发表时间2024年8月30日 (v1),2025年5月13日 (v2)
会议MLSys 2025
作者Jinghan Yao, Sam Ade Jacobs, Masahiro Tanaka, Olatunji Ruwase, Hari Subramoni, Dhabaleswar K. Panda
代码DeepSpeed PR #6462

1.2 摘要

具有长上下文能力的大语言模型(LLMs)对于自然语言处理和计算生物学中的复杂任务至关重要,如文本生成和蛋白质序列分析。然而,直接在极长上下文上训练LLMs需要大量的GPU资源和更多的内存,导致更高的成本和更大的复杂性。通过下游微调或适配引入长上下文能力的替代方法会带来显著的设计限制。

本文提出了全流水线分布式Transformer (Fully Pipelined Distributed Transformer, FPDT),用于以极高的硬件效率高效训练长上下文LLMs。对于GPT和Llama模型,与当前最先进的解决方案相比,在相同硬件上实现了16倍的序列长度提升。通过专用的序列块流水线设计,现在可以在仅4个GPU上训练具有200万序列长度的8B LLM,同时保持超过55%的MFU。FPDT与现有训练技术无关,并被证明在不同LLM模型上都能高效工作。

2.1 问题背景

随着GPT-4、Claude和Gemini等LLMs的快速发展,对扩展上下文窗口以处理更长输入序列的需求日益增长。长上下文能力对于以下应用至关重要:

  • 文档分析:处理整个法律文档或科学论文
  • 长篇内容生成:撰写书籍或详细报告
  • 对话式AI:维护连贯且上下文相关的长期对话
  • 复杂推理任务:医疗保健、气候和金融领域

2.2 现有方法的局限性

方法问题
Megatron-SP训练7B模型256K上下文需要32+ A100 80G GPU
DeepSpeed Ulysses训练1.2B GPT模型1M上下文需要64 A100 GPU
RoPE外推性能会随着序列长度增加而崩溃
FlashAttention虽然将内存复杂度从O(N²)降到O(N),但常数因子仍然很大

2.3 核心目标

FPDT旨在解决以下关键挑战:

  1. 内存瓶颈:随着序列长度增加,激活内存和中间缓冲区成比例增长
  2. 硬件效率:在有限GPU资源下实现超长上下文训练
  3. 通用性:与现有训练技术正交组合

3.1 整体架构

FPDT基于DeepSpeed Ulysses序列并行,结合以下关键技术:

graph TB
    A[输入序列] --> B[序列分块]
    B --> C[QKV投影]
    C --> D[Alltoall通信]
    D --> E[分布式注意力计算]
    E --> F[FFN层]
    F --> G[输出]

    H[GPU HBM] <--> I[CPU主机内存]
    E <--> I

3.2 关键组件

3.2.1 序列分块流水线

FPDT将输入序列分割为多个块(chunk),每个块的大小为原始序列的1/u:

符号含义
[b, s_local, h_global, d]非注意力操作的张量形状
[b, s_global, h_local, d]分布式注意力的张量形状
u分块数量
T_i第i个序列块

3.2.2 流水线调度流程

对于每个块 T_i:
  1. QKV投影: c_i -> q_i, k_i, v_i
  2. Alltoall通信: 分散头、收集序列
  3. 注意力计算: 使用因果掩码
  4. KV缓存: 将 k_i^, v_i^ 卸载到主机内存
  5. 前一块获取: 从主机内存获取之前的KV块
  6. 在线注意力更新: 逐步更新注意力输出

3.2.3 双缓冲设计

双缓冲流水线

双缓冲设计利用多个CUDA流,在反向传播期间将大部分卸载操作与注意力梯度计算重叠:

  • 计算流:执行注意力梯度计算
  • 通信流:执行数据预取和卸载
  • 重叠:几乎所有的预取操作与计算重叠

3.3 序列重排

序列重排

由于分块Alltoall操作,需要对序列进行重排以保持因果掩码的有效性:

  • 对角线因果掩码在每个块Alltoall操作后保持有效
  • NVLINK负载均衡通过数据布局实现
  • 标签相应重排,损失计算仍然匹配

4.1 内存分析与优化

4.1.1 Transformer块的内存占用

操作前向传播反向传播
HiddenNd2Nd
QKV投影3Nd6Nd
Alltoall4Nd8Nd
注意力4Nd8Nd
FFN4Nd-
其他操作3Nd-

关键发现:

  • QKV投影将内存占用增加3倍
  • Alltoall通信需要接收缓冲区,异步通信时需要6Nd内存
  • FlashAttention反向传播需要8Nd内存

4.1.2 内存优化策略

  1. 序列分块:将内存占用降低为原来的1/u
  2. 主机内存卸载:将KV缓存卸载到CPU
  3. 双缓冲重叠:隐藏通信延迟

4.2 与现有技术的正交组合

技术组合方式
DeepSpeed ZeRO-3分区所有参数、梯度和优化器状态
PyTorch FSDP完全分片数据并行
FlashAttention作为注意力计算的后端
激活检查点可选组合使用

4.3 与相关方法的对比

方法通信模式内存效率硬件要求
Megatron-SPAll-gather/Reduce-scatter中等高(32+ GPU)
Ring Attention多步通信中等中等
DeepSpeed UlyssesAll-to-all高中等(64 GPU)
MsT-中等(仅MLP)低
MEMOTensor并行中等中等
FPDTAll-to-all + 卸载极高极低(4 GPU)

5.1 最大上下文长度支持

最大上下文长度

Table 1: FPDT支持的最大上下文长度

模型规模A100 40G (1 GPU)A100 40G (2 GPU)A100 40G (4 GPU)A100 40G (8 GPU)A100 80G (4 GPU)A100 80G (8 GPU)A100 80G (16 GPU)A100 80G (32 GPU)
2.7B128K512K2M4M4M8M+8M+8M+
8B---1M2M4M8M+8M+
13B---256K512K3M4M8M+
30B-----1M3M4M
70B------1M4M

注:+表示只测试到该长度,-表示模型本身无法放入GPU内存。

5.2 性能对比

性能对比

5.2.1 关键性能指标

模型GPU数量最大序列长度MFU
GPT 2.7B2 GPU512K>55%
GPT 6.7B4 GPU1M>55%
Llama 8B4 GPU2M>55%
GPT 13B8 GPU3M>55%
Llama 70B32 GPU4M>55%

5.2.2 与现有方案对比

对比维度现有方案FPDT提升倍数
7B模型256K上下文32 GPU (Megatron-SP)4 GPU8x
1.2B模型1M上下文64 GPU (Ulysses)4 GPU16x
8B模型2M上下文不可行4 GPU-

5.3 序列块大小权衡

6.7B模型4GPU性能 8B模型4GPU性能

块大小的选择需要平衡:

  • 较小块:更低内存占用,更多通信开销
  • 较大块:更高内存占用,更少通信开销

5.4 收敛性评估

13B模型8GPU性能

FPDT在训练过程中保持了与标准训练相当的收敛性,损失曲线正常下降。

6.1 内存高效Transformer

方法技术内存复杂度
FlashAttention在线softmaxO(N)
低秩近似矩阵分解O(N)
核方法近似注意力O(N)
稀疏注意力选择性计算O(N)
局部+全局上下文混合注意力O(N)

6.2 长上下文训练方法

方法核心思想优势局限性
Megatron-SP张量并行+序列并行成熟稳定硬件要求高
BPT块并行Transformer内存高效需要仔细调参
Ring Attention环形通信可扩展性强依赖设备数量
DeepSpeed UlyssesAll-to-all通信通信效率高部署复杂
MsTMLP分块简单未解决注意力内存
MEMO主机内存卸载内存效率高需要整数规划

6.3 FPDT的定位

FPDT综合了多种技术的优势:

  • 继承FlashAttention的块计算思想
  • 采用DeepSpeed Ulysses的All-to-all通信
  • 引入主机内存卸载和流水线设计
  • 实现极低硬件要求下的超长上下文训练

7.1 主要贡献

  1. 端到端内存分析:识别了Transformer训练中的内存峰值,特别是QKV投影和Alltoall通信的内存开销

  2. 全流水线设计:基于DeepSpeed Ulysses设计了专用的序列块流水线,实现近零开销的训练流程

  3. 双缓冲机制:利用多个CUDA流重叠计算和通信,隐藏卸载延迟

  4. 极低硬件要求:

    • 8B模型可在4个GPU上训练200万序列
    • 70B模型可在32个GPU上训练400万序列
    • 相比现有方案实现16倍提升
  5. 通用性:与DeepSpeed ZeRO、PyTorch FSDP正交组合,适用于GPT、Llama等模型

7.2 技术亮点

特性描述
内存效率序列分块降低内存占用为1/u
计算效率保持>55%的MFU
通信效率双缓冲重叠通信与计算
通用性支持多种模型架构
易用性集成到DeepSpeed框架

7.3 应用场景

  • 超长文档处理:法律文档、科学论文、书籍
  • 长对话系统:多轮对话、客服系统
  • 基因组学:蛋白质序列分析
  • 代码理解:大型代码库分析

7.4 未来方向

论文提到的未来工作方向:

  • 进一步优化通信模式
  • 支持更多模型架构
  • 探索与MoE模型的结合
  • 优化推理阶段的长上下文支持

8.1 论文链接

8.2 代码资源

8.3 关键图表

图表描述路径
Figure 1MFU和最大上下文长度对比figures/ultra-long-context/intro.png
Figure 2DeepSpeed Ulysses分布式注意力figures/ultra-long-context/ulysses.png
Figure 4带卸载的分布式注意力设计figures/ultra-long-context/distributed_attention_with_offloading.png
Figure 5带获取和卸载的分布式注意力figures/ultra-long-context/distributed_attention_with_fo.png
Figure 6序列块重排figures/ultra-long-context/sequence_shuffle.png
Figure 7双缓冲设计figures/ultra-long-context/db_pipeline.png
性能对比时间对比figures/ultra-long-context/compare_time.png
6.7B性能4GPU配置figures/ultra-long-context/6.7b_4gpu.png
8B性能4GPU配置figures/ultra-long-context/llama_8b_4gpu.png
13B性能8GPU配置figures/ultra-long-context/13b_8gpu.png

8.4 相关论文

论文主题关系
FlashAttention (Dao et al., 2022)内存高效注意力基础技术
DeepSpeed Ulysses (Jacobs et al., 2023)序列并行基础架构
Megatron-SP (Korthikanti et al., 2023)序列并行对比方法
Ring Attention (Liu et al., 2023)环形通信对比方法
BPT (Liu & Abbeel, 2024)块并行相关工作
MsT (Luo et al., 2024)MLP分块对比方法
MEMO (Zhao et al., 2024)内存卸载对比方法

8.5 引用格式

@article{yao2024training,
  title={Training Ultra Long Context Language Model with Fully Pipelined Distributed Transformer},
  author={Yao, Jinghan and Jacobs, Sam Ade and Tanaka, Masahiro and Ruwase, Olatunji and Subramoni, Hari and Panda, Dhabaleswar K.},
  journal={arXiv preprint arXiv:2408.16978},
  year={2024}
}