跳转至

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() 确保所有进程完成恢复后才继续训练。