Appearance
实战第 4 步:评估与导出——项目交付
模型训好了,还差"最后一公里":在测试集上正式验收、把权重保存成可分发的文件、再写一个独立的推理脚本证明"换个进程也能用"。完成后你会得到完整可交付的项目结构。
测试集正式验收
训练脚本里每个 epoch 都在看测试集(教学简化;正式项目应留出独立验证集,测试集只在最后碰一次)。现在做最终评估,新建 evaluate_final.py:
python
# evaluate_final.py
import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
from model import DigitMLP
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.1307,), (0.3081,)),
])
device = torch.device("cuda" if torch.cuda.is_available()
else "mps" if torch.backends.mps.is_available()
else "cpu")
model = DigitMLP().to(device)
state = torch.load("weights/best.pt", map_location=device, weights_only=True)
model.load_state_dict(state)
model.eval()
test_loader = DataLoader(
datasets.MNIST("./data", train=False, download=True, transform=transform),
batch_size=256)
loss_fn = nn.CrossEntropyLoss()
total_loss, correct, seen = 0.0, 0, 0
with torch.no_grad():
for images, labels in test_loader:
images, labels = images.to(device), labels.to(device)
logits = model(images)
total_loss += loss_fn(logits, labels).item() * labels.size(0)
correct += (logits.argmax(1) == labels).sum().item()
seen += labels.size(0)
print(f"测试集: loss={total_loss/seen:.4f} acc={correct/seen*100:.2f}%")终端运行,预期输出:
text
测试集: loss=0.1008 acc=97.10%保存与加载:三种方式怎么选
python
# 方式 1(推荐日常使用):只存参数字典
torch.save(model.state_dict(), "model.pt")
model.load_state_dict(torch.load("model.pt", map_location="cpu", weights_only=True))
# 方式 2:续训存档——参数 + 优化器状态 + epoch 一起打包
torch.save({"epoch": 5, "model": model.state_dict(),
"opt": opt.state_dict()}, "checkpoint.pt")
# 方式 3:存整个模型对象(含类结构定义)
torch.save(model, "full_model.pt") # 加载时 model.py 的路径必须在 sys.path 中| 方式 | 优点 | 缺点 | 场景 |
|---|---|---|---|
| state_dict | 小、稳、跨代码版本灵活 | 需要模型类定义在手 | 交付、分享权重 |
| checkpoint | 可无损续训 | 较大 | 长训练中断恢复 |
| 整对象 | 加载省事 | pickle 耦合代码结构,有任意代码执行风险,类一改就废 | 只用于自己信任的本地实验 |
版本提醒:PyTorch 2.6 起
torch.load的weights_only默认值改为True(安全考虑,只反序列化张量)。加载自己训练的 state_dict 不受影响;加载老 checkpoint(含 numpy 等非张量对象)需显式weights_only=False并确认文件来源可信。永远不要torch.load来路不明的模型文件——等价于执行陌生人的代码。
独立推理脚本:单图预测
交付物最重要的是"别人拿到权重就能用"。新建 predict.py:
python
# predict.py —— 加载权重,对测试集任意一张图给出预测
import sys
import torch
import torch.nn.functional as F
from torchvision import datasets, transforms
from model import DigitMLP
def predict(image_tensor, weight_path="weights/best.pt"):
device = torch.device("cuda" if torch.cuda.is_available()
else "mps" if torch.backends.mps.is_available()
else "cpu")
model = DigitMLP().to(device)
model.load_state_dict(torch.load(weight_path, map_location=device, weights_only=True))
model.eval()
if image_tensor.dim() == 3: # 单张 [1,28,28] → 补 batch 维!
image_tensor = image_tensor.unsqueeze(0)
with torch.no_grad():
logits = model(image_tensor.to(device))
probs = F.softmax(logits, dim=1) # 推理展示时才需要概率
return int(logits.argmax(1)), probs[0].tolist()
if __name__ == "__main__":
idx = int(sys.argv[1]) if len(sys.argv) > 1 else 0
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.1307,), (0.3081,)),
])
test_set = datasets.MNIST("./data", train=False, download=True, transform=transform)
img, label = test_set[idx]
pred, probs = predict(img)
print(f"第 {idx} 张图: 真实标签={label} 预测={pred} 置信度={probs[pred]:.2%}")
print("各类别概率:", [f"{p:.3f}" for p in probs])终端运行:
powershell
python predict.py 7预期输出:
text
第 7 张图: 真实标签=2 预测=2 置信度=99.87%
各类别概率: ['0.000', '0.000', '0.999', '0.000', ...]注意两处细节:① 单张图 unsqueeze(0) 补 batch 维——就是张量那篇埋的伏笔;② 推理展示结果时才套 Softmax,训练时不套、报告时随便套,两者不冲突。
完整项目结构与运行方式
text
mnist-pytorch/
├── data/ # 自动下载的 MNIST
├── weights/
│ └── best.pt # 训练产出的最优权重(约 440KB)
├── model.py # DigitMLP 定义
├── train.py # 训练入口
├── evaluate_final.py # 测试集验收
├── predict.py # 推理入口
└── requirements.txt # torch>=2.1 torchvision>=0.16三步跑起来(激活 PyTorch 环境的终端):
powershell
pip install -r requirements.txt # 1. 装依赖
python train.py # 2. 训练,产出 weights/best.pt
python predict.py 42 # 3. 预测第 42 张测试图小结(实战篇完)
- 验收看测试集指标,保存用
state_dict,续训用 checkpoint。 weights_only是 PyTorch 2.6+ 的新默认,加载老文件要显式处理。- 交付标准:新环境
pip install -r requirements.txt+ 两条命令能复现结果。
你已经拥有一个完整训练并交付的 PyTorch 项目。接下来去进阶与最佳实践,看看真实项目和教学项目差在哪。