跳转至

fsdp_ulysses.py — FSDP + Ulysses 分片管理器

文件概述

处理 FSDP 和 Ulysses 序列并行之间数据重分片的管理器。

核心类

FSDPUlyssesShardingManager

class FSDPUlyssesShardingManager(BaseShardingManager):
    """FSDP + Ulysses 序列并行的数据分片管理

    问题:
    - FSDP 按 DP 维度分片数据
    - Ulysses 需要同一 SP 组内的所有 rank 拥有相同数据
    - 进入模型计算前需要在 SP 组内 AllGather 数据

    解决:
    - preprocess: AllGather 数据到 SP 组
    - postprocess: 按 SP rank 切分结果
    """

    def __enter__(self):
        """切换到模型特定的 SP 组"""
        self.prev_sp_group = get_ulysses_sequence_parallel_group()
        set_ulysses_sequence_parallel_group(self.device_mesh["sp"].get_group())

    def __exit__(self, ...):
        """恢复之前的 SP 组"""
        set_ulysses_sequence_parallel_group(self.prev_sp_group)

    def preprocess_data(self, data):
        """AllGather: 在 SP 组内收集完整数据"""
        group = self.device_mesh["sp"].get_group()
        all_gather_data_proto(data=data, process_group=group)
        return data

    def postprocess_data(self, data):
        """Split: 按 SP rank 切分数据"""
        sp_size = self.device_mesh["sp"].size()
        sp_rank = self.device_mesh["sp"].get_local_rank()
        data = data.chunk(chunks=sp_size)[sp_rank]
        return data

数据流示意

FSDP 分片后(每个 DP rank 持有不同数据):
  DP rank 0: [batch_0]
  DP rank 1: [batch_1]

Ulysses SP(SP 组内需要相同数据):
  ┌─ SP group 0 ─┐
  │ SP rank 0: [batch_0]  ← 持有完整数据
  │ SP rank 1: [batch_0]  ← AllGather 复制
  └──────────────┘

  ┌─ SP group 1 ─┐
  │ SP rank 0: [batch_1]
  │ SP rank 1: [batch_1]
  └──────────────┘

与其他模块的关系

  • 继承 base.py 的 BaseShardingManager
  • 被 fsdp_workers.py 中的 Actor/Critic 使用
  • 使用 verl/protocol.py 的 all_gather_data_proto

小结

解决了 FSDP 数据并行和 Ulysses 序列并行之间数据布局不匹配的问题,通过 AllGather 和 Split 实现透明的数据重分片。