跳转至

fsdp_model_merger.py — FSDP Checkpoint 合并器

源码路径:verl/model_merger/fsdp_model_merger.py

文件概述

这个文件实现了 FSDPModelMerger 类,用于将 FSDP(Fully Sharded Data Parallel)分布式训练产生的分片 checkpoint 合并为完整的 HuggingFace 格式模型。

什么是 FSDP?

FSDP 是 PyTorch 提供的分布式训练策略。在 FSDP 中,模型参数被 分片(shard) 到多个 GPU 上。每个 GPU 只保存一部分参数,训练时通过通信临时收集完整参数,计算完再释放。这样可以训练超过单卡显存的大模型。

保存 checkpoint 时,每个 rank(GPU 进程)保存自己持有的那一部分参数。因此要恢复完整模型,就需要把所有 rank 的分片按正确的方式拼接回来。

关键代码讲解

1. 获取 world_size

def _get_world_size(self) -> int:
    config_path = Path(self.config.local_dir) / "fsdp_config.json"
    with open(config_path) as f:
        config = json.load(f)
    world_size = config.get("world_size", None)
    return world_size

从 checkpoint 目录下的 fsdp_config.json 读取训练时使用的 GPU 数量(world_size)。这个信息决定了需要加载多少个分片文件。

2. 提取 Device Mesh 信息

def _extract_device_mesh_info(self, state_dict, world_size):
    pivot_key = sorted(list(state_dict.keys()))[0]
    weight = state_dict[pivot_key]

    if isinstance(weight, DTensor):
        device_mesh = weight.device_mesh
        mesh = device_mesh.mesh
        mesh_dim_names = device_mesh.mesh_dim_names
    else:
        mesh = np.array([world_size], dtype=np.int64)
        mesh_dim_names = ("fsdp",)

    return mesh, mesh_dim_names

DTensor 是 PyTorch 的分布式张量类型,它记录了张量的分片方式。这个方法从一个样本参数中提取 device mesh(设备网格)信息:

  • mesh:描述 GPU 的拓扑结构,比如 [8] 表示 8 个 GPU 组成一维网格
  • mesh_dim_names:每个维度的名称,比如 ("fsdp",) 表示纯 FSDP,("ddp", "fsdp") 表示 DDP + FSDP 混合

3. 计算分片配置

def _calculate_shard_configuration(self, mesh, mesh_dim_names):
    assert mesh_dim_names in (("fsdp",), ("ddp", "fsdp"))

    if "tp" in mesh_dim_names:
        total_shards = mesh.shape[-1] * mesh.shape[-2]
        mesh_shape = (mesh.shape[-2], mesh.shape[-1])
    else:
        total_shards = mesh.shape[-1]
        mesh_shape = (mesh.shape[-1],)

    return total_shards, mesh_shape

根据 mesh 信息计算需要合并的分片总数。目前支持: - 纯 FSDP(一维分片) - DDP + FSDP(数据并行 + 完全分片)

TP(Tensor Parallelism)支持在代码中预留了但尚未完整实现。

4. 加载并合并所有分片(核心方法)

def _load_and_merge_state_dicts(self, world_size, total_shards, mesh_shape, mesh_dim_names):
    model_state_dict_lst = [None] * total_shards

    def process_one_shard(rank, model_state_dict_lst):
        model_path = Path(self.config.local_dir) / f"model_world_size_{world_size}_rank_{rank}.pt"
        state_dict = torch.load(model_path, map_location="cpu", weights_only=False)
        model_state_dict_lst[rank] = state_dict
        return state_dict

    # 并行加载所有分片
    with ThreadPoolExecutor(max_workers=min(32, os.cpu_count())) as executor:
        futures = [executor.submit(process_one_shard, rank, model_state_dict_lst)
                   for rank in range(total_shards)]
        for future in tqdm(futures, desc=f"Loading {total_shards} FSDP shards"):
            future.result()

这里使用 线程池 并行加载多个分片文件,大幅加速 I/O。每个分片文件的命名格式为 model_world_size_{N}_rank_{i}.pt。

加载完成后,按 key 合并各分片的张量:

    for key in set(model_state_dict_lst[0].keys()):
        state_dict[key] = []
        for model_state_shard in model_state_dict_lst:
            tensor = model_state_shard.pop(key)
            if isinstance(tensor, DTensor):
                state_dict[key].append(tensor._local_tensor.bfloat16())
                # 记录 placement 信息用于后续合并
                placements = tuple(tensor.placements)
                if mesh_dim_names[0] in ("dp", "ddp"):
                    placements = placements[1:]  # 去掉 DP 维度的 replicate
                param_placements[key] = placements
            else:
                state_dict[key].append(tensor.bfloat16())

