你的 GPT-6 训练任务已经在 AI 训练集群上跑了 10 天,却突然崩溃了。屏幕熄灭,心跳加速。之前的进度难道全没了吗?要想恢复训练进度,必须提前规划。没人愿意眼睁睁看着几周的算力投入付诸东流,也没人愿意面对漫长停机所带来的压力。

这份实战指南将提供切实可行的步骤。首先,弄清楚哪些状态丢失了。然后,保存并恢复这些状态以继续训练。接着,调整学习率。最后,验证恢复是否成功。

理解被中断的训练会话

AI 训练集群故障的常见原因

训练任务中断的原因有很多,其中硬件故障最为常见。GPU 可能因过热而自动关机,内存模块可能出现错误,电源也可能在毫无预警的情况下失效。软件缺陷同样会导致崩溃。一个细微的代码错误,可能在运行数天后才破坏训练循环。网络问题也可能让节点与集群失去连接。无论是哪种故障,结果都一样:你的训练任务会在中途停止。

其影响远不止任务停止本身。你的集群会处于空闲状态,其他任务在队列中等待,团队的有效工作时间被浪费。这样的停机损失不只是算力成本,更会打乱研究节奏,拖慢项目进度。理解这些原因,有助于你为不可避免的中断做好准备。你无法阻止所有故障,但你可以掌控自己的应对方式。

训练停止时会丢失哪些状态?

很多工程师误以为只需要保存模型权重,这种想法往往会导致恢复失败。真正的续训,要求你恢复完整的训练状态。模型参数只是其中一部分,优化器状态同样至关重要。以 Adam 优化器为例,它保存了每个参数对应的动量和方差估计值。没有这些值,优化器的行为就会像刚开始训练一样。学习率调度器也保存了位置信息,它知道当前训练进行到衰减计划的哪一步。

随机数生成器(RNG)状态同样重要。训练过程中,数据打乱和数据增强都依赖随机性。每个 epoch 的数据顺序都会不同,图像变换会随机裁剪和翻转。这些操作都由 RNG 状态决定。如果不恢复它,数据管道就会生成不同的样本序列。这种不一致会带来细微但真实的训练差异,模型的收敛路径可能因此改变,验证结果也可能与崩溃前的日志不再一致。

一个完整的检查点应当同时保存这些组成部分:模型、优化器、调度器以及 RNG 状态。采用这种完整保存方式,恢复过程就会简单得多。你可以把所有内容恢复到故障前的精确状态,让训练仿佛从未中断。这样的准备,能够把原本可能是一场灾难的故障,变成一次小小的不便。你为完善检查点机制投入的时间,会在今后的每一次故障中获得回报。

通过检查点恢复训练进度

保存完整的模型与优化器状态

检查点必须捕获训练状态中的每一个关键部分。你不能只保存模型权重。优化器保存了每个参数的动量值,调度器记录了自己在学习率衰减曲线中的位置,随机数生成器决定了数据打乱的模式。每一个组件,都是平滑恢复训练不可或缺的一环。

你可以通过一条命令创建一个完整的检查点。字典结构能让所有内容保持清晰有序。保存当前 epoch 编号,以便知道从哪里重新开始;保存模型的状态字典,以还原网络结构中的参数;保存优化器的状态字典,以保留已经累积的梯度信息;保存调度器状态,以延续既定的学习率计划;记录 RNG 状态,以保证数据管道的可复现性。这个完整快照,能在任何故障后真正实现训练进度的恢复。

torch.save({
    'epoch': epoch,
    'model_state_dict': model.state_dict(),
    'optimizer_state_dict': optimizer.state_dict(),
    'scheduler_state_dict': scheduler.state_dict(),
    'rng_state': torch.get_rng_state()
}, 'checkpoint.pth')

随着模型规模增大,这个文件会迅速膨胀。一个简单经验法则是:若包含优化器状态,每个参数大约需要 12 字节。权重本身占 4 字节,优化器状态占 8 字节。这个估算规律可以帮助你在训练开始前就规划好存储容量。

存储检查点以实现快速恢复

你选择的存储介质,将直接决定恢复所需时间。像 HDFS 或 S3 这样的分布式文件系统提供更强的持久性,而本地 NVMe 则提供更高的速度。两者各有用途,也往往都需要。活跃训练过程通常写入并行文件系统,长期归档则转移到对象存储中,其每 TB 成本通常只有前者的十分之一。

