千卡集群训大模型,故障比代码先崩?聊聊分布式训练容错那点事

训过千卡集群的人都懂一个扎心事实:硬件故障不是"会不会发生",而是"多久发生一次"

Meta 用千卡集群训 OPT-175B,两个月崩了 105 次,最长一次稳定训练只撑了 2.8 天。训 Llama3-70B 的万卡集群更离谱——419 次中断,GPU 故障占了 58.7%。

一组节点挂掉,整个训练任务直接归零。几周的电费、算力、人力,说没就没。

所以今天聊聊:分布式训练容错到底怎么搞,才能让训练"崩了也能接着跑"


一、大模型训练都会遇到哪些故障

先盘点下常见"犯罪嫌疑人":

故障类型典型场景后果
硬件故障GPU卡死、节点掉电、电源异常训练中断、数据丢失
网络故障IB链路抖动、通信阻塞同步失败、性能骤降
软件异常框架崩溃、内存泄漏、进程挂死Rank失联、作业终止
数据故障数据分片丢失、数据不一致训练结果出错

传统方案就两板斧:

方案一:全量重启(Full Restart)。任意节点挂了,整个集群从头再来——从最后一个 checkpoint 重新加载。恢复时间长到离谱,算力全浪费。

方案二:定期 Checkpoint。每 N 个 iteration 存一次模型状态。问题是:Checkpoint 间隔内的训练进度全部丢失;大模型存一次盘本身就很慢,磁盘 I/O 还拖累训练吞吐。

这两个方案在千卡规模下基本是"能用,但憋屈"。


二、智韧容错:四大模块的联动体系

要真正解决"训练不中断",光靠单点优化不够。整个体系可以拆成四大模块:

故障处理

故障发现与隔离

根因分析

自动替换故障节点

局部重启恢复

计算容错

内存备份与恢复

异步持久化

远端内存读写

Checkpoint 秒级恢复

实时监控

任务日志聚合

故障告警推送

性能指标追踪

硬件状态监控

节点管理

节点健康检测

节点池维护

异常节点隔离

一句话总结:训练前筛掉坏节点,训练中盯着异常,出故障秒级恢复,修好后自动归队


三、节点检测:训练前先把"雷"排了

分布式训练最怕的是跑到一半某个节点崩了。所以训练前必须做一轮全面体检:

  1. NHC 检测:检查 CPU/GPU 状态、利用率、内存占用、残留进程。筛出明显异常的节点。
  2. Rccl 通信检测:用 AllReduce、AllGather 等通信模式跑模拟测试,验证节点间互联性能和通信稳定性。
  3. GEMM 计算检测:跑矩阵乘算子测试,确保所有计算节点性能一致,没"瘸腿"的。
  4. 分布式训练模拟:多机多卡跑一轮模拟训练,端到端验证。

检测完产出五份节点清单:正常节点、异常节点、备用节点、运行中节点、全量节点。后续训练只从正常/备用池里取节点。

通过

失败

通过

失败

通过

失败

通过

失败

全量节点

NHC检测

Rccl通信检测

异常节点池

GEMM计算检测

分布式训练模拟

正常节点池

训练就绪


四、Checkpoint 异步化:hyckpt 的设计思路

这是整个容错体系里最核心的一环。

4.1 痛点在哪

Megatron 框架原生 Checkpoint 流程是这样的:训练进程把模型数据从 GPU 显存拷出来,写磁盘,写完再继续训练。大模型动辄几百 GB 的 Checkpoint,磁盘 I/O 把训练卡住几十秒甚至几分钟。千卡集群里每来一次 Checkpoint,DCU 利用率就掉一截。

4.2 hyckpt 怎么解决

核心思路就三个字:异步化

hyckptd 服务端

训练进程

拷贝

写入共享内存后立即返回

异步

GPU 显存

CPU 内存/共享内存

训练继续

从共享内存异步读取

数据写入磁盘集群

多节点内存协同备份

具体流程:

  1. 训练进程保存 Checkpoint 时,先把数据从 GPU 显存拷到 CPU 内存/共享内存区域。
  2. 拷完立刻返回继续训练——不等待磁盘写入
  3. hyckptd 后台进程异步从共享内存读数据,慢慢写到磁盘集群。

4.3 多节点内存协同备份

单节点内存备份的问题是:如果这台节点直接挂了,内存里的 Checkpoint 数据也没了。

所以做了跨节点内存协同备份——每个节点的 Checkpoint 数据同时写到本节点内存和远端节点内存。这样任何单节点故障,数据都在别的节点内存里有备份。

节点4

节点3

节点2

节点1

互相备份

互相备份

