mindspore-ai / mindspore-ai/hyper-parallel

【RFC】支持DCP下get_optim_state_dict和set_optim_state_dict接口

Open
#679 0 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

Dominant language
Python
Stars
53
Forks
63
Avg merge
23h 45m
Merged PRs (30d)
63

Description

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

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 的具体实现细节,即可正确保存和加载分布式训练的检查点。

schema_version: 1
source: gitcode
gitcode_repo: mindspore/hyper-parallel
gitcode_issue: 240
source_url: https://gitcode.com/mindspore/hyper-parallel/issues/240

Contributor guide

No contributing guide indexed for this repository

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start by reading the existing get_optimizer_state_dict, set_optimizer_state_dict, and StateDictOptions entry points described in the issue, then trace how optimizer state, FQN mappings, FSDP parameters, and DCP flattening are currently handled. Done means the APIs support distributed optimizer save and load across FSDP, MPMD, resharding, partial-state edge cases, empty parameter groups, scheduler fields, and DTensor preservation, with tests covering the listed scenarios.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
distributed-systems
Issue type
Feature
Difficulty
5/5
Estimated time
Over a week
Activity status
Active
Clarity
Needs clarification
Newbie friendliness
25/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.