基于PyTorch的ResNet18胸部X光图像分类器3×3混淆矩阵构建及代码问题排查求助
修复代码问题并实现混淆矩阵
我帮你排查了代码里的几个关键问题,同时实现了你需要的3×3混淆矩阵功能,下面是详细的修正和实现内容:
一、代码问题修复
我找到了代码里的4个问题并逐一修正:
- 测试集路径错误:
test_dirs中viral的路径错误指向了test/normal,已修正为正确的test/viral路径 - 函数参数笔误:
show_preds函数中调用show_images时,误将preds写成了press,已修正 - 数据集加载逻辑混乱:原
ChestXRayDataset的__getitem__方法随机选择类别,导致样本加载顺序混乱、重复或遗漏,重新实现了按顺序遍历所有样本的逻辑 - 缺失混淆矩阵实现:添加了基于
sklearn和seaborn的混淆矩阵生成与可视化代码
二、完整修正后代码
import os import random import shutil import numpy as np import matplotlib.pyplot as plt import torch import torchvision from PIL import Image from sklearn.metrics import confusion_matrix import seaborn as sns class_names = ['normal', 'viral', 'covid'] root_dir = 'COVID-19 Radiography Database' source_dirs = ['NORMAL', 'Viral Pneumonia', 'COVID-19'] # 数据目录整理(仅第一次运行时需要) if os.path.isdir(os.path.join(root_dir, source_dirs[1])): os.mkdir(os.path.join(root_dir, 'test')) for i, d in enumerate(source_dirs): os.rename(os.path.join(root_dir, d), os.path.join(root_dir, class_names[i])) for c in class_names: os.mkdir(os.path.join(root_dir, 'test', c)) for c in class_names: images = [x for x in os.listdir(os.path.join(root_dir, c)) if x.lower().endswith('png')] selected_images = random.sample(images, 30) for image in selected_images: source_path = os.path.join(root_dir, c, image) target_path = os.path.join(root_dir, 'test', c, image) shutil.move(source_path, target_path) class ChestXRayDataset(torch.utils.data.Dataset): def __init__(self, image_dirs, transform): self.transform = transform self.class_names = ['normal', 'viral', 'covid'] # 整理所有样本路径和对应标签 self.samples = [] for class_idx, class_name in enumerate(self.class_names): image_paths = [os.path.join(image_dirs[class_name], fname) for fname in os.listdir(image_dirs[class_name]) if fname.lower().endswith('png')] print(f'Found {len(image_paths)} {class_name} examples') self.samples.extend([(path, class_idx) for path in image_paths]) def __len__(self): return len(self.samples) def __getitem__(self, index): image_path, label = self.samples[index] image = Image.open(image_path).convert('RGB') return self.transform(image), label train_transform = torchvision.transforms.Compose([ torchvision.transforms.Resize(size=(224, 224)), torchvision.transforms.RandomHorizontalFlip(), torchvision.transforms.ToTensor(), torchvision.transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) test_transform = torchvision.transforms.Compose([ torchvision.transforms.Resize(size=(224, 224)), torchvision.transforms.ToTensor(), torchvision.transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) train_dirs = { 'normal': '/content/gdrive/MyDrive/covidcalssifierdataset/COVID-19 Radiography Database/normal', 'viral': '/content/gdrive/MyDrive/covidcalssifierdataset/COVID-19 Radiography Database/viral', 'covid': '/content/gdrive/MyDrive/covidcalssifierdataset/COVID-19 Radiography Database/covid' } train_dataset = ChestXRayDataset(train_dirs, train_transform) test_dirs = { 'normal': '/content/gdrive/MyDrive/covidcalssifierdataset/COVID-19 Radiography Database/test/normal', 'viral': '/content/gdrive/MyDrive/covidcalssifierdataset/COVID-19 Radiography Database/test/viral', # 修复路径 'covid': '/content/gdrive/MyDrive/covidcalssifierdataset/COVID-19 Radiography Database/test/covid' } test_dataset = ChestXRayDataset(test_dirs, test_transform) batch_size = 6 dl_train = torch.utils.data.DataLoader(train_dataset, batch_size=batch_size, shuffle=True) dl_test = torch.utils.data.DataLoader(test_dataset, batch_size=batch_size, shuffle=False) # 测试集建议不shuffle,方便后续分析 print('Number of training batches', len(dl_train)) print('Number of test batches', len(dl_test)) def show_images(images, labels, preds): plt.figure(figsize=(8, 4)) for i, image in enumerate(images): plt.subplot(1, 6, i + 1, xticks=[], yticks=[]) image = image.numpy().transpose((1, 2, 0)) mean = np.array([0.485, 0.456, 0.406]) std = np.array([0.229, 0.224, 0.225]) image = image * std + mean image = np.clip(image, 0., 1.) plt.imshow(image) col = 'green' if preds[i] != labels[i]: col = 'red' plt.xlabel(f'{class_names[int(labels[i].numpy())]}') plt.ylabel(f'{class_names[int(preds[i].numpy())]}', color=col) plt.tight_layout() plt.show() # 查看训练集和测试集样本 images, labels = next(iter(dl_train)) show_images(images, labels, labels) images, labels = next(iter(dl_test)) show_images(images, labels, labels) # 初始化模型 resnet18 = torchvision.models.resnet18(pretrained=True) resnet18.fc = torch.nn.Linear(in_features=512, out_features=3) loss_fn = torch.nn.CrossEntropyLoss() optimizer = torch.optim.Adam(resnet18.parameters(), lr=3e-5) def show_preds(): resnet18.eval() images, labels = next(iter(dl_test)) outputs = resnet18(images) _, preds = torch.max(outputs, 1) show_images(images, labels, preds) # 修复参数错误 def train(epochs): print('Starting training..') for e in range(0, epochs): print('='*20) print(f'Starting epoch {e + 1}/{epochs}') print('='*20) train_loss = 0. val_loss = 0. resnet18.train() for train_step, (images, labels) in enumerate(dl_train): optimizer.zero_grad() outputs = resnet18(images) loss = loss_fn(outputs, labels) loss.backward() optimizer.step() train_loss += loss.item() if train_step % 20 == 0: print('Evaluating at step', train_step) accuracy = 0 resnet18.eval() for val_step, (images, labels) in enumerate(dl_test): outputs = resnet18(images) loss = loss_fn(outputs, labels) val_loss += loss.item() _, preds = torch.max(outputs, 1) accuracy += sum((preds == labels).numpy()) val_loss /= (val_step + 1) accuracy = accuracy/len(test_dataset) print(f'Validation Loss: {val_loss:.4f}, Accuracy: {accuracy:.4f}') show_preds() resnet18.train() if accuracy >= 0.95: print('Performance condition satisfied, stopping..') return train_loss /= (train_step + 1) print(f'Training Loss: {train_loss:.4f}') print('Training complete..') # 开始训练 train(epochs=20) # 生成并可视化混淆矩阵 def generate_confusion_matrix(model, dataloader, class_names): model.eval() all_labels = [] all_preds = [] with torch.no_grad(): # 关闭梯度计算,节省内存 for images, labels in dataloader: outputs = model(images) _, preds = torch.max(outputs, 1) all_labels.extend(labels.numpy()) all_preds.extend(preds.numpy()) # 计算混淆矩阵 cm = confusion_matrix(all_labels, all_preds) # 可视化混淆矩阵 plt.figure(figsize=(8,6)) sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=class_names, yticklabels=class_names) plt.xlabel('预测类别') plt.ylabel('真实类别') plt.title('胸部X光三分类混淆矩阵') plt.show() return cm # 训练完成后调用 confusion_matrix_result = generate_confusion_matrix(resnet18, dl_test, class_names) print('混淆矩阵结果:') print(confusion_matrix_result)
三、关键说明
- 数据集加载修正:重新实现了
ChestXRayDataset,将所有样本路径和标签提前整理成列表,确保__getitem__按顺序返回每个样本,避免了原逻辑的随机性问题 - 混淆矩阵实现:通过遍历测试集所有样本,收集真实标签和预测结果,使用
sklearn.metrics.confusion_matrix计算矩阵,再用seaborn.heatmap可视化,清晰展示三类样本的分类情况 - 额外优化:测试集的DataLoader设置
shuffle=False,方便后续分析样本的分类结果
内容的提问来源于stack exchange,提问作者abohamza.2000
相关产品推荐
相关产品推荐

