将下面代码保存到 chapter09_visualization.py 文件,运行后会训练 5 个 epoch,实时显示训练曲线,最后画出混淆矩阵和错误样本。
代码如下:
"""
第 9 章:可视化 —— 实时训练曲线 + 混淆矩阵 + 错误样本
功能整体说明:
基于PyTorch搭建手写数字MNIST卷积神经网络,在模型训练全过程实时绘制损失、准确率动态曲线;
训练结束后输出混淆矩阵分析各类数字识别分布情况,可视化展示被模型预测错误的图片样本,
全方位直观观测模型训练效果、分类偏向、识别缺陷。
运行方式:python chapter09_visualization.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
# 绘图可视化基础库
import matplotlib
import matplotlib.pyplot as plt
import numpy as np
# 指定matplotlib渲染后端为TkAgg,保证绘图弹窗正常弹出显示
matplotlib.use('TkAgg')
# 解决matplotlib图表中文显示乱码问题
plt.rcParams["font.family"] = ["Microsoft YaHei"]
# 解决坐标轴负号展示成方框异常
plt.rcParams["axes.unicode_minus"] = False
# ---------------------- 卷积神经网络模型定义 ----------------------
class MNIST_CNN(nn.Module):
"""
适配MNIST手写数字单通道灰度图的卷积神经网络
网络结构:2层卷积+最大池化提取图像特征 → 两层全连接层完成分类输出
输入尺寸:[batch_size, 1, 28, 28](MNIST标准28*28灰度手写数字图)
输出:10维向量,对应数字0~9的分类logits得分
"""
def __init__(self):
super().__init__()
# 第一层卷积:输入通道1(灰度图),输出32特征图,卷积核3×3,padding补0保持尺寸不变
self.conv1 = nn.Conv2d(1, 32, 3, padding=1)
# 2×2最大池化,特征图宽高压缩为原来1/2
self.pool1 = nn.MaxPool2d(2)
# 第二层卷积:输入32通道,输出64通道特征图,3×3卷积核
self.conv2 = nn.Conv2d(32, 64, 3, padding=1)
self.pool2 = nn.MaxPool2d(2)
# 两次池化后特征图尺寸7*7,64通道展平后接入全连接层
self.fc1 = nn.Linear(64 * 7 * 7, 128)
# 最终输出10个类别(数字0~9)
self.fc2 = nn.Linear(128, 10)
def forward(self, x):
"""模型前向传播逻辑"""
x = torch.relu(self.conv1(x)) # 卷积+ReLU激活
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):
"""
单个epoch训练一轮数据集
:param model: 待训练CNN模型
:param loader: 训练集DataLoader迭代器
:param loss_fn: 损失函数(交叉熵损失)
:param optimizer: 优化器Adam
:return: 当前轮次平均loss、整体训练集准确率
"""
model.train() # 切换模型至训练模式,启用dropout/bn训练逻辑(本网络无dropout/bn,规范写法)
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()
pred = outputs.argmax(dim=1) # 取概率最大下标作为预测类别
correct += pred.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):
"""
验证/测试集评估,关闭梯度计算节省显存与算力
:param model: 训练完成的模型
:param loader: 测试集DataLoader
:param loss_fn: 损失函数
:return: 测试集平均loss、测试集整体准确率
"""
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()
pred = outputs.argmax(dim=1)
correct += pred.eq(labels).sum().item()
total += labels.size(0)
return total_loss / len(loader), 100.0 * correct / total
# ---------------------- 可视化工具函数 ----------------------
def plot_confusion_matrix(cm, save_path=None):
"""
绘制分类混淆矩阵热力图,直观查看各类数字识别正误分布
:param cm: sklearn计算得到的混淆矩阵二维数组
:param save_path: 图片保存路径,不传则仅弹窗展示
"""
plt.figure(figsize=(8, 6))
plt.imshow(cm, cmap='Blues') # 热力图底色蓝色渐变
plt.colorbar(label='样本数') # 右侧色条,映射数值对应样本数量
plt.xlabel('预测标签')
plt.ylabel('真实标签')
plt.title('混淆矩阵')
# 在每个格子填入具体样本数量
for i in range(10):
for j in range(10):
# 根据格子深浅自动选用白色/黑色文字保证可读性
color = 'white' if cm[i][j] > cm.max() / 2 else 'black'
plt.text(j, i, str(cm[i][j]), ha='center', va='center', fontsize=8)
plt.xticks(range(10))
plt.yticks(range(10))
plt.tight_layout()
if save_path:
plt.savefig(save_path, dpi=150)
plt.show()
def show_misclassified(model, loader, max_show=10):
"""
筛选并可视化模型预测错误的手写数字样本
:param model: 训练好的模型
:param loader: 测试集加载器
:param max_show: 最多展示多少张错误图片
"""
model.eval()
misclassified = [] # 存放错误样本:(图像张量,真实标签,预测标签)
with torch.no_grad():
for images, labels in loader:
outputs = model(images)
preds = outputs.argmax(dim=1)
mask = ~preds.eq(labels) # 布尔掩码,标记预测错误的样本
if mask.any():
# 遍历所有错误下标
for idx in mask.nonzero(as_tuple=False).squeeze(1):
if len(misclassified) < max_show:
misclassified.append((
images[idx].squeeze(), # 去掉通道维度方便绘图
labels[idx].item(),
preds[idx].item()
))
if len(misclassified) >= max_show:
break
# 自动计算子图画布行列布局
n = len(misclassified)
cols = min(5, n)
rows = (n + cols - 1) // cols
fig, axes = plt.subplots(rows, cols, figsize=(cols * 3, rows * 3))
if n == 1:
axes = [axes]
else:
axes = axes.flatten()
for i, (img, true_lbl, pred_lbl) in enumerate(misclassified):
# 反向归一化还原图片原始像素区间,保证显示正常
img_disp = torch.clamp(img * 0.3081 + 0.1307, 0, 1)
axes[i].imshow(img_disp, cmap='gray')
axes[i].set_title(f'真实:{true_lbl} → 预测:{pred_lbl}', color='red')
axes[i].axis('off')
# 空白子图关闭坐标轴
for i in range(n, len(axes)):
axes[i].axis('off')
plt.suptitle('模型识别错误的样本', fontsize=14)
plt.tight_layout()
plt.show()
# ---------------------- 主训练入口流程 ----------------------
def main():
print("=" * 50)
print(" MNIST CNN 训练 + 可视化(第9章)")
print("=" * 50)
# 1、数据集加载与图像预处理
transform = transforms.Compose([
transforms.ToTensor(), # PIL图片转为张量,像素归一化至0~1
transforms.Normalize((0.1307,), (0.3081,)) # MNIST数据集全局均值、标准差标准化
])
# 加载官方MNIST手写数字数据集,自动下载
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) # Adam自适应优化器
# 开启matplotlib交互模式,实现训练过程动态刷新曲线图
plt.ion()
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 4))
num_epochs = 5
# 记录每一轮训练、测试的loss与准确率历史数据
history = {"train_loss": [], "train_acc": [], "test_loss": [], "test_acc": []}
# 3、循环迭代训练
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)
# 保存本轮指标至历史记录
for key, val in zip(
["train_loss", "train_acc", "test_loss", "test_acc"],
[train_loss, train_acc, test_loss, test_acc]
):
history[key].append(val)
# 实时刷新绘制损失曲线
epochs_list = range(1, epoch + 2)
ax1.clear()
ax1.plot(epochs_list, history["train_loss"], 'b-o', label='训练 Loss', markersize=4)
ax1.plot(epochs_list, history["test_loss"], 'r-s', label='测试 Loss', markersize=4)
ax1.set_xlabel('Epoch'); ax1.set_ylabel('Loss')
ax1.set_title(f'损失曲线(第{epoch+1}轮)'); ax1.legend(); ax1.grid(True, alpha=0.3)
# 实时刷新绘制准确率曲线
ax2.clear()
ax2.plot(epochs_list, history["train_acc"], 'b-o', label='训练 Acc', markersize=4)
ax2.plot(epochs_list, history["test_acc"], 'r-s', label='测试 Acc', markersize=4)
ax2.set_xlabel('Epoch'); ax2.set_ylabel('准确率 (%)')
ax2.set_title(f'准确率曲线(第{epoch+1}轮)'); ax2.legend(); ax2.grid(True, alpha=0.3)
plt.tight_layout()
plt.pause(0.3) # 短暂停顿,刷新画布
# 控制台打印本轮训练指标
print(f"Epoch {epoch+1}/{num_epochs} | "
f"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.2f}% | "
f"Test Loss: {test_loss:.4f} | Test Acc: {test_acc:.2f}%")
# 关闭交互绘图模式
plt.ioff()
# 训练结束,绘制完整版最终曲线并保存图片至本地
final_fig, (fax1, fax2) = plt.subplots(1, 2, figsize=(12, 4))
epochs_list = range(1, num_epochs + 1)
fax1.plot(epochs_list, history["train_loss"], 'b-o', label='训练 Loss', markersize=4)
fax1.plot(epochs_list, history["test_loss"], 'r-s', label='测试 Loss', markersize=4)
fax1.set_xlabel('Epoch'); fax1.set_ylabel('Loss'); fax1.set_title('损失曲线')
fax1.legend(); fax1.grid(True, alpha=0.3)
fax2.plot(epochs_list, history["train_acc"], 'b-o', label='训练 Acc', markersize=4)
fax2.plot(epochs_list, history["test_acc"], 'r-s', label='测试 Acc', markersize=4)
fax2.set_xlabel('Epoch'); fax2.set_ylabel('准确率 (%)'); fax2.set_title('准确率曲线')
fax2.legend(); fax2.grid(True, alpha=0.3)
plt.tight_layout()
plt.savefig("training_curves.png", dpi=150)
print("\n最终训练曲线已保存:training_curves.png")
plt.show()
# 4、计算混淆矩阵
model.eval()
all_preds, all_labels = [], []
with torch.no_grad():
for images, labels in test_loader:
outputs = model(images)
all_preds.extend(outputs.argmax(dim=1).tolist())
all_labels.extend(labels.tolist())
cm = confusion_matrix(all_labels, all_preds)
plot_confusion_matrix(cm)
# 5、展示识别错误的样本图片
show_misclassified(model, test_loader, max_show=10)
print("\n所有图表展示完毕!")
if __name__ == "__main__":
main()运行代码,开始训练,输出如下:
==================================================
MNIST CNN 训练 + 可视化(第9章)
==================================================
Epoch 1/5 | Train Loss: 0.1367 | Train Acc: 95.90% | Test Loss: 0.0424 | Test Acc: 98.55%
Epoch 2/5 | Train Loss: 0.0431 | Train Acc: 98.65% | Test Loss: 0.0375 | Test Acc: 98.81%
Epoch 3/5 | Train Loss: 0.0295 | Train Acc: 99.05% | Test Loss: 0.0342 | Test Acc: 98.88%
Epoch 4/5 | Train Loss: 0.0211 | Train Acc: 99.33% | Test Loss: 0.0331 | Test Acc: 99.07%
Epoch 5/5 | Train Loss: 0.0162 | Train Acc: 99.47% | Test Loss: 0.0352 | Test Acc: 98.93%
最终训练曲线已保存:training_curves.png
所有图表展示完毕!下图是训练工程中损失曲线和准确率曲线:

下图是混淆矩阵,其中对角线是预测正确的次数:

下图是总结出预测失败的样本:
