AI断点续训踩坑记录
训练跑着跑着就挂掉,原因五花八门:OOM、数据加载超时、集群维护重启都有。
我之前对断点续训比较随意:每几个 epoch 存一次 checkpoint,挂了就从最近那个接着跑。
断点续训要恢复什么
先搞清楚一件事:训练暂停后,如果要继续,到底需要恢复哪些东西?很多人第一反应就是"模型权重",但实际上这只是其中一部分。
以一个典型的 PyTorch 训练为例,完整的状态至少包括:
- 模型权重(
model.state_dict()) - 优化器状态(
optimizer.state_dict()),比如 Adam 的动量和方差缓存 - 学习率调度器状态(
scheduler.state_dict()),包括当前的 step 和参数 - 当前训练的 epoch 和 batch 索引
- 数据采样器的状态(RandomSampler 的 shuffle、SubsetRandomSampler 的索引)
- 生成任务的生成器状态,比如文本生成的 random seed
- 如果有混合精度训练,scaler 的状态
- 如果有 DDP 分布式训练,还需要恢复进程组的相关状态
少恢复任何一个,训练曲线都会出现明显的拐点。比如忘了恢复优化器状态,优化器的动量信息就丢了,相当于重新开始优化;如果数据采样器状态没恢复,batch 的顺序就乱了,这在某些敏感任务里会导致不可预测的差异。
所以 checkpoint 通常存整包状态,权重只是其中一块。PyTorch 官方建议的做法是这样:
def save_checkpoint(state_dict, filename):
torch.save({
'epoch': state_dict['epoch'],
'global_step': state_dict['global_step'],
'model_state_dict': state_dict['model_state_dict'],
'optimizer_state_dict': state_dict['optimizer_state_dict'],
'scheduler_state_dict': state_dict['scheduler_state_dict'],
'loss': state_dict['loss'],
'scaler_state_dict': state_dict.get('scaler_state_dict'),
'random_states': state_dict['random_states'],
}, filename)
这里有个细节:torch.save() 默认会使用 pickle 序列化,而 pickle 在不同 PyTorch 版本之间可能不兼容。所以如果训练和恢复的 PyTorch 版本不一致,可能会遇到反序列化失败。这个问题在长期训练中很常见——比如你开始训练时用的是 2.0,几个月后想恢复时环境已经升级到 2.1。解决方案要么是固定 PyTorch 版本,要么用更稳定的序列化格式,但后者会增加复杂度。
下面这张图总结了 checkpoint 包含的各个组件以及它们之间的关系:
要点:checkpoint 是一组状态的集合,缺任何一块恢复后曲线都可能拐一下——优化器动量和随机数 seed 最敏感。
checkpoint 的保存策略
checkpoint 保存的频率是个永恒的 trade-off:保存太频繁,磁盘 IO 和存储成本都受不了;保存太稀疏,崩溃后的回退成本又太高。
我之前用过几种策略:
固定间隔保存:比如每 1000 个 batch 保存一次。简单粗暴,但问题是不够灵活——如果前面的训练比较稳定、后面的阶段波动大,这种策略要么在前期浪费 IO,要么在后期不够安全。
基于时间保存:比如每小时保存一次。这种方式的好处是不管训练快慢,保证最多丢失一个小时的工作量。但如果训练很快,一小时可能已经跑了成千上万个 batch;如果训练很慢,一小时又可能只跑了几个 batch。
基于 epoch 保存:在每个 epoch 结束时保存。这是最常见的方式,但对于大 epoch 的任务(比如一个 epoch 要跑好几天),这种方式就不够用了。
混合策略:比如每个 epoch 保存一个"关键 checkpoint",同时在 epoch 内部每隔固定步数保存一个"临时 checkpoint"。关键 checkpoint 保留时间长(比如保留最近 10 个),临时 checkpoint 过期快(比如只保留最近 3 个)。这样既保证了细粒度的恢复点,又控制了存储成本。
最终我落地的是混合策略,大概这样:
def save_checkpoint_mixed(state_dict, is_epoch_end=False):
if is_epoch_end:
# 关键 checkpoint,保留更长时间
filename = f'checkpoint_epoch_{state_dict["epoch"]}.pth'
torch.save(state_dict, filename)
# 清理旧的关键 checkpoint,保留最近 10 个
clean_old_checkpoints('checkpoint_epoch_', keep=10)
else:
# 临时 checkpoint,保留时间短
step = state_dict['global_step']
if step % 1000 == 0:
filename = f'checkpoint_step_{step}.pth'
torch.save(state_dict, filename)
# 清理旧的临时 checkpoint,保留最近 3 个
clean_old_checkpoints('checkpoint_step_', keep=3)
这里有个坑点:清理旧 checkpoint 时要小心,不要删除正在使用中的那个文件。尤其是在异步保存的场景里,文件可能刚创建完就被清理了。我遇到过一次是因为清理脚本判断"文件修改时间早于 X 小时就删除",但异步保存的文件可能刚创建完就被判定为旧文件而被删掉。
异常捕获和自动恢复
checkpoint 机制有了,下一步就是遇到异常时怎么自动恢复。理想的情况是:训练进程挂掉后,监控系统能自动重启进程,然后检测到最近的 checkpoint,自动从那里继续。
但现实往往比这复杂。首先你得知道进程什么时候挂掉了——可能是因为 Python 异常,也可能是被 OOM killer 杀掉,也可能是网络断开。针对不同的挂起方式,捕获策略也不同。
Python 异常:这是最好处理的情况。用 try-except 把训练循环包起来,捕获所有异常,保存一个"紧急 checkpoint",然后退出。
try:
for epoch in range(start_epoch, num_epochs):
for batch_idx, (data, target) in enumerate(train_loader):
# 训练逻辑
pass
except Exception as e:
logging.error(f'Training failed with error: {e}')
save_checkpoint(state_dict, 'checkpoint_emergency.pth')
raise
OOM killer:这是最痛苦的。进程被杀掉时连 Python 异常都抛不出来,直接消失。针对这种情况,没有特别好的办法,只能通过外部监控。比如用 systemd 或者 supervisor 启动训练进程,设置重启策略;或者在训练脚本外再套一个监控脚本,定期检查进程是否存在。
一个实用的方式是在训练开始时创建一个"锁文件"或"心跳文件",定期更新它。监控脚本检查心跳文件,如果长时间没更新就认为训练挂了,然后重启。
网络异常:这种通常会在数据加载、模型保存等 IO 操作时抛出异常。比如保存 checkpoint 时 NFS 存储突然不可用了,这时候你得处理两个问题:一是 checkpoint 保存失败,二是训练进程是否继续。通常的策略是:如果 checkpoint 保存失败,先尝试保存到本地临时路径;如果连本地也失败了,再考虑终止训练。
def save_checkpoint_safe(state_dict, filename):
try:
torch.save(state_dict, filename)
except (IOError, RuntimeError) as e:
logging.error(f'Failed to save checkpoint to {filename}: {e}')
# 尝试保存到本地临时路径
local_path = '/tmp/' + os.path.basename(filename)
try:
torch.save(state_dict, local_path)
logging.info(f'Checkpoint saved to local: {local_path}')
except Exception as e2:
logging.error(f'Failed to save checkpoint locally: {e2}')
raise
这张图展示了异常检测与恢复的整体流程:
这张图展示了三种异常的处理路径:Python 异常可以直接捕获并保存紧急 checkpoint;OOM Killer 需要外部监控检测心跳;网络异常则采用降级保存策略。最终无论哪种异常,都会回到 checkpoint 检测和恢复的闭环中。
恢复后的对齐问题
checkpoint 恢复后,理论上训练应该"无缝"继续,但实际总会有一些对齐问题。
第一个问题是数据采样器的状态。PyTorch 的 DataLoader 默认会打乱数据,如果你没保存和恢复 RandomSampler 的状态,那恢复后的数据顺序就和原来不一样了。这会导致两个问题:一是训练曲线出现奇怪的跳变,二是某些对数据顺序敏感的任务(比如一些 NLP 任务的动态 batch 策略)会出现不稳定。
恢复数据采样器状态的方式是保存和恢复随机数生成器的状态:
def get_random_states():
states = {}
# Python 的 random
import random
states['random'] = random.getstate()
# NumPy 的 random
import numpy as np
states['numpy'] = np.random.get_state()
# PyTorch 的 random(用于 dropout、数据增强等)
states['torch'] = torch.random.get_rng_state()
# CUDA 的随机状态(如果用了 GPU)
if torch.cuda.is_available():
states['cuda'] = torch.cuda.get_rng_state_all()
return states
def set_random_states(states):
import random
import numpy as np
random.setstate(states['random'])
np.random.set_state(states['numpy'])
torch.random.set_rng_state(states['torch'])
if torch.cuda.is_available():
torch.cuda.set_rng_state_all(states['cuda'])
第二个问题是学习率调度器的状态。有些学习率策略(比如 cosine annealing)是基于 step 的,如果你只恢复了优化器但没恢复 scheduler,那学习率就会重置,导致训练曲线出现明显的拐点。
scheduler.load_state_dict(checkpoint['scheduler_state_dict'])
第三个问题是分布式训练的状态。DDP 场景下,模型和优化器恢复之外,还得把进程组状态对齐。常见做法是先销毁旧进程组,再重新 init:
def resume_ddp_training(checkpoint_path):
# 先销毁旧的进程组(如果存在)
if dist.is_initialized():
dist.destroy_process_group()
# 重新初始化进程组
dist.init_process_group(backend='nccl')
# 加载 checkpoint
checkpoint = torch.load(checkpoint_path, map_location='cpu')
# 恢复模型和优化器
model.load_state_dict(checkpoint['model_state_dict'])
optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
一些实际踩过的坑
说了这么多理论,还是来点实际的踩坑记录。
坑一:checkpoint 文件损坏
有一天发现某个 checkpoint 文件加载时报错,用 torch.load() 直接抛出异常。检查后发现是文件写入过程中磁盘满了,导致文件不完整。这个问题很隐蔽——你可能保存时没报错,但恢复时就失败了。
解决方式是:保存完 checkpoint 后,再加载一遍验证完整性;或者保存时用临时文件,保存成功后再重命名。
def save_checkpoint_verified(state_dict, filename):
# 先保存到临时文件
temp_filename = filename + '.tmp'
torch.save(state_dict, temp_filename)
# 验证完整性
try:
torch.load(temp_filename)
# 验证通过,重命名
os.rename(temp_filename, filename)
except Exception as e:
logging.error(f'Checkpoint verification failed: {e}')
os.remove(temp_filename)
raise
坑二:显存不够加载 checkpoint
这个问题在微调大模型时特别常见。假设你在 4 张 A100 上训练了一个 7B 模型,checkpoint 文件几十 GB。现在你想在一台只有 1 张 A100 的机器上恢复继续训练,结果显存根本不够。
解决方式是使用分片 checkpoint,或者只加载部分层。PyTorch 的 torch.save() 支持分片保存,加载时可以按需加载。
# 保存时使用分片
torch.save(state_dict, 'checkpoint.pth', _use_new_zipfile_serialization=True)
# 加载时可以只加载部分
checkpoint = torch.load('checkpoint.pth', map_location='cpu')
# 只加载模型权重,不加载优化器状态
model.load_state_dict(checkpoint['model_state_dict'])
坑三:不同框架之间的 checkpoint 兼容性
如果你想在不同的框架之间迁移模型(比如从 PyTorch 迁移到 JAX),checkpoint 文件的结构可能完全不兼容。这种情况下,你通常需要手动解析权重并转换格式,或者使用专门的迁移工具。
坑四:checkpoint 版本管理
训练时间一长,checkpoint 文件会堆成山。文件名里带上 epoch、step、loss,再维护一份 checkpoint_metadata.json,找"最低 loss 那份"会省很多时间。
{
"checkpoint_epoch_10.pth": {
"epoch": 10,
"global_step": 10000,
"loss": 0.234,
"timestamp": "2026-07-17T10:30:00",
"is_best": false
},
"checkpoint_epoch_15.pth": {
"epoch": 15,
"global_step": 15000,
"loss": 0.198,
"timestamp": "2026-07-17T15:45:00",
"is_best": true
}
}
这样你可以很容易找到"损失最低的那个 checkpoint"或者"最近 24 小时内的那个 checkpoint"。
最终落地的方案
整理一下,一个比较完整的断点续训方案大概是这样:
checkpoint 保存策略:混合策略,关键 checkpoint 每个 epoch 保存一次,临时 checkpoint 每 1000 步保存一次。关键 checkpoint 保留最近 10 个,临时 checkpoint 保留最近 3 个。
checkpoint 内容:包含模型权重、优化器状态、调度器状态、当前 epoch 和 step、损失、随机数状态、scaler 状态。
保存验证:保存完后立即加载验证完整性,使用临时文件+重命名的方式避免不完整文件。
异常捕获:捕获所有 Python 异常并保存紧急 checkpoint;使用外部监控系统检测进程挂起和心跳文件;IO 异常时尝试降级保存到本地。
恢复对齐:恢复时同步恢复随机数状态、数据采样器状态、学习率调度器状态;分布式训练时重新初始化进程组。
元数据管理:维护
checkpoint_metadata.json记录每个 checkpoint 的详细信息,方便后续查找和选择。
这套方案在实际使用中效果还不错。虽然不能完全避免训练崩溃,但至少能保证崩溃后的损失最小化。更重要的是,它给了我一种安全感——不用担心跑几天的训练突然挂掉,然后又要从头开始。
写在最后
断点续训管的是崩溃后的恢复成本:混合保存策略、保存后校验、心跳监控、随机数和 scheduler 状态对齐,这几项落地后,跑几天的训练挂了也不用从头来。
短任务或稳定集群,每个 epoch 存一次可能就够。我们这套是为长任务和不稳定环境准备的——最痛的是训练跑三天突然挂掉,然后对着日志发呆。
版权声明: 本文首发于 指尖魔法屋-AI断点续训踩坑记录(https://blog.thinkmoon.cn/post/280-ai-resume-training-crash-recovery-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。