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 包含的各个组件以及它们之间的关系:

graph TD A[Checkpoint 文件] --> B[模型权重] A --> C[优化器状态] A --> D[学习率调度器状态] A --> E[训练进度信息] A --> F[数据采样器状态] A --> G[随机数状态] A --> H[混合精度 Scaler] A --> I[分布式进程组状态] C --> C1[Adam 动量缓存] C --> C2[Adam 方差缓存] D --> D1[当前 step] D --> D2[调度器参数] E --> E1[当前 epoch] E --> E2[当前 batch 索引] E --> E3[当前 loss] F --> F1[RandomSampler 状态] F --> F2[数据顺序索引] G --> G1[Python random 状态] G --> G2[NumPy random 状态] G --> G3[PyTorch random 状态] G --> G4[CUDA random 状态] style A fill:#e1f5ff style B fill:#fff4e6 style C fill:#fff4e6 style D fill:#fff4e6 style E fill:#fff4e6 style F fill:#fff4e6 style G fill:#fff4e6 style H fill:#fff4e6 style I fill:#fff4e6

要点: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

这张图展示了异常检测与恢复的整体流程:

flowchart TD Start([训练启动]) --> LoadCheckpoint{检查<br/>checkpoint} LoadCheckpoint -->|存在| Resume[从 checkpoint 恢复状态] LoadCheckpoint -->|不存在| Init[初始化新训练] Resume --> TrainingLoop Init --> TrainingLoop TrainingLoop[训练循环] --> NormalBatch[正常批次训练] NormalBatch --> SaveCheckpoint{是否到<br/>保存时机?} SaveCheckpoint -->|是| Save[保存 checkpoint] SaveCheckpoint -->|否| CheckException Save --> CheckException CheckException{检测到异常?} CheckException -->|Python 异常| Emergency[保存紧急 checkpoint] CheckException -->|OOM Killer| Monitor[外部监控检测心跳] CheckException -->|网络异常| Fallback[降级保存到本地] CheckException -->|无异常| Continue[继续下一批次] Emergency --> Restart Fallback --> Continue Monitor --> HeartbeatCheck{心跳超时?} HeartbeatCheck -->|是| Restart HeartbeatCheck -->|否| Continue Restart([重启进程]) --> LoadCheckpoint Continue --> TrainingLoop style Start fill:#e1f5ff style Restart fill:#ffe1e1 style Emergency fill:#fff4e6 style Fallback fill:#fff4e6 style Monitor fill:#fff4e6 style LoadCheckpoint fill:#fff9e6 style CheckException fill:#fff9e6 style SaveCheckpoint fill:#fff9e6 style HeartbeatCheck fill:#fff9e6

这张图展示了三种异常的处理路径: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"。

最终落地的方案

整理一下,一个比较完整的断点续训方案大概是这样:

  1. checkpoint 保存策略:混合策略,关键 checkpoint 每个 epoch 保存一次,临时 checkpoint 每 1000 步保存一次。关键 checkpoint 保留最近 10 个,临时 checkpoint 保留最近 3 个。

  2. checkpoint 内容:包含模型权重、优化器状态、调度器状态、当前 epoch 和 step、损失、随机数状态、scaler 状态。

  3. 保存验证:保存完后立即加载验证完整性,使用临时文件+重命名的方式避免不完整文件。

  4. 异常捕获:捕获所有 Python 异常并保存紧急 checkpoint;使用外部监控系统检测进程挂起和心跳文件;IO 异常时尝试降级保存到本地。

  5. 恢复对齐:恢复时同步恢复随机数状态、数据采样器状态、学习率调度器状态;分布式训练时重新初始化进程组。

  6. 元数据管理:维护 checkpoint_metadata.json 记录每个 checkpoint 的详细信息,方便后续查找和选择。

这套方案在实际使用中效果还不错。虽然不能完全避免训练崩溃,但至少能保证崩溃后的损失最小化。更重要的是,它给了我一种安全感——不用担心跑几天的训练突然挂掉,然后又要从头开始。

写在最后

断点续训管的是崩溃后的恢复成本:混合保存策略、保存后校验、心跳监控、随机数和 scheduler 状态对齐,这几项落地后,跑几天的训练挂了也不用从头来。

短任务或稳定集群,每个 epoch 存一次可能就够。我们这套是为长任务和不稳定环境准备的——最痛的是训练跑三天突然挂掉,然后对着日志发呆。

版权声明: 本文首发于 指尖魔法屋-AI断点续训踩坑记录https://blog.thinkmoon.cn/post/280-ai-resume-training-crash-recovery-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!