检查点保存间隔需要仔细权衡。保存过于频繁,可以减少崩溃后需要重算的工作量;但保存过于频繁,也会增加写入时对 GPU 训练的阻塞。你必须在两者之间找到平衡。每 5 分钟保存一次通常过于激进,而对于持续数天、故障并不频繁的训练任务来说,每小时保存一次往往更合适。

检查点保存间隔应与任务的崩溃率(平均中断时间)成比例,而这个中断时间通常会随着所使用 GPU 数量的增加而缩短。

大规模训练任务在实践中通常遵循这一原则。一个 4050 亿参数的模型,如果平均 150 分钟会被中断一次,就会每 15 分钟保存一次检查点;一个 8000 亿参数的模型,则可能每 40 分钟保存一次。通常可将保存间隔设为预计故障间隔的大约十分之一,这样既能减少算力浪费,又不会带来过高额外开销。

现代大语言模型每次写入检查点时,往往需要写出 350–500 GB 的模型状态。为了尽量减少训练中断,写入过程必须在 5 分钟内完成。这意味着系统需要提供 1–2 GB/s 的持续突发写入带宽。因此,存储架构必须优先考虑突发性能,而不仅仅是长期吞吐能力。

像 TrainMover 这样的系统,进一步提升了恢复能力。替换机器会在真正发生故障前,预先完成 CUDA 内核编译并建立 NCCL 通信组。这种预热机制把检查点加载排除在关键恢复路径之外。其基于增量更新的设计,只更新离开节点与加入节点之间的连接,其他所有机器都无需修改。即使在 1,024 张 GPU 的规模下,这种架构也能将停机时间基本稳定在 20 秒以内。与此同时,由于准备状态保存在 CPU 内存或 NVMe 中,因此不会占用任何 GPU 显存,GPU 显存在最终切换前始终保持不变。

异步检查点是另一种有效的缓解策略。你可以先将数据从 VRAM 卸载到主机内存中,然后让 GPU 恢复训练,而由 CPU 在后台负责保存。对 85,000 次检查点的分析显示,这种方式可以将全局带宽控制在 1 TB/s 以下,而本地卸载速度可达到 50–200 GB/s。这样就能避免保存过程中 GPU 长时间停顿。

对于大规模训练而言,检查点期间 GPU 空转带来的成本每天可能超过 4,000 美元。优化停机时间,实际上就是在直接降低成本。你每减少一秒恢复时间,整个集群都会随之节省开支。因此,检查点策略应当与模型架构同等重视。你在稳健保存机制上投入的精力,会在之后的每一次故障中持续带来回报。

恢复 AI 训练集群状态

加载检查点并继续训练循环

恢复工作的起点,是一次简单的加载操作。你把之前保存的检查点文件重新读入内存,这一动作会恢复你先前保存的每一个组件。模型权重回到故障前的数值,优化器重新获得它的动量估计,调度器恢复到学习率衰减曲线中的原有位置。整个训练环境会回到崩溃前的精确状态。

checkpoint = torch.load('checkpoint.pth')
model.load_state_dict(checkpoint['model_state_dict'])
optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
scheduler.load_state_dict(checkpoint['scheduler_state_dict'])
start_epoch = checkpoint['epoch'] + 1

其中最关键的细节在最后一行:你必须从已保存 epoch 的下一个 epoch 开始,而不是从零重新开始。训练循环也需要相应调整,以便遵循这个起点。只需修改循环范围,使其从恢复后的 epoch 开始,就能避免重复计算已经完成的轮次。你的任务将从中断处精确续上,保留此前所有进度。

for epoch in range(start_epoch, total_epochs):
    train_one_epoch(model, train_loader, optimizer, scheduler)
    validate(model, val_loader)
    save_checkpoint(model, optimizer, scheduler, epoch)

像 fastai 这样的框架,还能把这个过程进一步简化。fit_one_cycle 函数会应用循环式学习率与动量策略。如果中断后处理不当,整个循环会从头开始,导致结果与未中断训练不同。fastai 通过 start_epoch 参数解决了这个问题。你只需重新实例化 learner,并用下一个 epoch 编号调用 fit_one_cycle。框架会自动加载此前为该 epoch 保存的文件,无需手动重新加载权重。这样,学习率和动量循环就能从准确位置继续执行,训练策略得以无缝延续。

重新建立数据加载器与随机种子

如果数据管道不能一致,仅仅恢复模型也没有意义。随机数生成器状态决定了数据加载器如何打乱样本顺序。每个 epoch 都会产生不同的顺序,而图像增强中的各种随机变换也都依赖 RNG 状态。

