☰
PyTorch断点续训实战:checkpoint保存与恢复的完整指南
2026/9/30 6:37:59 网站建设 项目流程

训练跑了一大半,几十个小时的算力眼看就要出结果,结果登录服务器一看:进程没了。这种事我遇到过太多次,很多时候真不是你代码写得不对,而是训练环境本身就充满不确定性——远程会话断了、显存被别的任务占满触发OOM、机房临时断电,随便哪个意外都能让你前面的辛苦全部清零。早年间跑小规模实验,硬扛着从头再来也就多花一晚上;后来训练周期超过一周,我彻底受不了了,老老实实把断点续训做进了所有训练脚本里。

这篇文章就从工程实战的角度,完整聊一聊PyTorch断点续训怎么做。核心思路是什么、checkpoint里该存哪些东西、torch.save和torch.load有哪些容易踩的坑、多卡切换怎么处理,以及恢复训练后怎么验证“续上了”而不是“跑偏了”。不管你刚入门还是已经跑过不少实验,本文都会给出一套可以直接抄的模板。

1. 断点续训的核心思路:不是“加载权重”这么简单

1.1 恢复训练与冷启动的本质区别

很多人第一次做断点续训,第一反应是:把model.state_dict()保存下来,下次加载回模型不就行了吗?这么想没错,但这只是“加载权重”,不是真正的“断点续训”。这两者的区别在训练中后期特别明显。

模型权重只是训练状态的一部分。现代优化器比如Adam、AdamW,内部维护着一阶动量(也就是前几次梯度的加权均值)和二阶动量(梯度平方的加权均值),这些统计量是训练过程中逐步累积出来的。如果你只恢复模型参数而让optimizer从零开始,Adam会以为这是新开始训练,更新幅度会突然变大,参数更新的方向也可能偏掉,loss要么跳变、要么震荡好一阵子才能缓过来。在训练任务比较重的场景下,这个“好一阵子”往往就是一两天的算力。

同样的道理也适用于学习率调度器。CosineAnnealingLR、ReduceLROnPlateau这类调度器内部记录了当前周期、历史最优指标等信息,如果你从第40个epoch恢复,却让scheduler以为又是第1轮,那学习率曲线会整体错位,后面的训练节奏完全乱掉。断点续训的核心,是“把整套训练状态完整恢复”,而不是只把模型权重捞回来。

1.2 checkpoint里到底该保存哪些内容

我自己常用的做法,是把checkpoint设计成一个字典,里面至少包含以下几类内容:

  • model:模型权重,用model.state_dict()取
  • optimizer:优化器完整状态,含各类动量统计量
  • scheduler:学习率调度器状态
  • epoch / global_step:训练进度,用于恢复训练循环和日志记录
  • best_metric:当前最优指标,用于判断后续是否刷新最优模型
  • args或config:本次实验的超参数和环境信息
  • 随机数状态:random、numpy、torch、torch.cuda各自的RNG状态

为什么连随机数状态都要存?因为DataLoader的shuffle、数据增强策略都会引入随机数。如果训练中断前数据集正好加载到第2000个batch,恢复训练后从头shuffle或者直接从第0个batch开始,都不是同一套数据顺序,数据分布和增强组合都不一样,指标波动就很难解释。保存并恢复RNG状态,能让训练过程在数据层面上尽量衔接上。

有人觉得RNG状态那么大一块,存下来有必要吗?我的看法是:如果你的实验不需要对指标波动做非常严格的对比,那可以省掉;但如果你在研究过程中要复现某一次训练,或者需要精确判断“某一个改动是否导致loss下降”,那RNG状态就是必需品,否则连实验对照都做不了。当然这也意味着每次保存的checkpoint文件会更大,磁盘代价需要自己权衡。

1.3 文件命名与保存策略

checkpoint文件的命名和保存策略,看起来是小事,实际影响很大。我早期只用了一个checkpoint.pth来存,每次覆盖,结果有一次训练到一半发现最优模型在第20个epoch,但文件已经被覆盖成第25个epoch的数据了,想回头拿第20个epoch的权重去跑测试,只能遗憾错过。