内存/Checkpoint

内存/Checkpoint

内存/Checkpoint

内存/Checkpoint

4.4 实测效果

在 llama-70b 模型上跑过对比,仅统计 PyTorch 接口耗时:

环境PyTorch 原生 loadhyckpt load
dp0 保存恢复35s ~ 75s0.3s ~ 11s

恢复时间从分钟级压到秒级,差了不止一个数量级。

在实际 256 卡集群上跑 llama3_70b 的完整恢复流程,优化前从故障到恢复进 Iter 需要约 15 分 39 秒,优化后压到 4 分 42 秒。其中 load-ckpt 阶段从 122 秒优化到 36 秒。


五、状态转换机制:训练故障的自动闭环

容错系统通过状态机驱动整个故障恢复流程:

正常

检测到故障

Init: 初始化节点和计算进程

Resolve_failure: 检测节点状态,补充备用节点

Start_training: 启动训练进程

Check_process: 核对实际运行进程

Monitor: 监控DCU状态

Stop_training: 终止训练

Clean_env: 清理残留进程

这套状态机实时感知 Megatron 训练进程的状态变化。一旦检测到进程异常退出或通信超时,自动走"停止→清理→替换节点→重启训练"的闭环,全程无需人工介入。


六、掉队检测:别让一颗老鼠屎坏了一锅粥

千卡集群里最容易忽视的问题是"掉队节点"——某个 GPU 没挂,但算得比别的 GPU 慢,拖慢整个同步步调。

检测逻辑

  1. 周期性收集每个 GPU 的 RTT(每个训练 iteration 完成时间)和估算吞吐量。
  2. 计算所有 GPU 的 RTT 平均值。
  3. 找出低于平均值 2 倍标准差的节点——这些就是掉队的。
  4. 自动触发告警。

同时监控损失曲线和吞吐量曲线,损失突然跳变或吞吐量异常下降时,结合日志数据快速定位问题节点。

训练吞吐量趋势示意12345678910训练 Iteration500450400350300250200150100500吞吐量 (TFLOPS)

七、产品对比:同类方案横向看

功能模块DLRoverFt_launcher (NVIDIA)国产方案
技术栈Torchrun + K8sTorchrun + Slurm进程管理 + mpirun
异步 Checkpoint支持支持支持
掉队检测不支持支持支持
节点池管理支持不支持支持
弹性扩缩容支持不支持不支持
局部重启不支持支持支持
故障定位不支持支持支持
节点预筛查支持不支持支持
计算错误检测不支持不支持不支持

几个观察:

  • DLRover 功能全面但依赖 K8s,接入复杂且在大规模场景下稳定性存疑。
  • Ft_launcher 聚焦故障检测和局部重启,但缺少节点管理能力,多节点运行还得靠 Slurm/mpirun。
  • 国产方案 在节点池管理、节点预筛查、轻量部署上有明显优势,更适配大规模国产加速卡集群的实际场景。

八、关键数据汇总

把全文最核心的性能指标拉一张表:

指标数据
单节点 llama3_8b 恢复耗时(内存)~1.25 分钟
单节点 llama3_8b 恢复耗时(磁盘)~1.25 分钟
4 节点 llama3_8b 恢复耗时~1.0 ~ 1.4 分钟
8 节点 llama3_70b 恢复耗时~1.6 ~ 1.9 分钟
32 节点 llama3_70b 恢复耗时~1.6 分钟
64 节点 llama3_70b 恢复耗时~2 分钟
hyckpt load vs PyTorch 原生 load0.3s~11s vs 35s~75s
256 卡完整故障恢复(优化后)~4 分 42 秒

可以看到:节点数从 8 扩到 64,恢复时间几乎没有线性增长——这说明内存协同备份 + 异步持久化的架构在扩展性上表现不错。


九、写在最后

分布式训练容错不是什么花活,是千卡以上规模训练必须过的坎

核心思想其实就三条:

  1. 训前排雷:NHC + Rccl + GEMM 三重检测,别让坏节点混进训练池。
  2. 训中兜底:hyckpt 异步 Checkpoint + 跨节点内存协同备份,故障恢复从分钟级压到秒级。
  3. 自动闭环:状态机自动检测→隔离→替换→重启,全程不需要人盯着。

目前这套方案在国产加速卡集群上已经跑过 64 节点 llama3_70b 的实测验证。后续方向包括适配更多训练框架(不仅限于 Megatron),以及补充计算错误检测等能力。


参考
[1] Meta, “OPT: Open Pre-trained Transformer Language Models”, 2022.
[2] Meta, “The Llama 3 Herd of Models”, 2024.

更多推荐