跳转至

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 维度之间进行数据重排,使得长序列训练成为可能。