Checkpointing 和 TrainJob 恢复
一个运行六小时的微调,最终总会遇到节点重启、Pod 被驱逐、抢占,或者在第 5 小时发生 OOM。没有 checkpoint,这些事件中的任何一个都会让你付出整个训练过程的代价。
Kubeflow Trainer v2 没有自己的 checkpoint API。TrainJob 或 TrainingRuntime 规范中没有任何内容会写入、查找或恢复 checkpoint。职责划分如下:
如果任意一列做错,失败模式都是一样的,而且是静默的:Job 重启后找不到任何 checkpoint,然后愉快地从 step 0 重新训练。
本指南介绍其工作机制。对于故意中断训练、把 GPU 还给 InferenceService 的特定场景,请参见 使用 Kueue 的可抢占 TrainJob,它建立在本文内容之上。
目录
前置条件checkpoint 必须包含什么选择能在 Pod 之外存活的存储让 runtime 感知 checkpoint提交一个可恢复的 TrainJob各类中断如何恢复崩溃或节点故障抢占有意暂停各框架的参数开关多节点运行验证它确实能恢复故障排查前置条件
checkpoint 必须包含什么
“恢复”并不等于“重新加载权重”。只保存 model weights 的 checkpoint 会让 optimizer 从零开始——momentum 缓冲区、learning-rate 调度位置以及 step 计数器都会丢失,loss 曲线会在恢复点处出现明显折角。
一个可恢复的 checkpoint 应包含:
- Model weights
- Optimizer state — Adam 的 moment 估计。通常约为模型大小的 ~2 倍(fp32),也是 checkpoint 中最大的一部分。
- LR scheduler state — 这样学习率调度才能继续,而不是重新 warm up。
- Step / epoch counter — 用于从哪里继续。
- RNG state 和 data-loader 位置 — 这样恢复后的运行不会以相同顺序再次看到同一批样本。
HuggingFace Trainer 的 checkpoint-N/ 目录已经包含了以上所有内容。如果你是手写训练循环,只保存 model.state_dict() 是最常见的“恢复”方式失效的原因,它会悄悄破坏训练过程。
选择能在 Pod 之外存活的存储
checkpoint 目录必须位于 PVC 上。emptyDir 会随着 Pod 一起消失,而 hostPath 或本地存储 class 会把 checkpoint 留在某个节点上,而重启后的 Pod 可能永远不会再调度到那个节点。
访问模式只取决于有多少个 trainer Pod 会同时写入:
ReadWriteOnce 的含义是“一次一个节点”,而不是“永远只有一个节点”。一个通过网络挂载的 RWO volume(如 Ceph RBD、EBS 及类似存储)会从旧节点卸载,并在重启后的 Pod 被调度到哪里时重新挂载,因此单节点运行即使恢复到不同节点也能正常继续。只有并发的多 Pod 访问才需要 RWX。
这不适用于 local-storage / topolvm / hostPath,它们在物理上固定在某一个节点上。
应用该 PVC(checkpoint-pvc.yaml):
容量按 checkpoint size x save_total_limit 再加上余量来规划。完整权重 checkpoint 通常比人们预期的大——一个带 optimizer state 的 7B 模型大约每个 checkpoint 80 GiB。LoRA/QLoRA adapter 只有几 MB,因此同一个 PVC 可以保存更多此类 checkpoint。
让 runtime 感知 checkpoint
checkpoint-trainingruntime.yaml 是一个可直接运行的 TrainingRuntime,完整模式已经配置好。关键部分如下:
以及脚本中真正执行恢复的两行:
这种自动检测机制使同一个 manifest 既能用于首次提交,也能用于之后的每次重启。硬编码 resume_from_checkpoint=True 会让第一次运行失败(因为那时还没有 checkpoint);硬编码路径则会把你永远绑定到某一个 checkpoint。
trainer replicatedJob 上的 trainer.kubeflow.org/trainjob-ancestor-step: trainer label 是必需的。没有它,controller 不会识别该 replicatedJob 是 trainer step,而你的 TrainJob 中所有 spec.trainer.* 字段——env、resourcesPerNode、image、command、numNodes——都会被静默丢弃:没有错误、没有事件、没有警告。TrainJob 会改用 runtime 的默认值运行,因此你以为分配了 GPU 和 CKPT_DIR 的 Job,实际上两者都没有。
请验证这些覆盖项是否真正生效,而不要想当然地相信它们:
应用该 runtime:
提交一个可恢复的 TrainJob
trainjob-resume.yaml 通过 podTemplateOverrides 挂载 checkpoint PVC,因此一个共享 runtime 可以服务多个运行,而每个运行都把 checkpoint 写到不同位置:
podTemplateOverrides 可以添加 volumes、volumeMounts、nodeSelector、tolerations 和 affinity,但它不能为 trainer 或 initializer 容器设置 env——validating webhook 会直接拒绝该 TrainJob:
环境变量应放在 spec.trainer.env 中。
各类中断如何恢复
崩溃或节点故障
有两层相互独立的机制会重启失败的 trainer,而具备 checkpoint 感知能力的脚本在这两种情况下都能正确恢复:
无论哪种情况,新进程都会执行相同的自动检测,并接着使用最新的 checkpoint-N。观察 JobSet 上的 restarts 递增:
一旦 maxRestarts 用尽,TrainJob 就会变为 Failed——因此它的值要足够高,以吸收临时性的节点波动,但也不要高到让真正损坏的脚本无限循环。
抢占
Kueue 会驱逐 Workload,而 Trainer v2 会在 quota 重新释放后重建 JobSet;trainer 会从 PVC 恢复。cohort/quota 的配置本身是一个独立话题——请参见 使用 Kueue 的可抢占 TrainJob。
有意暂停
spec.suspend 可以在不丢失训练状态的情况下暂停运行——例如临时把 GPU 借给更紧急的任务一小时:
恢复后的 Pod 是一个新的 Pod,具有新的名称,并且可能被调度到不同的节点。它的第一行日志应该是你的 [checkpoint] resuming from ...。
当一个被暂停的 TrainJob 恢复时,activeDeadlineSeconds 会重新开始计时——也就是说,每次暂停后,Job 都会重新获得完整的截止时间,而不是从创建时刻开始累计。
各框架的参数开关
所有基于 HuggingFace-Trainer 的栈都暴露相同的三个开关,因此上面的模式可以原样迁移:
根据你能够容忍的中断时间来选择 save_steps:风险最大的工作量最多就是 save_steps x seconds_per_step。如果每步 5 s,save_steps: 100 的风险大约是 8 分钟。保存过于频繁本身也有成本——完整权重 checkpoint 可能会让 GPU 停顿几十秒,因此在稳定集群上可以把值设大一些,在可抢占队列上则应设小一些。务必同时设置 save_total_limit,否则 PVC 会无限增长。
多节点运行
当 numNodes: > 1 时,PVC 必须是 RWX,并且会有两点变化:
- 对于标准的数据并行训练,只有 rank 0 应写入 checkpoint,否则每个 rank 都会争抢写同一路径。HF Trainer 已经处理了这一点;手写循环必须用
if rank == 0进行保护。 - 所有 rank 必须从同一个 checkpoint 恢复。 如果不同 rank 恢复到不同 checkpoint,会静默分叉,而不是崩溃。
对于分片策略(FSDP、DeepSpeed ZeRO),每个 rank 持有一部分 weights 和 optimizer state,把所有内容汇总到 rank 0 往往对大模型来说根本不可行。这正是 torch.distributed.checkpoint 的用途:每个 rank 并行写入自己的 shard,结果可以在不同的 world size 下重新加载——因此一个在 8 张 GPU 上 checkpoint 的运行可以恢复到 4 张 GPU 上。保存为单个聚合文件的 checkpoint 则无法以这种方式重新分片。
验证它确实能恢复
不要等到真正的故障发生后才发现问题。每个新的 runtime 都先执行一次这个演练:
如果第 3 步打印的是 no prior checkpoint, starting fresh,说明恢复有问题——真实抢占发生时也一样会有问题。
通过 grep trainer 日志中的 HF Trainer “Saving model checkpoint to” 这一行并不可靠:tqdm 进度条会覆盖它。应当检查 PVC 上是否存在 checkpoint-* 目录,上面的日志命令正是在做这件事。