已开启
【RFC】支持DCP下get_optim_state_dict和set_optim_state_dict接口 #240
zhangbuxue创建于  6月23日
zhangbuxue成员
6月23日 创建

一、背景:为什么需要这些接口?

1.1 核心问题:优化器使用参数 ID,而非 FQN

PyTorch 原生 optimizer.state_dict() 返回的结构如下:

{
    "state": {
        0: {"step": 10, "exp_avg": tensor(...), "exp_avg_sq": tensor(...)},
        1: {"step": 10, "exp_avg": tensor(...), "exp_avg_sq": tensor(...)},
    },
    "param_groups": [
        {"lr": 0.001, "betas": (0.9, 0.999), "params": [0, 1]}
    ]
}

这里的键是参数 ID(整数),而非参数名。参数 ID 是优化器内部按照 optim.param_groups 中参数的顺序分配的。

这在分布式训练中会引发严重问题:

问题 1:不同 rank 的参数 ID 不一致

在流水线并行(Pipeline Parallelism)等场景下,不同 GPU 上的优化器管理不同的参数子集。Rank 0 上的参数 ID 0 可能对应 layer1.weight,而 Rank 1 上的参数 ID 0 可能对应 layer2.weight。当所有 rank 同时保存检查点时,相同的参数 ID 指向不同的参数,导致冲突。

问题 2:FSDP 扁平化参数

FSDP(Fully Sharded Data Parallel)将多个原始参数扁平化(flatten)为一个 FlatParameter。优化器直接操作这个扁平参数,其状态也是针对扁平参数的。保存时需要将扁平参数的状态"反扁平化"(unflatten)回原始参数的状态。

问题 3:Resharding 需求

训练过程中可能需要:

  • 在不同数量的 GPU 之间恢复训练
  • 在不同并行策略之间切换
  • 从分布式检查点加载到单卡模型

这些都需要复杂的张量重分片逻辑。

1.2 与模型 state_dict 的差异

模型 state_dict 的键是参数名(如 layer1.weight),虽然会被并行包装修改(如加上 module. 前缀),但至少是字符串标识。而优化器状态使用整数 ID,完全无法跨 rank 对齐。

因此,get_optimizer_state_dict 的设计目标比 get_model_state_dict 更复杂:

  1. 需要将参数 ID 转换为规范 FQN(Fully Qualified Name)
  2. 需要处理 FSDP 扁平化参数的反扁平化
  3. 需要支持分片状态的聚合与重新分片
  4. 需要支持MPMD(Multiple Program Multiple Data,如流水线并行)场景

二、设计方案详解

2.1 核心设计原则

原则 1:参数 ID → FQN 转换

get_optimizer_state_dict 将优化器内部使用的参数 ID 转换为规范 FQN,与 get_model_state_dict 返回的键名一致。

# 原始 optimizer.state_dict() 的 "state" 部分
{0: {"step": 10, "exp_avg": ...}, 1: {"step": 10, "exp_avg": ...}}

# get_optimizer_state_dict 转换后
{"layer1.weight": {"step": 10, "exp_avg": ...},
 "layer1.bias": {"step": 10, "exp_avg": ...}}

这样,无论使用 DDP、FSDP 还是 TP,返回的优化器 state_dict 键名都一致。

原则 2:FSDP 扁平参数的反扁平化

FSDP 将多个参数扁平化为一个 FlatParameter,优化器状态也是针对这个扁平参数的。get_optimizer_state_dict 需要:

  1. All-gather 分片的优化器状态到完整状态
  2. Unflatten 扁平参数状态到各个原始参数的状态
  3. 重新分片(如果需要)到目标拓扑

FSDP 内部提供了 _unflatten_optim_state_communicate_optim_state 等函数来完成这些操作。

原则 3:支持 MPMD 的 Flatten 模式

在流水线并行等 MPMD 场景下,不同 rank 的 param_groups 结构不同。DCP 在保存时会将字典扁平化,但 param_groups 中的列表(如 params: [0, 1, 2])会导致键冲突。

例如,Rank 0 和 Rank 1 都有 param_groups.0.lr 这样的键,但对应的参数不同。

解决方案:引入 flatten_optimizer_state_dict 选项,将优化器状态进一步扁平化为每个参数一个键

# 标准格式(无法支持 MPMD)
{
    "state": {"layer1.weight": {"step": 10, "exp_avg": ...}},
    "param_groups": [{"lr": 0.1, "params": ["layer1.weight"]}]
}

# Flatten 格式(支持 MPMD)
{
    "state.layer1.weight.step": 10,
    "state.layer1.weight.exp_avg": tensor(...),
    "param_group.layer1.weight.lr": 0.1,
    "param_group.layer1.weight.betas": (0.9, 0.999),
}

