_state_dict_utils.py — 分布式 state_dict 工具函数¶
源码路径:
verl/third_party/torch/distributed/_state_dict_utils.py
文件概述¶
这个文件从 PyTorch 2.7.0 的 torch/distributed/_state_dict_utils.py 复制而来,提供了一组底层工具函数,用于在分布式训练环境中处理模型的 state_dict。核心功能包括:
- 收集(gather):将分布在多个 GPU 上的分片张量收集为完整张量
- 广播(broadcast):将 rank 0 的完整 state_dict 广播到所有 rank
- 复制(copy):在两个结构相同的 state_dict 之间复制数据
- CPU 卸载(offload):将 GPU 上的张量移到 CPU 以节省显存
- 展平/还原(flatten/unflatten):将嵌套字典展平为一级字典或反向操作
这些函数是 checkpoint/state_dict.py 的底层支撑。
背景知识¶
什么是 state_dict?¶
state_dict 是 PyTorch 中保存模型参数的标准格式,本质上是一个 dict[str, Tensor]。键是参数名(如 model.layers.0.self_attn.q_proj.weight),值是对应的张量。
分布式 state_dict 的挑战¶
在分布式训练中,参数可能以多种形式分布在不同 GPU 上: - ShardedTensor:PyTorch 的旧分片张量类型 - DTensor:PyTorch 的新分布式张量类型,记录了分片的 placement 信息 - 普通 Tensor:已经是完整的本地张量
获取或设置 state_dict 时,需要根据张量类型做不同的处理,这就是本文件各函数的核心工作。
关键代码讲解¶
1. 核心遍历函数 _iterate_state_dict()¶
def _iterate_state_dict(
iter_object: Any,
sharded_tensor_func: Callable,
dtensor_func: Callable,
tensor_func: Callable,
*,
pg: Optional[dist.ProcessGroup] = None,
device: Optional[torch.device] = None,
cpu_offload: bool = False,
companion_obj: Any = None,
ranks_only: tuple[int, ...] = (),
type_check: bool = True,
non_blocking: bool = True,
) -> dict[str, Any]:
这是整个文件的核心枢纽函数。它递归遍历 state_dict 的每个元素,根据元素类型调用不同的处理函数:
- 遇到 ShardedTensor → 调用
sharded_tensor_func - 遇到 DTensor → 调用
dtensor_func - 遇到 普通 Tensor → 调用
tensor_func - 遇到 dict/list/tuple → 递归遍历
- 遇到 基本类型(int, float, str)→ 原样返回
if isinstance(iter_object, ShardedTensor):
ret = sharded_tensor_func(iter_object, pg, device, companion_obj)
elif isinstance(iter_object, DTensor):
ret = dtensor_func(iter_object, pg, device, companion_obj)
elif isinstance(iter_object, torch.Tensor):
ret = tensor_func(iter_object, pg, device, companion_obj)
elif isinstance(iter_object, dict):
ret = {
key: _iterate_state_dict(value, sharded_tensor_func, dtensor_func, tensor_func, ...)
for key, value in iter_object.items()
}
# ... 类似地处理 list, tuple, 基本类型
companion_obj 参数的巧妙用法:当提供 companion_obj 时,函数会将处理后的结果复制到 companion_obj 对应位置,实现就地更新。这在 _copy_state_dict 和 _broadcast_state_dict 中被使用。
ranks_only 参数:如果指定了 rank 列表,只有这些 rank 会获得实际的 state_dict,其他 rank 得到空字典。这在 CPU offload 场景中很有用(只让 rank 0 持有完整数据)。
2. 收集 ShardedTensor _all_gather_sharded_tensor()¶
def _all_gather_sharded_tensor(
sharded_tensor: "ShardedTensor",
pg: Optional[dist.ProcessGroup] = None,
device: Optional[torch.device] = None,
) -> torch.Tensor:
world_size = dist.get_world_size(pg)
shards = sharded_tensor.local_shards()
dim_0_size = sharded_tensor.size()[0]
tensor_numel = sharded_tensor.size().numel()
chunk_size = math.ceil(dim_0_size / world_size) * tensor_numel // dim_0_size
# 本地分片可能需要 padding
local_tensor = shards[0].tensor.flatten()
num_padding = chunk_size - local_tensor.numel()
if num_padding > 0:
local_tensor = F.pad(local_tensor, [0, num_padding])
# all_gather 收集所有分片
tensor = torch.empty(chunk_size * world_size, dtype=local_tensor.dtype, device=pg_device)
dist.all_gather_into_tensor(tensor, local_tensor, group=pg)
# 去掉 padding 并 reshape
tensor = tensor.narrow(0, 0, tensor_numel).reshape(sharded_tensor.size())
return tensor
这个函数将分片到各 GPU 上的 ShardedTensor 收集为完整张量。流程:
1. 计算每个分片的 chunk_size(需要对齐,可能需要 padding)
2. 执行 all_gather 通信操作收集所有分片
3. 去掉 padding 并 reshape 为原始形状
3. 收集 state_dict _gather_state_dict()¶
def _gather_state_dict(state_dict, *, pg=None, device=None,
cpu_offload=False, ranks_only=(), type_check=True):
def dtensor_func(value, pg, device, companion_obj):
# 将所有分片收集为 Replicate
placements = [Replicate() for _ in value.placements]
value = value.redistribute(device_mesh=value.device_mesh, placements=placements)
value = value.to_local()
if isinstance(value, AsyncCollectiveTensor):
value = value.wait()
return value
def sharded_tensor_func(value, pg, device, companion_obj):
output_tensor = _all_gather_sharded_tensor(value, pg, device)
return output_tensor
return _iterate_state_dict(
state_dict, sharded_tensor_func, dtensor_func, _identity_func,
pg=pg, device=device, cpu_offload=cpu_offload, ranks_only=ranks_only)
这个函数是获取完整 state_dict 的高层接口。它对不同类型的分布式张量做不同处理: - DTensor:将 placement 改为全部 Replicate,相当于收集完整数据 - ShardedTensor:使用 all_gather 收集 - 普通 Tensor:不做处理(identity 函数)
DTensor 的 redistribute 方法会自动处理通信,将分片数据重新分布为完整副本。
4. 广播 state_dict _broadcast_state_dict()¶
def _broadcast_state_dict(
full_state_dict, local_state_dict, device, pg=None, strict=False, cpu_offload=False
):
ret = {}
if dist.get_rank() == 0:
for key, value in full_state_dict.items():
if not torch.is_tensor(value):
ret[key] = value
elif value.dim() == 0:
ret[key] = value.cpu()
else:
ret[key] = _TensorInfo(value.size(), value.dtype)
broadcast_list = [ret]
dist.broadcast_object_list(broadcast_list, src=0, group=pg)
ret = broadcast_list[0]
keys = []
for key, value in ret.items():
if not isinstance(value, _TensorInfo):
if key in local_state_dict:
local_state_dict[key] = value
continue
keys.append(key)
if len(keys) >= 1:
_broadcast_tensors(ret, local_state_dict, keys, device, pg)
keys.clear()
广播的流程:
1. Rank 0 先将 state_dict 的元信息(tensor 的 shape 和 dtype)广播给所有 rank
2. 然后逐个张量进行广播,避免一次性广播所有张量导致 OOM
3. 如果 local_state_dict 中有 DTensor,广播后会自动切分为对应的本地分片
为什么逐个广播? 因为大模型的参数总量可能达到数百 GB,一次性广播会导致显存不足。逐个广播可以及时释放中间结果。
5. 分发张量 _distribute_tensors()¶
def _distribute_tensors(local_state_dict, keys, device, pg=None):
for key in keys:
_local_state = local_state_dict.get(key, None)
if _local_state is None or torch.is_tensor(_local_state):
continue
local_state = _local_state[0] # DTensor
full_tensor = _local_state[1] # 广播得到的完整张量
shape, offset = compute_local_shape_and_global_offset(
full_tensor.shape, local_state.device_mesh, local_state.placements
)
slices = [slice(cur_offset, cur_offset + cur_shape)
for cur_shape, cur_offset in zip(shape, offset)]
if local_state.is_meta:
local_tensor = full_tensor[slices].detach().clone()
ret = DTensor.from_local(local_tensor, local_state.device_mesh,
local_state.placements, ...)
else:
ret = local_state
ret.to_local().copy_(full_tensor[slices])
local_state_dict[key] = ret
当需要将完整张量分发到各 rank 的 DTensor 中时,这个函数会: 1. 根据 DTensor 的 device_mesh 和 placements 计算每个 rank 应该持有的局部形状和全局偏移 2. 从完整张量中切出对应部分 3. 复制到本地 DTensor 中
6. 展平/还原 state_dict¶
def _flatten_state_dict(state_dict):
"""将嵌套字典展平为一级字典"""
# {'a': {'b': tensor}} → {'a.b': tensor}
flattened = {}
mappings = {}
def flat_copy(path, value):
new_fqn = ".".join(map(str, path))
flattened[new_fqn] = value
mappings[new_fqn] = path
_traverse_state_dict(state_dict, flat_copy)
return flattened, mappings
def _unflatten_state_dict(state_dict, mapping):
"""根据 mapping 将展平的字典还原为嵌套字典"""
nested = {}
for key, value in state_dict.items():
_set_element(nested, mapping[key], value)
return nested
这对函数用于将嵌套的 state_dict 展平为扁平字典,以及反向还原。在优化器 state_dict 处理中会用到(优化器的 state 通常是嵌套的)。
7. 创建 CPU state_dict _create_cpu_state_dict()¶
def _create_cpu_state_dict(state_dict, pin_memory=False, share_memory=False):
def tensor_func(obj, pg, device, _):
if share_memory:
t = torch.empty(*tuple(obj.size()), dtype=obj.dtype)
t = t.share_memory_()
if pin_memory:
# 手动注册 CUDA pinned memory
torch.cuda.cudart().cudaHostRegister(t.data_ptr(), t.numel() * t.element_size(), 1)
return t
elif pin_memory:
return torch.empty(*tuple(obj.size()), dtype=obj.dtype).pin_memory()
else:
return torch.empty(*tuple(obj.size()), dtype=obj.dtype)
return _iterate_state_dict(state_dict, _identity_func, dtensor_func, tensor_func, ...)
创建一个结构相同但数据为空的 CPU state_dict,支持: - pin_memory:锁页内存,加速 CPU-GPU 数据传输 - share_memory:共享内存,支持多进程间共享 - 两者同时启用时使用 CUDA API 手动注册锁页内存
核心类/函数列表¶
| 名称 | 类型 | 说明 |
|---|---|---|
_identity_func() |
函数 | 恒等函数,原样返回张量 |
_all_gather_sharded_tensor() |
函数 | all_gather 收集 ShardedTensor |
CompanionMismatch |
异常类 | companion_obj 结构不匹配时抛出 |
_iterate_state_dict() |
函数 | 核心遍历框架,按类型分发处理 |
_gather_state_dict() |
函数 | 收集分布式 state_dict 为完整版本 |
_offload_state_dict_to_cpu() |
函数 | 将 state_dict 卸载到 CPU |
_copy_state_dict() |
函数 | 在两个 state_dict 间复制数据 |
_create_cpu_state_dict() |
函数 | 创建空的 CPU state_dict |
_check_state_dict_similarity() |
函数 | 检查两个 state_dict 结构是否一致 |
_TensorInfo |
NamedTuple | 存储张量的 shape 和 dtype |
_broadcast_tensors() |
函数 | 广播张量并分发到本地 DTensor |
_distribute_tensors() |
函数 | 将完整张量分发到各 rank 的 DTensor |
_broadcast_state_dict() |
函数 | 从 rank 0 广播完整 state_dict |
_distribute_state_dict() |
函数 | 全量 state_dict 就地分发 |
_traverse_state_dict() |
函数 | 递归遍历 state_dict |
_flatten_state_dict() |
函数 | 展平嵌套字典 |
_unflatten_state_dict() |
函数 | 还原展平的字典 |
与其他模块的关系¶
- 被
checkpoint/state_dict.py大量导入和使用 - 被 verl 的 FSDP 训练逻辑间接使用
- 依赖 PyTorch 的
DTensor、ShardedTensor、dist等分布式基础设施
小结¶
_state_dict_utils.py 是一个底层工具文件,提供了分布式 state_dict 操作的全套基础设施。它的核心设计模式是通过 _iterate_state_dict 提供统一的遍历框架,然后通过传入不同的处理函数来实现收集、广播、复制等不同操作。这种策略模式(Strategy Pattern)使得代码高度可复用。
理解这个文件需要对 PyTorch 分布式训练中的 DTensor、ShardedTensor、ProcessGroup 等概念有基本了解。如果你是初学者,建议先了解 PyTorch 的分布式训练基础,再来阅读这部分代码。