ulysses.py — Ulysses 序列并行¶
文件路径: verl/utils/ulysses.py
文件概述¶
ulysses.py 实现了 DeepSpeed Ulysses 序列并行的工具函数。序列并行将一个长序列切分到多个 GPU 上并行计算注意力,然后通过 AllToAll 通信合并结果。
背景知识¶
序列并行 (Sequence Parallelism) 解决长序列注意力计算的内存瓶颈。Self-Attention 的内存复杂度是 O(seq_len^2),当序列很长时单个 GPU 放不下。
Ulysses 方法的核心思路: 1. 将序列维度切分到多个 GPU(每个 GPU 只持有 seq/N 的 token) 2. 注意力计算前,用 AllToAll 将 "序列分片 + 完整 head" 变为 "完整序列 + head 分片" 3. 各 GPU 独立计算自己负责的 head 的注意力 4. 计算后再用 AllToAll 变回来
核心函数详解¶
1. 进程组管理¶
_ULYSSES_SEQUENCE_PARALLEL_GROUP = None
def set_ulysses_sequence_parallel_group(group: dist.ProcessGroup):
global _ULYSSES_SEQUENCE_PARALLEL_GROUP
_ULYSSES_SEQUENCE_PARALLEL_GROUP = group
def get_ulysses_sequence_parallel_world_size(group=None) -> int:
group = get_ulysses_sequence_parallel_group() if group is None else group
return dist.get_world_size(group) if group else 1
2. AllToAll 通信¶
def gather_seq_scatter_heads(x, seq_dim, head_dim, unpadded_dim_size=0, group=None):
"""
[bsz, seq/n, h, ...] -> [bsz, seq, h/n, ...]
收集序列维度,分散 head 维度
"""
x = SeqAllToAll.apply(group, x, head_dim, seq_dim)
# 如果有 padding,去掉 padding
if unpadded_dim_size and unpadded_dim_size % sp_world != 0:
x = _unpad_tensor(x, seq_dim, padding_size)
return x
def gather_heads_scatter_seq(x, head_dim, seq_dim, group=None):
"""
[bsz, seq, h/n, ...] -> [bsz, seq/n, h, ...]
收集 head 维度,分散序列维度(反向操作)
"""
这两个函数是一对互逆操作,分别在注意力计算前后调用。
3. Padding 处理¶
def ulysses_pad(input_ids_rmpad, position_ids_rmpad=None, sp_size=1, pad_value=0):
"""确保序列长度是 sp_size 的倍数,不足的部分补 padding"""
pad_size = (sp_size - total_seq_len % sp_size) % sp_size
if pad_size > 0:
input_ids_rmpad = F.pad(input_ids_rmpad, (0, pad_size), value=pad_value)
return input_ids_rmpad, position_ids_rmpad, pad_size
序列必须能被 SP 并行度整除才能均匀切分,所以需要先 padding。
4. SeqAllToAll 自定义 autograd 函数¶
class SeqAllToAll(torch.autograd.Function):
@staticmethod
def forward(ctx, group, local_input, scatter_dim, gather_dim, async_op=False):
return all_to_all_tensor(local_input, scatter_dim, gather_dim, group, async_op)
@staticmethod
def backward(ctx, *grad_output):
# 反向传播时,交换 scatter 和 gather 维度
return (None, all_to_all_tensor(input_t, ctx.gather_dim, ctx.scatter_dim, ctx.group), ...)
通过 torch.autograd.Function,使 AllToAll 通信支持反向传播,梯度能正确传播。
数据流示意¶
GPU 0: [seq_0, seq_1] x [h0, h1, h2, h3]
GPU 1: [seq_2, seq_3] x [h0, h1, h2, h3]
gather_seq_scatter_heads (AllToAll)
↓
GPU 0: [seq_0, seq_1, seq_2, seq_3] x [h0, h1]
GPU 1: [seq_0, seq_1, seq_2, seq_3] x [h2, h3]
各 GPU 独立计算注意力
gather_heads_scatter_seq (AllToAll)
↓
GPU 0: [seq_0, seq_1] x [h0, h1, h2, h3]
GPU 1: [seq_2, seq_3] x [h0, h1, h2, h3]
核心函数/类列表¶
| 函数/类 | 说明 |
|---|---|
set/get_ulysses_sequence_parallel_group() |
管理 SP 进程组 |
gather_seq_scatter_heads() |
序列聚合 + head 分散 |
gather_heads_scatter_seq() |
head 聚合 + 序列分散 |
ulysses_pad() |
序列 padding |
ulysses_pad_and_slice_inputs() |
pad + 切分输入 |
slice_input_tensor() |
切分张量 |
SeqAllToAll |
可求导的 AllToAll |
Gather |
可求导的 AllGather |
gather_outputs_and_unpad() |
聚合输出并去 padding |
validate_ulysses_config() |
配置校验 |
与其他模块的关系¶
- 被 Actor/Critic 模型的注意力层在序列并行模式下调用
- 依赖
torch.distributed进行通信 validate_ulysses_config()确保 num_heads 能被 SP 并行度整除
小结¶
ulysses.py 实现了 DeepSpeed Ulysses 序列并行的所有通信原语,通过 AllToAll 操作在序列维度和注意力 head 维度之间进行数据重排,使得长序列训练成为可能。