现在我习惯在脚本里同时维护两类文件:

  • latest_checkpoint:记录最近一次保存的完整训练状态,用于断点续训
  • best_checkpoint:记录历史最佳指标对应的一次完整状态,用于模型评估和发布

latest文件用于“恢复”,best文件用于“最终取用”,两个逻辑分开,就不会出现“拿训练中断位置去当最优模型”的问题。另外,在磁盘充足的情况下,我还会保留最近1到2个历史checkpoint,防止latest文件写入一半出现损坏,至少能从上一个版本往回退几个epoch。

2. checkpoint相关的核心API与底层细节

2.1 torch.save与torch.load的正确打开方式

PyTorch里保存和加载的接口,核心就是torch.save和torch.load这两个函数,但用法上有很多讲究。以保存为例,推荐把整个训练上下文打包进一个字典:

checkpoint = { 'epoch': epoch, 'global_step': global_step, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'scheduler_state_dict': scheduler.state_dict(), 'best_metric': best_metric, 'config': config, 'random_state': random.getstate(), 'numpy_random_state': np.random.get_state(), 'torch_random_state': torch.random.get_rng_state(), 'cuda_random_state': torch.cuda.get_rng_state_all(), 'args': args, } torch.save(checkpoint, checkpoint_path)

这里有几个细节值得展开。第一,保存的是整个字典而不是单独的state_dict,这样恢复时能一次拿到所有上下文信息。第二,文件后缀用.tar或者.pth都行,内容本质都是pickle序列化后的字典,.tar后缀更符合很多人印象里的打包习惯。第三,我还会额外单独保存一份纯模型权重的model_final.pth或者model_best.pth,方便后续推理部署时只load模型,不需要碰优化器和调度器。

再说torch.load,最容易被忽略的是map_location参数和weights_only参数。如果你的checkpoint在GPU机器上保存,下次加载时机器没有GPU,或者GPU编号发生了变化,你需要用map_location='cpu'把存储映射到CPU上再取。如果明确要在多卡环境中加载,也可以通过map_location={'cuda:0': 'cuda:1'}做显式映射。

从PyTorch 2.6开始,torch.load的默认行为发生了变化,weights_only默认成了True,也就是说默认只允许加载基本数据类型和Tensor,不允许执行任意的Python类构造函数。这个改动是出于安全考虑,但很容易让老代码报错,尤其是当你的checkpoint字典里包含自定义类实例时。遇到这种报错,要么在加载时显式指定weights_only=False,要么改进存档结构,尽量只存tensor和纯基础类型。

2.2 optimizer与scheduler的状态恢复为什么关键

接着聊optimizer。optimizer.state_dict()内部对应每个参数保存了exp_avg、exp_avg_sq等键,分别对应一阶动量估计和二阶动量估计。如果这些信息丢了,优化器就得从头累积梯度统计量,模型参数即便没变,后续更新轨迹也会完全不一样。恢复时,optimizer.load_state_dict(optimizer_state)必须在optimizer定义之后、训练循环开始之前调用,并且顺序有讲究:先加载模型权重,再构造optimizer,再load optimizer状态,最后把optimizer里的状态张量移动到对应设备,这个顺序不能乱。

scheduler的恢复原理类似,但要注意的不是API问题,而是“恢复时机”。有些调度器会在每次调用step()时更新内部计数,比如CosineAnnealingLR需要知道自己当前处于第几个周期,加载完checkpoint后,如果又从epoch 0的循环开始跑,学习率就会从起点重新走一遍。好在scheduler_state_dict里已经包含了last_epoch这些信息,加载时一并恢复即可。

我自己踩过的一个坑是:加载了scheduler状态之后,在循环开头又手动调了一次scheduler.step()或者重新set last_epoch,等于把状态重置了,导致学习率曲线错乱。正确做法是:加载checkpoint后,直接用加载到的那一版epoch作为循环起点,不要在循环前额外触发任何更新操作。