这样每个参数的状态都有唯一的键,避免了跨 rank 的冲突。

2.2 get_optimizer_state_dict 的实现逻辑

def get_optimizer_state_dict(model, optimizers, *, options=None):
    """
    返回优化器的 state_dict,键名为规范 FQN。
    
    主要功能:
    1. 收集模型参数到 FQN 的映射
    2. 调用优化器的 state_dict() 获取原始状态
    3. 将参数 ID 转换为规范 FQN
    4. 处理 FSDP 扁平参数的反扁平化
    5. 根据 options 进行分片/聚合/卸载
    """

内部实现的关键步骤:

步骤 1:收集模型信息

# 收集所有参数的 FQN 映射
param_to_fqns = _get_param_to_fqns(model)
# 收集 FSDP 模块信息(用于反扁平化)
fqn_to_fsdp_param_info = _get_fqn_to_fsdp_param_info(model)

步骤 2:获取原始优化器 state_dict

# 调用优化器自身的 state_dict()
optim_state_dict = optimizer.state_dict()

步骤 3:参数 ID → FQN 转换

# 建立参数到参数 ID 的映射
param_to_param_key = _get_param_key_to_param(optim, model, ...)
# 将 state_dict["state"] 中的整数 ID 键替换为 FQN 键

步骤 4:FSDP 反扁平化(如果适用)

# 对于 FSDP 管理的参数,需要进行:
# 1. All-gather 分片的优化器状态
# 2. Unflatten 扁平参数状态到原始参数
# 3. 可选:重新分片到目标拓扑
if use_orig_params:
    state = _convert_state_with_orig_params(...)
else:
    state = _convert_state_with_flat_params(...)

步骤 5:处理 param_groups

# 将 param_groups 中的参数 ID 也替换为 FQN
# 如果启用 flatten_optimizer_state_dict,进一步扁平化

2.3 set_optimizer_state_dict 的实现逻辑

def set_optimizer_state_dict(model, optimizers, *, optim_state_dict, options=None):
    """
    将 FQN 键名的优化器 state_dict 加载到优化器中。
    
    主要功能:
    1. 验证 state_dict 的键名
    2. 将 FQN 转换回参数 ID(根据目标优化器的参数顺序)
    3. 处理 FSDP 扁平参数的重新扁平化
    4. 调用 optimizer.load_state_dict() 完成加载
    """

内部实现的关键步骤:

步骤 1:FQN → 参数 ID 反向映射

# 根据目标优化器的参数顺序,建立 FQN -> 参数 ID 的映射
# 这与保存时的映射可能不同(因为优化器可能重新初始化)

步骤 2:处理 FSDP 扁平参数

# 对于 FSDP 管理的参数,需要将原始参数状态重新扁平化
# 以匹配目标优化器中 FlatParameter 的结构

步骤 3:Split 和 Load

# _split_optim_state_dict: 将 FQN-based state_dict 拆分回 param ID-based
# 然后调用 optimizer.load_state_dict()

2.4 StateDictOptions 中与优化器相关的配置

@dataclass
class StateDictOptions:
    full_state_dict: bool = False          # 是否返回完整状态(所有 rank 聚合)
    cpu_offload: bool = False              # 是否卸载到 CPU
    flatten_optimizer_state_dict: bool = False  # 是否扁平化(用于 MPMD)
    strict: bool = True                    # 加载时是否严格匹配
    broadcast_from_rank0: bool = False     # 是否从 rank 0 广播
  • flatten_optimizer_state_dict=True:用于流水线并行等 MPMD 场景,避免 param_groups 键冲突
  • broadcast_from_rank0=True:在加载完整状态时,从 rank 0 广播到其他 rank
  • cpu_offload=True:将优化器状态卸载到 CPU,减少 GPU 内存占用

三、与模型 state_dict 的对比

特性 get_model_state_dict get_optimizer_state_dict
原始键类型 字符串(参数名) 整数(参数 ID)
并行包装影响 键名被添加前缀 参数 ID 顺序可能不同
FSDP 影响 参数被扁平化 状态需要反扁平化
核心转换 FQN 规范化 ID → FQN 转换
MPMD 支持 相对简单 需要 flatten_optimizer_state_dict
Resharding 张量重分片 状态张量重分片 + 参数 ID 重映射

四、典型使用模式

4.1 基本保存和加载

from torch.distributed.checkpoint.state_dict import (
    get_optimizer_state_dict, set_optimizer_state_dict, StateDictOptions
)
import torch.distributed.checkpoint as dcp

