训练断点恢复需要保存哪些状态?
简化版
训练断点不仅要保存模型权重,还要保存优化器、学习率调度器、AMP scaler、随机数状态和当前 step;否则只能加载参数,不能严格续训。分布式训练还需保证每个 rank 的数据游标与分片状态一致。
详细版
-
权重决定当前位置,优化器动量决定下一步方向,两者缺一都会改变轨迹。
-
定期断点、最佳指标权重和仅推理权重用途不同,应分开命名。
-
保存应先写临时文件再原子重命名,避免进程中断留下半个文件。
-
恢复后要核对 global step、学习率、数据顺序和首个 batch loss。
-
跨代码版本恢复需记录配置、代码提交和依赖版本。
完整版教学
一、断点是训练状态机的完整快照
Adam 的一阶、二阶矩会影响同一梯度对应的更新量,只恢复参数等价于突然换了一个新优化器。
调度器若从第 0 步重新开始,warmup 或衰减也会重复。
严格复现还依赖随机数与数据采样器状态。
只保存全局 seed 不够,因为训练途中各随机生成器已经消费了不同数量的随机数。
二、底层机制与公式
checkpoint = {model, optimizer, scheduler, scaler, epoch, global_step, RNG, sampler}
resume step: load all states -> set data cursor -> continue at global_step+1
三、带数字的推演
在第 20,000 步保存,学习率已从 1e-3 衰减到 2e-4。
若只载入权重,调度器重启后学习率回到 warmup 的 1e-5,后续曲线显然不是原训练的延续。
四、方案对比
| 方案/对象 | 核心特点 | 代价或边界 |
|---|---|---|
| 仅权重 | 迁移或推理 | 不能严格续训 |
| 完整断点 | 故障恢复 | 体积大且与代码结构耦合 |
| 最佳模型 | 按验证指标选取 | 未必包含最新进度 |
五、执行流程
训练进程 -> 写 tmp 完整状态 -> fsync/校验 -> 原子 rename
恢复进程 -> 校验配置 -> 加载各状态 -> 跳过已消费数据 -> 单步对账
六、边界条件与工程代价
多机 world size 改变时,优化器分片和 sampler 状态可能无法直接恢复,需要框架支持的重分片方案。
保存间隔是恢复时间与 I/O 成本的权衡;每 100 步保存很安全,却可能让共享存储成为训练瓶颈。
记忆钩子:记忆链是“参数定位置、动量定方向、调度定步幅、随机态定数据”。
七、常见误区与追问
-
误区:保存 model.state_dict 就能无缝续训。 优化器动量、调度器和随机状态都会丢失。
-
追问:为什么还要保存 AMP scaler? 动态 loss scale 决定溢出检测和是否跳过更新。
-
误区:设置相同 seed 就能恢复相同数据顺序。 还需恢复生成器已前进到的状态与 sampler 游标。
-
追问:如何防止断点文件损坏? 临时写入、校验完成后原子替换,并保留上一版本。
-
追问:恢复成功如何验收? 比较 step、LR、首批数据 ID、loss 与连续训练基线。
八、加强记忆
记忆链是“参数定位置、动量定方向、调度定步幅、随机态定数据”。
断点要能让下一步与未中断训练对上,而不只是让程序成功启动。