2.3 多卡与分布式场景的额外细节

如果训练时用到了DataParallel或者DistributedDataParallel,保存和加载时会遇到module前缀问题。DataParallel包装后的模型,state_dict里的key会带上module.前缀,直接保存下来的checkpoint加载到普通模型上,会报unexpected key或者missing key的错误。

我的建议是,保存checkpoint时不要把外层包装也算进去,而是先拿到model.module.state_dict()再保存:

if isinstance(model, torch.nn.DataParallel): model_state = model.module.state_dict() else: model_state = model.state_dict()

恢复训练时,如果把带module前缀的checkpoint加载到单卡模型上,也可以做一个自动清洗,把key里的module.前缀剔除。这个问题在跨机换卡、单卡多卡切换时非常常见,值得在脚本里做一层防御逻辑。

3. 完整实操流程:从保存到恢复的一整套代码

3.1 训练循环中保存checkpoint的标准写法

这一节给出一套可以直接套用的模板,基于常见的PyTorch训练代码改造,假设已经写好了train_one_epoch()和validate()两段函数。关键点是把保存逻辑封装成一个函数,并固定在每个epoch验证后调用。

import os import random import numpy as np import torch import torch.nn as nn def save_checkpoint(state, is_best, save_dir, epoch): os.makedirs(save_dir, exist_ok=True) latest_path = os.path.join(save_dir, 'latest_checkpoint.tar') torch.save(state, latest_path) if is_best: best_path = os.path.join(save_dir, 'best_checkpoint.tar') torch.save(state, best_path)

在每个epoch结束后,把完整状态写入latest文件;如果当前指标优于best_metric,再额外写入best文件。这里有一点值得强调:不要在epoch中间频繁保存,也不要只在训练结束才保存。中间保存太频繁会影响训练速度;只在结束时保存等于没有,因为训练常常在途中就挂了。固定每1个或每2个epoch验证后保存一次,性价比最高。

如果使用周期性保存(比如每2000个step存一次),我同样建议把global_step也存进checkpoint,这样后续可以和TensorBoard的步数对应起来。在大规模训练中,每一点训练时间都很宝贵,保存过程建议用单独的线程执行,避免阻塞训练主进程,尤其当checkpoint文件有几个GB的时候,同步保存对训练吞吐的影响是肉眼可见的。

3.2 训练中断的具体场景与恢复入口

在恢复训练时,脚本一般通过命令行参数来控制:

parser.add_argument('--resume', type=str, default=None, help='path to latest checkpoint')

如果resume参数不为空,就执行恢复逻辑,否则从头开始。恢复逻辑的核心代码大概长这样:

if args.resume: checkpoint = torch.load(args.resume, map_location='cpu') 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 global_step = checkpoint['global_step'] best_metric = checkpoint['best_metric'] if 'random_state' in checkpoint: random.setstate(checkpoint['random_state']) np.random.set_state(checkpoint['numpy_random_state']) torch.random.set_rng_state(checkpoint['torch_random_state']) if torch.cuda.is_available(): torch.cuda.set_rng_state_all(checkpoint['cuda_random_state']) print(f"Resumed from epoch {checkpoint['epoch']}, best_metric: {best_metric:.4f}") else: start_epoch = 0 global_step = 0 best_metric = float('inf')

这里我特意把map_location='cpu'写在torch.load里,然后再把模型移动到GPU。原因是:如果直接在GPU上加载,当机器环境变化或GPU index变化时容易报设备错误。先全部映射到CPU,再通过model.to(device)统一迁移,逻辑最稳。加载完还要把optimizer里的动量状态张量也一并to到设备上,因为optimizer里的状态变量默认跟着参数设备走。

关于start_epoch的取值,我习惯把保存的epoch当成“已完成”的轮次,恢复时从epoch+1开始。这个取决于你自己定义的语义,只要训练循环、日志记录和scheduler保持一致即可,但要在全项目里统一,别一半从+1开始、一半从当前epoch开始。

