简介:面向医学图像分类任务的肺炎胸片四分类数据集,适合深度学习初学者与医疗影像研究人员直接用于模型训练与验证。数据涵盖COVID(新型冠状肺炎)、Lung_Opacity(肺部浑浊)、Normal(正常)、Viral_Pneumonia(病毒性肺炎)四个类别,共21165张胸片图像,其中训练集16933张、测试集4232张,已按类别分目录存放,无需额外清洗即可输入深度学习框架。资源包总体积约743MB,共2000个文件,其中1998张PNG图片按训练集与测试集分目录存放,另附JSON类别字典与Python可视化脚本,后者可随机展示样本并保存预览图,便于快速核对数据分布。解压后data-train与data-test目录结构清晰,子文件夹名即类别名,配合类别映射文件能方便地完成多分类训练与结果分析。目前已有698人下载学习,是一套适合肺炎图像识别、医学影像分类等教学与实验场景的入门级实践数据。
1. 医学图像分类数据集实测:肺炎胸片4分类,先看懂再动手
一份肺炎胸片4分类的医学图像分类数据集拿到手,第一步不是开训练,而是先把标签、划分和预处理逐项摸清。这个数据集按正常、细菌性肺炎、病毒性肺炎和COVID-19四类组织胸片图像,是典型的医学图像识别与深度学习分类练习样本,也是研究图像识别算法在医疗场景下真实表现的常用载体。新手可以用它把整个图像分类流程从数据读取到指标评估完整跑一遍,熟手则建议重点看它的类别边界、患者级划分和类别不平衡处理,因为这三件事直接决定训练出来的模型能不能落地。
我很少拿到数据集就直接开训,先花半小时把目录结构、文件名规则和标签分布列出来,能少走一整个晚上的弯路。后面几章就按照我平时拆数据集的顺序来写:先讲标签和预处理,给出可直接运行的PyTorch训练骨架,再对比三个常用分类骨干网络,最后把真正容易翻车的地方集中成一份排查清单。
2. 拆解4分类数据集:目录结构、标签映射与预处理管线
2.1 四类标签怎么划定:边界不只在“有没有病”
决定四分类边界的是影像学上的可见差异。正常胸片没有明显局灶或间质病变;细菌性肺炎最常见表现为肺叶或肺段的实变,伴随空气支气管征;病毒性肺炎多表现为间质纹理增粗、磨玻璃影和网状影;COVID-19的特征则集中在双肺外周和胸膜下分布为主的磨玻璃影与实变。模型要学的不是“这是不是病”,而是这四类之间的影像学差异。很多新手直接把普通二分类网络改成四分类就跑,结果混淆矩阵里viral和covid互相错得离谱,原因就是没有先确认类别边界定义。
拿到数据集第一件事是看标签文件名和自带标签文件,确认四类具体是哪四类。有些版本的4分类把“其他肺部疾病”作为第四类,COVID-19单列;另一些则把COVID-19并入病毒性肺炎。标签定义不同,训练目标和评价口径完全不同。如果数据集给出的类别名不是normal、bacterial、viral、covid,就按实际目录名重写后面代码里的class_to_idx映射,其余流程不用变。处理数据集的顺序永远是先摸清语义,再写代码。
2.2 目录结构:train/val/test怎么分
这类胸片分类数据集的常见组织方式是图像按类别分子目录,train、val、test三套目录互相独立:
data/ ├── train/ │ ├── normal/ # 正常胸片 │ ├── bacterial/ # 细菌性肺炎 │ ├── viral/ # 病毒性肺炎 │ └── covid/ # COVID-19 ├── val/ │ ├── normal/ │ ├── bacterial/ │ ├── viral/ │ └── covid/ └── test/ ├── normal/ ├── bacterial/ ├── viral/ └── covid/划分比例常见的有8:1:1和6:2:2两种。如果是自己切数据,我一般用6:2:2,因为医学数据集本身噪声大,验证集太薄会导致早停频次和最终指标的方差变大,20%的验证数据能让每个epoch的评估结果稳定不少。更重要的是划分方式:不要按图片随机划分,要按患者划分。同一个患者可能有多张胸片,如果一张进了train、另一张进了val,模型在val上相当于提前见过了这个人,指标虚高,且很难在事后通过调参消除。后面第5章会专门讲怎么用患者ID重新划分,这里先记住结论。
2.3 预处理:灰度转RGB与归一化
胸片在文件层面是灰度图,但torchvision里常用的预训练模型输入是三通道。处理这个矛盾的常见做法是把灰度图复制成三通道,再用ImageNet的均值方差归一化:
from PIL import Image import numpy as np # 方式A:PIL直接转换 img = Image.open("chest_xray.jpg") img_gray = img.convert("L") # 强制转成单通道灰度 img_rgb = img_gray.convert("RGB") # 单通道复制成三通道 # 方式B:numpy复制通道,适合做像素级分析时使用 arr = np.array(img_gray) # shape: (H, W) arr_rgb = np.stack([arr, arr, arr], axis=-1) # shape: (H, W, 3)逻辑说明:convert("L")先把图片转成严格灰度,避免某些JPEG里残留的彩色通道信息干扰后续分布统计;convert("RGB")再把灰度值复制到三个通道,让输入shape满足预训练模型的3通道要求。方式B适合在需要自己写归一化逻辑、或需要做灰度直方图统计的时候使用。
参数说明:复制三通道只解决shape匹配,没有解决分布匹配。ImageNet预训练权重看到的是自然图像的统计特征,胸片是三通道同值的灰度分布,二者不完全一致。所以后面微调时训练轮数不能太少,一般会跑15到30个epoch,给底层卷积核足够时间适应胸片的灰度分布,而不是指望第一个epoch就收敛。
2.4 数据增强:哪些能用、哪些别碰
医学图像增强和自然图像增强有明显区别。胸片上病灶可能只占整张图像的百分之几,RandomResizedCrop这种先随机裁剪再缩放的增强,很容易把病灶区域裁掉,模型学到的全是正常组织纹理。ColorJitter里的大范围亮度、对比度扰动,也可能把磨玻璃影和正常灰度之间的微小差异抹平。在自然图像上越强的增强,在胸片上越可能是负优化。
我在这类数据集上常用的增强配置如下:
from torchvision import transforms IMG_SIZE = 224 train_transform = transforms.Compose([ transforms.Resize((IMG_SIZE, IMG_SIZE)), transforms.RandomRotation(degrees=10), transforms.RandomHorizontalFlip(p=0.5), transforms.RandomAffine(degrees=0, translate=(0.05, 0.05), scale=(0.9, 1.1)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) val_transform = transforms.Compose([ transforms.Resize((IMG_SIZE, IMG_SIZE)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ])参数说明:RandomRotation的10度是有意限制的,胸片有明确的上下方向,大的旋转会让模型学到错误的解剖方向,这在解剖结构强相关的医学图像里是常见坑。RandomAffine里的scale范围0.9到1.1模拟不同设备的放大倍率差异,translate的精度0.05模拟轻微的拍摄偏移。验证集上不做任何随机增强,只做Resize和Normalize,这是为了让每次验证指标完全可复现,比较不同实验时才有意义。
3. 用PyTorch训练肺炎胸片4分类模型:从数据加载到指标评估
3.1 自定义Dataset:读目录、建标签映射
有了上面的目录结构,用torch.utils.data.Dataset写一个能同时用于train和val的类。关键点在于:类别顺序要稳定,读取时要过滤非图片文件,路径和标签要一一对齐。
import os from PIL import Image from torch.utils.data import Dataset class PneumoniaDataset(Dataset): def __init__(self, root_dir, transform=None): self.root_dir = root_dir self.transform = transform # 对类别名排序,保证每次运行标签顺序一致 self.classes = sorted(os.listdir(root_dir)) self.class_to_idx = {name: idx for idx, name in enumerate(self.classes)} self.samples = [] for cls in self.classes: cls_dir = os.path.join(root_dir, cls) if not os.path.isdir(cls_dir): continue for fname in os.listdir(cls_dir): if fname.lower().endswith((".jpg", ".jpeg", ".png")): path = os.path.join(cls_dir, fname) self.samples.append((path, self.class_to_idx[cls])) def __len__(self): return len(self.samples) def __getitem__(self, idx): path, label = self.samples[idx] image = Image.open(path).convert("RGB") if self.transform: image = self.transform(image) return image, label逻辑说明:__init__里把类别目录名按字母序排序,normal、bacterial、viral、covid对应的标签索引依序是0、1、2、3。这只在目录名恰好是这四个单词时成立,如果数据集目录命名不同,排序后索引也会跟着变,训练前先打印train_ds.classes确认一下,否则后续评估的target_names全盘错位。
参数说明:过滤后缀是必须的,有些数据集目录里混着Thumbs.db或者其他非图片文件,不过滤会在训练一开始就报解码错误。convert("RGB")统一转三通道,避免某些单通道PNG直接进入模型时shape对不上。__getitem__里每次重新打开图片,瓶颈在磁盘IO,Windows上用num_workers=0或1更稳,Linux可以开到4以上。
3.2 数据增强流水线与DataLoader参数
加载部分把第2章的transform接进Dataset,再设置batch size和num_workers。batch size不是越大越好:胸片数据集总量不大,batch太大一个epoch就几百步,BN统计量不够平滑;batch太小梯度抖动又太大。
from torch.utils.data import DataLoader train_ds = PneumoniaDataset("data/train", transform=train_transform) val_ds = PneumoniaDataset("data/val", transform=val_transform) train_loader = DataLoader(train_ds, batch_size=32, shuffle=True, num_workers=4) val_loader = DataLoader(val_ds, batch_size=32, shuffle=False, num_workers=4)参数说明:shuffle=True只用于训练集,val_loader固定顺序,目的是让每个epoch的验证顺序一致,避免随机顺序干扰early stopping判断。batch_size选32是兼顾显存和梯度稳定性的折中,显存不够就降到16,每类样本数较多时也可以升到64,但要重新确认val指标没有明显变差。num_workers设置为4时,如果机器内存紧张或Windows系统容易报DataLoader worker错误,改成0最省事,代价是训练速度会慢一些。
3.3 类别不平衡:先用加权损失
胸片数据集的另一个常见问题是类别数量不均衡。正常胸片数量通常远多于细菌性、病毒性,如果直接CrossEntropyLoss,模型会倾向于把一切预测成正常类,整体准确率看着还行,实际对肺炎类别几乎没有判别力。处理方式是在损失函数里按类别样本数反比加权:
import numpy as np import torch import torch.nn as nn from sklearn.utils.class_weight import compute_class_weight # device 在训练脚本里用 torch.device("cuda" if torch.cuda.is_available() else "cpu") 得到 labels = [label for _, label in train_ds.samples] classes = np.array([0, 1, 2, 3]) weights = compute_class_weight("balanced", classes=classes, y=np.array(labels)) weights = torch.tensor(weights, dtype=torch.float).to(device) criterion = nn.CrossEntropyLoss(weight=weights)逻辑说明:compute_class_weight的balanced模式计算公式是n_samples除以类别数和类别频数的乘积,少数类样本少,权重就大。CrossEntropyLoss的weight参数会作用在归一化后的每个样本loss上,权重大的类别在反向传播里梯度放大,逼着模型在少数类上多花学习容量。
参数说明:如果训练后发现少数类recall还是偏低,例如viral类始终只有30%左右,可以把对应权重再乘一个1.5的缩放因子,把类别权重张量里的viral索引位单独放大。这种调整要配合验证集混淆矩阵来观察,别只盯着训练loss数值。
3.4 训练主循环与学习率设定
接下来是标准训练循环,优化器用AdamW配合L2正则。医学小数据集上weight_decay太大会压制模型容量,太小又控制不住噪声,我用1e-4起步。模型构建函数build_model在第4章给出完整实现,这里先按同一接口调用。
def train_one_epoch(model, loader, criterion, optimizer, device): model.train() total_loss, total_correct, total_num = 0.0, 0, 0 for images, labels in loader: images = images.to(device) labels = labels.to(device) optimizer.zero_grad() outputs = model(images) loss = criterion(outputs, labels) loss.backward() optimizer.step() total_loss += loss.item() * images.size(0) total_correct += (outputs.argmax(dim=1) == labels).sum().item() total_num += images.size(0) return total_loss / total_num, total_correct / total_num model = build_model(backbone="efficientnet_b0", num_classes=4) optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-4) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=30)逻辑说明:outputs.argmax(dim=1)在四分类输出上取最大概率对应的类作为预测结果,total_correct统计当前epoch内预测正确的图片数。最后返回平均loss和训练集准确率,方便打印日志和判断是否过拟合。
参数说明:lr=1e-4是微调预训练模型时的常见起点,太大容易在第一个epoch就把预训练权重冲坏,太小则需要更多轮数才能看到loss明显下降。CosineAnnealingLR的T_max设为30,对应计划训练30个epoch,学习率从1e-4余弦降到接近0,省去手动分段调整学习率的麻烦。如果训练到一半发现验证loss反弹,就提前停止,不要等完整30轮。
3.5 评估:准确率只是起点,要看每类召回率
四分类医学任务里accuracy容易被类别不平衡掩盖,必须同时看分类报告和混淆矩阵。这一步输出比训练循环本身更能说明模型问题:
from sklearn.metrics import classification_report, confusion_matrix def evaluate(model, loader, device): model.eval() all_preds, all_labels = [], [] with torch.no_grad(): for images, labels in loader: images = images.to(device) outputs = model(images) all_preds.extend(outputs.argmax(dim=1).cpu().numpy()) all_labels.extend(labels.numpy()) print(classification_report(all_labels, all_preds, target_names=["normal", "bacterial", "viral", "covid"])) print(confusion_matrix(all_labels, all_preds))逻辑说明:with torch.no_grad()关闭梯度计算,评估阶段推理提速并省显存。classification_report输出每个类别的precision、recall、f1,confusion_matrix输出4乘4的矩阵,能直接看到哪两类互相混淆。这里返回的all_labels和all_preds在项目收尾时还会再用来画图。
参数说明:target_names要与Dataset里的类别索引顺序一致,如果顺序错了,报告里的类名和数字对不上,比不打印还容易误导人。对胸片四分类,重点看bacterial和viral的recall是否明显低于normal,而不是先看整体accuracy。这一点在医学图像识别场景里几乎是默认共识。
4. 迁移学习选型:ResNet、DenseNet、EfficientNet怎么选
4.1 为什么小样本医学分类必须迁移学习
胸片分类数据集的规模通常只有几千到几万张,四分类情况下每类可能只有几百到几千张,从头训练一个CNN很容易过拟合,而且训练时间成本高。预训练模型在ImageNet上已经学到了边缘、纹理、形状这些通用视觉特征,迁移学习的本质是把这部分通用特征直接拿过来,只需要在胸片数据上微调高层语义特征。对于肺炎胸片识别,这是目前最稳的做法,比从零搭网络效果好很多,收敛速度也能快一半以上。
4.2 三个backbone的参数对比
三个模型在torchvision里都能用一行代码加载预训练权重,但各自特点差别不小:
| backbone | 参数量 | 输入分辨率 | 训练稳定性 | 适合场景 |
|---|---|---|---|---|
| ResNet50 | 约25.6M | 224×224 | 很稳定 | 显存充裕、追求稳妥基线 |
| DenseNet121 | 约8M | 224×224 | 稳定 | 小数据、小显存、特征复用 |
| EfficientNet-B0 | 约5.3M | 224×224 | 对学习率敏感 | 资源受限、快速实验 |
DenseNet121的Dense Block把每层输出都拼接到后续层,特征复用让它在参数更少的情况下保留足够表达能力,我在小规模胸片数据集上更常用它。EfficientNet-B0理论FLOPs最低,但它的MBConv结构对学习率和batch size更敏感,训练时如果发现loss来回震荡,通常不是网络设计问题,而是学习率没调对,需要配合warmup或更低的学习率才能稳住。
4.3 可切换backbone的模型构建代码
写一个build_model函数,用字符串参数切换三个骨架,避免每个实验单独改模型定义:
import torch.nn as nn from torchvision import models def build_model(backbone="efficientnet_b0", num_classes=4): if backbone == "resnet50": model = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V2) model.fc = nn.Linear(model.fc.in_features, num_classes) elif backbone == "densenet121": model = models.densenet121(weights=models.DenseNet121_Weights.IMAGENET1K_V1) model.classifier = nn.Linear(model.classifier.in_features, num_classes) else: model = models.efficientnet_b0(weights=models.EfficientNet_B0_Weights.IMAGENET1K_V1) model.classifier[1] = nn.Linear(model.classifier[1].in_features, num_classes) return model逻辑说明:三个模型改的是不同位置的分类头。ResNet50最后是fc层,DenseNet121是classifier线性层,EfficientNet-B0的classifier是一个Sequential容器,最后一个元素才是全连接层,所以要改classifier[1]而不是classifier。这是用同一个函数切换backbone时最容易写错的地方。
参数说明:weights=models.xxx是torchvision新版本推荐写法,旧版本的pretrained=True已经弃用。num_classes设为4时,新分类头的输出维度是4,与前面Dataset的标签索引范围一致。如果数据集类别改名了,这里num_classes要同步调整。
4.4 两段式微调与BN冻结
用预训练权重微调时,我习惯分两步走:先冻结骨干,只训练刚替换的分类头,跑5个epoch让分类头先收敛到大致可用的状态;再解冻全部参数,用更小的学习率微调整个网络。这样做的原因是新分类头随机初始化,如果一开始就直接全量微调,大梯度会反向传播到骨干网络,把预训练权重破坏掉。
def set_bn_eval(model): for module in model.modules(): if isinstance(module, nn.BatchNorm2d): module.eval()逻辑说明:冻结骨干后,torchvision预训练模型里的BatchNorm层仍然会更新running mean和running var。如果batch size比较小,BN统计量波动大,反而把已经稳定的底层特征搞乱。set_bn_eval让所有BN层进入eval模式,保持预训练统计量不变,只让分类头和少数残差分支的参数参与更新。
参数说明:第一段微调用lr=1e-3,只更新分类头参数,优化器里需要传入filter(lambda p: p.requires_grad, model.parameters())来排除冻结参数;第二段解冻骨干后把lr降到1e-4。如果机器显存允许,建议第一阶段就用4到8的batch size跑,第二阶段再切回32,因为冻结阶段参数量少,小batch也能快速收敛。
5. 肺炎胸片分类的避坑与排查:五个血泪经验
5.1 数据泄漏:同一患者出现在train和val
现象:在默认划分的val集上准确率有0.95,但一换到独立的test集只有0.85;或者训练时val曲线一直很漂亮,打印混淆矩阵却发现问题样本集中在某几个患者ID前缀下。
原因:数据集的图像文件是按图片随机划分的,同一患者的多张随访胸片被拆进了train和val。这类胸片常有连续拍摄的序列,相邻帧差异极小,模型实际记住了患者ID而不只是病灶特征。
解决:拿到数据先提取文件名中的患者ID,按患者分组重划。没有现成患者ID字段时,预处理阶段直接把文件路径切出一段作为ID。用GroupShuffleSplit实现:
import pandas as pd from sklearn.model_selection import GroupShuffleSplit df = pd.DataFrame({"path": paths, "label": labels, "patient_id": patient_ids}) gss = GroupShuffleSplit(n_splits=1, test_size=0.2, random_state=42) train_idx, val_idx = next(gss.split(df, groups=df["patient_id"]))逻辑说明:GroupShuffleSplit的groups参数接收患者ID数组,划分时保证同一个patient_id的所有行被分到同一侧。n_splits=1表示只需要切一次,test_size=0.2对应20%进验证集。这里直接next取出唯一的划分结果,比for循环遍历更简洁。
参数说明:patient_id需要从文件名里提取,比如文件名Rx_0234_20240105.jpg里的0234就是患者号,正则表达式提取后先做唯一性检查,确认没有重复才能当group字段用。如果文件名里没有患者号,就把完整文件名当group,这样至少避免同一文件被复制到不同目录时造成泄漏。
5.2 灰度图复制三通道的域偏移
现象:训练前几个epoch的loss下降很慢,最终准确率卡在一个不高不低的位置,怎么调学习率都突破不了。
原因:ImageNet预训练权重是在彩色自然图像上训练的,底层卷积核学习的是彩色边缘和纹理组合。胸片复制三通道后,三个通道完全一样,输入分布的统计特性和预训练分布不匹配,需要额外几个epoch才能让底层适应。
解决:默认接受复制三通道,但要给足微调轮数,我用30轮起步。如果实验环境允许,也可以使用在医学图像上预训练的权重,或在自己收集的大批量无标签胸片上做自监督预训练再微调。后者成本高,通常只有研究场景才值得做。
5.3 类别不平衡让模型只输出正常类
现象:整体准确率不低,打印classification_report后却发现normal的recall很高,bacterial、viral的recall只有30%左右;模型把大部分胸片都判成正常。
原因:CrossEntropyLoss对每个样本同等看待,大多数样本是正常类,梯度被正常类主导,少数类的错误对loss影响太小,模型选择躺平。
解决:用第3.3节的加权CrossEntropy。如果加权后少数类recall还是低,考虑Focal Loss,它按样本难易程度给权重,gamma=2、alpha=0.25是常见起手参数。Focal Loss核心实现如下:
import torch import torch.nn as nn import torch.nn.functional as F class FocalLoss(nn.Module): def __init__(self, gamma=2.0, alpha=None): super().__init__() self.gamma = gamma self.alpha = alpha def forward(self, logits, targets): ce = F.cross_entropy(logits, targets, reduction="none") p = torch.exp(-ce) loss = (1 - p) ** self.gamma * ce if self.alpha is not None: alpha = self.alpha[targets] loss = alpha * loss return loss.mean()逻辑说明:p是模型对正确类别的置信度,p越接近1说明样本越容易被分类,系数(1-p)^gamma越小,难样本的loss权重越大。alpha按类别加权,处理样本数量差异。两个机制叠加就是Focal Loss的核心。
参数说明:gamma增大会让难样本权重更明显,一般从2.0开始调,验证集viral的recall上不去就适当加大;alpha需要传入和类别数等长的张量,比如torch.tensor([0.4, 0.8, 1.0, 0.8]),这里torch.tensor的dtype默认float32,直接乘到loss上没有问题。
5.4 数据增强过猛把病灶抹平
现象:用了RandomResizedCrop和较强ColorJitter后,训练集loss下降变慢,val指标波动反而变大;偶尔出现val准确率比train还高的情况。
原因:胸片病灶通常只占图像几个百分点,RandomResizedCrop的随机裁剪可能把病灶切掉,模型只能学正常纹理。ColorJitter的高对比度扰动又可能让磨玻璃影和周围组织灰度差被压缩。这类增强在自然图像上是常规操作,在医学图像上属于负优化。
解决:按第2.4节的轻量增强配置执行。旋转不超过10度,scale范围0.9到1.1,translate控制在0.05以内,别用RandomResizedCrop。如果发现val指标波动大,先检查增强强度是不是太大,再考虑换模型,多半是增强配置的问题而不是网络容量的锅。
5.5 标签噪音与跨设备域差异
现象:训练曲线正常,但混淆矩阵里normal和viral互混特别多,且错误样本集中在某几个文件名前缀相同的批次里。
原因:数据集可能由多个来源拼接,不同医院或不同DR设备的胸片灰度分布差异很大;部分标签也可能在人工标注环节标错。这属于标签噪音和域差异的叠加,模型可能学到的是设备特征而不是病理特征。
解决:用训练好的模型在测试集上找置信度高的错误预测,逐个人工复核;画出各来源样本的灰度直方图,观察分布是否明显分群。如果确认存在域差异,按来源做分层划分,避免某个来源的样本只在train或val里单独出现。处理数据集时的这一层排查,往往比换模型结构带来的收益更明显。
6. 验证模型的最后一公里:混淆矩阵与患者级分区检查
训练结束不等于项目收工。对胸片4分类这种医学任务,我最后一步永远是用患者级分区重新验证一遍。做法很简单:从文件名中提取患者ID,用GroupShuffleSplit按患者重新切一次训练验证集,再训练同一个模型、看同一组指标。如果按患者划分的验证准确率比默认随机划分低2到3个百分点,说明原划分里存在数据泄漏;如果两个指标几乎一致,说明模型学的确实是病灶特征而不是患者特征。这个验证只需要半天时间,但能避免拿到一个看起来漂亮、实际上靠记忆患者ID撑起来的模型。
验证时只打印准确率不够,我把混淆矩阵直接画出来存图。第3.5节evaluate函数返回的all_labels和all_preds在这里直接复用:
import matplotlib.pyplot as plt from sklearn.metrics import ConfusionMatrixDisplay cm = confusion_matrix(all_labels, all_preds) disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=["normal", "bacterial", "viral", "covid"]) disp.plot(cmap="Blues") plt.savefig("confusion_matrix_patient_level.png", dpi=300)逻辑说明:ConfusionMatrixDisplay把sklearn的混淆矩阵包装成图,display_labels传入类别名,plot函数自动填充色块和数字。dpi=300保证后续复用清晰度,日常快速验证用150就够。
参数说明:如果四类类别名和实际索引顺序不符,这张图会误导后续判断。保存前先确认train_ds.classes的顺序,再决定display_labels怎么写。另外,图上数字的排列能直接暴露类间混淆模式,比如normal被大量错分成viral,就要回到第5.5节的域差异排查,而不是继续调参。
从那以后,我每次拿到医学图像分类数据集,都强制走一遍患者ID分组、灰度直方图核对、轻量增强、加权损失、患者级验证这五步。肺炎胸片4分类这个场景尤其如此,因为正常类和疾病类样本量悬殊,多设备来源又让域差异明显,跳过任何一步都可能把数据处理的痕迹误判成模型能力。数据集的目录结构、标签定义和划分方式拿到手先验一遍,再往训练阶段走,这是最省时间的顺序。希望帮到你。
本文还有配套的精品资源,点击获取