恢复这个状态其实很简单。只需把之前保存的 RNG 值重新传给 PyTorch,数据加载器就会生成与未崩溃时完全相同的样本序列。这种可复现性对于验证一致性非常重要。你的损失曲线能够在中断前后保持可比,梯度模式也会更加稳定可预测。

torch.set_rng_state(checkpoint['rng_state'])

你还必须恢复各个 worker 专属的随机种子。数据加载器通常会启用多个 worker 进程,而每个 worker 都维护着自己的随机状态。你需要在重新创建数据加载器实例之前,先设置这些种子。这样才能确保每个 worker 都复现出原本的数据打乱模式。至此,你的训练环境才算真正回到了故障前的状态。

最后的验证步骤,用于确认恢复是否真的成功。先跑几个 batch,并把损失值与崩溃前的日志进行对比。二者应当大致接近,梯度范数也应落在合理范围内。如果偏差明显,就说明恢复过程存在问题。尽早发现这些错误,能避免你在损坏状态下继续浪费数小时算力。

你的检查点加载策略,也会直接影响恢复时间。一个结构清晰、组织良好的检查点,能让你迅速定位到正确文件,并顺利完成加载,无需复杂解析。这样,你的任务通常可以在几分钟内恢复到完整运行状态。高效的恢复机制,会把一场本可能很严重的故障,变成一次短暂的暂停,让训练以最小停机、最大把握继续推进。

调整学习率以继续训练

学习率调度器保存着训练进展中的关键信息。它记录了当前在衰减序列中的位置。仅恢复模型和优化器,并不意味着恢复工作已经完成,调度器状态同样必须回到崩溃前的位置。如果你新建一个调度器,整个学习率计划就会被打乱,因为模型此时本应使用后期更低的学习率。

计算正确的 step 和 epoch

调度器通常根据 step 或 epoch 来决定当前应使用的学习率。你必须知道崩溃发生时所处的精确位置。可以使用以下公式:current_step = epoch * len(train_loader) + batch_index。这个计算结果给出了你在学习率计划中的精确 step 编号。然后,将 last_epoch 参数设置为这个 step 值,调度器就能从那个准确位置继续运行。学习率曲线将跨越中断点保持平滑,模型在后续每一步都能获得正确的学习率。这种做法可确保恢复过程不会破坏原本的训练节奏。

对于多周期调度策略来说,epoch 编号同样重要。余弦退火(cosine annealing)和 one-cycle 策略都依赖于完整周期长度。你的检查点中保存了最后一个已完成的 epoch,只需在此基础上加一,再传入调度器构造函数,调度器就能正确计算剩余衰减过程。这一步能让恢复过程避免拍脑袋猜测。

重新初始化调度器以保持连续性

训练崩溃后,不要新建一个全新的调度器。新的调度器会把学习率重置为初始值,而此时模型通常已经进入训练后期,理应使用更低学习率。突然回到较高学习率,可能导致梯度不稳定、损失暴涨,甚至权重发散。这样的错误,会让你之前的恢复努力前功尽弃,甚至迫使你再次重启训练。

如果在崩溃后重新创建一个新的调度器,你的学习率计划就会从头开始。这个突增可能让模型不稳定,并抵消数小时的训练成果。正确做法始终是恢复已保存的调度器状态。

正确方式是先创建调度器,再加载其保存的状态。scheduler.load_state_dict() 方法会恢复调度器内部所有参数,让它回到原本的 step 计数位置。这样,学习率就能沿着既定轨迹继续前进,仿佛故障从未发生。整个训练过程将无缝衔接,不错过任何一步。

妥善恢复调度器,可以彻底规避上述风险。你的训练动态得以延续,验证指标也保持可比。应当把这种做法应用到每一个检查点中。保存调度器状态几乎不需要额外成本,而忽略它的代价却可能非常高昂。正是这种对细节的重视,才能把一次崩溃从重大挫折,变成短暂暂停。

验证恢复是否成功

对损失和梯度进行合理性检查

在恢复检查点之后,先运行少量 batch。把此时的损失值与崩溃前日志中的数值进行比较。损失应与该 epoch 最后记录的值大致接近。如果匹配良好,就说明训练状态恢复正确;如果出现明显跳变,则意味着检查点加载流程存在问题。

