深度模型的保存,表面上看就是一行 torch.save 的事,但真正在生产里跑过一轮的人都知道,这里面藏着能让你加班到凌晨的坑。我自己就遇到过接手同事的项目,看到目录里躺着 model.ckpt、best.pth、checkpoint-epoch12.pth 三个文件,一时搞不清楚哪个是权重、哪个是断点、哪个能直接拿去推理,最后靠比对文件大小和用 torch.load 一个个打印 key 才把关系理清楚。这篇文章就围绕 ckpt 和 pth 这两个高频后缀,把深度模型保存这件事从"文件长什么样"讲到"工程上该怎么定策略",中间会把 state_dict、断点续训、map_location、参数不匹配排查这些实操细节全部铺开。不管你是刚跑完第一个 MNIST 的新手,还是已经在带团队做模型交付的老手,都能在这里找到能直接抄走的代码模板和避坑经验。
1. 先弄明白:ckpt 和 pth 到底是不是两种格式
1.1 后缀只是习惯,真正决定内容的是你存进去的对象
很多人第一反应是"ckpt 和 pth 是两种不同的文件格式",这个理解在 90% 的情况下是错的。PyTorch 的torch.save底层走的是 pickle 序列化加 zip 打包,文件名后缀写什么它根本不管,你写成.model、.bin、.weights都能正常读回来。官方文档的示例里习惯用.pt和.pth,而 TensorFlow 系的 checkpoint 机制习惯用.ckpt,PyTorch Lightning 为了和 TF 的习惯对齐,也把自己的断点文件默认命名为.ckpt。于是就有了"看到 ckpt 以为是 TF,看到 pth 以为是 PyTorch"这种经验主义判断,实际上并不可靠。
真正决定一个文件里装了什么,是你传给torch.save的那个对象。它可能是整个nn.Module实例,可能是一个state_dict字典(也就是"参数名到张量"的映射),也可能是一个包含了模型权重、优化器状态、当前 epoch、学习率调度器状态的复合字典。这三种东西虽然都能叫.pth,但加载方式和适用场景完全不同。所以我判断一个权重文件,从来不看后缀,而是先跑一段探测代码,把它里面的 key 打出来。
import torch ckpt = torch.load("model.pth", map_location="cpu", weights_only=True) print(type(ckpt)) if isinstance(ckpt, dict): print(list(ckpt.keys())[:20])如果打印出来是一堆layer1.0.conv1.weight这种名字,那它就是纯 state_dict;如果打印出来是epoch、optimizer、lr_scheduler、model这种顶层 key,那就是断点文件;如果打印出来直接是一个 Module 对象,那就是整包保存。三种情况对应的加载写法完全不一样,这也是后面所有坑的源头。
1.2 state_dict 和整个模型对象的本质区别
理解 ckpt 和 pth 的差异,绕不开 state_dict 这个概念。nn.Module的state_dict()返回的是一个有序字典,里面只有可学习的参数和注册过的 buffer(比如 BatchNorm 的 running_mean、running_var),不包含网络结构本身。这就意味着,光有 state_dict 你没有模型定义代码是跑不起来的,因为 PyTorch 不知道这些张量该往哪个层里塞。
而torch.save(model, path)走的是另一条路,它用 pickle 把整个 Module 对象连同它的类定义引用一起序列化。听起来很方便,问题是 pickle 保存的是"类的引用路径",比如mymodels.resnet.CustomResNet。你把文件拷到另一台机器,如果那个模块路径不存在、或者类定义改了、或者文件夹结构变了,加载直接报ModuleNotFoundError或者AttributeError。这种耦合在单人实验环境里还好,一旦进入多人协作或者模型交付环节,就是灾难。
我在团队里推的规则很明确:只保存 state_dict,永远不保存整个模型对象。理由有三条。第一,可移植,换机器换目录都不影响;第二,文件小,pickle 整个对象会把一些冗余的 Python 属性也带上;第三,安全,反序列化一个完整的类对象比反序列化一个纯张量字典的风险高得多。代价是你必须维护模型定义代码,但这个代价在工程上完全可以接受,因为模型结构本来就应该进版本管理。
1.3 不同生态下 ckpt 的含义差异
虽然 PyTorch 系也能随便叫 ckpt,但在 TensorFlow/Keras 的世界里,ckpt 是一个有明确协议的东西。Keras 的 ModelCheckpoint 回调默认生成的是一组文件:.ckpt-5.index、.ckpt-5.data-00000-of-00001,还可能带一个checkpoint文本文件记录最新的是哪一步。这一组文件必须放在一起,缺一个都读不了。看到这种多文件结构,基本可以确定是 TensorFlow 系的产物,和 PyTorch 的单文件 ckpt 不是一回事。
还有一种情况是 Hugging Face Transformers 保存出来的权重,通常是pytorch_model.bin或model.safetensors,配一份config.json。有些人为了统一命名,手动把它改名成.ckpt,这就更让人迷惑了。所以养成习惯:拿到陌生权重文件,先看它是单文件还是多文件组,再看它有没有配套的 config,最后用探测代码确认内容结构,三步走下来就不会认错。
| 来源生态 | 典型文件名 | 是否单文件 | 内部结构 |
|---|---|---|---|
| PyTorch 手写训练脚本 | model.pth / best.ckpt | 单文件 | 通常为 state_dict 或复合字典 |
| PyTorch Lightning | epoch=12-step=900.ckpt | 单文件 | 复合字典,含 hyper_parameters |
| TensorFlow / Keras | model.ckpt-5.index + .data | 多文件组 | 图变量,需配套 meta |
| Hugging Face | pytorch_model.bin / .safetensors | 单文件加 config | state_dict,key 命名有前缀 |
| TorchScript 导出 | traced_model.pt | 单文件 | 可执行图,非字典 |
这张表建议收藏,遇到陌生文件先对号入座,能省掉大量试错时间。
2. PyTorch 三种保存姿势的取舍逻辑
2.1 整包保存:方便是真方便,坑也是真坑
torch.save(model, "full_model.pth")这种写法在教程里出现频率极高,因为它加载时只要一行model = torch.load("full_model.pth")就完事,不需要你在加载侧再写一遍网络定义。对于做快速验证、写 demo、跑课设作业的场景,它确实省事。
但它的三个硬伤决定了它上不了生产。第一是环境耦合,我踩过一次特别典型的坑:实验机上用 PyTorch 1.12 保存的整包模型,换到只有 1.8 的服务器上加载,报了一长串AttributeError,原因是某些层在序列化时记录了旧版本的内部属性。第二是代码耦合,模型类只要改了构造函数签名、加了一个新参数,旧文件就可能加载失败,因为 pickle 在重建对象时会重新调用__init__。第三是 pickle 反序列化的安全风险,pickle 文件在加载时是可以执行任意代码的,来源不明的权重文件绝对不能直接 load,这一点后面第 4 章会展开讲。
还有一种情况是"保存了整个对象,但只想取权重",这时候你会写torch.save(model.state_dict(), ...),注意这已经切换成第二种姿势了。我见过不少人在同一个项目里两种混用,结果加载时一半报错一半正常,排查起来非常痛苦。所以定一条死规矩:一个项目只用一种保存姿势,团队内统一。
2.2 state_dict 才是工程上的标准姿势
标准做法是把模型和权重分开对待,结构定义写在代码里,权重单独存文件。保存端就一行:
torch.save(model.state_dict(), "best.pth")加载端需要两步,先实例化同结构的模型,再灌权重:
model = build_model(num_classes=10) # 结构必须和训练时完全一致 state = torch.load("best.pth", map_location="cpu", weights_only=True) model.load_state_dict(state) model.eval()这里load_state_dict默认是严格模式strict=True,要求两边的 key 集合完全一致,多一个少一个都会抛错。这个默认值其实是好事,它帮你把结构不匹配的问题在第一时间暴露出来,而不是悄悄加载了一半参数然后推理结果莫名其妙。如果你确实需要放宽,比如只加载骨干网络、分类头重新初始化,那就显式写strict=False,然后务必打印缺失和多余的 key 列表确认一遍,别糊里糊涂就往下跑。
map_location这个参数是我认为最被低估的一个。它的作用是告诉 PyTorch 把张量映射到哪个设备上。GPU 上保存的权重直接在没有 GPU 的机器上加载,不加map_location会直接报 CUDA 不可用;加了map_location="cpu"就能正常读进来。常见的还有map_location="cuda:1"指定卡号,或者map_location={"cuda:0": "cuda:1"}做设备重映射。多卡场景下这个参数用得特别多,后面细说。
2.3 断点续训要保存的东西远不止权重
如果你只保存 state_dict,训练中断之后重新开始,优化器的动量、自适应学习率的历史累积、学习率调度器走到第几步、AMP 的梯度缩放因子这些全部丢失。表面上看模型能继续训,实际上收敛轨迹已经变了,尤其是在训练后期,优化器状态的重要性不比权重低。
所以断点文件的正确打开方式是存一个复合字典。我常用的模板长这样:
def save_checkpoint(path, model, optimizer, scheduler, scaler, epoch, best_metric, cfg): torch.save({ "epoch": epoch, "model": model.state_dict(), "optimizer": optimizer.state_dict(), "scheduler": scheduler.state_dict() if scheduler else None, "scaler": scaler.state_dict() if scaler else None, "best_metric": best_metric, "config": cfg, }, path)注意这里用了model.state_dict()而不是model,外层字典的 key 我叫model,所以完整的结构是"字典里有 model、optimizer 等子字典"。这种嵌套结构在加载时容易写错,很多人会写成state_dict["state_dict"]或者state_dict["model_state_dict"],其实取决于你保存时怎么命名的。这也是为什么我在团队里要求保存的 key 名固定为model、optimizer、scheduler、epoch、best_metric五个,不允许自由发挥,谁改谁负责同步所有下游加载代码。
顺带说一句config这个字段。把训练配置(学习率、batch size、数据增强参数、类别数)一起存进 ckpt,好处是一年后你翻出这个文件还能复现实验,坏处是如果配置里塞了不可序列化的对象(比如数据集实例、lambda 函数)就会保存失败。我的做法是只存基础类型和列表字典,复杂对象转成字符串描述。
3. 一套可以直接复现的保存与加载实操流程
3.1 保存端:区分 best 和 last 两个文件
训练脚本里我通常维护两个文件:last.pth每个 epoch 覆盖写,best.pth只在指标刷新时写。这样做的原因很实际:断点续训要用 last,模型交付要用 best,两个用途分开,互不干扰。如果只存一个文件,你会在"继续训练"和"保留最好模型"之间反复纠结。
目录结构建议这样组织:
runs/exp_20240612/ ├── config.yaml ├── last.pth ├── best.pth └── log.txt每个实验一个时间戳目录,两个权重文件加一份配置。这个结构看起来朴素,但它解决了一个大问题:三个月后你回来看结果,不用去猜best.pth是哪次实验、什么参数。我见过太多人把十几个实验的 best.pth 全堆在一个目录里,最后靠文件修改时间排序,那画面太心酸了。
保存 best 的判断逻辑本身也有讲究。分类任务常用准确率或 F1,我一般用验证集 F1,因为它在类别不均衡时比准确率更能反映真实水平。检测类任务可能用 mAP,分割类用 mIoU。注意阈值方向要写对,if metric > best_metric和if metric < best_metric差了十万八千里,我在早期代码里犯过一次反号错误,结果保存下来的 best.pth 是整个训练过程中最差的那个,白白浪费了一晚上算力。
3.2 加载端:三种场景的写法对照
加载场景可以归成三类,写法各有侧重。
纯推理,只关心权重,其他一律不管:
model = build_model(num_classes=cfg["num_classes"]) state = torch.load("best.pth", map_location="cpu", weights_only=True) model.load_state_dict(state["model"] if "model" in state else state) model.eval() with torch.no_grad(): out = model(x)注意这里我加了一个"model" in state的判断,用来兼容"纯 state_dict"和"复合字典"两种情况。这种容错写法在接手别人项目时特别有用,能少写两行调试代码。生产环境我还是建议解析清楚再写死,别把容错逻辑留在线上。
继续训练,需要完整恢复优化器状态:
ckpt = torch.load("last.pth", map_location="cpu", weights_only=False) model.load_state_dict(ckpt["model"]) optimizer.load_state_dict(ckpt["optimizer"]) scheduler.load_state_dict(ckpt["scheduler"]) start_epoch = ckpt["epoch"] + 1这里weights_only=False是因为优化器状态里包含了 param_groups 这类结构,纯张量模式读不了。这一点非常关键,PyTorch 2.6 起torch.load的weights_only默认值变成了 True,如果你还在用旧代码加载断点文件,升级后会突然报错,提示反序列化被拒绝。解决办法就是显式传weights_only=False,同时确认文件来源可信。
迁移学习,加载骨干、换掉分类头:
state = torch.load("pretrained.pth", map_location="cpu", weights_only=True) missing, unexpected = model.load_state_dict(state, strict=False) print("missing:", missing) print("unexpected:", unexpected)strict=False会返回两个列表,missing 是你模型里有但权重文件里没有的 key,unexpected 反过来。迁移学习场景下你期望看到的 missing 应该正好是新的分类头,unexpected 应该正好是旧的分类头。如果 missing 里出现了骨干网络的层名,说明结构对不上,得回头检查。
3.3 断点续训的完整实现与恢复点选择
把前面几块拼起来,一个能用的训练循环大概是这个样子:
start_epoch = 0 best_metric = 0.0 resume_path = "runs/exp/last.pth" if os.path.exists(resume_path) and args.resume: ckpt = torch.load(resume_path, map_location="cpu", weights_only=False) model.load_state_dict(ckpt["model"]) optimizer.load_state_dict(ckpt["optimizer"]) if ckpt.get("scheduler"): scheduler.load_state_dict(ckpt["scheduler"]) if ckpt.get("scaler"): scaler.load_state_dict(ckpt["scaler"]) start_epoch = ckpt["epoch"] + 1 best_metric = ckpt.get("best_metric", 0.0) print(f"resumed from epoch {start_epoch}") for epoch in range(start_epoch, total_epochs): train_one_epoch(...) metric = evaluate(...) save_checkpoint("runs/exp/last.pth", model, optimizer, scheduler, scaler, epoch, max(best_metric, metric), cfg) if metric > best_metric: best_metric = metric torch.save({"model": model.state_dict(), "epoch": epoch, "metric": metric}, "runs/exp/best.pth")这段代码里有几个细节值得强调。第一,start_epoch = ckpt["epoch"] + 1的加一是必须的,我见过漏掉加一的写法,导致恢复后同一个 epoch 训了两遍,数据采样和调度器步数都会错位。第二,best_metric也要跟着恢复,否则恢复后第一次验证不管多差都会被判定为"新最优",直接把真正的好模型覆盖掉。第三,如果用了自动混合精度,scaler的状态必须一起存,因为它的缩放因子是动态累积的,重置后前几十步的梯度会不稳。
还有一个容易忽略的点是数据加载器的随机种子。断点续训的理想状态是"从断点处严格继续",但 DataLoader 的 shuffle 顺序、各种数据增强的随机数,如果不做种子管理,恢复后的数据流和中断前是不一样的。严格复现需要保存随机数生成器状态,通常在研究场景才这么干,工程上大家接受少量偏差。但如果你在做对复现性要求很高的对比实验,那就得把torch.get_rng_state()、numpy.random.get_state()也塞进 ckpt。
3.4 从 pth 到 pt、TorchScript、ONNX 的转换思路
经常有人问"pth 是不是从 pt 导出的",这个问题其实建立在"pt 和 pth 是两种格式"的误会上。在 PyTorch 的语境里,.pt和.pth从来都是同一种东西的两种叫法,官方示例里两者混用,不存在"导出"关系。真正存在转换关系的是另外几组:state_dict 到 TorchScript、state_dict 到 ONNX。
TorchScript 导出是为了脱离 Python 环境部署,生成的文件通常叫traced.pt,它是一个可以独立执行的图,不依赖你的模型定义代码:
model.eval() example = torch.randn(1, 3, 224, 224) traced = torch.jit.trace(model, example) traced.save("traced.pt")这里model.eval()是必须的,因为 trace 会把当时的计算路径固定下来,如果模型在训练模式下,Dropout 和 BatchNorm 的行为会被错误地烘焙进图里,推理结果全乱。这个坑非常隐蔽,模型不会报错,只是结果不对。
ONNX 导出用于跨框架部署:
torch.onnx.export( model, example, "model.onnx", input_names=["input"], output_names=["output"], dynamic_axes={"input": {0: "batch"}, "output": {0: "batch"}}, opset_version=12, )dynamic_axes这行很关键,不写的话导出的模型 batch 维度是写死的,部署时只能一张一张推理,吞吐量直接砍到十分之一。opset_version的选择取决于你的推理引擎支持到哪个版本,选高了跑不起来,选低了某些算子不支持,我一般从 11 或 12 起步试。
4. 踩坑实录:ckpt 与 pth 相关的常见问题排查
4.1 加载报错速查表
下面这张表是我这些年攒下来的报错对照,基本覆盖了九成以上的加载问题。
| 报错信息关键词 | 大概率原因 | 处理方式 |
|---|---|---|
| Missing key(s) in state_dict | 结构不一致,或键名前缀不同 | 打印两边 key 对比,检查是否多了module.前缀 |
| Unexpected key(s) in state_dict | 权重文件比模型多参数 | 确认是否加载了错误文件,或用 strict=False |
| size mismatch for ... | 某层维度对不上,常见于类别数改动 | 检查输出层维度,必要时裁掉重训 |
| Attempting to deserialize object on a CUDA device | 无 GPU 环境加载 GPU 权重 | 加 map_location="cpu" |
| Weights only load failed | 新版默认 weights_only=True | 断点文件显式传 weights_only=False |
| ModuleNotFoundError | 整包保存且类路径失效 | 改用 state_dict 保存,重训或重建结构 |
| RuntimeError: Error(s) in loading state_dict | 复合字典没取对层级 | 打印顶层 keys,确认取 ckpt["model"] |
这张表我用得很频繁,基本看到报错第一行就能定位方向。
4.2 key 不匹配的定位方法:打印对比是唯一正解
面对 size mismatch 或者 missing key,最有效的手段就是把两边的 key 都打出来做对比,别靠猜。
model_keys = set(model.state_dict().keys()) ckpt_keys = set(ckpt["model"].keys()) print("only in model:", sorted(model_keys - ckpt_keys)[:10]) print("only in ckpt:", sorted(ckpt_keys - model_keys)[:10])如果发现 ckpt 里的 key 全都带module.前缀,说明权重是用nn.DataParallel或DistributedDataParallel训练时保存的,包了一层。这时候有两个办法,一是加载时手动去前缀:
new_state = {k.replace("module.", "", 1): v for k, v in ckpt["model"].items()} model.load_state_dict(new_state)二是用load_state_dict的strict=False配合key重映射。我更推荐第一种,因为改完的字典可以继续当普通 state_dict 用,逻辑干净。
反过来,如果模型本身被 DDP 包了而权重没前缀,加载前也要相应处理。这里有个经验:保存的时候统一存model.module.state_dict(),也就是剥掉 DDP 外壳,这样存出来的文件在任何场景下都能用,不必每次加载都做前缀清洗。这个习惯我从第一次遇到 DDP 权重问题之后就一直保持。
4.3 文件体积、安全与版本兼容
模型文件的安全问题值得单独说。pickle 格式在反序列化时可以触发任意代码执行,这已经不是理论风险了。所以第一条纪律是:只加载你自己训练或者可信来源的权重文件。网上随手下载的权重文件,先看来源,能用 safetensors 格式的就用 safetensors,它只存张量不支持代码执行,天然安全。
PyTorch 也在这方面做了收紧,weights_only=True就是只允许加载张量、基础类型和少量安全对象。从 2.6 开始这个参数默认变成 True,是好事,但会打破一批老代码。我的处理方式是:推理加载用默认的 True,断点续训显式写 False,并且把这条规则写进项目的 README。
体积方面,一个 50M 参数的模型,float32 权重差不多 200MB,如果断点文件里还带了优化器状态(Adam 会存两份动量),体积直接变成三倍。所以断点文件动辄几个 G 是正常的,磁盘规划要提前算好。我一般会保留最近三个 last.pth 加全部 best.pth,老的用脚本自动清理,不然跑一个月的实验能把盘塞满。
4.4 多卡、EMA 与混合精度下的特殊处理
多卡训练时的保存策略有两个主流选择。一是只在 rank 0 上保存,配合torch.distributed.barrier()保证其他进程不写文件,这样避免多个进程同时写同一个文件导致损坏。二是每个 rank 存自己的分片,用于后续并行加载大模型,这是近年来大模型训练的常见做法,但需要配套的加载逻辑。
if dist.get_rank() == 0: torch.save({"model": model.module.state_dict()}, "ckpt.pth") dist.barrier()barrier()不能省,否则 rank 0 还在写文件的时候其他进程可能已经进下一轮训练,损毁文件的风险是真实存在的。
EMA(指数移动平均)权重的情况也常见。做深度伪造检测模型时,比如基于 Xception 骨干的检测网络,很多人会同时维护原始权重和 EMA 权重两套参数。保存时要把两套都存下来,因为验证阶段通常用 EMA 权重评估,而继续训练要用原始权重更新梯度。命名上建议用model和ema_model两个 key 明确区分,别都叫state_dict。
混合精度的坑前面提过一次,这里再补一个:scaler的状态如果没保存,续训后缩放因子会从默认值重新开始。在前几千步里,梯度缩放可能偏小导致下溢,loss 会突然抖动。判断方法很简单,续训后前 100 步的 loss 曲线如果出现明显跳变,基本就是这个原因。
5. 真实项目里怎么定保存策略
5.1 从三个实际模型类型的保存需求说起
不同模型对保存策略的需求差别挺大,我用三个具体场景来说明。
第一个是深度伪造检测模型,典型结构是 Xception 骨干加一个二分类头。这类任务的特点是数据不平衡严重,指标波动大,所以 best 的判定不能只看单次验证结果,我一般会在验证集上跑多次或用滑动平均。另外这类模型经常要做跨数据集测试,权重文件的可移植性要求高,所以必须是纯 state_dict 加一份 config,坚决不能整包保存。
第二个是深度平衡模型(Deep Equilibrium Model),它的特点是前向过程是求解一个不动点,本身不带很多中间层参数,权重的存储结构反而比较轻。但它有个特殊点:训练时需要保存求解器的迭代次数和收敛阈值等状态,否则续训时数值行为会变。所以这类模型的 ckpt 里除了常规字段,还要额外记录求解器配置。
第三个是深度循环模型,比如各种 RNN、状态空间模型。它们的隐藏状态在网络内部流转,一般不作为参数保存,但如果你的实现里把某些状态做成了 buffer(比如某些变体中的初始状态),那就会进 state_dict,加载时必须保证 buffer 也对得上。我遇到过 buffer 维度不一致导致的 size mismatch,排查了半天才发现是序列长度配置变了。
5.2 命名规范与版本管理的一些约定
混乱的文件名是排查成本的最大来源。我现在的命名规范是这样的:
{实验名}_{数据集}_{骨干网}_{指标值}_{epoch}.pth比如dfdetect_ffpp_xception_f1-0.923_ep18.pth。指标值放文件名里,好处是一眼能看出好坏,不用加载。缺点是每次刷新 best 都要重命名,所以实践上 best.pth 用固定名,同时软链接指向带指标名的文件,兼顾便利和可读性。
版本管理方面,权重文件不要进 git,用 git-lfs 也尽量避免,因为大文件会拖慢仓库。我的做法是权重存独立的存储路径,git 里只保留一个weights_manifest.json,记录每个实验的权重路径、指标、训练命令。这个小文件几十行,但它让整个项目可追溯。
顺带说一句"保存本地模型配置失败"这类问题。它通常不是保存逻辑本身的错,而是配置里塞了不可序列化的东西,比如数据集对象、文件句柄、lambda 函数。写配置的时候坚持"只用基础类型和容器类型"这条原则,基本就不会遇到。如果确实需要存复杂对象,转成字符串描述或者只存它的构造参数。
5.3 保存频率与性能开销的权衡
每个 epoch 都保存断点文件的代价经常被低估。一个 3GB 的断点文件写盘,在普通机械盘上要好几秒,如果训练一个 epoch 只要 20 秒,那保存开销就占了 20% 以上,而且写盘 IO 会阻塞训练进程。优化方式有几个:一是只在满足条件时保存,比如每 N 个 epoch 或指标刷新时;二是先写临时文件再原子重命名,避免写一半崩溃导致文件损坏;三是用后台线程异步写盘,训练继续跑。
tmp = path + ".tmp" torch.save(obj, tmp) os.replace(tmp, path) # 原子操作,避免半截文件os.replace这一步看着不起眼,但它能防止断电或者进程被杀导致 last.pth 变成损坏文件。我在一次集群任务被抢占之后就加上了这个习惯,代价几乎为零,收益是文件永远可用。
6. 一些容易被问到的细节问题
6.1 关于 weights_only 的取舍
前面反复提到这个参数,值得单独梳理一下判断标准。纯推理加载 state_dict,用weights_only=True,安全且够用。加载含优化器状态的断点,用weights_only=False,因为优化器状态字典里有非张量结构。加载整包模型对象,必然要 False,但这条路径我建议直接放弃。
还有一个细节:如果你用的是较老版本 PyTorch,weights_only参数可能根本不存在,那说明版本在 2.0 之前,加载行为一直是宽松模式。升级版本时要留意这个行为变化,最好在升级前先把项目里的torch.load调用点全部列出来,逐个确认该传什么。
6.2 CPU 与 GPU 权重互转的实际影响
从 GPU 保存的权重加载到 CPU,除了要map_location="cpu",还有一点要注意:张量的 dtype 和布局不变,所以内存占用和 GPU 显存占用是一样的量级。一个大模型加载到 CPU 内存可能直接把内存吃满,这时候可以考虑用torch.load(..., mmap=True)做内存映射,按需读取,能显著降低峰值内存。
反向的 CPU 到 GPU 加载不需要特殊处理,load 完之后调用model.to("cuda")就行。但要注意map_location="cuda"和先 load 到 CPU 再 to cuda 的区别:前者在加载过程中就分配显存,如果显存不够会直接失败且可能留下碎片;后者更可控。我一般倾向于先加载到 CPU 再搬,出问题好排查。
6.3 权重文件损坏的识别与预防
权重文件损坏在长时间训练里不算罕见,尤其是写到一半进程被杀、或者是网络存储抖动。识别方法很直接,加载时如果报UnpicklingError或者zipfile.BadZipFile,基本就是文件坏了。预防手段就是前面说的原子写加上临时文件。
如果只是损坏了尾部,有时候还能抢救,用torch.load加mmap=True可能读到部分内容,但这个方法不保证成功,只适合应急。真正靠谱的还是保留多个历史断点,别只留一个文件。我的习惯是 last 文件保留最近三个轮转,这样即使最新一个坏了,损失也就一个 epoch。
6.4 跨框架权重迁移的现实难度
有时候需要把一个 PyTorch 的 pth 迁到别的框架,或者反过来。这件事的难度取决于层的对应关系。卷积、全连接、BatchNorm 这类标准层基本能一一对应,但涉及自定义算子、特殊的 padding 方式、不同的默认初始化,就会出现权重对得上但结果对不上的情况。
我的建议是做数值对齐验证:构造一个固定输入,在两边分别跑前向,逐层对比输出。哪一层开始出现明显偏差,问题就在那一层附近。这个流程我在做跨框架部署时走过好几次,虽然麻烦,但比盲目试错高效得多。
最后分享一个我自己用下来最省事的习惯。不管项目大小,我都会在训练脚本旁边放一个inspect_ckpt.py,功能就一个:传进任意权重文件,打印它的顶层结构、每个子字典的 key 数量和前几个 key 名、以及所有张量的 dtype 和 shape 概览。这个小脚本我用了好几年,它不能解决任何问题,但它能在三十秒内让你搞清楚手上这个文件到底是什么。踩过太多次"以为它是纯权重、结果是复合字典"的坑之后,我现在的原则是:任何陌生的 ckpt 或 pth,先 inspect,再动手。多花这半分钟,往往能省下半夜的排查时间。