测试集评估 + 混淆矩阵

🎉摘要:本文通过PyTorch实现MNIST手写数字识别CNN模型,演示测试集评估、混淆矩阵分析及过拟合检查。运行代码可查看训练/测试损失、准确率,并生成混淆矩阵与分类报告,帮助理解模型性能。

把下面代码保存为 chapter08_validation.py,运行代码观察测试集评估和混淆矩阵结果。

代码如下:

"""
第 8 章:验证、过拟合与评估 —— 测试集评估 + 混淆矩阵

运行方式:python chapter08_validation.py
"""
import torch
import torch.nn as nn
import torch.optim as optim
from torchvision import datasets, transforms
from torch.utils.data import DataLoader
from sklearn.metrics import confusion_matrix, classification_report
import numpy as np


# 网络定义
class MNIST_CNN(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(1, 32, 3, padding=1)
        self.pool1 = nn.MaxPool2d(2)
        self.conv2 = nn.Conv2d(32, 64, 3, padding=1)
        self.pool2 = nn.MaxPool2d(2)
        self.fc1 = nn.Linear(64 * 7 * 7, 128)
        self.fc2 = nn.Linear(128, 10)

    def forward(self, x):
        x = torch.relu(self.conv1(x))
        x = self.pool1(x)
        x = torch.relu(self.conv2(x))
        x = self.pool2(x)
        x = x.view(x.size(0), -1)
        x = torch.relu(self.fc1(x))
        x = self.fc2(x)
        return x


# 训练函数
def train_one_epoch(model, loader, loss_fn, optimizer):
    model.train() # 训练模式
    total_loss, correct, total = 0.0, 0, 0
    for images, labels in loader:
        optimizer.zero_grad()
        outputs = model(images)
        loss = loss_fn(outputs, labels)
        loss.backward()
        optimizer.step()
        total_loss += loss.item()
        correct += outputs.argmax(dim=1).eq(labels).sum().item()
        total += labels.size(0)
    return total_loss / len(loader), 100.0 * correct / total


# 评估函数
@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)
    return total_loss / len(loader), 100.0 * correct / total


@torch.no_grad()
def get_all_predictions(model, loader):
    """收集所有预测结果"""
    model.eval()
    all_preds, all_labels = [], []
    for images, labels in loader:
        outputs = model(images)
        all_preds.extend(outputs.argmax(dim=1).tolist())
        all_labels.extend(labels.tolist())
    return all_labels, all_preds


# 主程序
def main():
    print("=" * 50)
    print("  MNIST CNN 训练 + 测试集评估(第8章)")
    print("=" * 50)

    # 1.数据准备
    transform = transforms.Compose([
        transforms.ToTensor(),
        transforms.Normalize((0.1307,), (0.3081,))
    ])
    train_dataset = datasets.MNIST("./data", train=True, download=True, transform=transform)
    test_dataset = datasets.MNIST("./data", train=False, download=True, transform=transform)
    # 训练数据集
    train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True)
    # 测试数据集
    test_loader = DataLoader(test_dataset, batch_size=64, shuffle=False)

    # 2.模型、损失函数、优化器
    model = MNIST_CNN()
    loss_fn = nn.CrossEntropyLoss()
    optimizer = optim.Adam(model.parameters(), lr=0.001)

    # 3.训练(每个 epoch 都做验证)
    num_epochs = 5
    print("\n开始训练(每轮同时看训练集和测试集表现):")
    print(f"{'Epoch':<8}{'Train Loss':<14}{'Train Acc':<12}{'Test Loss':<14}{'Test Acc':<12}")
    print("-" * 60)

    history = {"train_loss": [], "train_acc": [], "test_loss": [], "test_acc": []}

    for epoch in range(num_epochs):
        train_loss, train_acc = train_one_epoch(model, train_loader, loss_fn, optimizer)
        test_loss, test_acc = evaluate(model, test_loader, loss_fn)

        history["train_loss"].append(train_loss)
        history["train_acc"].append(train_acc)
        history["test_loss"].append(test_loss)
        history["test_acc"].append(test_acc)

        print(f"{epoch+1:<8}{train_loss:<14.4f}{train_acc:<12.2f}{test_loss:<14.4f}{test_acc:<12.2f}")

    # 混淆矩阵分析
    print("\n" + "=" * 50)
    print("  混淆矩阵分析")
    print("=" * 50)

    true_labels, pred_labels = get_all_predictions(model, test_loader)
    cm = confusion_matrix(true_labels, pred_labels)

    # 打印混淆矩阵
    print("\n混淆矩阵(行=真实标签, 列=预测标签):")
    print("预→ " + " ".join(f"{i:4d}" for i in range(10)))
    for i in range(10):
        print(f"真{i}: " + " ".join(f"{cm[i][j]:4d}" for j in range(10)))

    # 找最常见的混淆
    errors = []
    for i in range(10):
        for j in range(10):
            if i != j and cm[i][j] > 0:
                errors.append((i, j, cm[i][j]))
    errors.sort(key=lambda x: x[2], reverse=True)

    print("\n模型最容易混淆的 Top 5:")
    for real, pred, count in errors[:5]:
        print(f"  真实数字 {real} → 被误认为 {pred}:{count} 次")

    # 分类报告
    print("\n分类报告:")
    print(classification_report(true_labels, pred_labels, digits=4))

    # 每个数字的准确率
    print("每个数字的识别准确率:")
    for i in range(10):
        correct = cm[i][i]
        total = cm[i].sum()
        print(f"  数字 {i}: {correct}/{total} = {100*correct/total:.2f}%")


