实时训练曲线 + 混淆矩阵 + 错误样本

🎉摘要:本文提供Python代码实现MNIST CNN模型训练,通过Matplotlib实时绘制训练/测试损失与准确率曲线,并生成混淆矩阵及展示错误分类样本,帮助深度学习初学者直观理解模型性能。

将下面代码保存到 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

所有图表展示完毕!

下图是训练工程中损失曲线和准确率曲线:

image.png

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

image.png

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

image.png



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