Ulysses 序列并行用于百万标记上下文训练
Hugging Face 已将 Ulysses 序列并行(Snowflake AI Research 的 Arctic 长序列训练协议的一部分)集成到 Accelerate、Transformers 和 TRL 库中。此集成使开发者能够通过在多 GPU 上分布注意力计算,克服注意力机制的二次方内存增长,从而在数十万甚至数百万标记的序列上训练大语言模型(LLM)。
长序列训练的挑战
标准 Transformer 注意力在 FLOPs 和内存上都随序列长度 $n$ 二次增长($O(n^2)$)。虽然 FlashAttention 将内存使用降低到 $O(n)$,但 $O(n^2)$ 的计算需求仍然存在。对于超过 32k 标记的序列,训练通常会超出单个 GPU 的内存容量,这就需要一种将序列本身拆分到多个设备上的方法,而不能仅依赖数据并行。
Ulysses 序列并行的工作原理
Ulysses 序列并行(SP)通过在 GPU 之间划分序列维度和注意力头来分配注意力计算。其流程如下:
- 序列分片:将输入序列在 $P$ 个 GPU 上切分,每个 GPU 持有本地的标记块。
- QKV 投影:每个 GPU 为其本地块计算 query、key、value 投影。
- All-to-All 通信:通过 all-to-all 集体操作重新分配数据,使每个 GPU 持有所有序列位置,但仅针对一部分注意力头。
- 本地注意力:GPU 使用 FlashAttention 或 SDPA 对分配的头部计算注意力。
- All-to-All 通信:第二次 all-to-all 操作将数据恢复为序列分片格式。
- 输出投影:每个 GPU 为其本地序列块计算输出投影。
通信复杂度
Ulysses 在每个注意力层需要两次 all-to-all 操作,每个 GPU 的通信量为 $O(n \cdot d / P)$(其中 $n$ 为序列长度,$d$ 为隐藏维度,$P$ 为并行度)。这比 Ring Attention 的每 GPU $O(n \cdot d)$ 通信量更高效,后者在 $P-1$ 跳之间串行传输。
生态系统集成
Accelerate
Accelerate 通过 ParallelismConfig 类和 DeepSpeed 集成实现 Ulysses。关键参数包括 sp_size(用于序列并行的 GPU 数量)和 sp_backend(必须设为 "deepspeed")。当调用 accelerator.prepare() 时,系统会自动使用 UlyssesSPAttentionHF 注册模型,并用 UlyssesSPDataLoaderAdapter 包装 dataloader。
Transformers Trainer
Transformers 的 Trainer 通过 TrainingArguments.parallelism_config 处理 Ulysses 集成。它自动完成 dataloader 包装、序列分片以及加权损失聚合,确保即使标记在各 rank 上分布不均,梯度仍然正确。
TRL SFTTrainer
TRL 的 SFTTrainer 为监督微调添加了优化,例如 packing 功能以减少填充浪费。它要求 pad_to_multiple_of 等于 sp_size,以保证序列可整除。启用 Ulysses 时,SFTTrainer 还会自动管理预移位标签。
Ulysses 与 Ring Attention 的对比
| 方面 | Ulysses (DeepSpeed) | Ring Attention (FSDP2) |
|---|---|---|
| 并行方式 | 注意力头划分 | 基于环的 KV 交换 |
| 后端 | DeepSpeed ZeRO | PyTorch FSDP2 |
| 注意力支持 | FlashAttention 2/3、SDPA | 仅 SDPA |
| 通信方式 | 每层两次 all-to-all |
P2P 环形通信 |
| 每 GPU 通信量 | $O(\text{total_seq} \times \text{hidden} / \text{sp_size})$ | $O(\text{total_seq} \times \text{hidden})$ |
| 头数约束 | num_heads >= sp_size |
无 |
性能基准
Hugging Face 使用 Qwen3-4B 在 Gutenberg 英文数据集上,配合 H100 80GB GPU,对 Ulysses SP 进行了基准测试。
内存降低
使用 sp=4 在相同序列长度下将每 GPU 内存降低了 3.3 倍。这使得从基准的 8K 标记(DP=4)可以扩展到 96K 标记(SP=4),仍然保持在 80GB 内存限制内。128K 标记时会出现 OOM(内存溢出)情况。
吞吐量
随着序列长度增长,吞吐量提升,因为二次方的注意力计算主导了通信开销。在 64K 标记时,sp=4 达到 13,396 标记/秒,是 8K 基准的 3.7 倍。
实施最佳实践
- 序列可整除:使用
pad_to_multiple_of确保序列长度能被sp_size整除。 - 注意力后端:在 Ampere GPU 上使用 FlashAttention 2,在 Hopper GPU 上使用 FlashAttention 3。
- 内存优化:将 Ulysses 与 DeepSpeed ZeRO Stage 3 结合,并设置环境变量
PYTORCH_ALLOC_CONF=expandable_segments:True以降低碎片化。 - 2D 并行:根据 GPU 数量平衡
sp_size与dp_shard_size,以优化最大序列长度或更高吞吐量。 - 额外内核:使用 Liger‑Kernel 的
FusedLinearCrossEntropy与TiledMLP,进一步在损失计算和大矩阵运算期间降低工作内存。