fsdp2_utils.py — 提供 FSDP2(Fully Sharded Data Parallelism v2) 模型...¶
文件路径: verl/experimental/fully_async_policy/fsdp2_utils.py
文件概述¶
提供 FSDP2(Fully Sharded Data Parallelism v2) 模型的分片参数保存和恢复功能。与 megatron_utils.py 类似,但专门处理 PyTorch 原生的 FSDP2 分布式训练场景。
FSDP2 使用 DTensor(分布式张量)将模型参数分片存储在多个 GPU 上,每个进程只持有参数的一部分。
关键代码讲解¶
1. 分片保存到 CPU¶
def fsdp2_sharded_save_to_cpu(model):
"""
每个进程只保存自己负责的 DTensor 分片到 CPU。
"""
cpu_sharded_state = {}
global_spec = None
for param_name, param in model.named_parameters():
if not isinstance(param, DTensor):
# 非分片参数(如 BatchNorm 的 running_mean)
cpu_tensor = param.detach().cpu()
cpu_sharded_state[param_name] = (cpu_tensor, None)
continue
# 记录全局分片规则
if global_spec is None:
global_spec = param._spec
# 提取本地分片 -> 移到 CPU
local_cpu_tensor = param._local_tensor.detach().cpu()
cpu_sharded_state[param_name] = (local_cpu_tensor, param._spec)
return cpu_sharded_state, global_spec
2. 从 CPU 恢复分片¶
def fsdp2_sharded_load_from_cpu(model, cpu_sharded_state, target_spec):
"""
每个进程将自己的 CPU 分片恢复到对应 GPU。
"""
# 验证 device_mesh 一致性
assert current_device_mesh == target_spec.device_mesh
for param_name, param in model.named_parameters():
local_cpu_tensor, saved_spec = cpu_sharded_state[param_name]
if isinstance(param, DTensor):
# 验证分片策略一致
assert saved_spec.placements == target_spec.placements
# 恢复到 GPU
target_device = param._local_tensor.device
param._local_tensor.copy_(local_cpu_tensor.to(target_device))
else:
param.data.copy_(local_cpu_tensor.to(param.device))
# 进程同步
dist.barrier()
与 Megatron Utils 的对比¶
| 特性 | megatron_utils | fsdp2_utils |
|---|---|---|
| 框架 | Megatron-Core | PyTorch FSDP2 |
| 参数类型 | DDP buffers | DTensor |
| 保存方式 | 全量复制 | 分片保存(每个进程只保存自己的分片) |
| 同步 | 不需要 barrier | 需要 dist.barrier() |
| PyTorch 版本 | 无要求 | >= 2.6 |
核心类/函数列表¶
| 名称 | 类型 | 说明 |
|---|---|---|
fsdp2_sharded_save_to_cpu() |
函数 | 分片保存到 CPU |
fsdp2_sharded_load_from_cpu() |
函数 | 从 CPU 恢复分片到 GPU |
与其他模块的关系¶
- 与
megatron_utils.py功能对等,针对不同的分布式训练框架 - 被 Worker 的
save_model_to_cpu / restore_model_from_cpu方法调用
小结¶
fsdp2_utils.py 为 FSDP2 训练模式提供了高效的参数快照能力。每个进程只保存/恢复自己负责的参数分片,避免了不必要的 AllGather 通信,节省内存和带宽。dist.barrier() 确保所有进程完成恢复后才继续训练。