mindspore-ai / mindspore-ai/hyper-parallel

【RFC】hyper parallel全流程断点续训

Open
#646 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

HyperParallel 断点续训特性文档

0. 基本信息

项目 内容
特性名称 HyperModels 断点续训(Checkpoint Save / Resume,text_trainer 全量)
代码位置 hyper_models/components/checkpoint/hyper_models/trainer/callbacks/checkpoint_callback.py
配套入口 hyper_models/trainer/base.py::BaseTrainer._init_callbacks(注册回调);hyper_models/trainer/text_trainer.py::TextTrainer.train(驱动训练循环,消费 start_epoch/start_step
开发分支 master(单提交 d5d168ce "support hyper resume training" 落地)
适用后端 仅 PyTorch —— hyper_models 全包没有 MindSpore 分支,直接依赖 torch,不经 hyper_parallel 那套 platform PT/MS 抽象层
已验证设备 PyTorch NPU/HCCL 4 卡(tp_size=2 × dp_shard_size=2)人工端到端实测;PyTorch CPU(pytest 单进程,DCP 用内存假后端 mock)UT 全量
参考实现对齐 存储原语对齐 torch.distributed.checkpoint(经 hyper_parallel.core.distributed_checkpoint 封装);extra_state 六桶拆分(progress / scheduler / dataloader / RNG)是本仓自定义设计,无直接对齐的上游实现
真实业务用法 examples/training_demo(Qwen3-30B-A3B,默认 8 卡 FSDP2,checkpoint.restore_from=LATEST 续训)
当前阶段 功能已合入单个提交,尚无自动化 ST,按整体转测

本期转测对象是以下组件构成的断点续训链路:

断点续训
├── CheckpointingConfig                          YAML checkpoint: 块对应的类型化配置
├── CheckpointerCallback                         策略层:何时存、存什么、怎么把恢复的 payload 映射回 trainer 运行时对象
├── CheckpointerBase / DistributedCheckpointer    存储层:payload 怎么落盘 / 读回,基于 DCP
├── DataLoader(components/data/dataloader.py)   可续读的 stateful dataloader
└── initialize_optimizer_state                   load 前给 optimizer“热身”,兼容 ChainedOptimizer

1. 背景

训练可能因为抢占、故障、主动分段跑等原因中断,需要能从中断点继续训练,而不是从头开始。这和“加载预训练权重”(model.pretrained_model_name_or_path 之类首次建模用的路径,见 _transformers/checkpoint_loader.py)是两件不同的事:断点续训要恢复的是完整训练状态——模型参数、优化器动量、LR scheduler 进度、数据读取位置、随机数状态——缺一个都可能导致“看起来在正常训练、loss 曲线也说得过去,但实际收敛到了另一个地方”这种静默错误。这也是 tests/hyper_models/trainer/test_checkpoint_callback.py 顶部注释直接点名的风险。

设计上延续“策略与存储解耦”的思路(checkpointer.py:18-29 docstring 明确写了这个划分):

  • CheckpointerCallback 只管策略——存不存、存什么、多久存一次、怎么把读回来的 payload 应用到 trainer 的运行时对象上;
  • CheckpointerBase(当前唯一实现 DistributedCheckpointer)只管存储——一个 dict 怎么写到目录、再读回来,存储格式可以独立于策略演化。

相关资料:

  • docs/guide/distributed_checkpoint.md(DCP 底层原语:hyper_parallel/core/distributed_checkpoint/
  • hyper-parallel断点续训方案.md(本仓库内,基于 examples/training_demo 的端到端人工实测记录;本文档 3.4 节引用其 §4 作为真实示例)

2. 本期目标与非目标

2.1 本期目标
  1. CheckpointingConfig 的字段可通过 YAML checkpoint: 块配置,并被 CheckpointerCallback 消费(除已声明但保留给 HF 导出的 4 个字段外,见 2.2)。
  2. CheckpointerCallbacksave_steps / save_epochs / on_train_end 三个时机存盘,_last_saved_step 去重同一 step 不重复存。
  3. save_ckpt 只管写、restore_from 只管读,两者正交,4 种组合都有确定行为(含“只读不写”和“两者都关、不注册 callback”)。
  4. 保存的 payload 覆盖 6 个状态桶:model / optimizer / global_step+epoch / lr_scheduler / train_dataloader / rng_state(CPU + device + Python 三路)。
  5. 加载严格区分“model 缺 key 必须报错”与“optimizer 缺 key 允许缺省”(_ModelStrictLoadPlanner),PEFT 场景下放宽 model 完整性要求。
  6. extra_state 支持 per-rank 文件与 DCP 内嵌两种布局,加载侧自动探测,不要求当前配置与写入时一致。
  7. DataLoadercomponents/data/dataloader.py)按 (epoch, batches_consumed) 续读;拓扑(dp_world_size/batch_size)变化时安全降级为从头读并告警,而不是错位续读。
  8. Optimizer 恢复前做状态“热身”(initialize_optimizer_state),兼容普通 optimizer 与本仓 ChainedOptimizer 组合包装器。
  9. 支持同步 / 异步(is_async)存盘,异步场景下不允许两次保存重叠,on_train_end 强制同步收尾。
2.2 本期非目标

以下内容不在本次断点续训转测范围内,但测试时需要知道它们的存在,避免和本特性混淆:

模块 / 字段 位置 说明
model_save_format / save_consolidated / staging_dir / best_metric_key CheckpointingConfig 已声明字段,但 CheckpointerCallback 目前不消费,保留给“合并导出 HuggingFace 格式权重”功能(config.py:30-34 docstring 明确写明)
hyper_parallel.trainer.callbacks.base.CheckpointCallback / SafetensorsExportCallback hyper_parallel/trainer/ 另一套更早的续训实现,类名只差一个 “er”(CheckpointCallback vs 本文档的 CheckpointerCallback),配置形状(args.checkpoint / args.train.checkpoint,字段是 output_dir/load_path/save_hf_weights)、落盘布局(optimizer_rank{R}.pt 等独立文件)都不同,服务的是 hyper_parallel.trainer.base 那条旧链路,与本文档覆盖的 hyper_models.trainer.text_trainer 互不调用。测试也分属 tests/torch/trainer/tests/ut/trainer/tests/hyper_models/trainer/ 两棵不同的树,转测时不要把前者的通过当成本文档功能的验证证据(详见 8.1)
HF safetensors 合并导出、离线格式转换 由其它组件负责,不在 CheckpointerCallback 的存 / 读路径内
Pipeline Parallel 下的续训 AcceleratorConfig.pp_size 存在,但断点续训与 PP 组合未见专门验证
Muon 优化器 + tp_size>1 + dp_shard_size>1 hyper_parallel/platform/torch/fully_shard/param.py 与续训无关的独立 bug(_logical_global_size 计算错误),会在第一次 optimizer.step() 就崩溃,阻塞了这个具体拓扑下的续训验证,详见第 7 节

也不承诺:

  • 续训前后训练配置(模型结构、并行度、max_steps)发生变化时的正确性——restore_from 只负责恢复“状态”,不做全局拓扑一致性校验(DataLoader.load_state_dict 里 dp_world_size/batch_size 不匹配会探测并降级,但那是 dataloader 自己的保护,不是全局校验);
  • 多 epoch + 多次中断的组合场景(见 7.1);
  • checkpoint 目录的跨集群 / 跨存储介质迁移;
  • 存盘耗时、异步保存对训练吞吐影响等性能指标。

3. 断点续训语义

3.1 CheckpointerBase / CheckpointerCallback 契约
class CheckpointerBase(ABC):
    @abstractmethod
    def save(self, path, state, *, global_step, save_async=False) -> None: ...
    @abstractmethod
    def load(self, path, state, *, strict_model=True, extra_state_skeleton=None) -> Dict[str, Any]: ...
    def maybe_wait_for_async_save(self) -> None: ...        # 默认 no-op
    def find_latest_checkpoint(self, checkpoint_dir) -> Optional[str]: ...  # 默认 NotImplementedError

约束:

  • 不能直接实例化 CheckpointerBase()ABC + abstractmethod)。
  • 当前唯一注册实现是 "dcp"DistributedCheckpointercheckpointer.py:131-138),经 build_checkpointer(ckpt_manager="dcp", extra_state_per_rank=...) 构造。
  • CheckpointerCallback 不直接碰磁盘,只持有一个 CheckpointerBase 实例(self.checkpointercheckpoint_callback.py:104-106),调用它的 save / load / maybe_wait_for_async_save / find_latest_checkpoint——这就是“策略与存储解耦”的具体体现,存储后端可以独立换掉而不动 CheckpointerCallback
