Skip to content

实战第 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.loadweights_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 项目。接下来去进阶与最佳实践,看看真实项目和教学项目差在哪。