3.3 训练循环的衔接与指标连续性验证

恢复训练后,训练循环要能从start_epoch直接往下跑,关键点在于dataloader的shuffle逻辑和batch的衔接。如果你在脚本里使用seed加shuffle,那么DataLoader的随机数序列和全局RNG有关。保存RNG状态后,即使重新创建DataLoader,只要在创建之前恢复了RNG状态,理论上能接上原来的数据顺序。

如果不想处理随机状态,也不想严格复现原始顺序,另一种可接受的做法是接受“重新shuffle”带来的微幅扰动,只要把global_step计数对齐,loss曲线在断点处可能会有小幅跳变,但对模型整体趋势影响不大。对大多数线上训练任务来说,这种方案完全够用,毕竟RNG状态会让checkpoint文件大不少。

恢复之后一定要做的验证动作是:先跑一个batch,确认模型前向、反向、优化器step都能正常执行,再看loss数量级是否与中断前一致。如果loss突然高出好几个数量级,多半是model权重和optimizer状态没有对齐,或者学习率被意外重置;如果loss比中断前低了很多,也别高兴太早——有可能是验证指标记录出了问题,或者加载错了文件,手动核对一下epoch编号更稳妥。

3.4 中断恢复后的日志与监控设计

断点续训不只是代码层面的事情,实验管理和监控同样重要。我在项目里习惯把训练日志写到独立log文件中,并且记录每条日志对应的epoch和global_step。这样恢复训练后,可以把前后日志拼接起来,在TensorBoard上看到连续的曲线,而不是曲线中间断了一截。

具体做法是,除了传统的print输出,我在save_checkpoint时额外写一条marker日志:

logger.info('SAVE CHECKPOINT at epoch {} step {}, best={:.4f}'.format( epoch, global_step, best_metric))

恢复训练后,通过grep这个marker就能快速定位上次保存位置。做好这些记录之后,即使checkpoint因为磁盘故障丢失,你也至少知道训练到哪一步,可以从最近的实验日志和单独保存的model权重做二次恢复。

4. 常见问题快查:断点续训的坑我都踩过

4.1 checkpoint加载时报“Missing key”或“Unexpected key”

这个问题九成九来自DataParallel、DDP的module前缀,或者模型结构定义不一致。处理方法是做一层key归一化清洗:

state_dict = checkpoint['model_state_dict'] new_state_dict = {} for k, v in state_dict.items(): if k.startswith('module.'): k = k[7:] new_state_dict[k] = v model.load_state_dict(new_state_dict)

还有一个常见原因是恢复训练前改动过模型结构,比如改了输出层维度、增加了一个分支。这个没有特别完美的解法,除非你明确知道改动点,否则建议谨慎处理:结构变了,优化器状态也可能对不上,要么选择从旧checkpoint中提取能对应上的层,要么直接在修改结构后重新训练,别硬续。依赖安装差异、PyTorch版本差异也可能导致模型定义里某些类无法加载,这类问题排查时先看完整报错堆栈。

4.2 恢复训练后loss不降反升,通常是什么原因

第一个要排查的就是optimizer状态是否恢复成功。如果只恢复了model权重的state_dict,optimizer里的动量全被重新初始化,Adam更新规则中分母(二阶动量开方)几乎为零,有效步长会被放大,loss就可能往上跳。第二种常见原因是学习率调度器被重置,比如scheduler的last_epoch和当前循环不同步,导致学习率曲线错位。第三种是随机数状态没恢复,数据增强方式出现明显差异。如果这些都排查过还是不对,可以把恢复后的学习率先降一个量级,跑几十个step,再把学习率调回脚本设定值,给optimizer一个平滑过渡的机会。

4.3 旧版PyTorch的checkpoint在新版本环境加载失败

