PyTorch 原生 optimizer.state_dict() 返回的结构如下:
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 中参数的顺序分配的。
optim.param_groups
这在分布式训练中会引发严重问题:
问题 1:不同 rank 的参数 ID 不一致
在流水线并行(Pipeline Parallelism)等场景下,不同 GPU 上的优化器管理不同的参数子集。Rank 0 上的参数 ID 0 可能对应 layer1.weight,而 Rank 1 上的参数 ID 0 可能对应 layer2.weight。当所有 rank 同时保存检查点时,相同的参数 ID 指向不同的参数,导致冲突。
layer1.weight
layer2.weight
问题 2:FSDP 扁平化参数
FSDP(Fully Sharded Data Parallel)将多个原始参数扁平化(flatten)为一个 FlatParameter。优化器直接操作这个扁平参数,其状态也是针对扁平参数的。保存时需要将扁平参数的状态"反扁平化"(unflatten)回原始参数的状态。
FlatParameter
问题 3:Resharding 需求
训练过程中可能需要:
这些都需要复杂的张量重分片逻辑。
模型 state_dict 的键是参数名(如 layer1.weight),虽然会被并行包装修改(如加上 module. 前缀),但至少是字符串标识。而优化器状态使用整数 ID,完全无法跨 rank 对齐。
module.
因此,get_optimizer_state_dict 的设计目标比 get_model_state_dict 更复杂:
get_optimizer_state_dict
get_model_state_dict
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 键名都一致。
FSDP 将多个参数扁平化为一个 FlatParameter,优化器状态也是针对这个扁平参数的。get_optimizer_state_dict 需要:
FSDP 内部提供了 _unflatten_optim_state、_communicate_optim_state 等函数来完成这些操作。
_unflatten_optim_state
_communicate_optim_state
在流水线并行等 MPMD 场景下,不同 rank 的 param_groups 结构不同。DCP 在保存时会将字典扁平化,但 param_groups 中的列表(如 params: [0, 1, 2])会导致键冲突。
param_groups
params: [0, 1, 2]
例如,Rank 0 和 Rank 1 都有 param_groups.0.lr 这样的键,但对应的参数不同。
param_groups.0.lr
解决方案:引入 flatten_optimizer_state_dict 选项,将优化器状态进一步扁平化为每个参数一个键:
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 的冲突。
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,进一步扁平化
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()
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
broadcast_from_rank0=True
cpu_offload=True
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)
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")
# 保存时使用 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)
# 收集完整优化器状态到 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)
有用户报告 get_optimizer_state_dict 会调用 _init_optim_state,该函数内部执行了 step() 且假设 lr=0 不会改变状态,但实际上某些优化器(如 AdamW)在 lr=0 时仍会修改状态,导致后续行为不一致。
_init_optim_state
step()
lr=0
如果 state_dict 中缺少某些参数的状态(如微调时新增可训练参数),set_optimizer_state_dict 会抛出 KeyError,而原生的 optimizer.load_state_dict() 可以成功。
KeyError
optimizer.load_state_dict()
当优化器包含没有参数的参数组时,set_optimizer_state_dict 可能导致 optim.step() 报错,因为参数组信息被错误处理。
optim.step()
initial_lr
param_groups 中的 initial_lr 等字段在保存/加载过程中需要被正确保留,否则学习率调度器可能无法正常工作。
加载后,优化器状态中的张量应该保持为 DTensor 类型(如果原始状态是 DTensor),而不是被转换为普通张量。
DTensor
get_optimizer_state_dict 和 set_optimizer_state_dict 是 PyTorch DCP 中处理优化器状态的关键适配层,其设计解决了以下核心问题:
这些 API 与 get_model_state_dict / set_model_state_dict 共同构成了 DCP 的高层封装,使得用户无需了解底层 FSDP/DDP/TP 的具体实现细节,即可正确保存和加载分布式训练的检查点。
set_model_state_dict
一、背景:为什么需要这些接口?
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 需求
训练过程中可能需要:
这些都需要复杂的张量重分片逻辑。
1.2 与模型 state_dict 的差异
模型 state_dict 的键是参数名(如
layer1.weight),虽然会被并行包装修改(如加上module.前缀),但至少是字符串标识。而优化器状态使用整数 ID,完全无法跨 rank 对齐。因此,
get_optimizer_state_dict的设计目标比get_model_state_dict更复杂:二、设计方案详解
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需要: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 广播到其他 rankcpu_offload=True:将优化器状态卸载到 CPU,减少 GPU 内存占用三、与模型 state_dict 的对比
get_model_state_dictget_optimizer_state_dictflatten_optimizer_state_dict四、典型使用模式
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_dict和set_optimizer_state_dict是 PyTorch DCP 中处理优化器状态的关键适配层,其设计解决了以下核心问题:flatten_optimizer_state_dict扁平化这些 API 与
get_model_state_dict/set_model_state_dict共同构成了 DCP 的高层封装,使得用户无需了解底层 FSDP/DDP/TP 的具体实现细节,即可正确保存和加载分布式训练的检查点。