上周我们集群又上演了一次"心跳骤停":一个跑满14天的千卡训练任务,因为一张H800的NVLink报错中断。定位用了半小时,恢复数据用了两小时,丢掉了整整三个检查点周期的训练进度。这种事情在万卡级训练里不是黑天鹅,而是家常便饭。断点续训和故障恢复做得好不好,直接决定了你号称"十万卡集群"是拿来给甲方看的PPT,还是真能跑出稳定训练曲线的基础设施。
这篇文章想把这件事讲透。我会从万卡集群故障的数学期望讲起,拆解检查点机制从"同步全量落盘"到"异步分层快照"的演进逻辑,梳理故障检测、坏节点隔离、热备接管和弹性恢复的完整链路,最后聊聊我们在恢复一致性上踩过的坑和故障演练的心得。内容面向分布式训练工程师、平台开发者和大模型训练运维,目标是让你看完之后能自己设计一套断点续训方案,而不是只会调一个torch.save。
1. 万卡级训练的真实故障图景与可靠性账本
1.1 一张卡挂掉很简单,难的是集群里"必然"有人挂
很多人第一次接触万卡集群时,第一反应是"多挂几张卡,跑得飞快"。但真正落地之后你会发现,万卡集群的核心矛盾不是算力,而是可靠性——或者说,是怎么在不可靠的硬件海洋里,让训练任务像一条船一样不沉没。
单张GPU的年度故障率在成熟数据中心里大约在1%到3%之间,听起来很低对吧?但把这个数字乘以一万张卡,每天24小时运行,结果就很不一样了。我做了一个简单的估算:假设单卡年故障率2%,折算到每天约为0.0055%。看似微不足道,但一万张卡叠加,任务连续运行7天,至少发生一次硬件故障的概率大约是1减去"所有卡都健康"的概率,算下来接近七成。如果任务要连续跑30天,这个概率几乎等于1。
更别提故障不只有显卡本身。我看过我们平台过去一年的故障工单统计,GPU Xid错误、显存ECC报错、NVLink链路退化、节点宕机、IB交换机端口飘移、存储写满、文件系统卡IO……种类多到让人怀疑机房下面是不是埋了个倒霉蛋。真正的工程挑战不是"某张卡坏了怎么办",而是"在每一轮训练周期内,几乎必然有卡会坏,你要怎么让训练任务视而不见"。
1.2 故障类型分布与MTBF估算
为了把问题说得更具体,我把常见故障归了归类。下面的表格来自我们内部平台的统计,不同集群可能有差异,但大方向一致:
| 故障类别 | 典型表现 | 占比(平台统计) | 检测手段 |
|---|---|---|---|
| GPU硬件故障 | Xid报错、显存ECC错误、NVLink降速 | 35%左右 | 设备驱动日志、健康探针 |
| 节点级故障 | 内核panic、掉电、内存错误、CPU过热 | 25%左右 | 心跳超时、带外管理 |
| 网络故障 | IB链路闪断、丢包率上升、拥塞 | 20%左右 | NCCL超时、链路监控 |
| 存储故障 | 写满、IO hang、文件损坏 | 15%左右 | 检查点落盘超时 |
| 软件/框架问题 | 分布式锁失控、集合通信卡死、OOM | 5%左右 | 训练心跳停滞 |
算完MTBF(平均故障间隔)之后,你就明白为什么行业里会有"万卡训练必修课"这种说法。假设单卡MTBF按比较乐观的5年算,一万张卡的集群MTBF就是5年除以10000,折合约4.4小时。也就是说,这个集群平均每四个多小时就会挂掉一张卡。你一个训练任务跑一天,中途不遇到任何硬件问题的概率简直可以忽略不计。
所以,我们在设计训练平台时一直跟团队强调一个观点:断点续训不是训练框架的附加功能,而是基础设施的生存底线。没有它,万卡训练就是在赌运气。
1.3 断点续训的本质是什么
断点续训本质上解决的是三个问题:第一,把"训练到第几步"的状态可靠地保存下来;第二,在故障发生后快速找回状态并继续训练;第三,保证恢复出来的训练过程和中断前保持基本一致,不会出现数据重复、状态错乱、收敛漂移。
这里最核心的概念是"训练状态快照"。它不只是模型权重,还包括优化器状态、学习率调度器、混合精度缩放、RNG随机数生成器状态、数据加载器的读取位置,甚至集合通信库的一些内部状态。保存的东西越完整,恢复出来的训练越"无缝"。
但"完整"是有代价的。一个万亿参数的模型,光模型权重用BF16保存就是2TB,加上Adam优化器的动量和方差(FP32各一份),轻松超过8TB。这种体量的快照,如果还用传统的全量同步落盘方式,训练会被存储写带宽卡死。我们后面章节会详细拆解怎么用分层、增量、异步的思路来解决。
2. 训练状态快照里到底有什么:从模型参数到数据指针
2.1 模型参数和优化器状态是主角
先看主角。模型参数是神经网络的权重,这部分数据一定不能丢,丢了就得重训。优化器状态同样重要——Adam优化器会为每个参数维护一阶动量(m)和二阶动量(v),也就是梯度均值和梯度平方均值,它们决定了下一步更新方向和步长。如果模型参数恢复了但优化器状态丢了,训练虽然能跑,但优化器相当于失忆了,量级和自适应率都要重新积累,收敛速度和稳定性会受到明显影响。
以一个700B稠密模型为例,换算成存储量:
| 状态类型 | 精度与格式 | 每参数占比 | 700B模型总大小 |
|---|---|---|---|
| 模型权重(主副本) | BF16,2字节 | 2B | 1.4TB |
| 模型权重(FP32主副本,用于更新) | FP32,4字节 | 4B | 2.8TB |
| Adam一阶动量 | FP32,4字节 | 4B | 2.8TB |
| Adam二阶动量 | FP32,4字节 | 4B | 2.8TB |
| 合计 | — | — | 约9.8TB |
这里还没算上梯度版本和通信缓冲。所以你会发现,在大模型训练里,优化器状态往往比模型参数还要占空间。这也是为什么DeepSpeed ZeRO和Megatron的分布式优化器都把优化器状态切到各卡上,而不是每卡存一份全量副本。
2.2 容易被忽略的"隐性状态"
除了模型和优化器,还有四类状态平时不起眼,恢复的时候缺了就会出幺蛾子。
第一是数据加载器的读取进度。你训练到第10000步,对应的是第4个epoch的第178个batch(全局采样顺序)。如果恢复时不清不楚地从第0步开始重放数据,模型会看到大量已经见过的数据,训练曲线出现诡异的突变。分布式场景下,每个rank的数据切片位置还要和全局采样策略对上,否则还可能漏掉一部分数据。
第二是随机数生成器状态。Dropout、数据增强、随机遮挡这些操作都依赖RNG。如果你的训练进程从检查点恢复时RNG没有同步,那么同一份数据的增强方式可能和中断前不一致。多数情况下这一点不会致命,但对结果一致性要求严格的实验会很困扰。
第三是学习率调度器和EMA。很多训练脚本把LR按step衰减,恢复后LR如果从初始值重新开始,等于前功尽弃。EMA影子权重如果没保存,中途恢复会导致后续评估指标一直对不齐。
第四是混合精度训练里的Loss Scaler。AMP(自动混合精度)训练中,动态损失缩放因子会根据梯度溢出情况自动调整。这个因子丢失后重新初始化,初期会频繁出现梯度溢出跳过更新,收敛曲线会留下痕迹。
2.3 检查点不是一个大文件:分片存储与全局视图
早期大家写PyTorch训练脚本,常用的是torch.save(model.state_dict(), "ckpt.pt"),把整个状态字典集中写到一个文件里。这在单卡、小模型时代没问题,但到了万卡时代,集中式保存有两个致命问题。
第一是容量和带宽。9.8TB的检查点如果集中写到某个节点,单机NVMe写带宽撑死几个GB/s,写一次要一小时,训练早停麻了。第二是单点问题。检查点文件如果只放在一个地方,存储节点挂掉,整个训练进度全部报销。
所以现在主流的分布式检查点方案是分片存储 + 全局元数据。每个rank只保存自己负责的那部分模型分片和优化器分片,写入各自关联的存储位置。元数据文件(也就是"plan")记录所有分片的组织方式、版本号、状态摘要、各分片的存放位置。恢复时先读元数据,再并行加载各个分片,最后通过全局视图重建训练状态。
PyTorch的torch.distributed.checkpoint就是按这个思想设计的。它不再要求每个rank保存完整的state_dict,而是自动把状态按张量切分规则保存到各rank的storage中。写代码的形式大致是:
from torch.distributed.checkpoint import FilePlanner, save, load from torch.distributed.checkpoint.default_planner import DefaultSavePlanner state_dict = { "model": model.state_dict(), "optimizer": optimizer.state_dict(), "rng_state": rng_state, "dataloader_state": dataloader_state, } save( state_dict=state_dict, storage_writer=FileStorageWriter("/ckpt/latest"), planner=DefaultSavePlanner(), )核心思想是:你的代码不关心分片细节,框架负责把大张量切成每个rank的独立文件。但要真用于生产,还得自己处理版本管理、原子提交换名、损坏校验等环节。日常跑通demo很简单,推向万卡稳定运行是另一回事。
3. 检查点机制的关键演化:从同步阻塞到异步分层
3.1 第一代方案的痛点:训练停等存储
我最早做分布式训练时,大家保存检查点的方式非常简单粗暴:每个epoch结束,调用一次torch.save,训练进程阻塞在那里,等所有参数序列化并写完才继续训练。那时候模型几百MB,写一次几秒钟,没人觉得这是个问题。
但把时间线拉到今天,这个方案在大规模训练里完全不可用。假设你的训练集群每5秒钟产生一个8TB的全量状态,如果你每隔N步同步保存一次,意味着训练要停下来等存储写完8TB数据。以并行文件系统20GB/s的写带宽为例,光写就400秒,训练损失一大截。如果你为了减少保存次数把保存间隔拉长,故障恢复时丢失的训练进度又会变大。这个矛盾在万卡时代被放大到极致。
用一张表来对比三种主流策略:
| 策略 | 保存频率 | 训练停滞时间 | 恢复丢失进度 | 典型适用 |
|---|---|---|---|---|
| 同步全量 | 每N步阻塞保存 | 高(分钟级) | 可控但训练浪费 | 小模型/单机 |
| 异步全量 | 每N步后台保存 | 低 | 可控 | 中型集群 |
| 异步分层+增量 | 高频快照+低频全量 | 极低 | 极低 | 万卡大模型 |
3.2 异步落盘与内存双缓冲
异步检查点解决的是"阻塞训练"的问题。核心思路是:训练进程把状态快照拷贝到一块独立的CPU内存缓冲区,然后马上继续训练;后台线程负责把缓冲区里的数据编码、压缩、写到存储。
这里有一个很关键的设计细节——双缓冲。我们平台是用两块CPU内存轮流当缓冲区。训练进程写第0块缓冲区时,后台线程正在把第1块缓冲区刷到存储;下一秒交换角色。这样训练进程只需要等待一次memcpy的时间,通常几秒到几十秒,而不是等到落盘完成。对训练曲线来说,这种停顿几乎可以忽略。
但要小心,异步写盘也有自己的坑。如果训练进程在缓冲区还没写完时就崩溃,这段时间的检查点数据就丢了。所以严格来说,异步检查点要让"训练第N步"和"检查点第N步"之间相差一个缓冲周期。你在恢复时接受的进度,最多落后一个保存周期,这个差距就是RPO(恢复点目标)。
3.3 增量快照与多层分片:大模型时代的核心思路
全量检查点再怎么做异步,8TB这个量级的落盘始终有物理极限。既然增量数据量和训练步长非跟模型规模强相关,那能不能只保存"变化的部分"?
这个方向在大型语言模型训练里特别有意义。大家观察到一个特性:Transformer模型里,不同层参数的变化速率差异非常大。Embedding层和最后的输出层梯度变化剧烈,中间层相对平缓。如果对所有层用同一频率保存检查点,既浪费带宽,又无法做到高频保护。
于是有了分层快照的思路。简单说,把模型的参数分成若干个参数组,每组有自己的保存频率。例如:
- 高频组:Embedding、最后的FFN层、LayerNorm参数,每500步保存一次增量。
- 中频组:靠近输出的Transformer block层,每1000步保存。
- 低频组:靠近输入的层,每2000步保存一次。
- 全量基线:每10000步或每2小时保存一次完整状态,作为恢复的锚点。
恢复时,先加载最近的基线全量,再叠加各组最近一次增量。这个思路有点像数据库的"全量备份+binlog增量恢复"。字节的MegaScale、各大厂的万卡训练平台都有类似实现,落地效果是:在保持恢复精度不丢的前提下,把检查点写入量降了一个数量级。
不过增量方案对元数据管理要求很高。每一层需要有对应的版本号、时间戳、校验信息,恢复时才能知道"这一层用哪个增量片段叠加到哪一版全量上"。谁负责维护这些一致性?答案是元数据服务器。我们用的是把检查点管理信息存到高可用KV存储里,每次保存时更新一个全局的"状态目录"。
4. 故障发生后:检测、隔离与重新编排的三段式恢复链
4.1 心跳、健康探针与超时判定
断点续训不只是保存和恢复,更关键的是有一套可靠的"故障发现-故障隔离-重新编排"链路,让训练任务在坏节点被踢掉之后还能继续跑。先讲检测。
每个节点上跑一个健康探针Agent,作用有三个:一是定期上报GPU健康状态(读取Xid错误、显存ECC计数、NVLink链路速度);二是维护训练进程的心跳;三是响应控制面的探活请求。控制面如果连续多次心跳超时,就会把这个节点标记为可疑状态,触发进一步的诊断。
GPU故障检测比较特殊。很多时候GPU不是直接崩掉,而是先出现Xid错误或者ECC纠错频次上升,性能开始劣化,但进程还没死。这种"亚健康"状态最坑人。你如果不理它,训练到后面会因为NVLink重传率上升而整体卡顿;如果你立刻把它踢了,又显得过于敏感。我们的经验是设置一个"前置通知"机制:设备Agent发现GPU健康指标异常时,先上报事件并预判,等故障真正影响通信时再把节点踢出训练集群。
4.2 坏节点剔除与热备接管
故障确认后,控制面要做的事是安全剔除,而不是直接杀掉整个任务。如果直接kill,集合通信库(比如NCCL)里面还没有处理故障节点的逻辑,其余rank会一直等一个永远不会回来的rank,形成Hang住的状态,训练进程变成僵尸。
正确的顺序是:先让训练框架进入"故障处理模式",暂停正常迭代;然后触发一次保存——有些平台支持直接从内存状态做一次快速快照,不用等完整落盘;保存结束后,把包含故障节点的rank集合从通信组里摘除;最后重新初始化通信组,继续训练。
热备节点的概念也很重要。万卡集群里通常会预留1%到2%的节点作为热备池,专门用于替换故障节点。这些节点不跑训练任务,但预置了完整的镜像、驱动和依赖环境。当坏节点被踢掉后,调度器从热备池分配节点,把检查点数据拉过来,新节点只需要重新搭好训练进程就能加入。
热备节点池是成本换可靠性的典型做法。一万张卡的集群,预留100-200张卡热备,看起来浪费,但对比一次故障中断导致的算力损失,这笔账很划算。
4.3 弹性重算 vs 固定世界大小恢复
这里要区分两种恢复策略,它们在工程实现上有本质差异。
固定世界大小恢复:训练的rank数量保持不变,比如1024卡训练,挂了4张,就从热备节点补4张,world size不变。这种方案对训练曲线最友好——数据并行分片、张量并行分片、流水线并行的stage分配都不变,恢复后训练行为和中断前完全一致。
弹性训练:训练进程容忍rank数量变化,挂了4张卡,就用剩下的1020张继续跑,通过梯度累积步数调整等效batch size,保持全局batch size不变。这背后的思想是"不依赖热备节点、不浪费算力",但实现复杂得多,数据分片要动态重算,通信拓扑要重排。
在实际生产环境里,我们更倾向于固定世界大小恢复,因为它对调度器和训练框架的改动更小,恢复一致性更好。弹性训练大多用在资源受限、不想预留热备节点的场景。如果预算允许,我建议优先做固定世界大小恢复。
5. 恢复后的"无缝感"从哪来:一致性校验与数据迭代对齐
5.1 数据Loader恢复:别让模型"跑回头路"
很多人做断点续训时只保存模型和优化器,结果恢复后训练曲线出现不正常的重复波动。这大概率是数据Loader没存对导致的。
分布式训练的数据流通常是:全局数据集被均匀切分给各个rank,每个rank通过一个分布式采样器维护自己的样本索引。在恢复时,除了要恢复每个rank的当前step,还要恢复它的采样器状态——当前epoch、当前batch偏移、shuffle时用的RNG种子。
如果采样器没恢复,训练会从第0步重新开始遍历数据。模型已经把前面的数据学过了,再学一遍,损失函数会显得"乱跳"。大数据集上可能一两万个样本之后才察觉,但小数据集上几分钟就能观察到训练曲线异常。
我们团队的做法是在检查点里专门存一个DataLoaderState字典,包括:
dataloader_state = { "epoch": current_epoch, "batch_index": current_batch_index, "shuffle_seed": shuffle_rng_state, "sampler_rank_offset": sampler_rank_offset, "consumed_samples": global_consumed_samples, }恢复时,把采样器重新set到这个位置。如果你用的是自定义数据读取逻辑,也建议把文件读取偏移、缓存状态一并保存。这个动作看起来琐碎,但它是保证"恢复后训练曲线平滑接上"的基础。
5.2 学习率调度器、EMA、混合精度缩放器的状态恢复
这部分经常被忽略,但影响也不小。先说学习率调度器。常见的策略是Warmup后线性衰减到某个值,如果你恢复时LR重置到初始值,训练步长会突然跳到之前的状态,学习率莫名其妙变大或变小,收敛曲线会出现一段不可预测的震荡。解决办法是把调度器的step数恢复,让LR函数接着跑。
EMA影子权重也是类似。很多大模型训练采用指数移动平均作为最终评估版本,如果EMA状态没保存,恢复出来的模型在评估指标上会出现一段时间的不稳定。侵入式做法是在每个训练循环里,把EMA权重视为普通状态一并导出;如果想省事,也可以只在保存时导出EMA,而梯度更新时照常。
再强调一下AMP的Loss Scaler。动态缩放因子如果丢失,训练恢复后会连续出现溢出导致的跳过更新,表现为loss曲线突然多了一段水平平台。解决办法也同样直接:把GradScaler的state_dict写进检查点。
这些"隐性状态"的保存,其实是一个训练框架成熟度的试金石。能把它们都吃到检查点里,你的恢复才会真正"无缝"。
5.3 恢复结果的验证方法
到底怎么证明恢复是成功的?我们的验证分三层。
第一层是数值一致性。训练任务恢复后,跑固定的几步,对比恢复前这几步的loss曲线和梯度范数。如果曲线平滑衔接,没有突变,基本可以判断状态恢复正确。为了这个目标,我们需要在训练日志里自动记录每个step的loss、lr、grad norm等信息,恢复时拉出来对比。
第二层是确定性对齐。如果在完全相同的环境和输入下,从同一检查点恢复,多次运行得到的loss应该完全一致(前提是禁用了非确定性操作)。我们会在CI里跑这个测试:保存检查点,重启进程,执行若干step,对比两次运行的tensor数值。有差异就说明RNG或数据迭代状态没对齐。
第三层是端到端benchmark。真正到万卡级别,把故障恢复后的训练跑几个epoch,对比总的收敛曲线和评估指标。如果整体指标和没有中断过的对照组基本一致,才算过关。
6. 故障注入、踩坑清单与恢复演练
6.1 真实的踩坑经历
断点续训方案上线初期,我们几乎每周都能从故障演练里挖出新的坑。挑几个有代表性的说说。
坑一:检查点写了一半,任务就崩了。这是最普遍的问题。训练进程正在写检查点文件时,如果节点掉电或者进程被kill,很容易留下一个半截文件。恢复时加载到这个损坏文件,反序列化直接报错。我们的解决思路是"先写影子文件再原子rename":每个checkpoint先写入临时路径,全部写完并通过校验后再rename成正式文件。这样任何时刻正式路径下都是完整可用的版本。
坑二:磁盘满了,保存静默失败。检查点保存是异步的,如果落盘失败,训练进程可能不会立刻感知。我们曾经遇到过磁盘写满后,保存线程一直重试,训练进程浑然不觉,等到故障发生想恢复时才发现"最近的checkpoint是两小时前的"。现在的方案是保存线程每次完成都会上报元数据,控制面检查检查点时间戳,一旦超过预设阈值自动告警。
坑三:多副本目录不一致。检查点如果同时写本地NVMe和远端并行文件系统,两处文件的更新时序需要一致性。如果控制面读到的是旧版本的元数据,恢复时可能加载到过期状态,白白浪费训练时间。我们后来统一用元数据服务器做主时间戳,文件目录只作为存储介质,不再承担"哪个是最新"的判断。
6.2 故障注入手段与混沌演练
断点续训方案不是写出来就能用的,必须通过故障注入演练反复验证。我们在演练中用的手段包括:
- 进程级故障注入:用
kill -9随机杀掉一个训练进程,模拟节点宕机或框架崩溃。 - GPU故障注入:通过驱动层的Xid注入工具模拟GPU报错,或者人为把某张卡的NVLink链路断开。
- 存储故障注入:把存储挂载点短暂摘掉,或让目录权限变成只读,模拟存储不可用。
- 网络故障注入:用TC工具做网络延迟、丢包,模拟IB链路劣化。
- 节点掉电模拟:直接通过带外管理接口把节点硬关机,这是最接近真实的演练方式。
演练频率上,一个星期至少做一次随机故障注入。每次演练之后都要复盘三个指标:RPO(丢了多少步训练进度)、RTO(恢复耗时多久)、MTR(平均修复时间,包括故障定位和人工介入时间)。
6.3 可落地检查清单
最后整理一份我在项目上线和演练时反复用来检查的清单,你可以直接复制到自己的运维手册里:
- 检查点是否包含模型参数、优化器状态、RNG、数据Loader索引、LR调度器、EMA、Loss Scaler?
- 检查点是否使用分片方式写入,元数据是否单独维护?
- 是否有影子文件+原子rename机制,避免半截文件?
- 保存失败是否有明确告警,检查点滞后是否有时间阈值报警?
- 故障检测是否覆盖GPU亚健康状态,NCCL超时参数是否调优?
- 故障剔除时是否会先暂停训练、触发内存快照,再进行通信组重排?
- 是否有热备节点池,替换节点后的环境是否预置完整?
- 恢复是否经过数值一致性验证,而不是只判断"进程起来了"?
- 是否定期做故障注入演练,RPO、RTO是否在预算范围内?
- 企业级方案里是否有高可用的控制面做编排调度,避免控制面单点故障?
这套检查单看着条目很多,但每一项背后都对应过一次线上事故。很多团队把断点续训想得太简单,以为"保存模型就能恢复",直到在一次真实的万卡故障面前才意识到,真正要解决的是一个从硬件、驱动、框架到调度、存储、元数据的全栈问题。
我个人在反复折腾这些系统后的最大体会是:断点续训不是某个组件的事,而是一套需要端到端设计的工程系统。从发现故障的那一秒开始,到训练曲线重新平滑前进为止,中间每一环都会咬你一口。把这一整套链路都踩实了,你的万卡集群才配叫"高可用"。