PyTorch的checkpoint本质上是pickle序列化的,和具体类定义强相关。如果你从PyTorch 2.0切到2.6,直接torch.load通常问题不大,但如果涉及自定义类,或者从旧版本换到新版本后类的导入路径变了,就会出问题。我建议在保存checkpoint之外,额外保存一份纯tensor格式的模型权重(model.state_dict()),这样即使整个checkpoint因为版本问题无法加载,也可以用纯权重文件重新构建训练状态,最多丢失优化器状态。

这里说的版本迁移,也同样适用于环境迁移场景。很多人在线上训练用A环境,离线调试用B环境,如果两个环境里PyTorch和CUDA版本差很多,推荐的做法是尽量保持一致。断点续训最怕的就是中途换环境,换完之后所有状态都“好像”加载成功了,实际上数值精度、随机数分布早就不一样了,最后的模型效果也很难复现。

4.4 checkpoint写入一半断电导致文件损坏

训练中断本来就很闹心,如果checkpoint文件本身还损坏,那就更难受了。我的做法是“先写临时文件,再原子替换”:

tmp_path = checkpoint_path + '.tmp' torch.save(state, tmp_path) os.replace(tmp_path, checkpoint_path)

torch.save先写入.tmp结尾的文件,全部写完再用os.replace替换成正式文件。os.replace在绝大多数文件系统上是原子操作,即使中途断电,要么有完整的旧文件,要么有完整的新文件,不会出现只有半截数据的情况。再配合保留最近2个历史版本,即使latest损坏,也可以用上一个历史版本恢复,最多损失三两个epoch的进度。

4.5 多卡切换单卡,或者反过来,如何处理

多卡切单卡最容易出现的就是module前缀和参数尺寸不一致的问题。另一个隐患是DDP的梯度同步机制恢复后无法原样重现,但这不是大问题,只要优化器状态恢复正确,训练依然能继续。如果你从单卡checkpoint恢复多卡训练,推荐在加载后用broadcast把模型参数同步到所有进程,避免因为初始参数不一致导致多卡训练发散。脚本里最好把世界大小(world_size)也记录在checkpoint中,恢复时能做一致性校验。

5. 一些值得补充的工程化建议

5.1 把恢复训练做成命令行的一等公民

我的训练脚本一开始没好好设计resume参数,导致每次想恢复训练都要改代码、改硬编码路径,非常容易出错。后来我统一在入口处增加args.resume,所有恢复逻辑都集中在脚本开头,只在resume参数有值时执行。这不仅是断点续训本身的需求,也是一种工程规范,让脚本能直接复用到不同项目、不同数据集上。脚本里同时建议增加--local_rank这类参数,方便后续在多卡环境下统一管理。

5.2 记录环境版本与依赖到checkpoint里

为了防止过了两个月你自己都忘了当时用的PyTorch版本,建议在保存checkpoint时把环境信息也写进去:

checkpoint['pytorch_version'] = torch.__version__ checkpoint['cuda_version'] = torch.version.cuda if torch.cuda.is_available() else None checkpoint['python_version'] = sys.version

平时看着没什么用,一旦出现版本问题,这些信息就是排查的第一手线索。尤其对于跨机器、跨团队协作的项目,别人拿到你的checkpoint后,先看版本信息再决定怎么加载,能省掉很多折腾时间。

5.3 定时同步checkpoint到远端或云盘

我有个教训:模型训练到第40个epoch,本地服务器磁盘突然故障,幸好checkpoint放在另一块盘上,才没有全丢。自那以后,我用一个简单的定时任务,每小时把latest_checkpoint同步到另一个节点或云盘。这个习惯成本很低,但关键时刻能救命。如果条件有限,至少要在关键节点(比如每个10倍epoch)把best_checkpoint手动拉一份到本地备份,避免单点故障。


最后再分享一个我自己用下来的习惯:每次恢复训练后,我都会把加载到的epoch、best_metric、学习率等关键信息打印出来,确认一遍再进入训练循环。这个习惯看起来简单,实际上帮我排掉了大量低级错误。断点续训本质上没有什么高深的技术,难的是把各种边界情况都想到、把状态完整保存下来。希望这篇内容能帮你避开这些坑,把训练中断的成本降到最低。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询