3.2 存 / 读路径的正交约定

save_ckpt 只控制写路径,restore_from 只控制读路径,两者独立(config.py:36-39 注释):

save_ckpt restore_from 是否注册 CheckpointerCallback 语义
True None 只存不读,从头训练
True "LATEST" / 具体路径 边读边存,典型续训
False "LATEST" / 具体路径 只读不存:从这个 checkpoint 起步,但不再写新的(比如 branch 出一个新实验)
False None 两者都关,base.py:720-726 直接不 append CheckpointerCallback,连每步的判断开销都不产生

第四种组合是否注册 callback 由 BaseTrainer._init_callbacks 在初始化阶段一次性决定,不是在每个 hook 里各自检查一次开关。

3.3 global_step(epoch, step) 换算

保存的是单一计数器 global_step;恢复时要把它换算回训练循环用的 (start_epoch, start_step)checkpoint_callback.py:214-227, 380-382):

steps_per_epoch = len(train_dataloader) 或 train_steps(无限长 loader 兜底)
start_epoch     = global_step // steps_per_epoch
start_step      = global_step % steps_per_epoch

train() 的外层循环是 for epoch in range(start_epoch, num_train_epochs),内层是 for _ in range(start_step, train_steps)base.py:940,952 / text_trainer.py:282,298text_trainer.pybase.py 逻辑一致)——内层 range 的上界用的是跑的总步数 train_steps 而非按 epoch 的步数,真正的 epoch 边界靠 dataloader 耗尽时的 StopIteration 打断,start_step 只在恢复后的第一个 epoch 生效,之后每个 epoch 结束都会把它复位为 0(base.py:967)。

恢复完成后,_last_saved_step 会被同步为刚读回的 global_stepcheckpoint_callback.py:384-387):如果恢复后没有新的训练发生(比如直接对齐到 max_steps),on_train_end 就不会把刚读进来的 checkpoint 又原样写一遍。

3.4 真实示例:4 卡 tp=2 × dp_shard=2 的存盘与恢复