5. 按 Placement 策略合并张量

def _merge_by_placement(self, tensors, placement):
    if placement.is_replicate():
        return tensors[0]           # 复制的,取任意一个即可
    elif placement.is_partial():
        raise NotImplementedError   # 部分聚合,暂不支持
    elif placement.is_shard():
        return torch.cat(tensors, dim=placement.dim).contiguous()  # 沿分片维度拼接

DTensor 的 Placement 类型决定了如何合并: - Replicate(复制):每个 rank 持有相同的完整副本,取任意一个 - Shard(dim)(分片):参数沿 dim 维度被切分,需要用 torch.cat 拼接

最终的合并逻辑:

    for key in sorted(state_dict):
        if key in param_placements:
            placements = param_placements[key]
            if len(mesh_shape) == 1:
                # 1-D: 纯 FSDP
                shards = state_dict[key]
                state_dict[key] = self._merge_by_placement(shards, placements[0])
            else:
                raise NotImplementedError("FSDP + TP is not supported yet")
        else:
            state_dict[key] = torch.cat(state_dict[key], dim=0)

6. 合并主流程 merge_and_save()

def merge_and_save(self):
    world_size = self._get_world_size()
    rank_zero_state_dict = self._load_rank_zero_state_dict(world_size)
    mesh, mesh_dim_names = self._extract_device_mesh_info(rank_zero_state_dict, world_size)
    total_shards, mesh_shape = self._calculate_shard_configuration(mesh, mesh_dim_names)
    merged_state_dict = self._load_and_merge_state_dicts(
        world_size, total_shards, mesh_shape, mesh_dim_names)

    if self.config.operation == "test":
        self._validate_state_dict(merged_state_dict)
    elif self.config.operation == "merge":
        self.save_hf_model_and_tokenizer(merged_state_dict)
        if self.config.hf_upload:
            self.upload_to_huggingface()

7. 验证合并结果 _validate_state_dict()

def _validate_state_dict(self, state_dict):
    auto_model_class = self.get_transformers_auto_model_class()
    hf_model = auto_model_class.from_pretrained(self.config.test_hf_dir, torch_dtype=torch.bfloat16)
    hf_state_dict = hf_model.state_dict()

    for key in hf_model_keys:
        assert hf_state_dict[key].shape == state_dict[key].shape
        assert hf_state_dict[key].dtype == state_dict[key].dtype
        torch.testing.assert_close(hf_state_dict[key], state_dict[key], atol=1e-6, rtol=1e-6)

验证合并后的 state_dict 与参考 HuggingFace 模型在 key、shape、dtype、数值上完全一致。

核心流程图

_get_world_size()                    读取 fsdp_config.json
        │
        ▼
_load_rank_zero_state_dict()         加载 rank 0 的 checkpoint
        │
        ▼
_extract_device_mesh_info()          从 DTensor 提取分片信息
        │
        ▼
_calculate_shard_configuration()     计算分片数量和 mesh 形状
        │
        ▼
_load_and_merge_state_dicts()        并行加载所有分片 → 按 placement 合并
        │
        ├── operation == "test"  → _validate_state_dict()
        │
        └── operation == "merge" → save_hf_model_and_tokenizer()
                                        │
                                        └── upload_to_huggingface() (可选)

核心类/函数列表

名称 类型 说明
FSDPModelMerger 类 FSDP checkpoint 合并器
_get_world_size() 方法 从配置文件获取 GPU 数量
_load_rank_zero_state_dict() 方法 加载 rank 0 的 checkpoint
_extract_device_mesh_info() 方法 提取 DTensor 的设备网格信息
_calculate_shard_configuration() 方法 计算分片配置
_merge_by_placement() 方法 根据 placement 策略合并张量
_load_and_merge_state_dicts() 方法 加载并合并所有分片
merge_and_save() 方法 合并主流程
_validate_state_dict() 方法 验证合并结果的正确性

与其他模块的关系

  • 继承自 BaseModelMerger,使用其 save_hf_model_and_tokenizer() 等公共方法
  • 依赖 PyTorch 的 DTensor、Placement、Shard 等分布式张量类型
  • 使用 ThreadPoolExecutor 并行 I/O 加速

小结

FSDPModelMerger 的核心工作是将 FSDP 训练时分片到各 GPU 的参数重新拼接为完整的模型。关键挑战在于理解 DTensor 的 Placement 策略(Replicate vs Shard),然后根据策略决定是取一份副本还是沿特定维度拼接。代码通过并行 I/O 提升了加载效率,并支持验证合并结果的正确性。