You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在PyTorch中绘制混淆矩阵?附MNIST-MLP代码求助

给MNIST MLP代码添加混淆矩阵绘制功能

步骤1:添加必要的库导入

在现有导入代码后补充:

from sklearn.metrics import confusion_matrix
import matplotlib.pyplot as plt
import numpy as np

步骤2:将模型移动到指定设备(GPU/CPU)

定义完model后添加一行代码,让模型使用你指定的计算设备:

model = model.to(device)

步骤3:修改训练和测试函数,适配设备并收集分类数据

修改train函数

在data = data.view(-1, 784)前添加,将训练数据移到指定设备:

data, target = data.to(device), target.to(device)

替换原test函数,新增混淆矩阵逻辑

把原有的test()函数替换为以下代码,该函数会同时完成测试评估和混淆矩阵绘制:

def test_and_generate_confusion_matrix():
    model.eval()
    test_loss = 0
    correct = 0
    all_targets = []
    all_preds = []
    with torch.no_grad():
        for data, target in test_loader:
            # 将测试数据移到指定设备
            data, target = data.to(device), target.to(device)
            data = data.view(-1, 784)
            output = model(data)
            test_loss += criterion(output, target).item()
            pred = output.data.max(1, keepdim=True)[1]
            correct += pred.eq(target.data.view_as(pred)).sum()
            # 收集所有真实标签和预测结果(移回CPU处理)
            all_targets.extend(target.cpu().numpy())
            all_preds.extend(pred.squeeze().cpu().numpy())

    test_loss /= len(test_loader.dataset)
    print('Test set: Average loss: {:.4f}, Accuracy: {}/{} ({:.0f}%) '.format(
        test_loss, correct, len(test_loader.dataset),
        100. * correct / len(test_loader.dataset)))
    
    # 计算混淆矩阵
    cm = confusion_matrix(all_targets, all_preds)
    
    # 可视化混淆矩阵
    plt.figure(figsize=(10, 8))
    plt.imshow(cm, interpolation='nearest', cmap=plt.cm.Blues)
    plt.title('MNIST MLP 混淆矩阵')
    plt.colorbar()
    # 设置刻度为0-9的数字
    tick_marks = np.arange(10)
    plt.xticks(tick_marks, [str(i) for i in range(10)], rotation=45)
    plt.yticks(tick_marks, [str(i) for i in range(10)])
    
    # 为每个单元格添加数值标签
    thresh = cm.max() / 2.
    for i, j in np.ndindex(cm.shape):
        plt.text(j, i, format(cm[i, j], 'd'),
                 horizontalalignment="center",
                 color="white" if cm[i, j] > thresh else "black")
    
    plt.tight_layout()
    plt.ylabel('真实标签')
    plt.xlabel('预测标签')
    plt.show()

步骤4:替换原测试调用

将代码末尾的test()替换为:

test_and_generate_confusion_matrix()

完整修改后的代码

import torch
from torch import nn
import torch.nn.functional as F
from torchvision import datasets, transforms
import time
from sklearn.metrics import confusion_matrix
import matplotlib.pyplot as plt
import numpy as np

# enable GPU
if torch.cuda.is_available():
    device = torch.device('cuda')
else:
    device = torch.device('cpu')
    
print('Using PyTorch version:', torch.__version__, ' Device:', device)

# Build a simple MLP to train on MNIST
model = nn.Sequential(
    nn.Linear(784, 128),
    nn.ReLU(),
    nn.Linear(128, 256),
    nn.ReLU(),
    nn.Linear(256, 512),
    nn.ReLU(),
    nn.Linear(512, 10),
    nn.LogSoftmax(dim=1)
)
# 将模型移到指定设备
model = model.to(device)

# Load the training data
train_loader = torch.utils.data.DataLoader(
    datasets.MNIST('data', train=True, download=True,
                     transform=transforms.Compose([
                            transforms.ToTensor(),
                            transforms.Normalize((0.5,), (0.5,))
                        ])),
    batch_size=64, shuffle=True)    

