Appearance
Q131 · 容错与长期训练
假设一个模型要在两张 GPU 上训练几天。已经完成第 400 次参数更新时,系统保存了一份可恢复的记录;训练继续到第 460 次更新后,一台机器突然故障。工程师最想知道的不是“有没有一个模型文件”,而是:重启后能从第几步继续,下一批该读哪些数据,学习率和优化器内部状态是否接得上,两张卡看到的训练进度是否一致?
长期训练会遇到进程崩溃、GPU 或节点失联、内存不足、存储故障、网络中断、任务被调度系统终止以及人为暂停。容错的目标是在这些事件后,能从最近一个完整且一致的检查点继续,并清楚知道哪些进度必须重做。它不意味着训练永不中断,也不意味着恢复后一位小数都与无故障运行相同。PyTorch 的容错分布式训练教程以故障后重启进程、从已保存快照继续为基本流程。
先解释术语
| 术语或记号 | 直白解释 | 本文例子 |
|---|---|---|
| GPU | 执行模型计算的加速设备 | GPU A 与 GPU B 一起训练 |
| 训练步 / step | 本文特指一次已经完成的优化器参数更新,不是读入一个样本 | 第 400 步表示第 400 次更新完成 |
| 模型权重 | 模型已经学到、会用于预测的参数 | 第 400 步时模型各层的数值 |
| 优化器状态 | 帮助下一步更新参数的内部记录;依算法而不同 | Adam 一类方法保存的历史统计 |
| 学习率调度器 | 决定不同时刻采用多大学习率的规则及其当前进度 | 第 401 步应使用的学习率 |
| 检查点 / checkpoint | 保存某个训练边界上足以继续运行的状态集合 | 已提交的第 400 步快照 |
| 数据游标 | 下一批数据应该从哪里读的进度信息 | 第 400 步完成后,每张卡下一批的索引或位置 |
| 采样器 | 决定样本顺序,以及多张卡各读哪一部分数据的组件 | A、B 各自要读不同样本 |
| 随机数状态 | 随机过程“走到了哪里”;单记一个最初种子通常不够 | 数据打乱、随机增强等操作的后续序列 |
| rank / 进程编号 | 分布式训练里区分工作进程的编号 | A 和 B 是两个不同 rank |
| 分片 | 大模型状态拆开,由不同进程分别持有 | A、B 各存一部分权重或优化器状态 |
| 提交 | 确认这份检查点的所有必需部分都写完且可用 | 第 400 步快照可供重启选择 |
| 恢复点目标 | 故障后最多愿意损失多少已做但未保存的进度 | 例子里最多重做约 60 步 |
这里的“第 400 步”只是教学设定,不是框架默认保存周期。“检查点”也不等于一个固定格式的单文件:单机可以是一个文件,分片训练通常由多个文件和描述它们的元数据构成。PyTorch Distributed Checkpoint(DCP)文档说明分布式检查点可由多个 rank 并行保存,并产生每个 rank 对应的文件。
为什么只留模型权重无法无缝继续训练
第 400 步后,模型权重确实能用于推理;但继续训练还需要知道怎样更新它。以有历史统计的优化器为例,如果重启时只装回模型权重、把优化器重新初始化,第 401 步的更新方向和幅度可能与原计划不同。PyTorch 的通用检查点教程明确区分“只用于推理的模型状态”与“用于恢复训练的通用检查点”,后者至少还应考虑优化器状态和已完成的训练进度。PyTorch:Saving and Loading Models
下面把本例在第 400 步更新完成、下一批尚未开始这个边界应保存的内容列清。选择这样的边界,可以避免把“更新做了一半、梯度累积还没完成”当作一个已提交步骤。
| 保存内容 | 为什么需要 | 漏掉后可能怎样 |
|---|---|---|
| 模型权重和必要缓冲区 | 把已学到的模型带回第 400 步 | 从旧模型或随机初始值重来 |
| 优化器状态 | 延续更新过程里的历史记录 | 第 401 步更新与原运行不同 |
| 已完成步数、学习率调度器状态 | 知道接下来是第 401 步及应采用的学习率 | 学习率突然回到初始阶段,或重复计数 |
| 数据集版本、数据游标与采样器状态 | 确定各进程下一批读什么、哪些样本已处理 | 大量重复或跳过样本;不同 rank 分配错乱 |
| 用到的随机数生成器状态 | 尽量延续打乱顺序、随机增强等过程 | 恢复后抽到不同样本或增强结果 |
| 混合精度缩放器状态(若使用) | 延续自动调整梯度缩放的状态 | 与之前不同的缩放或跳过更新行为 |
| 训练配置、代码与数据版本标识 | 确认检查点与当前程序和数据是否兼容 | 看似恢复,实际用了另一套模型或数据 |
并非每个框架都会自动保存表中所有状态。具体清单取决于训练脚本和优化器:例如没有学习率调度器,就无需它的状态;如果在梯度累积的中途强行保存,还须处理尚未提交的梯度及累积计数,远比在完整更新边界保存复杂。PyTorch 的混合精度教程也明确建议在迭代开始前或缩放器更新完成后保存其状态。PyTorch:Automatic Mixed Precision 的保存与恢复