以下记录取自 hyper-parallel断点续训方案.md §4,是本特性目前唯一一次贴近生产拓扑的人工端到端验证(Ascend 910B3 × 4,基于 examples/training_demo 改出的 tp_size=2, dp_shard_size=2 配置):

阶段一:跑到 step 15 主动停止--training.max_steps=15,比 kill 进程更可控可复现):

Rank0 Start training. Start step: 0. Train steps: 15. Start epoch: 0. Train epochs: 1.
...
Saving checkpoint: global_step=10, epoch=0, dir=.../global_step_10, ...   # save_steps=10 命中
...
Saving checkpoint: global_step=15, epoch=0, dir=.../global_step_15, ...   # on_train_end 强制补存

阶段二:恢复--checkpoint.restore_from=LATEST):

Resolved LATEST checkpoint: .../global_step_15
Read extra_state embedded in the DCP state dict
Checkpoint loaded successfully: path=.../global_step_15, global_step=15, start_epoch=0, start_step=15
Rank0 Start training. Start step: 15. Train steps: 25. Start epoch: 0. Train epochs: 1.

Training:  60%|██████    | 15/25 ... loss=6.90772, lr=3.09e-05   # 从 15 接着走,lr 曲线没有重置
...
Training: 100%|██████████| 25/25 ... loss=6.90778, lr=0

start_step=15 与保存时的 global_step=15 一致;LR 从中断处的 3.09e-05 继续按 cosine 衰减到 0(而不是从 warmup 重新开始),完整验证了 save → 中断 → resume → 续跑完成的整条链路。


4. 总体设计

4.1 架构与数据流
用户配置
  YAML checkpoint: 块 → CheckpointingConfig
        |
        v
BaseTrainer._init_callbacks()
  save_ckpt=true 或 restore_from 非 None 时才 append CheckpointerCallback(3.2)
        |
        v
CheckpointerCallback(策略层,trainer/callbacks/checkpoint_callback.py)
  on_train_begin → _load_checkpoint(起步前恢复一次;模型/优化器/dataloader 均已构建完毕)
  on_step_end    → save_steps 命中时 _save_checkpoint
  on_epoch_end   → save_epochs 命中时 _save_checkpoint(去重 _last_saved_step)
  on_train_end   → 未存过的最后一步强制同步存 + wait_for_pending_save
        |
        v
checkpointer.save(path, state, global_step=, save_async=)
checkpointer.load(path, state, strict_model=, extra_state_skeleton=)
  CheckpointerBase 抽象接口(3.1),当前唯一实现 DistributedCheckpointer
        |
        v
DistributedCheckpointer(存储层,components/checkpoint/dcp_checkpointer.py)
  save:extra_state 落盘(per-rank 或塞进 payload)→ dcp_save/dcp_async_save → _finalize_checkpoint 写 latest 指针
  load:探测 extra_state 布局 → dcp_load(planner=_ModelStrictLoadPlanner) → _read_extra_state
        |
        v
hyper_parallel.core.distributed_checkpoint(DCP 原语)
  save / async_save / load / StandardLoadPlanner
        |
        v
磁盘
  checkpoint_dir/
  ├── latest_checkpoint_iteration.txt        只存一个数字:最新 checkpoint 的 global_step
  └── global_step_N/
      ├── .metadata                          DCP 完整性标记,写完所有 rank 的 shard 才会落这个文件
      ├── <各 rank 的 shard 文件>
      └── extra_state/extra_state_rank_{R}.pt   仅 save_extra_state_per_rank=true 时存在

设计原则:

  • CheckpointerCallback 不重复实现存储逻辑,落盘/读回全部委托给 checkpointer
  • .metadata 是完整性的唯一判据:写到一半被打断的目录没有这个文件,恢复时会被自动跳过,不会读到半成品(dcp_checkpointer.py:232-247)。
  • 存和读互斥:save() / load() 开头都先 maybe_wait_for_async_save()dcp_checkpointer.py:402,454),保证同一个 checkpointer 实例上不会有两个 DCP 操作重叠。
