如何在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
相关产品推荐
相关产品推荐

