第8章:验证、过拟合与评估指标

🎉摘要:训练集准确率高不代表模型好,本文详解测试集评估的重要性,分析过拟合现象及Early Stopping、Dropout等解决方法,并介绍混淆矩阵、精确率、召回率和F1分数,帮助深度评估模型泛化能力。

这一章的目标

  • 理解为什么需要验证/测试集来评估模型

  • 理解过拟合(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}%")

过拟合(Overfitting)—— 初学者最该警惕的问题

过拟合 = 模型在训练集上表现很好,但在测试集上表现很差。

就像学生把课本例题的答案背下来了,但题目稍微变化就不会做了。

怎么判断过拟合?看这两条曲线

现象说明
训练 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()

点击查看完整的 测试集评估 + 混淆矩阵 示例代码。

  

说说我的看法
全部评论(
没有评论
关于
本网站专注于 Java、数据库(MySQL、Oracle)、Linux、软件架构及大数据等多领域技术知识分享。涵盖丰富的原创与精选技术文章,助力技术传播与交流。无论是技术新手渴望入门,还是资深开发者寻求进阶,这里都能为您提供深度见解与实用经验,让复杂编码变得轻松易懂,携手共赴技术提升新高度。如有侵权,请来信告知:hxstrive@outlook.com
其他应用
公众号