4.2 CheckpointerCallback(策略层)
Hook 触发条件 行为
on_train_begin 回调被注册就会跑 打印 checkpoint 配置日志;_load_checkpoint():解析 restore_from → 视情况给 optimizer “热身” → 经 checkpointer.load 读取 model/optimizer/extra_state → 写回 trainer 各运行时对象
on_step_end save_steps>0global_step % save_steps==0 若本 step 未存过则 _save_checkpoint;去重靠 _last_saved_step
on_epoch_end save_epochs>0(epoch+1) % save_epochs==0 若本 step 未被 on_step_end 存过则 _save_checkpoint,否则只打日志跳过(checkpoint_callback.py:137-147
on_train_end save_ckpt=trueglobal_step>0 且未存过 强制同步存盘(force_sync=True,忽略 is_async);随后 wait_for_pending_save() 排空在途的异步保存(checkpoint_callback.py:149-168

restore_from 的解析(_resolve_restore_pathcheckpoint_callback.py:274-299):

行为
None 不恢复,从头训练
"LATEST"(大小写不敏感) 交给 checkpointer.find_latest_checkpoint;找不到时 warning 并从头训练,不报错
具体路径 校验目录存在,否则 FileNotFoundError
4.3 DistributedCheckpointer(存储层)与 extra_state 双布局

extra_state_per_rank(构造参数,来自 save_extra_state_per_rank)决定保存时的布局,加载侧自动探测,不要求和保存时一致:

  • True:每个 rank 单独 torch.save 一个 extra_state/extra_state_rank_{R}.pt,永远正确。
  • False:内嵌进 DCP payload,DCP 会对相同 FQN 的条目跨 rank 去重,“每个 rank 恢复的是谁的副本”取决于去重胜出的是谁——只有 rank 间完全一致的状态才适合这样存(见 7.2)。

内嵌布局的哨兵值机制(dcp_checkpointer.py:57-62, 461-464, 494-508):extra_state 骨架里的 global_step 字段会被先强制置成 -1_UNRESTORED_STEP,真实 checkpoint 不可能出现负数 step),再交给 dcp_load

  • checkpoint 里真的有内嵌 extra_state-1 被真实值覆盖,正常返回;
  • 没有(比如实际是分 rank 存的、或者干脆没存训练状态)→ -1 原样留下来,被识别出来直接 raise FileNotFoundError,而不是静默地当成“成功恢复到 step 0”。

find_latest_checkpointdcp_checkpointer.py:319-335)优先读 latest_checkpoint_iteration.txt 指针文件(一次 I/O 即可回答);指针缺失、损坏、或指向的目录不完整时,退化为扫描 checkpoint_dir 下所有 global_step_*,取 .metadata 存在(即完整)的最大 step。指针发布是 _finalize_checkpoint 里的 barrier - rank0 写 - barrier 三段式(dcp_checkpointer.py:366-385):前一个 barrier 保证指针不会发布到还在写的 checkpoint 上,后一个 barrier 保证发布完成前没有 rank 抢跑去读。

