理解为什么需要验证/测试集来评估模型
理解过拟合(Overfitting) 是什么、怎么发现它
学会使用混淆矩阵和分类报告深度分析模型表现
第 7 章我们在训练集上算准确率,达到了 97%+。但问题是:这些图模型已经"看过"了。就像一个学生把月考卷子背下来,重新做当然满分,但换一套卷子就不一定了。
所以我们需要一个模型从没见过的测试集来真正评估它。
# 测试集评估函数
@torch.no_grad() # 评估时不需要计算梯度,省内存、加速
def evaluate(model, loader, loss_fn):
"""在测试集上评估模型,返回平均损失和准确率"""
model.eval() # 切换到评估模式
total_loss, correct, total = 0.0, 0, 0
for images, labels in loader:
# 向前传播,计算预测值
outputs = model(images)
# 计算损失
loss = loss_fn(outputs, labels)
total_loss += loss.item()
correct += outputs.argmax(dim=1).eq(labels).sum().item()
total += labels.size(0)
avg_loss = total_loss / len(loader)
accuracy = 100.0 * correct / total
return avg_loss, accuracy📌 @torch.no_grad() 装饰器
它告诉 PyTorch:"下面这段代码我不需要梯度,不用构建计算图"。这能节省大量内存,让评估更快。评估/推理时一定要加。
在每个 epoch 结束后,既看训练集表现,也看测试集表现:
for epoch in range(num_epochs):
# 训练
train_loss, train_acc = train_one_epoch(...)
# 评估
test_loss, test_acc = evaluate(model, test_loader, loss_fn)
print(f"Epoch {epoch+1}: "
f"Train Loss={train_loss:.4f}, Train Acc={train_acc:.2f}% | "
f"Test Loss={test_loss:.4f}, Test Acc={test_acc:.2f}%")过拟合 = 模型在训练集上表现很好,但在测试集上表现很差。
就像学生把课本例题的答案背下来了,但题目稍微变化就不会做了。
| 现象 | 说明 |
| 训练 Loss 持续下降,测试 Loss 也下降 | ✅ 正常学习 |
| 训练 Loss 还在降,但测试 Loss 开始上升 | ⚠️ 过拟合开始 |
| 训练准确率接近 100%,测试准确率远低于它 | ❌ 严重过拟合 |
📌 术语:泛化能力(Generalization)
模型在"没见过的数据"上表现良好的能力。训练的目标不是让模型背下训练集,而是让模型拥有泛化能力。 这是深度学习最核心的追求。
| 方法 | 说明 | 初学者友好度 |
| Early Stopping(早停) | 当测试 Loss 不再下降时就停止训练 | ⭐⭐⭐⭐⭐ |
| Dropout | 训练时随机"关掉"一部分神经元 | ⭐⭐⭐⭐ |
| 数据增强(Data Augmentation) | 把训练图旋转、平移、缩放,变出更多数据 | ⭐⭐⭐ |
| 减少模型参数 | 缩小网络,让它没那么容易"背答案" | ⭐⭐⭐⭐ |
| 增加训练数据 | 更多的数据 = 更不容易过拟合 | ⭐ |
对于 MNIST,过拟合不是大问题(数据量大、任务简单),但你可能会在第 10 个 epoch 后看到测试准确率不再涨了,训练准确率却继续涨 —— 这就是过拟合的苗头。
假设你的模型说"我有 95% 的准确率",听起来不错。但万一那 5% 的错误全集中在数字 "8" 上呢?那就意味着模型完全认不出 8。
混淆矩阵能告诉你:模型在哪些数字上容易搞混。
📌 术语:混淆矩阵(Confusion Matrix)
一个 N×N 的表格(N=类别数),行是"真实标签",列是"预测标签"。
对角线上的数字 = 预测正确的数量
非对角线 = 预测错误的数量(比如真实是3但预测成5的数量)
看着这个表,你一眼就能发现模型最不擅长识别哪个数字、最容易把哪个数字搞混。如下图:
上图中,“正确”单元格就是预测的正确值。其他区域是预测错误的值,可以通过单元格数量一眼看出预测错误主要集中在哪些值上面。
下面是混淆矩阵使用的关键代码,可以运行的代码查看本章最后。部分代码如下:
from sklearn.metrics import confusion_matrix, classification_report
import numpy as np
import torch
@torch.no_grad()
def get_predictions(model, loader):
"""
收集整个测试集的真实标签和预测标签
:param model: 训练好的 PyTorch 分类模型
:param loader: 测试集 DataLoader
:return: all_labels(真实标签列表), all_preds(模型预测标签列表)
"""
# 关闭梯度计算,节省显存、加速推理,推理阶段不需要反向传播
model.eval()
# 定义两个空列表,分别存储全部真实标签、全部预测标签
all_preds, all_labels = [], []
# 遍历测试集每一个batch
for images, labels in loader:
outputs = model(images)
preds = outputs.argmax(dim=1)
all_preds.extend(preds.tolist())
all_labels.extend(labels.tolist())
return all_labels, all_preds
# 运行函数,得到全部真实标签、模型预测标签
true_labels, pred_labels = get_predictions(model, test_loader)
# 计算混淆矩阵
# 行:真实标签,列:预测标签;对角线为预测正确样本
cm = confusion_matrix(true_labels, pred_labels)
print("混淆矩阵(行=真实标签,列=预测标签):")
print(cm)
# 输出分类报告:每一类精确率、召回率、F1、样本数量;digits=4保留4位小数
print("\n分类报告:")
print(classification_report(true_labels, pred_labels, digits=4))📌 术语:精确率、召回率、F1
精确率(Precision):模型说"这是3"的图片中,真正是3的比例。精度高 = 不乱猜。
召回率(Recall):所有真正的3中,模型找出了多少。召回高 = 不漏掉。
F1 分数:精确率和召回率的调和平均。越高越好,满分是 1.0。
对 MNIST,这三个指标一般在 0.97~0.99 之间。
找出模型最容易搞混的数字对:
# 找错误最多的数字对
# 存放混淆错误元组,格式:(真实标签i,预测标签j,混淆样本数量)
errors = []
# 遍历混淆矩阵所有行 i:真实标签;列 j:预测标签
for i in range(10):
for j in range(10):
# i != j:排除预测正确(对角线样本);cm[i][j] > 0:只保留存在混淆的类别对
if i != j and cm[i][j] > 0:
# 将真实类别、预测类别、混淆数量存入列表
errors.append((i, j, cm[i][j]))
# 按照混淆样本数量从大到小排序(第三个元素count降序)
errors.sort(key=lambda x: x[2], reverse=True)
print("最常见的5个混淆:")
# 取出Top5混淆情况循环打印
for real, pred, count in errors[:5]:
# real:真实数字,pred:模型误判数字,count:发生次数
print(f" 数字{real} → 被误认为{pred}:{count}次")你会发现 MNIST 上最常见的混淆是 4↔9、3↔5、7↔1……这些数字长得确实有点像,人有时候也会看错。
测试集是最终的"评判官",模型在测试集上的表现才代表真实水平
过拟合 = 训练好、测试差。发现后可以用早停、Dropout 等方法解决
混淆矩阵能告诉你模型具体在哪些数字上出错,比只看准确率有用得多
评估函数要加 @torch.no_grad() 和 model.eval()
点击查看完整的 测试集评估 + 混淆矩阵 示例代码。