已关闭
【RFC】hyper parallel全流程断点续训 #328
zhangbuxue创建于 8月11日关闭于 5 天前
8月11日 添加了label:RFC
8月11日 修改标题为 “【RFC】hyper parallel全流程断点续训”,原标题为“hyper parallel全流程断点续训”
8月11日 修改了issue 的描述
8月11日 修改了issue 的描述
8月12日 将 zhangbuxue 设为负责人
8月12日 修改了issue 的描述
8月17日 修改了issue 的描述
8月17日 修改了issue 的描述
8月17日 修改了issue 的描述
8月17日 修改了issue 的描述
8月17日 修改了issue 的描述
8月17日 修改了issue 的描述
8月17日 修改了issue 的描述
8月17日 修改了issue 的描述
8月17日 修改了issue 的描述
5 天前 issue状态由 TODO 改变为 DONE
5 天前 关闭了 issue
HyperParallel 断点续训特性文档
0. 基本信息
text_trainer全量)hyper_models/components/checkpoint/、hyper_models/trainer/callbacks/checkpoint_callback.pyhyper_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"落地)hyper_models全包没有 MindSpore 分支,直接依赖torch,不经hyper_parallel那套platformPT/MS 抽象层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续训)本期转测对象是以下组件构成的断点续训链路:
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-29docstring 明确写了这个划分):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”)。model/optimizer/global_step+epoch/lr_scheduler/train_dataloader/rng_state(CPU + device + Python 三路)。_ModelStrictLoadPlanner),PEFT 场景下放宽 model 完整性要求。extra_state支持 per-rank 文件与 DCP 内嵌两种布局,加载侧自动探测,不要求当前配置与写入时一致。DataLoader(components/data/dataloader.py)按(epoch, batches_consumed)续读;拓扑(dp_world_size/batch_size)变化时安全降级为从头读并告警,而不是错位续读。initialize_optimizer_state),兼容普通 optimizer 与本仓ChainedOptimizer组合包装器。is_async)存盘,异步场景下不允许两次保存重叠,on_train_end强制同步收尾。2.2 本期非目标
以下内容不在本次断点续训转测范围内,但测试时需要知道它们的存在,避免和本特性混淆:
model_save_format/save_consolidated/staging_dir/best_metric_keyCheckpointingConfigCheckpointerCallback目前不消费,保留给“合并导出 HuggingFace 格式权重”功能(config.py:30-34docstring 明确写明)hyper_parallel.trainer.callbacks.base.CheckpointCallback/SafetensorsExportCallbackhyper_parallel/trainer/CheckpointCallbackvs 本文档的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)CheckpointerCallback的存 / 读路径内AcceleratorConfig.pp_size存在,但断点续训与 PP 组合未见专门验证tp_size>1+dp_shard_size>1hyper_parallel/platform/torch/fully_shard/param.py_logical_global_size计算错误),会在第一次optimizer.step()就崩溃,阻塞了这个具体拓扑下的续训验证,详见第 7 节也不承诺:
max_steps)发生变化时的正确性——restore_from只负责恢复“状态”,不做全局拓扑一致性校验(DataLoader.load_state_dict里 dp_world_size/batch_size 不匹配会探测并降级,但那是 dataloader 自己的保护,不是全局校验);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_ckptrestore_fromCheckpointerCallbackTrueNoneTrue"LATEST"/ 具体路径False"LATEST"/ 具体路径FalseNonebase.py:720-726直接不 appendCheckpointerCallback,连每步的判断开销都不产生第四种组合是否注册 callback 由
BaseTrainer._init_callbacks在初始化阶段一次性决定,不是在每个 hook 里各自检查一次开关。3.3
global_step↔(epoch, step)换算保存的是单一计数器
global_step;恢复时要把它换算回训练循环用的(start_epoch, start_step)(checkpoint_callback.py:214-227, 380-382):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 进程更可控可复现):阶段二:恢复(
--checkpoint.restore_from=LATEST):start_step=15与保存时的global_step=15一致;LR 从中断处的3.09e-05继续按 cosine 衰减到 0(而不是从 warmup 重新开始),完整验证了 save → 中断 → resume → 续跑完成的整条链路。4. 总体设计
4.1 架构与数据流
设计原则:
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(策略层)on_train_begin_load_checkpoint():解析restore_from→ 视情况给 optimizer “热身” → 经checkpointer.load读取 model/optimizer/extra_state → 写回 trainer 各运行时对象on_step_endsave_steps>0且global_step % save_steps==0_save_checkpoint;去重靠_last_saved_stepon_epoch_endsave_epochs>0且(epoch+1) % save_epochs==0on_step_end存过则_save_checkpoint,否则只打日志跳过(checkpoint_callback.py:137-147)on_train_endsave_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 并从头训练,不报错FileNotFoundError4.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:extra_state→-1被真实值覆盖,正常返回;-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 六个状态桶
modelmodel.state_dict()(PEFT 下只含requires_grad=True的参数,_model_state_dict)model.load_state_dict(strict=not is_peft)optimizersave_optimizer=truestate_dict()initialize_optimizer_state热身,再 DCP 填充 +optimizer.load_state_dictcheckpoint_callback.py:342-353)global_step/epochsave_train_state=truestate.global_step、state.epochtrainer.state,并据此推出start_epoch/start_step(3.3)extra_statelr_schedulerstate_dict()checkpoint_callback.py:392-404)_as_list/_unwrap_single)train_dataloaderdata_iterator.state_dict(),否则train_dataloader.state_dict()trainer.train_dataloader.load_state_dict(...),仅当 dataloader 具备该方法rng_statetorch_cpu(torch.get_rng_state())、torch_device(get_device_rng_state())、python(random.getstate())torch.set_rng_state/set_device_rng_state/random.setstatetorch_device为None,直接跳过4.5 与其它模块的交互
hyper_parallel.core.distributed_checkpoint)save/async_save/load/StandardLoadPlannerDistributedCheckpointer不重复实现存储原语,只加一层_ModelStrictLoadPlanner做 model 严格 / optimizer 宽松的分级校验ChainedOptimizer(hyper_parallel/core/optimizer/optimizer.py)initialize_optimizer_state用getattr(optimizer, "chained_optimizers", None) or [optimizer]鸭子类型兼容.state在ChainedOptimizer上不存在,历史上因此崩过一次(3.4 节),现已修复fully_shard)state_dict(),与并行拓扑正交tp_size>1+dp_shard_size>1会在optimizer.step()就崩(无关 bug,见 2.2)_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()SkipDTensorDispatchinitialize_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_ckptTruecheckpoint_dir"./checkpoints"save_steps0(关闭)save_epochs1is_asyncFalseis_peftFalsesave_optimizerTruesave_train_stateTruesave_extra_state_per_rankTrueextra_state落盘布局,见 4.3restore_fromNoneNone/"LATEST"/ 具体路径restore_optimizerTruerestore_train_stateTruemodel_save_format"safetensors"save_consolidated"final""none"/"final"/"every"),见 2.2staging_dirNonebest_metric_key"default"真实示例(
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, TrainerStatebuild_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_steps/save_epochs/on_train_end触发 + 去重)tests/hyper_models/trainer/test_checkpoint_callback.py)save_ckpt/restore_from正交开关(4 种组合)LATEST解析(指针文件 + 扫描兜底 + 跳过不完整目录)extra_state双布局(per-rank / 内嵌)自动探测ChainedOptimizer兼容)strict_model=False)state_dict/load_state_dict,含有/无DistributedSampler两种路径、拓扑变化探测)tests/components/test_dataloader.py),限单进程场景async_save,多次保存排队等待)is_async=true场景无自动化 STtp=2 × dp_shard=2,见 3.4 节;无自动化 STtp_size>1+dp_shard_size>1续训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-212docstring)。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.pysave_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.pyset_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.pyhyper_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
restore_from="/no/such/dir"FileNotFoundError(checkpoint_callback.py:298)restore_from="LATEST",checkpoint_dir为空.metadatasave_optimizer=false,save_train_state=false存盘后restore_train_state=true去读FileNotFoundError,信息含 "no training state"(dcp_checkpointer.py:499-506)"extra_state"键_UNRESTORED_STEP触发"model.xxx"keyRuntimeError,信息含 "missing model key"(dcp_checkpointer.py:112-115)is_peft=true时同样删掉一个 keysave_ckpt/restore_from四种组合dp_world_size=2,恢复dp_world_size=4B. 存盘触发与去重
save_steps周期存盘save_steps=N,跑 2N+1 步save_epochs周期存盘save_epochs=1,不设save_stepssave_steps倍数on_epoch_end打日志跳过on_train_end补存max_steps提前于任何周期边界结束on_train_end不重复存save_steps=save_epochs=0on_train_end(若save_ckpt=true)补一次_last_saved_step已在恢复时同步,on_train_end不重写(见 3.3)C. 六桶恢复正确性
allcloseexp_avg/exp_avg_sq恢复后与保存前一致且非零ChainedOptimizer恢复global_step/epoch→start_epoch/start_stepsteps_per_epoch=5, global_step=8→start_epoch=1, start_step=3DistributedSampler)dp_world_size>1时跳过已消费 batch,不重不漏dp_world_size==1,私有Generator(seed+epoch)复现同一路 shuffleextra_state双布局往返save_extra_state_per_ranktrue/false 各跑一遍都能正确恢复D. 拓扑与并行组合(需要 NPU 多卡,当前均为转测补齐重点)
tp_size=2, dp_shard_size=2,4 卡,完整走通 save→中断→resume→续训(3.4 节场景固化为 ST)tp_size=1, dp_shard_size=N,对齐examples/training_demo默认 8 卡配置is_async=true)optimizer.step()报错(2.2 已知缺陷),不算本文档验收范围,仅确认失败模式没有变化num_train_epochs>=2,验证是否符合“该 epoch 从头重放”的既有告警语义(7.1),而非默默错位E. 平台矩阵
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.state_dict()["state"]里任意参数的exp_avg是否非零,或直接对比恢复前后是否allclose。lr是否等于中断前最后一步的lr。preset_pt/dummy 索引),记录中断前最后几条样本索引,恢复后确认不重复出现。restore_train_state前后,恢复后同一步的 dropout mask / 初始化噪声是否符合预期(开启时应确定性延续)。extra_state,确认抛出的是RuntimeError/FileNotFoundError,而不是静默训出一条“看起来正常”的 loss 曲线。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正确热身。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),对训练主循环零开销。TextTrainer(组合BaseTrainer而非继承它)与BaseTrainer驱动的其它 trainer 共享同一套on_train_begin/on_step_end/on_epoch_end/on_train_end回调分发,断点续训行为一致。9.3 明确报错(必须 fail-closed)
restore_from指向不存在的具体目录FileNotFoundErrorFileNotFoundError,信息含 "no training state"extra_state实际未写入(哨兵值未被覆盖)FileNotFoundError,信息含 "extra_state"RuntimeError,信息含 "missing model key"以下场景刻意不报错,是设计选择而非疏漏(见 7.5):
restore_from="LATEST"但目录下无任何 checkpointdp_world_size/batch_size)变化state_dict/load_state_dict9.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 回归通过。hyper_parallel.trainer.callbacks.base.CheckpointCallback(旧栈,见 2.2)的测试结果算作本文档功能的转测通过条件。