Skip to content

大模型分布式训练故障恢复策略与实践

内容整理自学习笔记,仅供面试备考参考;不构成录用、培训或考试承诺。

1. 为什么需要故障恢复机制?

大规模分布式训练场景中,集群规模庞大,芯片设备、主机、网络等组件均会不定期出现故障。若想在故障后继续训练,需要从最近保存的 checkpoint(ckpt)恢复并 resume training。

故障带来的时间开销不可避免,但可以尽量降低影响。

2. 如何获取最优的 Checkpoint 存储间隔?

假设均匀同步存储 ckpt,故障随机发生在 ckpt 间隔区间内,集群时间损失可定义如下:

集群时间损失 = ckpt 存储耗时 + 故障期望次数 × 恢复训练耗时 × (ckpt interval / 2 + 恢复训练耗时)

通过对导数求零,可以根据集群环境得到最优的 ckpt interval。需要注意的是,ckpt interval 通常远大于 1。

3. Checkpoint 能否实现异步或部分掩盖?

3.1 完全异步的困境

异步存储 ckpt 的最大问题是设备内存踩踏

  • 如果在另一个 stream 中做 D2H(Device to Host)数据拷贝,同时模型训练继续运行
  • 所有参数的 D2H 操作尚未完成时,下一个 step 已经开始更新参数或优化器状态
  • 后续未完成的拷贝操作会获取错误的数据

由于 ckpt 存储时间不可控(无法确定是否小于下一个 step 的执行时间),内存踩踏问题不可避免。完全异步的方案不可行。

3.2 部分掩盖的可行方案

方案一(训练脚本侧): 在下一次更新参数或优化器状态之前,强制等待 ckpt 存储完成,尽可能 overlap 计算和 IO。

方案二(框架侧): 类似 CUDA 在 H2D non-blocking 操作后,有数据依赖时强制加 sync point——框架侧可以在 D2H 拷贝后,有数据写操作时强制添加 sync point。

4. 断点续训 / 临终遗言是否可行?

4.1 可行性分析

绝对可行,但有一定条件限制。

大模型训练通常是 DP/TP/PP 多维并行的场景,任意节点都可能故障。关键在于整网参数的完整性

场景是否可临终存储说明
每个 PP stage 内 TP Group 完整可以框架侧捕获分布式 error,做临终参数存储
有 PP stage 内 TP Group 不完整不可以整网参数不完整,无法保存
故障发生在参数/优化器状态更新时不可以参数一致性无法保证

实践经验: 框架侧开发不会很难,需要结合 rank 编排做定制研发。

4.2 最佳实践

基于训练框架对深度学习框架做深度定制,是实现断点续训/临终遗言的最佳路径。

5. 故障恢复策略总结

策略优点缺点
同步 ckpt简单可靠存储时训练阻塞
异步 ckpt不阻塞训练内存踩踏风险,不可行
部分掩盖减少 IO 等待需要框架/脚本侧定制
临终遗言ckpt interval 趋近于 0依赖整网参数完整性