if __name__ == "__main__":
    main()

运行代码,输出如下:

==================================================
  MNIST CNN 训练 + 测试集评估(第8章)
==================================================

开始训练(每轮同时看训练集和测试集表现):
Epoch   Train Loss    Train Acc   Test Loss     Test Acc    
------------------------------------------------------------
1       0.1288        96.07       0.0391        98.63       
2       0.0411        98.73       0.0331        98.91       
3       0.0270        99.15       0.0323        98.83       
4       0.0204        99.32       0.0276        99.02       
5       0.0158        99.49       0.0350        98.86       

==================================================
  混淆矩阵分析
==================================================

混淆矩阵(行=真实标签, 列=预测标签):
预→    0    1    2    3    4    5    6    7    8    9
真0:  978    0    0    0    0    0    0    1    1    0
真1:    0 1134    0    0    0    1    0    0    0    0
真2:    5    1 1024    0    1    0    0    0    1    0
真3:    0    0    6  998    0    2    0    0    4    0
真4:    0    2    1    0  962    0    1    0    2   14
真5:    2    0    0    6    0  880    2    1    1    0
真6:    2    2    0    0    1    1  948    0    3    1
真7:    0   10    9    1    1    1    0 1002    1    3
真8:    2    1    2    0    0    2    0    0  965    2
真9:    1    3    0    0    3    4    0    2    1  995

模型最容易混淆的 Top 5:
  真实数字 4 → 被误认为 9:14 次
  真实数字 7 → 被误认为 1:10 次
  真实数字 7 → 被误认为 2:9 次
  真实数字 3 → 被误认为 2:6 次
  真实数字 5 → 被误认为 3:6 次

分类报告:
              precision    recall  f1-score   support

           0     0.9879    0.9980    0.9929       980
           1     0.9835    0.9991    0.9913      1135
           2     0.9827    0.9922    0.9875      1032
           3     0.9930    0.9881    0.9906      1010
           4     0.9938    0.9796    0.9867       982
           5     0.9877    0.9865    0.9871       892
           6     0.9968    0.9896    0.9932       958
           7     0.9960    0.9747    0.9853      1028
           8     0.9857    0.9908    0.9882       974
           9     0.9803    0.9861    0.9832      1009

    accuracy                         0.9886     10000
   macro avg     0.9887    0.9885    0.9886     10000
weighted avg     0.9887    0.9886    0.9886     10000

每个数字的识别准确率:
  数字 0: 978/980 = 99.80%
  数字 1: 1134/1135 = 99.91%
  数字 2: 1024/1032 = 99.22%
  数字 3: 998/1010 = 98.81%
  数字 4: 962/982 = 97.96%
  数字 5: 880/892 = 98.65%
  数字 6: 948/958 = 98.96%
  数字 7: 1002/1028 = 97.47%
  数字 8: 965/974 = 99.08%
  数字 9: 995/1009 = 98.61%

  

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