跳转至

_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。核心功能包括:

  1. 收集(gather):将分布在多个 GPU 上的分片张量收集为完整张量
  2. 广播(broadcast):将 rank 0 的完整 state_dict 广播到所有 rank
  3. 复制(copy):在两个结构相同的 state_dict 之间复制数据
  4. CPU 卸载(offload):将 GPU 上的张量移到 CPU 以节省显存
  5. 展平/还原(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 的分布式训练基础,再来阅读这部分代码。