把下面代码保存为 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%