图里的手提箱代表逻辑上同一个一致的检查点,不表示所有文件必须合成一个物理文件;图只画了四类核心信息,调度器、混合精度等按实际训练配置补充。假设故障发生在第 460 次参数更新后、下一份检查点提交前,那么第 401~460 步已做过却没有持久记录,重启后会重做这 60 步;并不是可以拿着内存里的第 460 步权重接着做第 461 步。
数据位置和随机状态为什么会决定恢复质量
训练不是一遍遍读同一条记录。程序可能打乱样本、给图片随机裁剪、把大数据流切成不同分片。若模型和优化器回到第 400 步,但数据读取器接着故障前第 461 步的位置走,模型的训练轨迹就出现第 401~460 步参数更新没留下、样本却被当成已读的缺口。反过来,若数据从开头重来,可能多次训练同一批记录。
因此在第 400 步要记录能够重建下一个样本位置和顺序的信息:至少明确数据版本、当前轮次、打乱种子、各 rank 的采样位置;复杂数据管道还可能需要它自己的可保存状态。PyTorch 的分布式采样器文档要求在每个训练轮次开始、创建数据迭代器前设置轮次,才能让不同轮次的打乱按预期变化。TorchData 的 StatefulDataLoader 则提供保存/装回迭代位置的接口;其文档特别说明它会处理同一 rank 内多个加载 worker 的状态,但不会自动跨 rank 汇总,分布式协调仍由训练系统负责。PyTorch:DistributedSampler · TorchData:Stateful DataLoader
随机状态也有层次。只写“最初种子是 42”,并不能表明第 400 步时随机序列已经被消耗到哪里;若要尽量重现后续采样、dropout 或数据增强,需保存实际使用的随机数生成器状态,涉及 Python、NumPy、PyTorch CPU/GPU 及数据加载 worker 时逐项考虑。即使这些都处理了,不同硬件、库版本或非确定性算子仍可能导致数值不完全一致;容错目标通常首先是可正确继续训练,要求位级一致时要把运行环境和确定性条件单独验证。PyTorch:Reproducibility · PyTorch:Numerical accuracy
对流式数据尤其要小心:数据源可能在训练期间发生追加、修改或删除。只有“第几个样本”而没有数据集版本、分片清单和可重现读取顺序,恢复时的“同一个位置”也许已经对应另一条记录。能否做到严格不重不漏,取决于数据管道提供的可重放能力;模型文件本身无法保证这一点。
两张卡怎样保存出一致的检查点
先区分两种常见布局。数据并行时,A、B 可能各有一份相同模型、各读不同数据并同步梯度;在满足该框架同步前提下,模型权重可以只由一个进程写,但每个进程的数据位置、随机状态等可能各不相同,不能因此只保留 A 的全部训练状态。PyTorch 的 DDP 教程展示了由一个进程保存同一模型、其他进程在保存完成后再加载的例子,同时强调等待保存完成和正确映射设备。PyTorch:DDP 保存与加载检查点
如果采用模型或优化器分片,A、B 各掌握状态的一部分,单拿 A 的文件就不完整。保存流程需要各 rank 对齐到同一已完成训练边界,然后并行写各自负责的部分;协调方确认必要的分片和元数据都写好,才能把这份检查点标为“可恢复”。恢复时也要把这些部分重新装入相应进程。PyTorch DCP 提供并行保存/加载以及装载时重新分片的能力,但模型、优化器以外的数据游标与随机状态仍需训练应用按其数据管道设计。PyTorch:Distributed Checkpoint · DCP 入门教程
这里的一致包含两个层面:一是 A、B 的快照都对应同一个已完成步数,不能把 A 的第 400 步和 B 的第 399 步拼起来;二是模型、优化器、调度器、数据位置等属于同一个逻辑时刻。实际实现可采用“先写到新版本目录→核验文件与元数据→最后发布可用标记”的提交协议。若第 500 步写到一半机器断电,恢复逻辑应该忽略这份未完成记录,退回已确认的第 400 步。具体的原子性依赖存储和框架实现,不能仅因调用了保存函数就假定远端文件已安全提交。
异步保存可让写盘与训练部分重叠,但需要先拿到不会被后续更新改写的状态副本,并确认后台写入成功后再宣告检查点可恢复。PyTorch DCP 对异步保存的说明包括先把状态暂存到训练安全的副本、随后在后台写入,以及等待完成结果;异步调用返回本身不等于持久化已成功。PyTorch:DCP async_save
故障后从第 400 步继续的具体顺序
- 发现故障并停止旧进程组。 A、B 的分布式通信有依赖,一台失联时让另一台独自继续可能卡在通信操作或导致状态不一致。以 PyTorch torchrun 的容错方案为例,故障时会重启相关训练进程;它负责重新启动,检查点内容和保存频率仍取决于训练脚本。PyTorch:Fault-tolerant Distributed Training
- 查找最近的完整检查点。 例如发现第 500 步文件不全,而第 400 步已有完成标记且可校验,就选第 400 步;不能只看文件名“step500”便相信它可恢复。
- 验证兼容性并重建进程。 检查模型结构、训练配置、数据版本、分布式拓扑与检查点格式。拓扑变化时可能需要框架支持重新分片;数据版本变化则要明确是否允许继续同一实验。
- 按组件装回状态。 各 rank 装回其应有的模型与优化器状态、调度器、混合精度状态,并恢复各自的数据游标和随机状态;确认全体进程看到同一个已完成步数 400。
- 从下一步执行。 第 401 次更新使用对应的数据批次与学习率。训练若已在故障前做到第 460 步,也要把 401~460 的未提交进度重新做一遍。若恢复策略无法精确重放数据顺序,要如实记录偏差,不能声称“无缝完全一致”。
- 核对第一段恢复日志。 检查各 rank 的步数、学习率、数据批次、梯度/损失是否在合理范围;新检查点成功提交后,再把它纳入可恢复列表。
上面的步骤是设计流程,不是可直接运行的特定框架代码;不同分布式栈的具体 API 和状态清单有所差异。
多久保存一次:重做损失与保存开销的取舍
检查点保存得太稀,失败时要重做很多训练;保存得太密,I/O、网络带宽和存储空间又会拖慢训练。用本例假设:每步计算耗时 5 秒,每 100 步保存一次且一次同步保存耗时 20 秒。一个 100 步计算窗口是 500 秒,保存额外占 20 秒,粗略相当于计算时间的 4%;若第 400 步保存后在第 460 步故障,约 60 步、300 秒的计算要重做。若每 20 步保存,单次故障最多损失约 20 步计算,但每 100 秒计算就花 20 秒保存,粗略是 20% 的额外保存时间。
这些数都是假设算例,没有把压缩、异步重叠、网络拥塞、恢复装载耗时和实际故障频率算进去。真正选间隔时应量测检查点的完整大小、同步/异步保存耗时、恢复耗时、平均故障间隔和可接受的进度损失;最好还规定保留几个历史检查点,防止最新一份损坏后毫无退路。
怎样做一次真正的失效演练
先在可控的小规模训练里保存第 400 步完整检查点,记录各 rank 的下一批数据标识、当前学习率和模型/优化器摘要。让训练再走 60 步,故意终止其中一个进程,模拟在第 460 步后发生故障。重启时只允许选“已提交”的第 400 步快照,确认每个 rank 的步数都回到 400,并在重新执行第 401 步时检查读到的批次和学习率是否符合预期。最后观察能否继续超过 460 步并成功写入下一份可恢复检查点。
再测两条失败路径。第一,保存新检查点时故意中断一个 rank 或破坏一个分片:加载器必须拒绝不完整的快照,回退到上一份完整版本,而非默默拼接新旧状态。第二,只恢复模型、故意不恢复优化器或数据游标:检查训练日志是否能检测出学习率/优化器状态重置、重复或跳过数据。若这些异常没有监控信号,单靠“训练进程重新跑起来了”不能算通过容错测试。
还应在演练记录里区分两种目标:连续训练目标要求任务能从合理状态继续且结果可用;严格重现目标进一步要求相同输入和环境下后续轨迹尽可能一致。后者可能受非确定性计算、拓扑变化和数据源变化限制,应单独验证,不应当作每个检查点自动保证的性质。
面试里怎样回答
“长期训练会遇到机器、网络、进程或调度故障,容错的基本做法是在完整参数更新边界定期保存检查点,故障后从最近一份完整检查点重启。检查点不能只有模型权重,还要包括优化器、学习率进度、已完成步数,以及让下一批数据接得上的采样器和随机状态;用了混合精度或分片训练,还要保存对应状态。分布式场景要保证各 rank 的快照来自同一步,并在所有必要分片写完后才标记可用。比如两卡训练第 400 步提交检查点、第 460 步故障,就从第 400 步恢复、重做之后未提交的 60 步。保存间隔是 I/O 开销与故障重做成本的取舍,最后要用杀进程和破坏半写入快照的演练证明真的能恢复。”
如果追问“只保存模型能不能继续”,应区分“拿来推理或重新开始微调”与“尽量延续同一训练轨迹”:前者可能够用,后者还缺优化器、数据和训练进度。若追问“随机种子保存了为什么仍不完全复现”,应说明一个初始种子不等于第 400 步的实际随机状态,且跨硬件/版本仍可能有数值差异。
资料依据
- PyTorch:Saving and Loading Models、Automatic Mixed Precision:恢复训练的模型、优化器与缩放器状态。
- PyTorch:Fault-tolerant Distributed Training、Distributed Checkpoint:进程重启、分布式快照和重新分片。
- PyTorch:DistributedSampler、TorchData:Stateful DataLoader:数据顺序与迭代状态。
- PyTorch:Reproducibility、Numerical accuracy:随机和数值一致性的限制。