不知道你有没有过这种经历:照着教程复制粘贴跑通了一个手写数字识别,准确率看着还挺像回事,然后教程结束,模型也成了摆设。最尴尬的是,过了两天想把训练时保存的 .pth 文件重新加载出来看一眼,结果各种报错;或者同组同事问你要这个模型,反问一句“这文件怎么用啊”,你才发现自己只会跑那个训练脚本,连模型怎么单独调用都说不清楚。
这种状态太常见了。市面上大量资料把“网络结构设计”和“训练调参”讲得很重,但真正进入实际工作后,每天反复操作的东西反而不是造新网络,而是模型文件的使用与修改、保存与读取这一整条链路。模型怎么从磁盘里恢复?怎么跑一次前向推理?别人给你的预训练权重改一个分类任务要怎么接?为什么 torch.save(model) 保存出来的模型换个环境就加载失败?这篇文章就把这些问题一次说透。内容以 PyTorch 为主,因为它是目前入门深度学习最友好的框架之一,也顺带提一遍 ONNX、checkpoint 这些实际工程里绕不开的东西。适合已经跑通过一个最小训练脚本、但对“模型文件生命周期”一头雾水的朋友,也适合刚看完理论课还不敢动手的初学者直接照着抄作业。
1. 先别急着追新网络,把模型的“生命周期”打通
1.1 模型的使用、修改、保存、读取其实是同一件事
我见过很多初学者把“训练出一个模型”当成终点,仿佛模型训练完就自动能用了。但实际上,训练只是整个链路里的一个环节。模型的真实生命周期是这样的:定义网络结构 → 喂数据训练 → 得到参数 → 把参数保存到磁盘 → 下次从磁盘读回来 → 做前向推理 → 为了适配新任务修改结构或参数 → 重新保存 → 重新加载。
你会发现,使用、修改、保存、读取这四个动作,在整个生命周期里是反复穿插的。它们从来不是四件独立的事,而是一条完整的传送带。只学会“定义网络”和“跑训练”,等于传送带只装了两节,工件走到一半就掉地上了。
所以在往下看之前,我建议你把目标缩小一点:不要贪多,不要今天追 Transformer、明天追扩散模型,先老老实实把一个已经训练好的 CNN 模型,从加载到推理、再到修改后重新保存这条路走通。这条路走通之后,你再去看任何 “加载预训练模型做微调”的教程,都会觉得它们只是把这条路里的某个环节替换了一下而已。
1.2 模型文件里到底装了什么
很多人对 .pth、.pt 这类文件没有概念,总觉得它像一个神秘的压缩包。其实如果你用 PyTorch 保存一个 state_dict,里面的本质就是一个 Python 字典,字典的 key 是网络中每一层的参数名,比如 features.0.weight、fc.bias,value 是 PyTorch 的 Tensor。你可以直接把它读出来打印,完全不用怕。
打个不那么严谨但有帮助的比方:网络结构是一张“图纸”,state_dict 是图纸上标注的每一个零件参数。光有参数没有图纸,你不知道它怎么组装;光有图纸没有参数,你造出来的东西没有实际能力。所以加载模型时,你永远要做两件事:先创建对应结构的网络,再把参数灌进去。model = SimpleCNN(); model.load_state_dict(torch.load("xxx.pt")) 这行代码之所以能成立,前提就是网络类和训练时完全一致。只要类名改了、某个层删了、输出维度变了,立刻就会报尺寸不匹配。
理解这一点后,你就能明白为什么很多人推荐只保存 state_dict,而不是直接 torch.save(model) 整个模型。直接保存整个模型虽然省事,但它除了参数之外还把“图纸”的定义也序列化了,一旦调用环境里类定义路径变了,就会因为反序列化失败而打不开。
1.3 实操主线:一个能随时修改的 CNN
为了让后面所有操作都落在具体代码上,后面统一用 PyTorch 写一个结构非常简洁的 CNN,在 MNIST 手写数字数据集上做演示。不需要 GPU,CPU 跑几分钟就能完成训练,方便你在没有显卡的电脑上也能完整复现整个流程。
先给出模型定义:
python复制import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
class SimpleCNN(nn.Module):
def __init__(self, num_classes=10):
super().__init__()
self.features = nn.Sequential(
nn.Conv2d(1, 16, 3, padding=1),
nn.ReLU(),
nn.MaxPool2d(2),
nn.Conv2d(16, 32, 3, padding=1),
nn.ReLU(),
nn.MaxPool2d(2),
)
self.fc = nn.Linear(32 * 7 * 7, num_classes)
def forward(self, x):
x = self.features(x)
x = x.view(x.size(0), -1)
return self.fc(x)
这个结构没什么特别的,两个卷积池化组加上一个全连接分类层。唯一值得留意的地方是 num_classes 被放到构造参数里了,这对我后面演示“修改模型适配新任务”非常重要。如果你把一个任务的分类数量写死在网络代码里,后面再想迁移到另一个类别数量的项目,就得改文件重写类,容易牵扯出一堆维护问题。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 用一个小 Demo 训练出“第一个可复用模型”
2.1 训练脚本不要追求花哨,但一定要有保存动作
很多速成教程会在训练环节写一大堆学习率衰减、早停、分布式训练的东西,对新手来说信息量太大。这里我不展开调参,训练代码只做一件事:把模型参数稳定存到磁盘。因为你后续所有关于“读取”“修改”的操作,都必须先有一个真实存在的权重文件作为输入。
python复制def train():
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.1307,), (0.3081,))
])
train_set = datasets.MNIST("./data", train=True, download=True, transform=transform)
train_loader = DataLoader(train_set, batch_size=64, shuffle=True)
model = SimpleCNN(num_classes=10)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
criterion = nn.CrossEntropyLoss()
for epoch in range(3):
model.train()
total_loss = 0.0
for images, labels in train_loader:
optimizer.zero_grad()
outputs = model(images)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
total_loss += loss.item()
avg_loss = total_loss / len(train_loader)
print(f"epoch={epoch + 1}, avg_loss={avg_loss:.4f}")
torch.save(model.state_dict(), "mnist_cnn.pt")
if __name__ == "__main__":
train()
这里的数据集第一次运行时需要联网下载到 ./data 目录。如果下载慢,可以手动从任意渠道下载后解压到对应目录,MNIST 官方数据目录格式很固定,网上搜一下就能找到位置。
2.2 训练轮数、优化器和精度之间的朴素关系
训练里我用了 3 个 epoch。有些刚从理论转过来的人会问:为什么是 3?是不是越多越好?这个问题没有标准答案。如果把一个模型比作学生做题,epoch 就是学生把习题册从头到尾做了几遍。做得太少,知识还没巩固,欠拟合;做得太多,可能把答案背下来而不是学会规律,在标准数据集上会过拟合。更关键的是,对于 MNIST 这类简单任务,CNN 在很短时间内就能达到很高精度,继续加 epoch 收益很小,所以拿来演示模型保存读取时,跑 3 个 epoch 完全够用。
实际项目里到底跑多少轮,建议直接用两条曲线来判断:训练集 loss 和验证集 loss。只要验证集 loss 还在稳步下降,就可以继续训练;验证集 loss 开始回升、训练集 loss 还在下降,大概率就是过拟合了,应当停止。这比照着别人的“20 epoch”“50 epoch”盲目复制靠谱得多。
2.3 简单验证保存结果是否正常
训练结束后,目录下会出现一个 mnist_cnn.pt 文件。你可以用几行代码快速确认里面到底存了什么:
python复制state = torch.load("mnist_cnn.pt", map_location="cpu")
print(type(state))
print(list(state.keys()))
正常你会看到输出是一个 collections.OrderedDict,key 类似 features.0.weight、features.0.bias、features.3.weight、features.3.bias、fc.weight、fc.bias。看到这些就说明权重确实落盘了。
这里先别往下走,我特别建议你做一个动作:故意把 mnist_cnn.pt 复制一份改名成 mnist_cnn_backup.pt,后面很多折腾都能从容恢复。新手最容易犯的错误就是反复用同一个文件名保存,一旦改了代码,旧权重被覆盖,想回到“还能用的版本”就再也找不回来了。
3. 模型的使用:只会跑训练脚本不等于会推理
3.1 训练脚本和推理脚本要分家
我见过不少人把推理逻辑直接塞在训练循环里:训练完一个 epoch,立刻拿测试集算一下准确率,然后说“我会用模型了”。这话得打折扣。部署场景下,推理往往对应着一个独立的服务、一个批处理脚本或者一个 Web 接口,它不负责计算 loss,更不负责反向传播。你要做的是:加载一个已经存在的权重文件,给它一张图片,让它输出预测结果。
在实际工程里,训练和推理分属两份代码是最基本的纪律。训练代码里通常还包含数据增强、标签采样、loss 计算等逻辑,这些东西对推理毫无意义,甚至会成为错误来源。
3.2 最小推理代码:model.eval() 与 no_grad()
下面给出完整的单张图片推理流程:
python复制import torch
from PIL import Image
def predict_image(image_path, model, device="cpu"):
# 1. 读取图片并做与训练一致的预处理
image = Image.open(image_path).convert("L") # 转灰度
image = image.resize((28, 28))
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.1307,), (0.3081,))
])
input_tensor = transform(image).unsqueeze(0) # 增加 batch 维度
input_tensor = input_tensor.to(device)
# 2. 模型进入评估模式
model.eval()
# 3. 推断阶段不计算梯度
with torch.no_grad():
logits = model(input_tensor)
pred = logits.argmax(dim=1).item()
return pred
# 加载模型
device = "cuda" if torch.cuda.is_available() else "cpu"
model = SimpleCNN(num_classes=10)
model.load_state_dict(torch.load("mnist_cnn.pt", map_location=device))
model.to(device)
print(predict_image("test_digit.png", model, device))
这段代码里有三个关键点,新手非常容易忽略。
第一个是 model.eval()。这个操作会切换模型中某些层的运行模式。比如 Dropout 层在训练时会随机丢弃神经元,测试时应当关闭;BatchNorm 层在训练时会用当前 batch 的统计量,在评估时则用训练阶段累积的全局统计量。如果你不调用 eval(),模型在推理时可能仍带有随机性或使用错误的统计量,导致同一张图每次预测结果不一样。很多人说“我用训练好的模型预测不准”,先检查一下自己是不是漏了这行。
第二个是 torch.no_grad()。在推理阶段,我们不需要反向传播,也就不需要保存计算图。如果不加这个上下文管理器,每一次前向都会额外占用大量显存或内存,连续推理时很容易把显存跑爆。它不会影响预测结果,纯粹是一个性能优化,但养成本能习惯后,能帮你避开很多“内存告警”问题。
第三个是图像预处理必须和训练时一致。训练阶段做了灰度、缩放、归一化,推理阶段也必须用一模一样的系数。很多应用里,训练时用的是 Normalize((0.485, 0.456, 0.406), ...) 这类三通道参数,推理时复制过来却只喂了一张单通道图,就会因为 Tensor 形状不匹配而报错。更隐蔽的是缩放尺寸:训练时是 28×28,推理时如果直接拿原图 224×224 送进去,模型结构内部的展平维度会算错,报错提示可能会让你一头雾水。
3.3 给同事用的推理函数要怎么写
这里再给一个经验之谈:当别人找你要模型,不要只丢一个 .pt 文件过去。最负责任的做法是把推理函数封装好并附带一段极简示例,告诉对方输入是什么格式、输出是什么含义、预处理有什么要求。
我自己习惯把 predict_image 这类函数收进一个 inference.py,然后在文件顶部用三行注释写清楚:
python复制# 输入:单张灰度图片路径
# 预处理:resize到28x28,Normalize(0.1307, 0.3081)
# 输出:0-9之间的整数
这些小细节在关键时候能救命。很多开发者拿到模型后第一件事不是看 paper,而是看你给他的代码怎么调用。如果你的模型文件越大,越应该把这层“使用说明”写清楚,否则你迟早会被大量重复的“怎么加载”“喂什么数据”问题淹没。
4. 模型的修改:换分类头、冻结参数、选择性加载
4.1 两类最常见的修改需求
“修改模型”在真实项目里一般有两层含义。第一层是改网络结构,比如模型原本输出 10 类,现在你的业务只需要 5 类,那么最后那个全连接层就要换掉;第二层是改参数更新方式,比如你想让模型在保留原有特征提取能力的同时,只训练新加的那部分结构,就需要使用“冻结”技巧。
无论哪一层需求,都要理解一个核心事实:模型参数的 shape 是由结构决定的。你把最后输出改成 5,那 fc.weight 的形状就从 [10, 1568] 变成 [5, 1568],旧权重没办法直接整体导入。正因为如此,“修改”和“加载”总是绑定出现的。
4.2 换掉最后一层,同时保留前面特征层的权重
先用一个最常见场景做演示:假设我原来的模型在 MNIST 上训练过,现在要在一个 5 类分类任务上做迁移学习。我的处理方法是新建一个输出为 5 的 SimpleCNN 实例,然后加载旧权重时忽略最后一层。
python复制# 新任务只有5类
model = SimpleCNN(num_classes=5)
# 读取旧模型的权重
old_state = torch.load("mnist_cnn.pt", map_location="cpu")
# 旧权重里的 fc.weight/fc.bias 形状不匹配,先剔除
old_state.pop("fc.weight", None)
old_state.pop("fc.bias", None)
# strict=False 允许部分key不存在
model.load_state_dict(old_state, strict=False)
完成后,model.features 各层用的是 MNIST 上训练出来的预训练权重,model.fc 则是一个随机初始化的新分类头。接下来在这个 5 类数据上继续训练即可。
为什么用 pop 而不是直接用 strict=False?因为 strict=False 并不能容忍“同名 key 但 shape 不一致”的情况。只要某个 key 存在且尺寸对不上,PyTorch 依然会报 size mismatch。只有把不匹配的 key 从待加载字典里删掉,才能真正跳过它。还有一点,如果新旧任务的图片通道数、分辨率不一样,前面 features 的权重也可能会不匹配。比如原来是三通道 RGB 图片,现在换成单通道灰度图,第一个卷积层的输入通道就必须从 3 改成 1,那时仅仅 pop 最后一层就不够了,需要连 features.conv0.weight 也一起处理。这类问题没有万金油方法,只能靠打印 state_dict 里的 key 和 shape 去逐层比对。
4.3 冻结参数:让迁移学习不把旧知识冲掉
模型结构改好了,接下来往往是希望只训练新分类头而不要大幅度调整前面的卷积层。原因很简单:卷积层学到的是边缘、纹理这些通用特征,新任务往往数据量很小,如果从头训练这些层,不仅容易过拟合,还可能把预训练模型里积累的通用能力冲掉。
python复制for name, param in model.named_parameters():
if name.startswith("features."):
param.requires_grad = False
optimizer = torch.optim.Adam(
filter(lambda p: p.requires_grad, model.parameters()),
lr=1e-3
)
关键点是:requires_grad = False 之后,model.parameters() 仍然会把所有参数返回给优化器。如果直接把所有参数交给 Adam,它虽然不会更新那些 requires_grad=False 的参数,但会额外记录状态,浪费内存。所以构造优化器时,我习惯用 filter(lambda p: p.requires_grad, model.parameters()) 过滤一遍。
调试时怎么确认冻结生效了呢?可以用一个很朴素的方法:打印 named_parameters(),观察冻结层参数是否带 .requires_grad=False。或者在训练跑通后,保存前后各打印一次 model.fc.weight.mean() 和 model.features[0].weight.mean(),正常情况下 fc 的数值会变化,features 的数值在冻结时不会变化。
4.4 打印模型结构,修改前先摸清“家底”
说了这么多,最终要落到具体实现上。在你动手修改某个模型前,我强烈建议先做一次“家底盘点”,把每一层名称、参数 shape 都打印出来:
python复制model = SimpleCNN(num_classes=10)
for name, param in model.state_dict().items():
print(f"{name}: {tuple(param.shape)}")
比如输出可能是:
text复制features.0.weight: (16, 1, 3, 3)
features.0.bias: (16,)
features.3.weight: (32, 16, 3, 3)
features.3.bias: (32,)
fc.weight: (10, 1568)
fc.bias: (10,)
看到 fc.weight 是 10 行、1568 列,你就知道它把前面展平后的 1568 维特征映射到 10 个类别。假如换成一个 5 分类数据集,这个张量必须变成 (5, 1568)。在改结构的时候,把这个打印结果作为镜子对照,能避免大量盲改代码。
5. 保存与读取:别只会 torch.save(model)
5.1 为什么推荐保存 state_dict 而不是整个模型
很多教深度学习的代码喜欢写 torch.save(model, "model.pt"),看起来是少写了一行 state_dict(),实际埋了不少坑。torch.save(model) 序列化的是整个模型对象,它依赖保存时所处环境的类定义。当你把文件拷到另一台机器,或者把自己的代码目录结构调整过,反序列化时就会因为找不到原来那个类的定义路径而报错。更尴尬的是,如果对方环境里 PyTorch 版本不同,有些旧版本保存的对象可能无法在新版本中加载。
我的做法是尽量只保存 model.state_dict()。这样做有几个明显好处:文件更小,因为不携带类定义和框架缓存;加载灵活,因为我可以选择加载到任何结构一致的新模型实例;也更容易做模型融合、迁移学习、跨框架转换这类操作。
5.2 checkpoint:把优化器状态和训练进度一起留下来
如果你只保存模型参数,断点续训时还有一个隐藏问题:Adam 这类优化器内部保存了动量、二阶梯度估计等状态。如果只把模型权重加载回来,重新建一个全新的 optimizer,学习过程会丢失之前的“惯性”,模型虽然能继续训练,但前几百步可能相当于重新适应。对于大模型训练,这个代价无法接受。
所以正规的训练脚本都会保存一个完整 checkpoint,里面不只有模型参数,还有 optimizer 状态、当前 epoch、当前最佳指标等。举个例子:
python复制checkpoint = {
"model_state_dict": model.state_dict(),
"optimizer_state_dict": optimizer.state_dict(),
"epoch": epoch + 1,
"best_acc": best_acc,
"config": {"lr": lr, "batch_size": batch_size}
}
torch.save(checkpoint, "checkpoint_epoch3.pt")
读取并恢复训练的套路是:
python复制ckpt = torch.load("checkpoint_epoch3.pt", map_location="cpu")
model = SimpleCNN(num_classes=10)
model.load_state_dict(ckpt["model_state_dict"])
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
optimizer.load_state_dict(ckpt["optimizer_state_dict"])
start_epoch = ckpt["epoch"]
best_acc = ckpt["best_acc"]
很多新手不理解为什么“已经训练好的模型继续训练,loss 反而比上一次低时高”。多数情况下,不是因为代码错了,而是 optimizer 状态没保存。你从头创建了一个没有历史状态的优化器,相当于把它以前积累的调整方向清零了。
5.3 想跨框架部署?用 ONNX 导出
PyTorch 生态虽好,但真正部署到推理引擎、边缘设备或异构平台时,不能要求每台机器都装一个 PyTorch。这时候就需要一个中间表示,ONNX 是目前最通用的选择之一。它相当于一个跨框架的“通用模型格式”,PyTorch 训练好的模型可以导出为 .onnx,再由 ONNX Runtime、TensorRT、OpenVINO 这类推理引擎加载。
用 ONNX 导出的核心操作是给模型一个“哑输入”,让 PyTorch 沿着这个输入走一遍前向,记录计算图。
python复制model.eval()
dummy_input = torch.randn(1, 1, 28, 28)
torch.onnx.export(
model,
dummy_input,
"mnist_cnn.onnx",
input_names=["input"],
output_names=["output"],
dynamic_axes={
"input": {0: "batch"},
"output": {0: "batch"}
},
opset_version=12
)
dynamic_axes 的作用是告诉导出工具,batch 维度是可变的。如果不设置,导出的模型输入 batch 只能固定为 1 或你训练时的 batch size。设置这项的好处是,无论你之后一次推理一张图片还是一百张图片,都不需要重新导出。
加载 ONNX 模型用 ONNX Runtime,不需要再依赖 PyTorch:
python复制import onnxruntime as ort
import numpy as np
sess = ort.InferenceSession("mnist_cnn.onnx", providers=["CPUExecutionProvider"])
input_name = sess.get_inputs()[0].name
# input_numpy 的形状是 (1, 1, 28, 28),dtype=np.float32
result = sess.run(None, {input_name: input_numpy})
pred = np.argmax(result[0], axis=1).item()
这里要注意输出已经是概率或 logits 的 numpy 数组,不再是 PyTorch 的 Tensor。在实际工程中,这个部署链路往往比在同一个框架里反复折腾更常见。
5.4 不同框架的文件格式不能混用
稍微延伸一下。你可能会接触到各种后缀的模型文件:.pt、.pth、.h5、.pb、.onnx、.tflite、.weights,它们分别对应 PyTorch、Keras/TensorFlow、ONNX、TensorFlow Lite、Darknet 等不同生态。后缀名本身不是硬性标准,关键要看里面用什么方式序列化的对象。你不能拿 torch.load 去读一个 .h5 文件,也不能拿 Keras 的 load_model 去直接读 PyTorch 的 state_dict。跨框架前先转成中间格式,这是一个通用规则。
6. 高频报错与排查经验
6.1 常见的报错、原因和解决方式速查
实际动手过程中,报错才是最让人长记性的老师。下面我把最容易遇到的几类问题整理成了一张速查表,方便你以后直接翻:
| 报错/表现 | 常见原因 | 处理方式 |
|---|---|---|
size mismatch for fc.weight |
模型输出类别数或某层结构改了,但还在用旧权重 | 打印 state_dict 比对 key 和 shape,剔除不匹配参数后加载 |
RuntimeError: Attempting to deserialize object on a CUDA device |
当前机器没有 GPU,或加载时没指定 map_location | torch.load(path, map_location="cpu") |
Expected all tensors to be on the same device |
模型在 GPU,输入数据还在 CPU,或反过来 | 确保 model.to(device) 和 input.to(device) 一致 |
AttributeError: Can't get attribute 'SimpleCNN' |
直接 torch.save(model) 后,当前环境类名或代码路径不一致 |
尽量使用 state_dict 保存,或保持类定义一致 |
| 推理结果和训练时差别巨大 | 缺少 model.eval(),或预处理和训练不一致 |
加上 model.eval(),严格比对灰度、缩放、归一化参数 |
多 GPU 训练后权重 key 多了 module. 前缀 |
使用了 DataParallel 或 DistributedDataParallel |
加载前去除 module. 前缀,或保存 model.module.state_dict() |
6.2 尺寸不匹配的完整排查过程
很多人第一次看到 size mismatch 就慌,其实排查思路非常固定。第一步,打印当前模型的 state_dict 形状;第二步,打印待加载文件的 state_dict 形状;第三步,找到哪个 key 不一样。
写一个每次都会用到的对比工具函数:
python复制def compare_state_dict(model, ckpt_state):
model_state = model.state_dict()
model_keys = set(model_state.keys())
ckpt_keys = set(ckpt_state.keys())
print("只存在于模型:", model_keys - ckpt_keys)
print("只存在于权重文件:", ckpt_keys - model_keys)
for key in model_keys & ckpt_keys:
if tuple(model_state[key].shape) != tuple(ckpt_state[key].shape):
print(f"形状不一致: {key}, 模型 {tuple(model_state[key].shape)}, 文件 {tuple(ckpt_state[key].shape)}")
运行之后,如果输出只有最后一层形状不一致,说明你只是改了分类数量,按前面 pop 的方式处理即可。如果前面的卷积层也不一致,那就不是简单的“微调最后一层”能解决的事了,需要重新检查输入图片尺寸、通道数,以及网络结构是否和保存权重时完全一致。
6.3 训练中断后从断点续跑的实操
很多训练任务一跑就是几小时甚至几天,如果中途断电、被杀进程,训练进度全丢,心态很容易崩。避免这个问题的最好方法,就是在每个 epoch 结束时都保存一个带编号的 checkpoint,同时保留一个 best_model.pt。
实操建议的目录结构长这样:
text复制checkpoints/
├── mnist_epoch01.pt
├── mnist_epoch02.pt
├── mnist_epoch03.pt
└── best_model.pt
每次保存 best_model.pt 前,先比较当前验证集指标和历史最佳值,只有更好的时候才覆盖保存。这样就算你代码改出问题、跑崩了,至少还有一份历史最佳可用。同时,不要只保留一个文件。每隔几个 epoch 保留一个带编号的 checkpoint,这样可以回到任意中间状态,尤其在后续想要做模型融合或对比实验时会非常有用。
我见过一些人,习惯把所有模型都存成同一个名字,下次训练前直接把旧文件覆盖掉。短时间内看没什么,一旦你想回退到昨天的版本,傻眼了。所以哪怕嫌麻烦,至少要在文件命名里带上 epoch 或时间戳。
6.4 一个小习惯:保存后立刻加载验证
最后分享一个让我少踩很多坑的习惯:任何模型保存之后,立刻在一个干净的脚本里重新加载并跑一次预测,确认整个流程闭环。不要等到第二天、不要等到部署现场,才第一次尝试读文件。
所谓“能保存不算数,能读回来并输出正确结果才算数”。如果保存的时候类定义和加载时候的类定义有细微差异,最快的发现时机就是保存结束后的 5 分钟内。等这个习惯养成之后,你对“保存与读取”这个事情基本上就不会再恐慌了。任何复杂的模型文件,在你眼里都会回归成一个普通字典加一张网络图纸的组合:结构对得上就加载,对不上就打印、比对、剔除、选择性加载,总之没有真正解决不了的问题。