# 保存
optim_sd = get_optimizer_state_dict(model, optimizer)
dcp.save({"optimizer": optim_sd}, checkpoint_id="checkpoint")

# 加载
optim_sd = get_optimizer_state_dict(model, optimizer)  # 获取空结构
dcp.load({"optimizer": optim_sd}, checkpoint_id="checkpoint")
set_optimizer_state_dict(model, optimizer, optim_state_dict=optim_sd)

4.2 使用 Stateful 包装器

from torch.distributed.checkpoint.stateful import Stateful

class AppState(Stateful):
    def __init__(self, model, optimizer):
        self.model = model
        self.optimizer = optimizer

    def state_dict(self):
        model_sd = get_model_state_dict(self.model)
        optim_sd = get_optimizer_state_dict(self.model, self.optimizer)
        return {"model": model_sd, "optim": optim_sd}

    def load_state_dict(self, state_dict):
        set_model_state_dict(self.model, state_dict["model"])
        set_optimizer_state_dict(self.model, self.optimizer, state_dict["optim"])

# DCP 自动调用 Stateful 接口
dcp.save({"app": AppState(model, optimizer)}, checkpoint_id="checkpoint")
dcp.load({"app": AppState(model, optimizer)}, checkpoint_id="checkpoint")

4.3 流水线并行(MPMD)场景

# 保存时使用 flatten 模式
opts = StateDictOptions(flatten_optimizer_state_dict=True)
optim_sd = get_optimizer_state_dict(model, optimizer, options=opts)
dcp.save({"optimizer": optim_sd}, checkpoint_id="checkpoint")

# 加载时同样使用 flatten 模式
opts = StateDictOptions(flatten_optimizer_state_dict=True)
optim_sd = get_optimizer_state_dict(model, optimizer, options=opts)
dcp.load({"optimizer": optim_sd}, checkpoint_id="checkpoint")
set_optimizer_state_dict(model, optimizer, optim_state_dict=optim_sd, options=opts)

4.4 完整状态保存(用于推理或迁移)

# 收集完整优化器状态到 rank 0
opts = StateDictOptions(full_state_dict=True, cpu_offload=True)
optim_sd = get_optimizer_state_dict(model, optimizer, options=opts)
if dist.get_rank() == 0:
    torch.save(optim_sd, "full_optim.pt")

# 加载时从 rank 0 广播
opts = StateDictOptions(full_state_dict=True, broadcast_from_rank0=True)
optim_sd = torch.load("full_optim.pt") if dist.get_rank() == 0 else None
set_optimizer_state_dict(model, optimizer, optim_state_dict=optim_sd, options=opts)

五、已知问题与注意事项

5.1 get_optimizer_state_dict 可能修改优化器状态

有用户报告 get_optimizer_state_dict 会调用 _init_optim_state,该函数内部执行了 step() 且假设 lr=0 不会改变状态,但实际上某些优化器(如 AdamW)在 lr=0 时仍会修改状态,导致后续行为不一致。

5.2 set_optimizer_state_dict 不支持部分加载

如果 state_dict 中缺少某些参数的状态(如微调时新增可训练参数),set_optimizer_state_dict 会抛出 KeyError,而原生的 optimizer.load_state_dict() 可以成功。

5.3 空参数组(Empty Param Group)问题

当优化器包含没有参数的参数组时,set_optimizer_state_dict 可能导致 optim.step() 报错,因为参数组信息被错误处理。

5.4 需要保留 initial_lr 等字段

param_groups 中的 initial_lr 等字段在保存/加载过程中需要被正确保留,否则学习率调度器可能无法正常工作。

5.5 DTensor 类型保持

加载后,优化器状态中的张量应该保持为 DTensor 类型(如果原始状态是 DTensor),而不是被转换为普通张量。


六、总结

get_optimizer_state_dictset_optimizer_state_dict 是 PyTorch DCP 中处理优化器状态的关键适配层,其设计解决了以下核心问题:

问题 解决方案
优化器使用参数 ID,无法跨 rank 对齐 将参数 ID 转换为规范 FQN
FSDP 扁平化参数 反扁平化/重新扁平化优化器状态
分片状态需要聚合 All-gather + 重新分片
MPMD(流水线并行)键冲突 flatten_optimizer_state_dict 扁平化
不同并行策略的参数 ID 差异 FQN 作为统一标识

这些 API 与 get_model_state_dict / set_model_state_dict 共同构成了 DCP 的高层封装,使得用户无需了解底层 FSDP/DDP/TP 的具体实现细节,即可正确保存和加载分布式训练的检查点。

likedislike
Zzhangbuxue成员
6月23日 修改了issue 的描述
Lliviageng
7月22日 关联了pull request:feat: support DCP optimizer state dict in HyperParallel