4.4 六个状态桶
写入条件 内容 恢复方式 备注
model 总是写 model.state_dict()(PEFT 下只含 requires_grad=True 的参数,_model_state_dict DCP 原地填充 + 显式 model.load_state_dict(strict=not is_peft) DTensor/FSDP2 分片状态,每 rank 只落自己的 shard
optimizer save_optimizer=true 各 optimizer(可能是 list)的 state_dict() load 前先 initialize_optimizer_state 热身,再 DCP 填充 + optimizer.load_state_dict 数量与当前 optimizer 数不一致时按较短者配对,多出的保持初始状态(仅 warning,checkpoint_callback.py:342-353
global_step / epoch save_train_state=true state.global_stepstate.epoch 直接写回 trainer.state,并据此推出 start_epoch/start_step(3.3) 内嵌于 extra_state
lr_scheduler 同上 各 scheduler(可能是 list)的 state_dict() 数量不匹配时同 optimizer 的截断 + warning 策略(checkpoint_callback.py:392-404 支持单/多 scheduler(_as_list/_unwrap_single
train_dataloader 同上 优先 data_iterator.state_dict(),否则 train_dataloader.state_dict() trainer.train_dataloader.load_state_dict(...),仅当 dataloader 具备该方法 非 stateful dataloader 时告警“本 epoch 将从头重放”
rng_state 同上 torch_cpu(torch.get_rng_state())、torch_device(get_device_rng_state())、python(random.getstate()) 分别 torch.set_rng_state / set_device_rng_state / random.setstate CPU-only 环境下 torch_deviceNone,直接跳过
4.5 与其它模块的交互
模块 支持的交互 不支持 / 注意
DCP(hyper_parallel.core.distributed_checkpoint save/load 全部依赖它的 save / async_save / load / StandardLoadPlanner DistributedCheckpointer 不重复实现存储原语,只加一层 _ModelStrictLoadPlanner 做 model 严格 / optimizer 宽松的分级校验
ChainedOptimizerhyper_parallel/core/optimizer/optimizer.py initialize_optimizer_stategetattr(optimizer, "chained_optimizers", None) or [optimizer] 鸭子类型兼容 直接访问 .stateChainedOptimizer 上不存在,历史上因此崩过一次(3.4 节),现已修复
TP / FSDP2(fully_shard 参数已是 DTensor / 分片状态,checkpoint 直接存这个分片后的 state_dict(),与并行拓扑正交 Muon + tp_size>1 + dp_shard_size>1 会在 optimizer.step() 就崩(无关 bug,见 2.2)
PEFT _model_state_dict 只存 trainable 参数;strict_model=False 放宽 _ModelStrictLoadPlanner 的 model 完整性检查 恢复时 model.load_state_dict(..., strict=False),基座冻结权重不会被当成“缺失”
BackgroundPrefetcher / HyperItertrainer/base.py _collect_extra_state 优先取 data_iterator.state_dict()——后台预取线程已经比训练 step 实际消费的 batch 更超前,用 iterator 在“取出当前 batch 那一刻”捕获的快照,而不是 dataloader 的实时状态,才能对齐“刚训完这一步”的数据位置 未启用预取(use_background_prefetcher=false)时两者等价,直接退化为 train_dataloader.state_dict()
SkipDTensorDispatch initialize_optimizer_state 用它包住热身用的 optimizer.step(),并显式 no_skip={torch.zeros_like} 热身只是为了创建 exp_avg/exp_avg_sq 等状态张量,不能真的改变参数值,所以同时把 lr/weight_decay 清零,finally 里再恢复
异步保存 is_async=true 时走 DCP 自己的 async_savemaybe_wait_for_async_save 保证同一 checkpointer 实例上两次保存 / 一次读不会重叠 on_train_end 的最后一次保存永远同步(force_sync=True),不受 is_async 影响

5. 对外接口

5.1 YAML 配置

CheckpointingConfig 全部字段(config.py:24-73):

字段 默认值 含义
save_ckpt True 只控制写路径,不影响读(见 3.2)
checkpoint_dir "./checkpoints" 保存根目录
save_steps 0(关闭) 每 N 个 optimizer step 存一次
save_epochs 1 每 N 个 epoch 存一次
is_async False 是否异步存盘
is_peft False 是否只持久化可训练(adapter)权重
save_optimizer True 是否保存 optimizer 桶
save_train_state True 是否保存 progress/scheduler/dataloader/RNG 桶
save_extra_state_per_rank True extra_state 落盘布局,见 4.3
restore_from None None / "LATEST" / 具体路径
restore_optimizer True 是否恢复 optimizer
restore_train_state True 是否恢复 step/epoch/scheduler/dataloader/RNG
model_save_format "safetensors" 保留字段,见 2.2
save_consolidated "final" 保留字段("none"/"final"/"every"),见 2.2
staging_dir None 保留字段,见 2.2
best_metric_key "default" 保留字段,见 2.2

真实示例(examples/training_demo/train.yaml:155-178):

checkpoint:
  save_ckpt: true
  checkpoint_dir: ./outputs/training_demo/checkpoints
  save_steps: 10
  save_epochs: 1
  is_async: false
  save_optimizer: true
  save_train_state: true
  save_extra_state_per_rank: false

  # Resume with:  bash examples/training_demo/run.sh --checkpoint.restore_from=LATEST
  # or point at one directory: --checkpoint.restore_from=./outputs/training_demo/checkpoints/global_step_10
  restore_from: LATEST
  restore_optimizer: true
  restore_train_state: true

--dotted.path=valuehyper_models.config.manager 的 CLI override 语法,可覆盖 YAML 里任意字段。

5.2 编程接口
from hyper_models.components.checkpoint import (
    CheckpointingConfig, CheckpointerBase, build_checkpointer, CHECKPOINTER_REGISTRY,
)
from hyper_models.trainer.callbacks import CheckpointerCallback, TrainerState
  • build_checkpointer(ckpt_manager="dcp", **kwargs):从 CHECKPOINTER_REGISTRY 按名字取实现并构造(checkpointer.py:40-53);CheckpointerCallback.__init__ 是当前唯一调用方(checkpoint_callback.py:104-106),普通用户不需要直接构造。
  • CheckpointerCallback(trainer):正常由 BaseTrainer._init_callbacks 按 3.2 的规则自动注册;测试代码可以直接构造它对着一个满足鸭子类型的 trainer 对象跑(tests/hyper_models/trainer/test_checkpoint_callback.py 的做法)。
  • 扩展点:新增一个存储后端只需要实现 CheckpointerBase.save/load,再 @CHECKPOINTER_REGISTRY.register("your_name")CheckpointerCallback 不用改。

6. 当前支持矩阵

能力 状态 备注
六桶 save/load 编排(save_steps/save_epochs/on_train_end 触发 + 去重) 支持 UT 全覆盖(tests/hyper_models/trainer/test_checkpoint_callback.py
save_ckpt/restore_from 正交开关(4 种组合) 支持 UT 覆盖全部组合
LATEST 解析(指针文件 + 扫描兜底 + 跳过不完整目录) 支持 UT 覆盖,含指针损坏/指向缺失目录/忽略无关目录
extra_state 双布局(per-rank / 内嵌)自动探测 支持 UT 覆盖两种布局往返 + 哨兵值缺失报错
optimizer 状态热身(含 ChainedOptimizer 兼容) 支持 UT 覆盖 + NPU 实测发现并修复过一次真实 bug(3.4 节)
PEFT 部分权重 checkpoint(strict_model=False 支持 UT 覆盖
DataLoader 续读(state_dict/load_state_dict,含有/无 DistributedSampler 两种路径、拓扑变化探测) 支持 UT 覆盖(tests/components/test_dataloader.py),限单进程场景
异步保存(真正的 DCP async_save,多次保存排队等待) 支持 UT 覆盖(mock 后端);NPU 实机 is_async=true 场景无自动化 ST
与 TP/FSDP2 并行拓扑组合下的续训 端到端人工验证通过 4 卡 tp=2 × dp_shard=2,见 3.4 节;无自动化 ST
多 epoch + 多次中断的组合续训 未覆盖 已知边界,见 7.1
Muon 优化器 + tp_size>1 + dp_shard_size>1 续训 不支持 无关 bug 阻塞,见 2.2 / 7.4
CUDA/NCCL 未作为验收设备 本期按 NPU/HCCL 实机 + CPU mock UT

7. 风险与限制

7.1 多 epoch 恢复的 start_epoch/start_step 双路径风险

start_epoch/start_stepcheckpoint_callback.py:380-382)是用 global_step // steps_per_epoch 推导出来的,而不是直接读 dataloader 自己保存的真实 _epoch/_batches_consumed——这是两条独立路径,只在 len(train_dataloader) 全程严格不变时才会一直吻合。DataLoader.set_epoch()dataloader.py:138-148)只有在推导出来的 epoch 恰好等于 dataloader 自己记的 _epoch 时才会保留 _batches_consumed;一旦不等,会静默清零 _batches_consumed,导致这个 epoch 内一部分样本重复训练、另一部分被跳过。已知触发条件:dp_world_size==1 且续训时更换了拓扑或 batch_size(load_state_dict 会先探测拓扑,见 7 节表格),或配置了不支持 state_dict/load_state_dict 的 dataloader。单 epoch 训练不会触发;多 epoch + 多次中断的组合目前没有专门覆盖。

7.2 save_extra_state_per_rank=false 要求状态 rank 间一致

内嵌布局下 DCP 会对相同 FQN 的条目跨 rank 去重,“每个 rank 恢复的是谁的副本”取决于去重胜出的是谁(dcp_checkpointer.py:199-212 docstring)。extra_state 里的 rng_state 理论上并不满足“rank 间完全一致”这个前提——各 rank 的 RNG 流通常不同。需要精确恢复每个 rank 各自 RNG 状态时,save_extra_state_per_rank=true 是唯一保真的选项。

7.3 ChainedOptimizer 状态初始化的鸭子类型依赖

initialize_optimizer_stategetattr(optimizer, "chained_optimizers", None) 判断是否为组合优化器(3.4 节的历史 bug 已修复)。如果未来出现第三种优化器包装形态(既不是普通 Optimizer 也不暴露 chained_optimizers),这里会重新踩坑,建议转测时补一个“未知包装器类型”的用例把这个假设显式钉住。

7.4 Muon + TP>1 + FSDP2 dp_shard>1

与断点续训无关的独立问题(_logical_global_size 计算错误,hyper_parallel/platform/torch/fully_shard/param.py:417),会在第一次 optimizer.step() 就崩溃,阻塞了这个具体拓扑下的续训验证。3.4 节的实测因此换用了 AdamW 绕开。

7.5 数量不匹配只告警、不报错

optimizer / lr_scheduler 数量与 checkpoint 记录不一致时(checkpoint_callback.py:342-353, 392-404),只按较短列表配对并打 warning,多出的部分保持初始状态——这是有意的宽松策略(允许恢复后调整优化器组合),但也意味着配置写错导致的数量不匹配不会 fail-closed,容易被忽略掉一条 warning 日志就继续跑。


8. 验证设计与当前结果

8.1 已有覆盖(开发侧)

UT(CPU,不启分布式,DCP 用内存假后端 mock):

文件 覆盖
tests/hyper_models/trainer/test_checkpoint_callback.py save 节奏(step/epoch/train_end,含去重)、四种 save_ckpt × restore_from 组合的注册决策、六桶按需写入、optimizer 热身(含 noop 场景)、六桶完整 round-trip(含 start_epoch/start_step 推导)、extra_state 双布局往返、model/extra_state 缺失的 fail-closed、PEFT 部分权重、“恢复后不重写”幂等性、restore_train_state=false 只读权重、LATEST 解析(指针命中/回退扫描/跳过不完整目录/指针损坏兜底/忽略无关目录)、异步保存排空
tests/components/test_dataloader.py DP 分片、set_epoch 转发给 DistributedSamplerstate_dict 记录消费位置、resume 后精确续读不重不漏、无 sampler(单机)场景下用私有 Generator 复现 shuffle、跨 epoch 计数器复位、拓扑变化时丢弃旧位置并告警、非法 dp_world_size/dp_rank 校验

不属于本文档覆盖范围、但名字容易混淆的邻近测试(见 2.2):

文件 实际覆盖对象 提醒
tests/torch/trainer/test_checkpoint_callback.pytests/torch/trainer/_test_checkpoint_callback.pytests/ut/trainer/callbacks/test_checkpoint_callback.pytests/ut/trainer/test_checkpoint_callback_config.py hyper_parallel.trainer.callbacks.base.CheckpointCallback(旧栈,注意类名无 “er”) 转测时不要把这几个文件的通过当作本文档功能的验证证据

无 ST:目前仓库里没有针对 hyper_models.trainer.text_trainer + CheckpointerCallback 这条链路的自动化多卡 / NPU 测试用例。

人工端到端记录hyper-parallel断点续训方案.md 记录了一次完整的 4 卡 Ascend 910B3 实机验证(详见 3.4 节),是目前唯一一次贴近生产拓扑的真实验证,尚未固化成 ST。转测建议优先把它变成自动化的 8.2 节 D1 用例。

8.2 转测建议用例
A. 接口与 fail-closed
ID Feature Description Expectation
A1 恢复路径不存在 restore_from="/no/such/dir" FileNotFoundErrorcheckpoint_callback.py:298
A2 LATEST 但目录下无任何 checkpoint restore_from="LATEST"checkpoint_dir 为空 warning,从头训练,不报错
A3 LATEST 指针指向不完整目录 手工写坏指针指向的 step 缺 .metadata 回退扫描,warning,取次新的完整 checkpoint
A4 权重-only checkpoint 去恢复 train_state save_optimizer=false,save_train_state=false 存盘后 restore_train_state=true 去读 FileNotFoundError,信息含 "no training state"(dcp_checkpointer.py:499-506
A5 内嵌 extra_state 实际未写入 手工删掉 payload 里的 "extra_state" 同 A4,哨兵值 _UNRESTORED_STEP 触发
A6 model 权重缺 key 手工删掉 saved state_dict 里的一个 "model.xxx" key RuntimeError,信息含 "missing model key"(dcp_checkpointer.py:112-115
A7 PEFT 模型缺 base 权重 key is_peft=true 时同样删掉一个 key 正常恢复,不报错
A8 save_ckpt/restore_from 四种组合 见 3.2 表 callback 是否注册符合表格
A9 optimizer/scheduler 数量不一致 checkpoint 存 2 个,当前 run 只有 1 个(或反之) warning + 按较短者配对,不崩溃(见 7.5)
A10 dataloader 拓扑变化 存盘 dp_world_size=2,恢复 dp_world_size=4 告警并丢弃保存位置,从该 epoch 第一条样本重新读
B. 存盘触发与去重
ID Feature Description Expectation
B1 save_steps 周期存盘 save_steps=N,跑 2N+1 步 恰好在 step N、2N 各存一次
B2 save_epochs 周期存盘 save_epochs=1,不设 save_steps 每个 epoch 结束存一次
B3 同一 step 被两个触发点同时命中 step 恰好是 epoch 最后一步且是 save_steps 倍数 只存一次,on_epoch_end 打日志跳过
B4 on_train_end 补存 max_steps 提前于任何周期边界结束 强制同步存最后一步
B5 on_train_end 不重复存 最后一步恰好刚被存过 不再存
B6 两个周期都关 save_steps=save_epochs=0 训练期间不存,仅 on_train_end(若 save_ckpt=true)补一次
B7 恢复后立即收尾不重复写 恢复到 checkpoint 后没有新训练发生 _last_saved_step 已在恢复时同步,on_train_end 不重写(见 3.3)
C. 六桶恢复正确性
ID Feature 说明
C1 model 权重恢复 恢复后与保存前 allclose
C2 optimizer 动量恢复(AdamW) 训练几步后存盘,exp_avg/exp_avg_sq 恢复后与保存前一致且非零
C3 optimizer 未训练过就存盘再恢复 全零 state 不报错(容忍无 grad 的参数)
C4 ChainedOptimizer 恢复 Muon/AdamW 拆分或 decay/no-decay 拆分,每个子 optimizer 都被正确热身
C5 global_step/epochstart_epoch/start_step 例如 steps_per_epoch=5, global_step=8start_epoch=1, start_step=3
C6 lr_scheduler 恢复 续训 loss/lr 曲线延续,不从 warmup 重新开始(3.4 节实测口径)
C7 dataloader 位置恢复(有 DistributedSampler dp_world_size>1 时跳过已消费 batch,不重不漏
C8 dataloader 位置恢复(无 sampler) dp_world_size==1,私有 Generator(seed+epoch) 复现同一路 shuffle
C9 RNG 三路恢复 CPU + device + Python 恢复后下一次随机操作与未中断的参考轨迹一致
C10 extra_state 双布局往返 save_extra_state_per_rank true/false 各跑一遍都能正确恢复
C11 PEFT 权重恢复 只有 adapter 参数被保存/恢复,基座权重保持初始加载值
D. 拓扑与并行组合(需要 NPU 多卡,当前均为转测补齐重点)
ID Feature Description
D1 TP+FSDP2 组合续训 tp_size=2, dp_shard_size=2,4 卡,完整走通 save→中断→resume→续训(3.4 节场景固化为 ST)
D2 纯 FSDP2(无 TP) tp_size=1, dp_shard_size=N,对齐 examples/training_demo 默认 8 卡配置
D3 异步保存(is_async=true 训练中频繁触发 save_steps,不阻塞训练;下一次保存前等待上一次完成;进程退出前完成收尾
D4 Muon + TP>1 + FSDP2>1 预期在 optimizer.step() 报错(2.2 已知缺陷),不算本文档验收范围,仅确认失败模式没有变化
D5 多 epoch 跨越 + 二次中断 num_train_epochs>=2,验证是否符合“该 epoch 从头重放”的既有告警语义(7.1),而非默默错位
E. 平台矩阵
后端 设备 必跑
PyTorch CPU,pytest 单进程(mock DCP) A 类全部 + B 类全部 + C1/C3/C4/C5/C6/C9/C10/C11
PyTorch CPU,dataloader 专项(无 DCP 依赖) C7/C8(已有覆盖,见 8.1)
PyTorch NPU/HCCL,4 卡 D1(已人工验证一次,待固化)、D3、C2/C7 的多卡版本
PyTorch NPU/HCCL,8 卡 D2,对齐 examples/training_demo 默认拓扑
8.3 如何验证(对应测试常问问题)
  1. 六桶是否都在:存盘后看 global_step_N/ 目录,.metadata 存在即完整;save_extra_state_per_rank=true 时看 extra_state/extra_state_rank_{R}.pt 是否每个 rank 都有;为 false 时改看 DCP payload 里是否有 extra_state key。
  2. 恢复是否正确接续:对比日志行 Checkpoint loaded successfully: ... global_step=%s, start_epoch=%s, start_step=%s 与保存时的 Saving checkpoint: global_step=%s, epoch=%s 是否吻合;start_step 应等于 global_step % steps_per_epoch
  3. optimizer 动量是否真的恢复:恢复后打印 optimizer.state_dict()["state"] 里任意参数的 exp_avg 是否非零,或直接对比恢复前后是否 allclose
  4. LR 是否接续而非重新 warmup:看恢复后第一步打印的 lr 是否等于中断前最后一步的 lr
  5. dataloader 是否跳过已消费样本:数据可枚举时(如 preset_pt/dummy 索引),记录中断前最后几条样本索引,恢复后确认不重复出现。
  6. RNG 是否真的生效:对比开关 restore_train_state 前后,恢复后同一步的 dropout mask / 初始化噪声是否符合预期(开启时应确定性延续)。
  7. fail-closed:故意删掉 model 的某个 key / 整个 extra_state,确认抛出的是 RuntimeError/FileNotFoundError,而不是静默训出一条“看起来正常”的 loss 曲线。
  8. 与 TP/FSDP2 的交互:确认恢复后模型参数仍是预期的 DTensor/FSDP2 分片(而非不小心变成普通 tensor),grad_norm/loss 量级与未中断的参考 run 可比。

9. 验收标准

9.1 功能验收
  • CheckpointingConfig 全部消费字段(除 2.2 声明的 4 个保留字段外)在 CheckpointerCallback 中生效。
  • save_steps/save_epochs/on_train_end 三个触发点行为符合 4.2 表,_last_saved_step 去重生效。
  • save_ckpt/restore_from 四种组合的 callback 注册与读写行为符合 3.2 表。
  • 六个状态桶按 save_optimizer/save_train_state 精确控制是否写入;恢复后 start_epoch/start_step/optimizer 动量/lr_scheduler 进度/dataloader 位置/RNG 与保存时一致。
  • extra_state 两种布局(save_extra_state_per_rank true/false)均可正确保存与恢复,加载侧自动探测。
  • ChainedOptimizer 与普通 torch.optim.Optimizer 都能被 initialize_optimizer_state 正确热身。
  • PEFT(is_peft=true)只保存/恢复可训练参数,不要求基座权重存在于 checkpoint。
9.2 兼容性验收
  • 不配置 checkpoint: 块(全部用默认值)时,行为等价于 save_ckpt=True, save_steps=0, save_epochs=1, restore_from=None——默认每个 epoch 存一次、不主动恢复,不应影响现有训练脚本。
  • save_ckpt=Falserestore_from=None 时,CheckpointerCallback 完全不注册(base.py:720-726),对训练主循环零开销。
  • TP/FSDP2/PEFT 与断点续训正交组合可用,不要求为了支持续训而改变并行拓扑配置方式。
  • TextTrainer(组合 BaseTrainer 而非继承它)与 BaseTrainer 驱动的其它 trainer 共享同一套 on_train_begin/on_step_end/on_epoch_end/on_train_end 回调分发,断点续训行为一致。
9.3 明确报错(必须 fail-closed)
场景 异常
restore_from 指向不存在的具体目录 FileNotFoundError
权重-only checkpoint 被要求恢复 train_state FileNotFoundError,信息含 "no training state"
内嵌 extra_state 实际未写入(哨兵值未被覆盖) FileNotFoundError,信息含 "extra_state"
model state_dict 缺任意 key(非 PEFT) RuntimeError,信息含 "missing model key"

以下场景刻意不报错,是设计选择而非疏漏(见 7.5):

场景 行为
restore_from="LATEST" 但目录下无任何 checkpoint warning + 从头训练
optimizer/lr_scheduler 数量与 checkpoint 不一致 warning + 按较短者配对
optimizer 缺某参数的 state(未参与过反向) 静默保留零初始化值
dataloader 拓扑(dp_world_size/batch_size)变化 warning + 该 epoch 从头重放
dataloader 不支持 state_dict/load_state_dict warning + 跳过该桶恢复
9.4 一致性口径
  • 单元测试口径:权重/动量恢复用 torch.allclose(atol=1e-6),状态字典逐 key 相等比较(tests/hyper_models/trainer/test_checkpoint_callback.py 的 round-trip 用例)。
  • 端到端口径:不是“和单卡参考比数值”,而是看“中断前最后一步”和“恢复后第一步”打印的 loss/lr 是否落在同一条轨迹上(3.4 节实测:中断前 lr=3.09e-05,恢复后第一步同为 3.09e-05 并继续按 cosine 衰减)。DataLoader.__iter__ 是“重新读并丢弃已消费前缀”(dataloader.py:116-136),不是“从内存快照续跑”,严格意义上比特级可复现要求 RNG + 数据顺序完全对齐,属于第 7 节已知边界,不作为本次验收的强制口径。
9.5 转测完成定义
  • tests/hyper_models/trainer/test_checkpoint_callback.py + tests/components/test_dataloader.py 全量 UT 回归通过。
  • 8.2 节 D1(TP+FSDP2 4 卡)在 NPU 实机固化为自动化 ST,不再依赖人工记录。
  • 8.2 节 A 类 fail-closed 场景中,至少 A4/A5/A6 三个“恢复失败必须报错”的场景在 NPU 实机上复验一次。
  • 7.1 节“多 epoch 恢复边界语义”的行为(而非修复)有明确文档记录,不要求本期修复。
  • 不把 hyper_parallel.trainer.callbacks.base.CheckpointCallback(旧栈,见 2.2)的测试结果算作本文档功能的转测通过条件。

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

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 with tests/hyper_models/trainer/test_checkpoint_callback.py and the entry points BaseTrainer._init_callbacks and TextTrainer.train. Compare the documented save, restore, dataloader, optimizer, RNG, and asynchronous paths with the existing implementation, then use examples/training_demo for the documented end-to-end scenario. Done means the full checkpoint-resume flow is tested and any unsupported cases remain explicitly identified.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
distributed-systems, testing-qa
Issue type
Feature
Difficulty
5/5
Estimated time
Over a week
Activity status
Active
Clarity
Mostly clear
Newbie friendliness
35/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.