# Load the test data
test_loader = torch.utils.data.DataLoader(
    datasets.MNIST('data', train=False, transform=transforms.Compose([
                            transforms.ToTensor(),
                            transforms.Normalize((0.5,), (0.5,))
                        ])),
    batch_size=64, shuffle=True)


# Define the optimizer
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)

# Define the loss function
criterion = nn.NLLLoss()

# Train the model
def train(epoch):
    model.train()
    for batch_idx, (data, target) in enumerate(train_loader):
        # 将数据移到指定设备
        data, target = data.to(device), target.to(device)
        data = data.view(-1, 784)
        optimizer.zero_grad()
        output = model(data)
        loss = criterion(output, target)
        loss.backward()
        optimizer.step()
        if batch_idx % 100 == 0:
            print('Train Epoch: {} [{}/{} ({:.0f}%)]\tLoss: {:.6f}'.format(
                epoch, batch_idx * len(data), len(train_loader.dataset),
                100. * batch_idx / len(train_loader), loss.item()))

# Test the model and generate confusion matrix
def test_and_generate_confusion_matrix():
    model.eval()
    test_loss = 0
    correct = 0
    all_targets = []
    all_preds = []
    with torch.no_grad():
        for data, target in test_loader:
            # 将数据移到指定设备
            data, target = data.to(device), target.to(device)
            data = data.view(-1, 784)
            output = model(data)
            test_loss += criterion(output, target).item()
            pred = output.data.max(1, keepdim=True)[1]
            correct += pred.eq(target.data.view_as(pred)).sum()
            # 收集所有真实标签和预测结果(移回CPU处理)
            all_targets.extend(target.cpu().numpy())
            all_preds.extend(pred.squeeze().cpu().numpy())

    test_loss /= len(test_loader.dataset)
    print('Test set: Average loss: {:.4f}, Accuracy: {}/{} ({:.0f}%) '.format(
        test_loss, correct, len(test_loader.dataset),
        100. * correct / len(test_loader.dataset)))
    
    # 计算混淆矩阵
    cm = confusion_matrix(all_targets, all_preds)
    
    # 可视化混淆矩阵
    plt.figure(figsize=(10, 8))
    plt.imshow(cm, interpolation='nearest', cmap=plt.cm.Blues)
    plt.title('MNIST MLP 混淆矩阵')
    plt.colorbar()
    # 设置刻度为0-9的数字
    tick_marks = np.arange(10)
    plt.xticks(tick_marks, [str(i) for i in range(10)], rotation=45)
    plt.yticks(tick_marks, [str(i) for i in range(10)])
    
    # 为每个单元格添加数值标签
    thresh = cm.max() / 2.
    for i, j in np.ndindex(cm.shape):
        plt.text(j, i, format(cm[i, j], 'd'),
                 horizontalalignment="center",
                 color="white" if cm[i, j] > thresh else "black")
    
    plt.tight_layout()
    plt.ylabel('真实标签')
    plt.xlabel('预测标签')
    plt.show()

start = time.time()

# main
if __name__ == '__main__': 

    # Run the training loop
    # This is the loop you have to time
    for epoch in range(1, 10):
        train(epoch)

        
end = time.time()
print(end - start)

test_and_generate_confusion_matrix() 
   
# Save the model
torch.save(model.state_dict(), "mnist_mlp.pt")

说明

  • 新增代码会在测试完成后生成可视化混淆矩阵,清晰展示每个数字类别的正确/错误分类情况。
  • 若使用GPU训练,需通过.cpu()将显存中的数据移回CPU,再转换为numpy数组处理。
  • 混淆矩阵的颜色深浅对应分类样本数量,单元格内的数字为具体分类数。

内容的提问来源于stack exchange,提问作者Sakib Mohammad

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.17 05:21:14