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

基于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)

三、关键说明

  1. 数据集加载修正:重新实现了ChestXRayDataset,将所有样本路径和标签提前整理成列表,确保__getitem__按顺序返回每个样本,避免了原逻辑的随机性问题
  2. 混淆矩阵实现:通过遍历测试集所有样本,收集真实标签和预测结果,使用sklearn.metrics.confusion_matrix计算矩阵,再用seaborn.heatmap可视化,清晰展示三类样本的分类情况
  3. 额外优化:测试集的DataLoader设置shuffle=False,方便后续分析样本的分类结果

内容的提问来源于stack exchange,提问作者abohamza.2000

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.29 16:23:12