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 提升了加载效率,并支持验证合并结果的正确性。