分布式训练如何保存、重分片并精确恢复 Checkpoint?
先用自己的话答,再看参考说法
60 秒是练习上限,不是必须凑满。原理题讲清因果,设计题讲清约束,项目题只讲真实证据。
别照着背。参考说法只用于对照;直接回答问题,简单题说清楚就收尾。项目题只用自己的经历和数字。
参考内容当前已显示;开始口述后会暂时隐藏。
面试时怎么答
先定义“精确恢复”包含什么:模型权重、优化器、学习率步数、Scaler、各类 RNG 和数据游标都不能缺。多 Rank 保存时要先写临时分片,校验完成后原子发布 Manifest。面试官若问换 World Size,说明需要可重分片格式;若问是否真的恢复正确,用强杀演练检查 loss 连续和样本不重不漏。
可以这样答:
分布式 Checkpoint 不能只保存模型权重,还要保存优化器状态、学习率步数、混合精度 Scaler、随机数状态和数据加载位置。各 Rank 先写临时分片,全部校验成功后再原子发布 Manifest,避免读取半份快照。若恢复时 World Size 改变,格式还要支持重分片。可靠性应通过故障注入验证恢复前后 loss 连续、样本不重复也不跳过。
核心回答
大模型 Checkpoint 应让各 Rank 并行保存本地参数/Optimizer 分片,避免先在单 Rank 聚合完整状态造成显存、主存和网络峰值。Checkpoint 由数据分片与全局元数据共同组成,只有全部必需对象写完并通过校验后,才原子发布一个完成标记或 Manifest;恢复端只加载已提交版本。
精确续训还要保存 LR Scheduler、Grad Scaler、全局 Step、RNG、数据游标和训练配置。若恢复到不同 World Size,加载器需要按全局张量布局读取旧分片并重新分发,而不是假设“旧 Rank 文件对应新 Rank”。
展开说明
- 格式:逻辑 State Dict 应尽量与物理分片解耦,元数据描述每个张量的全局形状、Dtype、偏移和分片位置。
- 提交:先写临时目录/唯一前缀,校验所有分片,再发布 Manifest;对象存储上不能依赖目录 Rename 等同于本地原子操作。
- 兼容:参数重命名、模型结构、Optimizer 类型或并行布局变化需要显式迁移,不能静默跳过不匹配状态。
- 保留:按最近 N 个、里程碑与最佳指标分层保留;删除前确认没有运行正在引用。
RNG 与数据游标决定“接下来看到什么”,Optimizer State 决定“接下来怎样更新”。只恢复模型权重是 Warm Start,不是精确 Resume。
版本边界:PyTorch Distributed Checkpoint API、State Dict 规范和 Planner 能力会演进;跨框架或跨大版本恢复前应锁定格式版本,并用小规模迁移演练验证。
工程实践
定期做恢复演练:在已知 Step 强制终止,换相同和不同 World Size 恢复,核对首个 Batch ID、LR、RNG 抽样、Optimizer 统计和后续 Loss。监控保存耗时、训练暂停、总字节、失败分片、校验失败与恢复 RTO,并把 Checkpoint 写入与训练进度日志关联。
常见追问
- 为什么不能只让 Rank 0 保存完整模型? 超大模型聚合可能让 Rank 0 OOM,并把网络和存储带宽串行化;并行分片保存能扩展,但必须有全局元数据和一致提交。
- 如何避免加载到“半个 Checkpoint”? 数据先写到未发布版本,所有分片成功并校验后再写完成 Manifest;恢复端只枚举有有效 Manifest 的版本。
- 换 World Size 后怎样恢复? 按全局张量坐标读取旧 Shard,再根据新布局重分片;同时重建通信组和数据 Sampler,不能复用旧 Rank 编号映射。
一句话复习
分布式 Checkpoint 要把逻辑状态与物理分片解耦,用一致提交保证完整性,并保存足够状态让新 World Size 可重分片恢复。
参考资料
评论与补充
评论会直接显示在这道题下面。可以写自己的答法、继续追问或指出错误,不需要 GitHub 账号,也不会跳转到 Issue;内容会公开,请勿填写个人隐私、公司机密或受保密约束的材料。
正在连接站内评论服务…
正在加载评论…
还没有评论,你可以先写下自己的理解或追问。