本文目录导读:

“断点续训”是深度学习和模型训练中非常实用的机制,它允许你暂停一段耗时很长的训练过程(比如因为断电、Out of Memory、或者只是下班了),并在之后从中断的地方重新开始,而不是从头再来。
核心思想
它的核心逻辑是:定期保存模型的状态快照(Checkpoint / 检查点),当训练中断后,加载最近的快照,恢复模型、优化器、学习率调度器等所有状态,继续训练。
为什么需要断点续训?
- 节省时间和成本:大模型训练可能需要数天甚至数周,一次中断可能意味着成千上万的算力成本(GPU/TPU时间)白白浪费,断点续训是救命的。
- 容错能力:训练环境不稳定(云服务器可能重启、网络故障、内存溢出、硬件故障等),断点续训是保证训练最终能够完成的必要手段。
- 超参数调整:如果你在训练中途想调整学习率或其他参数,可以先保存当前状态(checkpoint),调整代码中的参数,然后从该checkpoint恢复训练,而无需从头开始。
- 资源管理:可以灵活地利用非工作时间或低负载时间段进行训练,并在资源被抢占时(如共享GPU集群)从容保存。
具体实现机制
通常在深度学习框架(如 PyTorch、TensorFlow/Keras、JAX)中,断点续训需要保存和恢复三大部分的信息:
需要保存的内容(Checkpoint 文件)
一个完整的 Checkpoint 通常是一个字典或文件,包含以下关键组件:
- 模型参数(Model State Dict):模型的权重和偏置 (
model.state_dict())。 - 优化器状态(Optimizer State Dict):优化器的动量和缓存(如 Adam 的
exp_avg,exp_avg_sq),这是恢复训练效果的关键,如果不恢复优化器状态,恢复后的第一次更新可能会产生较大波动,导致训练不稳定。 - epoch / step 计数器:当前训练进行到的轮次(epoch)和迭代步(iteration/step),用于正确继续下一个 epoch 的 DataLoader,以及恢复学习率调度器的状态。
- 学习率调度器状态(Scheduler State Dict):如果使用了动态学习率(如 StepLR、CosignAnnealing),需要保存当前的 epoch/step 和优化器状态,以便调度器继续计算正确学习率。
- 损失值/其他监控指标:当前验证集的最佳 Loss 或 Accuracy,用于后续 Early Stopping 和模型选择。
- 随机数生成器状态(可选):如果希望训练结果完全可复现,应保存 Python、NumPy、PyTorch 的随机种子状态。
保存策略(Saving Policy)
不是每个 step 都保存,否则文件会爆炸,常用策略:
- 固定间隔保存:每 N 个 epoch 或每 M 个全局步(global step) 保存一次,比如每 1000 steps 或每 1 epoch 保存。
- 最佳模型保存:仅在验证集上达到新的最高性能(如最低损失、最高准确率)时保存,通常会保存“最佳模型”的单独文件。
- 滚动保存 + 最佳保存:保留最近的 K 个 Checkpoint(例如最近5个),同时保留一个始终指向历史最佳的那个,这样做可以避免磁盘爆满,同时允许回滚到任何历史检查点。
恢复流程(Resume / Restore)
当训练中断重新启动时,脚本需要:
- 检查是否存在 Checkpoint 文件:通常是一个特定路径(如
./checkpoints/checkpoint_epoch_10.pt)或一个指向最新 Checkpoint 的符号链接(如./checkpoints/latest.pt)。 - 加载 Checkpoint:如果有,加载它到内存。
- 重建模型、优化器、调度器实例:使用相同的架构、配置。
- 加载状态:将保存的状态字典分别加载到
model.load_state_dict()、optimizer.load_state_dict()、scheduler.load_state_dict()。 - 设置起始 epoch/step:从已保存的 epoch/step 开始循环 DataLoader(注意:DataLoader 本身不会记住上次读到了哪个样本,需要手动或通过
DistributedSampler来跳过已处理的数据,或直接从step开始迭代)。 - 继续训练:开始后续的训练循环。
代码示例 (PyTorch 风格)
import torch
import os
# 参数
checkpoint_dir = './checkpoints'
model = MyModel()
optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)
scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=30, gamma=0.1)
start_epoch = 0
best_loss = float('inf')
def save_checkpoint(state, filename='checkpoint.pth.tar'):
"""保存checkpoint"""
# 确保目录存在
os.makedirs(checkpoint_dir, exist_ok=True)
filepath = os.path.join(checkpoint_dir, filename)
torch.save(state, filepath)
print(f"Checkpoint saved to {filepath}")
def load_checkpoint():
"""加载最新的checkpoint,返回起始epoch和状态"""
latest_path = os.path.join(checkpoint_dir, 'latest.pth.tar')
if os.path.isfile(latest_path):
print(f"Loading checkpoint from {latest_path}")
checkpoint = torch.load(latest_path)
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 # 从下一轮开始
best_loss = checkpoint['best_loss']
# 如果有,也加载随机种子
# torch.set_rng_state(checkpoint['rng_state'])
print(f"Resumed from epoch {checkpoint['epoch']}")
return start_epoch, best_loss
else:
print("No checkpoint found, starting from scratch.")
return 0, float('inf')
# --- 主训练循环 ---
if __name__ == '__main__':
start_epoch, best_loss = load_checkpoint()
for epoch in range(start_epoch, NUM_EPOCHS):
train_one_epoch(model, optimizer, train_loader)
val_loss = validate(model, val_loader)
# 更新学习率调度器 (按epoch或step更新)
scheduler.step()
# 构建保存的状态字典
state = {
'epoch': epoch,
'model_state_dict': model.state_dict(),
'optimizer_state_dict': optimizer.state_dict(),
'scheduler_state_dict': scheduler.state_dict(),
'best_loss': best_loss,
# 'rng_state': torch.get_rng_state() # 如果需要随机状态
}
# 1. 保存最新的checkpoint (用于恢复)
save_checkpoint(state, 'latest.pth.tar')
# 2. 保存最佳模型 (用于推理)
if val_loss < best_loss:
best_loss = val_loss
save_checkpoint(state, 'best_model.pth.tar')
print(f"Epoch {epoch+1}/{NUM_EPOCHS}, Loss: {val_loss:.4f}, Best Loss: {best_loss:.4f}")
常见问题与注意事项
- DataLoader 的循环问题:直接
for epoch in range(start_epoch, ...)是正确的方式,但如果你的 DataLoader 使用了DistributedSampler,需要为每个 epoch 设置sampler.set_epoch(epoch),否则数据 shuffle 顺序会不一致(但通常不影响最终结果,只是影响复现性),更精细的做法是保存 DataLoader 的内部迭代器状态(不常见,比较复杂)。 - 分布式训练 (DDP):在 DDP 中,每个进程(GPU)都需要保存自己的 Checkpoint,一般做法是只让一个进程(通常是 rank 0)来负责保存,恢复时,所有进程加载相同的 Checkpoint(包含 model state dict),并且各自恢复优化器状态(如果优化器也是分布式的,则只从 rank 0 同步)。
- 文件系统与原子性:保存 Checkpoint 时,如果写入一半断电,文件会损坏,常见的做法是:先写入一个临时文件(如
checkpoint.pth.tmp),然后使用os.rename()(在 Linux/Unix 上通常是原子操作)重命名为最终文件名。 - 磁盘空间:频繁保存或保存大模型时,磁盘可能很快填满,务必实现滚动删除(如只保留最近 K 个 checkpoint)或压缩存储。
- 后处理——导出推理模型:恢复训练用的 Checkpoint 通常包含了优化器状态和额外信息(体积较大)。推理时,只需要加载模型参数(可能还需要经过
model.eval()),断点续训的 Checkpoint 和用于推理的模型文件是两回事,最佳实践是:在 save_checkpoint 后,额外导出一个只包含模型参数的小文件(model.pth),供部署使用。
| 组件 | 是否必须保存 | 作用 |
|---|---|---|
| 模型参数 | ✅ 必须 | 恢复模型权重,这是训练的核心 |
| 优化器状态 | ✅ 强烈建议 | 恢复动量/缓存,保证训练稳定和收敛 |
| 当前 epoch/step | ✅ 必须 | 知道从何处继续,skip 已处理的 epoch |
| 学习率调度器状态 | ✅ 强烈建议 | 恢复学习率变化轨迹 |
| 最佳验证指标 | ✅ 推荐 | 用于 Early Stopping 和模型选择 |
| 随机种子 | ❌ 可选 | 如果需要完全复现结果 |
断点续训是工业级训练管线的标配能力。 实现它不仅是为了防止意外中断,更是为了提高训练灵活性(超参数调整、资源调度)。