mindspore-ai / mindspore-ai/hyper-parallel
【RFC】hyper parallel全流程断点续训
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 本期目标
CheckpointingConfig的字段可通过 YAMLcheckpoint:块配置,并被CheckpointerCallback消费(除已声明但保留给 HF 导出的 4 个字段外,见 2.2)。CheckpointerCallback按save_steps/save_epochs/on_train_end三个时机存盘,_last_saved_step去重同一 step 不重复存。save_ckpt只管写、restore_from只管读,两者正交,4 种组合都有确定行为(含“只读不写”和“两者都关、不注册 callback”)。- 保存的 payload 覆盖 6 个状态桶:
model/optimizer/global_step+epoch/lr_scheduler/train_dataloader/rng_state(CPU + device + Python 三路)。 - 加载严格区分“model 缺 key 必须报错”与“optimizer 缺 key 允许缺省”(
_ModelStrictLoadPlanner),PEFT 场景下放宽 model 完整性要求。 extra_state支持 per-rank 文件与 DCP 内嵌两种布局,加载侧自动探测,不要求当前配置与写入时一致。DataLoader(components/data/dataloader.py)按(epoch, batches_consumed)续读;拓扑(dp_world_size/batch_size)变化时安全降级为从头读并告警,而不是错位续读。- Optimizer 恢复前做状态“热身”(
initialize_optimizer_state),兼容普通 optimizer 与本仓ChainedOptimizer组合包装器。 - 支持同步 / 异步(
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"→DistributedCheckpointer(checkpointer.py:131-138),经build_checkpointer(ckpt_manager="dcp", extra_state_per_rank=...)构造。 CheckpointerCallback不直接碰磁盘,只持有一个CheckpointerBase实例(self.checkpointer,checkpoint_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,298,text_trainer.py 与 base.py 逻辑一致)——内层 range 的上界用的是跑的总步数 train_steps 而非按 epoch 的步数,真正的 epoch 边界靠 dataloader 耗尽时的 StopIteration 打断,start_step 只在恢复后的第一个 epoch 生效,之后每个 epoch 结束都会把它复位为 0(base.py:967)。
恢复完成后,_last_saved_step 会被同步为刚读回的 global_step(checkpoint_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>0 且 global_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=true 且 global_step>0 且未存过 |
强制同步存盘(force_sync=True,忽略 is_async);随后 wait_for_pending_save() 排空在途的异步保存(checkpoint_callback.py:149-168) |
restore_from 的解析(_resolve_restore_path,checkpoint_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_checkpoint(dcp_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_step、state.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_device 为 None,直接跳过 |
4.5 与其它模块的交互
| 模块 | 支持的交互 | 不支持 / 注意 |
|---|---|---|
DCP(hyper_parallel.core.distributed_checkpoint) |
save/load 全部依赖它的 save / async_save / load / StandardLoadPlanner |
DistributedCheckpointer 不重复实现存储原语,只加一层 _ModelStrictLoadPlanner 做 model 严格 / optimizer 宽松的分级校验 |
ChainedOptimizer(hyper_parallel/core/optimizer/optimizer.py) |
initialize_optimizer_state 用 getattr(optimizer, "chained_optimizers", None) or [optimizer] 鸭子类型兼容 |
直接访问 .state 在 ChainedOptimizer 上不存在,历史上因此崩过一次(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 / HyperIter(trainer/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_save,maybe_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=value 是 hyper_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_step(checkpoint_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_state 靠 getattr(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 转发给 DistributedSampler、state_dict 记录消费位置、resume 后精确续读不重不漏、无 sampler(单机)场景下用私有 Generator 复现 shuffle、跨 epoch 计数器复位、拓扑变化时丢弃旧位置并告警、非法 dp_world_size/dp_rank 校验 |
不属于本文档覆盖范围、但名字容易混淆的邻近测试(见 2.2):
| 文件 | 实际覆盖对象 | 提醒 |
|---|---|---|
tests/torch/trainer/test_checkpoint_callback.py、tests/torch/trainer/_test_checkpoint_callback.py、tests/ut/trainer/callbacks/test_checkpoint_callback.py、tests/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" |
FileNotFoundError(checkpoint_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/epoch → start_epoch/start_step |
例如 steps_per_epoch=5, global_step=8 → start_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 如何验证(对应测试常问问题)
- 六桶是否都在:存盘后看
global_step_N/目录,.metadata存在即完整;save_extra_state_per_rank=true时看extra_state/extra_state_rank_{R}.pt是否每个 rank 都有;为false时改看 DCP payload 里是否有extra_statekey。 - 恢复是否正确接续:对比日志行
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。 - optimizer 动量是否真的恢复:恢复后打印
optimizer.state_dict()["state"]里任意参数的exp_avg是否非零,或直接对比恢复前后是否allclose。 - LR 是否接续而非重新 warmup:看恢复后第一步打印的
lr是否等于中断前最后一步的lr。 - dataloader 是否跳过已消费样本:数据可枚举时(如
preset_pt/dummy 索引),记录中断前最后几条样本索引,恢复后确认不重复出现。 - RNG 是否真的生效:对比开关
restore_train_state前后,恢复后同一步的 dropout mask / 初始化噪声是否符合预期(开启时应确定性延续)。 - fail-closed:故意删掉 model 的某个 key / 整个
extra_state,确认抛出的是RuntimeError/FileNotFoundError,而不是静默训出一条“看起来正常”的 loss 曲线。 - 与 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_ranktrue/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=False且restore_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
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- 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