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 实现透明的数据重分片。