你还需要检查梯度范数。这些数值衡量的是权重更新的幅度。把它们与崩溃前的日志进行对照,若结果相近,通常表示恢复成功;若差异明显,则说明某处出了问题。例如,调度器状态可能没有正确恢复,或者随机数生成器与原始环境不一致。这些细微的不匹配,都会导致训练结果不稳定。

如果损失值或梯度数值出现大幅跳变,就说明恢复过程出现了错误。检查点文件可能已损坏,也可能内容不完整。这种快速检查能帮助你尽早发现问题,避免在错误状态下继续训练数小时,从而节省宝贵算力。

建议采用对照式验证方法。提前记录崩溃前某一步的损失值和梯度范数,恢复后使用相同数量的 batch 重新跑一遍,再记录新的结果,并进行逐项对比。两组数值应当高度接近。整个过程只需要几分钟,却能在后面节省大量时间。

监控恢复初期是否出现发散

恢复后的最初几个 step,往往最能说明问题。理想情况下,训练应当和中断前完全一致地继续进行,损失曲线保持原有下降趋势,梯度模式也应维持稳定。

你需要特别留意一些警告信号。比如梯度爆炸,会表现为参数更新突然异常增大;NaN 值则意味着出现了数值不稳定。这两种情况通常都源于恢复错误。如果不及时处理,训练不会自行回到正常状态,反而会继续浪费资源。因此,必须尽早发现。

如果观察到发散现象,应立即停止任务,不要抱着侥幸心理继续运行。请检查检查点文件是否完整,确认各个组件是否都已正确加载,并核对 RNG 状态是否与原始环境一致。

一次干净、正确的恢复,不会出现发散。损失曲线会在各个 epoch 之间平滑延续,梯度范数保持稳定可预测,学习过程继续正常推进。这意味着你已经成功从故障中恢复,任务能够按计划继续运行,停机阶段正式结束,训练进度也得以延续。

至此,整个验证流程完成。此后,你就可以对模型后续结果保持足够信心。

现在,你已经理解了完整的恢复流程:保存全部状态——模型、优化器、调度器和 RNG;实施稳健且间隔合理的检查点机制;精确恢复训练环境;正确延续学习率计划。这些步骤,能够把服务器故障从灾难性事件,变成可控的短暂中断。

只要准备充分,你的训练任务就能在崩溃后继续存活。前期在检查点机制上的投入,会在未来节省无数小时。恢复时间将从数天缩短到几分钟。你的训练会从中断处准确接续,不浪费算力,也不丢失进度。

你保存下的每一个 epoch,都是对停机风险的保险;你做出的每一次学习率修正,都是对训练稳定性的保障;你写下的每一个检查点,都是对模型轨迹的守护。

有了这份指南,你就能更从容地面对下一次服务器故障,因为你知道训练进度是安全的,接下来的道路也是清晰的。请在下一次崩溃发生之前,就把这些实践落实到位,而不是事后补救。

常见问题

训练过程中应多久保存一次检查点?

你应根据预期崩溃率来设置检查点间隔。比如一个任务平均每 150 分钟中断一次,那么每 15 分钟保存一次会比较合理。这种做法既能减少重算损失,也不会带来过多额外开销。最终可行的间隔,还取决于你的存储写入速度。

如果我只保存模型权重,会发生什么?

你会丢失关键的优化器状态,包括 Adam 的动量值;学习率调度器会重置到初始阶段;随机数生成器会生成不同的数据打乱顺序。这些缺失都会导致训练结果不一致。因此,若想实现真正的恢复,必须始终保存完整的状态字典。

我可以在不同数量的 GPU 上恢复训练吗?

可以,你可以在不同 GPU 数量的配置上恢复训练。检查点中保存的模型参数和优化器状态可以跨配置迁移。不过,你需要相应调整 batch size 和学习率。数据加载器的分发方式可能会改变,但保存的 epoch 和模型状态依然有效。

我怎么判断恢复是否成功?

先运行几个 batch,并把损失值与崩溃前的日志进行比较;梯度范数也应当尽量接近。同时观察恢复初期是否出现 NaN 值或梯度爆炸。任何显著偏差,都说明恢复流程存在错误,应立刻停止并排查,避免继续浪费算力。

检查点会拖慢训练吗?

检查点机制确实会带来一定开销,但现代系统已经能把这部分成本压得很低。传统存储方式通常会占用 5%–10% 的训练时间,而高速 NVMe 分层存储可将其降到 1% 以下。异步检查点还能把数据先卸载到主机内存,让 GPU 在保存过程中继续工作。