☰
医学图像分类实战:微生物数据集与CNN/YOLOv5训练指南
2026/10/7 23:08:11 网站建设 项目流程

简介:这是一份面向医学图像分类与目标检测任务的微生物识别数据集,包含阿米巴、眼虫属、水螅、草履虫等8类显微图像,数据规模适中,适合深度学习初学者与科研人员快速训练CNN分类网络或YOLOv5分类模型。数据已按文件夹整理为训练集和测试集,训练集共630张图片、测试集共150张图片,以jpg/png/jpeg图片为主,另有类别字典json文件和一个可视化show脚本,可以快速查看每个类别的样本分布。整个压缩包约102MB,共792个文件,目录结构清晰,能够直接接入PyTorch、TensorFlow或YOLOv5等常见训练流程,免去自行标注和划分数据的麻烦,也方便按需调整训练集与测试集比例。目前已有177人学习,该数据集可用于微生物识别课题、课程设计、算法对比实验或作为入门分类项目的练手数据,整体规范且开箱即用。

1. 医学图像分类数据集:8种微生物图像识别,训练测试划分好的开箱方案

做图像分类的同学应该都有过这种经历:模型结构选好了,环境配好了,结果卡在数据集上。要么是公开数据集太大,下载半天发现类别对不上;要么是数据没划分,得自己写脚本按比例切,切完还要担心验证集和训练集有没有图片重复。这个8种微生物图像识别数据集,就是冲着这个痛点来的:训练集630张、测试集150张,已经按目录结构划分好,附带类别字典JSON文件,解压就能直接喂给YOLOv5分类分支或者CNN分类网络。

数据集覆盖阿米巴、眼虫属、水螅、草履虫等8个类别,属于医学显微图像里比较典型的微生物形态。单类样本量不大,但正因为小,反而适合用来跑通分类训练全流程:数据加载、标签映射、模型训练、指标评估。对刚接触分类任务的初学者来说,拿它练手比直接上ImageNet级别的数据集要友好得多;对熟手来说,它最值钱的地方在于划分好的目录结构和JSON类别字典,省掉了数据工程里最琐碎的一步。

2. 从文件目录到训练管线:读透数据集的结构与字典

拿到这个数据集,第一步不是直接开始训练,而是先在本地把目录结构完整盘一遍。很多人在这一步跳过,结果训练到一半发现路径写错或者类别索引对不上,回来排查的时候又浪费一两个小时。这个数据集的设计很直接:根目录下只有一个data文件夹,里面放着train和test两个子目录,每个子目录里按类别分别存放图片,类别信息同时维护在一份JSON字典文件里。

2.1 目录结构解析:训练集与测试集的存放逻辑

我习惯用tree命令先把结构打出来,只看前两层就行,不需要把每张图片都列出来。这个数据集解压之后的结构应该是这样的:

├── data │ ├── train │ │ ├── amoeba │ │ ├── euglena │ │ ├── hydra │ │ └── ... │ └── test │ ├── amoeba │ ├── euglena │ ├── hydra │ └── ... ├── class_dict.json └── show.py

各目录的含义如下:

  • data/train:训练集目录,存放630张图片,按类别分子文件夹
  • data/test:测试集目录,存放150张图片,划分方式和训练集保持一致
  • class_dict.json:类别字典文件,维护类别名称与标签索引的映射关系
  • show.py:可视化脚本,用于随机展示数据集中的样本图片

注意训练集和测试集的子目录名称必须一致,这是后续加载数据时能直接复用目录名作为标签的前提。之前遇到过有些数据集训练集用英文名、测试集用中文名,或者train目录里写Amoeba、test目录里写amoeba,大小写不一致导致标签映射错位,这个数据集没有这个问题。

2.2 类别字典JSON:标签映射的单一数据源

JSON字典是整个数据集的"黑匣子"开关。训练脚本读取类别顺序时,必须以这个文件为准,不能靠手数目录数量来猜。里面的内容格式大概是这样的:

{ "0": "amoeba", "1": "euglena", "2": "hydra", "3": "paramecium", "4": "stentor", "5": "volvox", "6": "yeast", "7": "diatom" }

这个映射关系决定了模型输出的类别序号对应什么微生物。训练时类别顺序就是模型最后全连接层的输出维度顺序,比如0对应"amoeba"(阿米巴),那模型预测输出索引0就意味着它认为这张图是阿米巴。如果后续你想增加类别或者调整顺序,只改这个JSON不够,对应目录结构也要同步调整,两边必须严格一致。

加载JSON并生成训练标签,常见做法是这样:

import json import os from torch.utils.data import Dataset from PIL import Image class MicrobeDataset(Dataset): def __init__(self, data_dir, class_dict_path, transform=None): self.transform = transform with open(class_dict_path, 'r', encoding='utf-8') as f: self.class_dict = json.load(f) # 根据字典的值(类别名)反查索引 self.class_to_idx = {name: int(idx) for idx, name in self.class_dict.items()} self.samples = [] for class_name in os.listdir(data_dir): class_path = os.path.join(data_dir, class_name) if not os.path.isdir(class_path): continue idx = self.class_to_idx[class_name] for img_name in os.listdir(class_path): self.samples.append((os.path.join(class_path, img_name), idx)) def __len__(self): return len(self.samples) def __getitem__(self, index): img_path, label = self.samples[index] img = Image.open(img_path).convert('RGB') if self.transform: img = self.transform(img) return img, label

这段代码的核心逻辑是:先把JSON里的映射反转成类别名 -> 索引,然后遍历data_dir下的所有子文件夹,把每个文件夹内的图片路径和对应标签配对成样本列表。__getitem__里返回的是PIL图像和整数标签,后续交给DataLoader时,需要在这里接入transform做预处理和增强。

这里有两个参数值得注意:class_dict_path建议传绝对路径,避免相对路径在不同工作目录下解析出错;transform参数控制在外部传入,数据集类本身不做图像预处理,这样可以在训练脚本里灵活切换训练集和测试集的不同增强策略。

2.3 可视化脚本:训练前先确认数据没跑偏

数据加载写完之后,别急着训练,先用自带的可视化脚本试一试。这个show.py做的事情本质上是随机抽几张图,把类别名打在图上展示出来,确认图片和标签确实对得上。如果你习惯自己写可视化,可以用更简单的办法:

import matplotlib.pyplot as plt from torchvision.utils import make_grid from torch.utils.data import DataLoader dataset = MicrobeDataset('data/train', 'class_dict.json') loader = DataLoader(dataset, batch_size=16, shuffle=True) images, labels = next(iter(loader)) grid = make_grid(images, nrow=4, padding=4) plt.imshow(grid.permute(1, 2, 0)) plt.title('Microbe Training Samples') plt.axis('off') plt.show()

跑一遍这个脚本,重点看两件事:一是类别数量是不是和JSON里一致,二是每张图的微生物形态是否明显可区分。如果类别之间形态过于相似,后面的训练需要更大的分辨率或者更强的数据增强,这一点在模型选型时就要想清楚。

提示:目录结构和JSON字典文件是配套使用的,不要单独改动其中一项。如果你把训练集里的某类子目录改名,一定要同步改JSON里对应的类别名和索引。

3. 把数据集喂给分类网络:两条训练路线的参数选择与脚本拆解

数据准备好之后,接下来是实际训练环节。这个数据集可以同时用于传统CNN分类网络和YOLOv5的分类分支,两者在数据加载方式上有区别:CNN分类网络需要自己写Dataset类并在训练循环里迭代,YOLOv5则直接把数据集放在指定目录结构下跑train.py就行。两条路线各有适合的场景,下面的参数设置可以参考。

3.1 ResNet分类训练:迁移学习与关键参数

对于只有630张训练图片的小数据集,直接从头训练一个深度网络效果往往不理想,这也容易给自己加阻力。常见做法是使用ImageNet预训练权重做迁移学习,微调最后一层即可。用ResNet18做示例,训练脚本的核心部分长这样:

import torch import torch.nn as nn import torch.optim as optim from torchvision import models, transforms from torch.utils.data import DataLoader # 预训练权重加载 model = models.resnet18(pretrained=True) num_features = model.fc.in_features model.fc = nn.Linear(num_features, 8) # 8个微生物类别 # 设置损失函数和优化器 criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=0.0005) # 训练集与测试集的预处理策略 train_transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomHorizontalFlip(p=0.5), transforms.RandomRotation(15), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) test_transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) train_dataset = MicrobeDataset('data/train', 'class_dict.json', train_transform) test_dataset = MicrobeDataset('data/test', 'class_dict.json', test_transform) train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=4) test_loader = DataLoader(test_dataset, batch_size=32, shuffle=False, num_workers=4)

这里的关键参数说明:

  • pretrained=True:使用ImageNet预训练权重初始化网络。虽然微生物图像和ImageNet的自然图像分布差异不小,但底层边缘、纹理特征仍然可迁移,能明显加快收敛速度
  • fc层输出设置为8:对应类别字典里定义的8个微生物类别。如果类别数变了,这里必须同步修改
  • lr=0.0005:迁移学习场景下,全连接层是随机初始化的,其余层是预训练的。这个学习率对预训练层偏大,但对新初始化的分类头偏小。经验做法是分类头用0.001、骨干网络用0.0001,这里取一个折中值
  • RandomRotation(15):微生物图像没有明确的"正"方向,15度的随机旋转可以提升模型对拍摄角度的鲁棒性
  • Normalize参数使用ImageNet的均值和标准差,这和预训练权重的初始化统计分布保持一致

训练循环本身没有特别之处,每个epoch结束后在测试集上计算准确率,保存验证集上最佳模型:

best_acc = 0.0 for epoch in range(30): model.train() running_loss = 0.0 for images, labels in train_loader: optimizer.zero_grad() outputs = model(images) loss = criterion(outputs, labels) loss.backward() optimizer.step() running_loss += loss.item() * images.size(0) model.eval() correct = 0 total = 0 with torch.no_grad(): for images, labels in test_loader: outputs = model(images) _, predicted = torch.max(outputs, 1) total += labels.size(0) correct += (predicted == labels).sum().item() acc = 100.0 * correct / total print(f'Epoch {epoch+1}, Loss: {running_loss/len(train_dataset):.4f}, Acc: {acc:.2f}%') if acc > best_acc: best_acc = acc torch.save(model.state_dict(), 'best_model_resnet18.pth')

测试集150张,准确率的波动会比较明显,可能出现某个epoch已经达到98%、下一个epoch掉到94%的情况。建议保存最高精度的模型权重作为最终产物。另外,整个训练过程不需要设置早停,跑完全部epoch再挑最优即可,因为小数据集的验证集波动大,早停阈值不好卡。

3.2 YOLOv5分类模式:数据格式直接对接训练命令

如果不想手写训练脚本,YOLOv5的分类训练模式会更省事。它的分类任务直接支持目录结构的数据集,运行classify/train.py脚本即可。YOLOv5要求的数据组织和这个数据集天然匹配:

python classify/train.py \ --model yolov5s-cls.pt \ --data data \ --epochs 50 \ --img 224 \ --batch 32 \ --lr 0.01 \ --save-period 10 \ --name microbe_cls

命令里的参数含义如下:

  • --model yolov5s-cls.pt:YOLOv5官方提供的分类预训练模型权重,yolov5s是轻量版本,630张图的训练集在普通显卡上几分钟就能跑一个epoch
  • --data data:直接指向刚才解析过的data目录,YOLOv5会自动识别data/train和data/test两个子目录作为分类训练和验证数据
  • --img 224:输入图像分辨率,和CNN路线保持一致
  • --lr 0.01:YOLOv5分类任务默认的初始学习率
  • --save-period 10:每10个epoch保存一次权重,防止最后几轮过拟合导致最优权重丢失

训练完成后,推理用classify/predict.py,对单张图片预测类别:

python classify/predict.py --weights runs/train/microbe_cls/weights/best.pt --source test.jpg

输出结果会直接显示预测类别名和置信度,因为YOLOv5会自动解析训练集目录名作为类别标签。使用这个数据集时,整个训练流程只需要在启动前关注一遍目录有没有放对位置,剩下的步骤比较顺畅。

3.3 两条路线的适用边界

从实践角度给两条路线做个对比选择参考:

  • CNN + ResNet18适合需要精细控制训练过程的场景,比如自定义数据增强、调整分类头结构、观察特征图输出。代码在手,改起来灵活
  • YOLOv5分类模式只适合跑通整个流程,类别数、输入分辨率都固定了,想加个标签平滑得自己改源码,调试成本不小

从个人项目经验来说,如果目标是把这份数据集作为基准测试来做实验对比,优先选择CNN路线;如果目标是快速验证某个检测模型在分类任务上的baseline,YOLOv5分类模式更直接。

4. 微生物图像分类的避坑指南:数据与训练的五条踩坑记录

数据量小的分类项目,跑通容易,但每个环节都有容易踩坑的地方。以下五条是实践中比较典型的翻车案例,每条按照"现象→原因→解决"来梳理。

4.1 反序列化类别字典时出现乱码

现象:用json.load()读取class_dict.json文件时,类别名出现中文乱码,比如"阿米巴"变成\u963f\u7c73\u5df4。

原因:JSON文件本身的编码格式不是UTF-8。Windows环境下用记事本另存为时可能默认保存成ANSI编码,Python的open()函数默认按UTF-8解码,遇到ANSI编码文件就会出错。

解决:读取时显式指定编码格式。

with open('class_dict.json', 'r', encoding='utf-8') as f: class_dict = json.load(f)

如果已经出现乱码,先检查文件编码:file class_dict.json,在Linux下可以直接看出编码格式。Windows下用VSCode打开,右下角会显示当前编码。把文件另存为UTF-8编码再使用。这里也建议把JSON文件的BOM头去掉,有些编辑器保存UTF-8时会自动加BOM,Python解析时会报unexpected BOM错误。

提示:如果这个数据集是某博主配套源码分发的,注意整个工程的文件编码是否统一,源码里读取JSON的地方有没有显式指定编码,不指定的话默认跟随系统,Windows上大概率出问题。

4.2 训练集和测试集类别顺序不一致导致标签错位

现象:训练阶段准确率很高(95%以上),但推理阶段对新图片的预测结果完全不对,比如把眼虫属预测成阿米巴。

原因:训练脚本里手动指定了类别顺序,比如classes = ['阿米巴', '眼虫属', ...],而测试集目录顺序或JSON字典顺序与此不同,导致模型训练时的索引和推理时的索引对应不上。数据加载时如果分别遍历train和test目录里的os.listdir,两次返回的目录顺序可能不同,标签索引就会错位。

解决:整个项目只信任class_dict.json这一个文件作为类别标签的来源,不要用os.listdir的返回值顺序作为标签。训练和推理时都要通过字典反查类别名,确保索引映射完全一致。

4.3 图片格式大小写混乱导致读取失败

现象:数据加载时报错,提示Image.open()无法识别文件格式,但用图片查看器打开发现图片本身没有损坏。

原因:数据集里部分图片的扩展名混用大小写,JPEG和jpg混在一起。部分Windows环境下的NTFS文件系统对大小写不敏感,但Linux和macOS的文件系统严格区分,代码里如果写死了处理jpg后缀,遇到JPEG后缀就会漏掉或者读取失败。

解决:在Dataset的初始化阶段统一处理图片格式,兼容大小写和后缀差异。

valid_extensions = ('.jpg', '.jpeg', '.png', '.JPG', '.JPEG', '.PNG') for img_name in os.listdir(class_path): if img_name.endswith(valid_extensions) and not img_name.startswith('.'): self.samples.append((os.path.join(class_path, img_name), idx))

4.4 过拟合速度过快,验证集指标大幅波动

现象:训练集准确率在第10个epoch就接近100%,但测试集准确率只有70%-80%,且每个epoch之间波动超过5%。

原因:训练集只有630张,模型容量相对数据量过大。ResNet18在ImageNet上超过100万张图的规模才能充分训练,在600多张图上很快就记住了训练集的具体像素分布,但没学会泛化。

解决:在train_transform中加大数据增强力度,比如增加RandomResizedCrop(裁剪缩放比设到0.6-1.0)、加大旋转角度到30度、适当调整颜色抖动。另一个做法是把预训练冻结的层数增加,只微调最后两层。数据增强和迁移学习本质上是把ImageNet上学会的特征保留住,避免在小数据集上做剧烈扰动。

4.5 图像尺寸过小导致Resize后变形

现象:训练时transforms.Resize((224, 224))能跑通,但推理阶段对某些输入图片预测异常,可视化后发现图片被严重拉伸。

原因:微生物图像本身可能接近正方形,但显微成像时图幅比例不一致,统一Resize成224×224会破坏原始比例。微生物的外形特征是分类的重要线索,比如水螅是长条形、草履虫是椭球形,拉伸后会损失形态信息。

解决:用Resize配合CenterCrop或RandomResizedCrop的组合,先等比例缩放短边到224,然后中心裁剪得到224×224:

transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), ])

这样可以保留微生物的相对位置和形态比例,减少几何畸变带来的分类损失。

5. 模型验证与指标解读:混淆矩阵、可视化与推理边界

训练完成拿到权重之后,验证工作不能只看一个准确率数字。150张测试图,走一遍测试集能反映出不少问题。准确率只是一个宏观指标,更细的类别混淆和典型误判方向还得靠混淆矩阵和可视化来定位。

5.1 混淆矩阵分析:定位类别间误判

用scikit-learn生成混淆矩阵,代码不复杂:

import numpy as np import matplotlib.pyplot as plt import seaborn as sns from sklearn.metrics import confusion_matrix y_true = [] y_pred = [] model.eval() with torch.no_grad(): for images, labels in test_loader: outputs = model(images) _, predicted = torch.max(outputs, 1) y_true.extend(labels.numpy()) y_pred.extend(predicted.numpy()) cm = confusion_matrix(y_true, y_pred) class_names = list(class_dict.values()) plt.figure(figsize=(10, 8)) sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=class_names, yticklabels=class_names) plt.xlabel('Predicted') plt.ylabel('True') plt.tight_layout() plt.show()

混淆矩阵读出来的信息量比较大。比如草履虫和眼虫属看起来都有长条形的轮廓,互相误判的概率就会偏高。如果对角线的数字明显高于非对角线,说明模型学到了有区分度的特征;如果某两类混淆严重,一方面可以收集更多这两类的样本,另一方面可以在数据增强里针对这两类的差异特征做强化,比如草履虫有口沟结构、眼虫属有眼点,增加小区域裁剪的增强可以放大这些局部细节。

5.2 错误样本可视化:看模型为什么错

把预测错的样本带着真实标签和预测标签打出来:

import torch import matplotlib.pyplot as plt from torchvision.utils import make_grid model.eval() misclassified_samples = [] with torch.no_grad(): for images, labels in test_loader: outputs = model(images) _, predicted = torch.max(outputs, 1) for i in range(len(labels)): if predicted[i] != labels[i]: misclassified_samples.append((images[i], labels[i], predicted[i])) # 展示前8个错误样本 for i in range(min(8, len(misclassified_samples))): img, true_label, pred_label = misclassified_samples[i] plt.subplot(2, 4, i + 1) img = img.permute(1, 2, 0) # CHW -> HWC plt.imshow(img.numpy()) plt.title(f'True: {class_names[true_label]}\nPred: {class_names[pred_label]}', fontsize=8) plt.axis('off') plt.tight_layout() plt.show()

看错误样本时重点关注图像的成像质量:显微图像的染色差异、气泡遮挡、背景杂质,这些非微生物本身的干扰因素,往往比微生物形态本身的相似性更容易导致误判。实验记录里把每张错分图像的干扰因素标注下来,如果大量错误集中在同类干扰上,下一步就有方向了。

5.3 模型推理部署与推理边界

验证之后是推理落地,把模型封装成推理函数,输入图片路径,输出类别名和置信度:

def predict_image(img_path, model, class_dict, device='cpu'): model.eval() model.to(device) transform = test_transform img = Image.open(img_path).convert('RGB') img_tensor = transform(img).unsqueeze(0).to(device) with torch.no_grad(): output = model(img_tensor) prob = torch.softmax(output, dim=1) confidence, pred_idx = torch.max(prob, dim=1) # 根据类别字典反查类别名 idx_to_class = {int(idx): name for name, idx in class_dict.items()} return idx_to_class[pred_idx.item()], confidence.item() result = predict_image('data/test/paramecium/Image_63.jpeg', model, class_dict) print(f'预测类别: {result[0]}, 置信度: {result[1]:.2f}')

推理边界方面,有几个点需要留意。confidence反映的是模型软max输出分布,与真实正确概率有差距,尤其在训练样本代表性不足的类别上,高置信度也可能出错。图像加载编码上,如果遇到相机直出的RAW格式图片,需要兼容性处理,比如先统一转换;本数据集最好统一转成RGB三通道的JPG再送入网络。

6. 数据增强实验:在小数据集上稳定提升精度的实操配置

对于630张的训练集来说,数据增强策略对最终精度的影响是决定性的。微生物图像有一些独特性质:旋转不变性很强,微生物在视野中朝哪个方向都有;尺度变化明显,不同样本在同倍率下大小不一;拍摄条件差异大,光照、染色、背景干净程度都不稳定。因此数据增强策略不能照搬ImageNet的标准模板,需要针对这些特性做调整。

第一组推荐配置以几何扰动为主,适合作为基线增强:

train_transform = transforms.Compose([ transforms.RandomResizedCrop(size=224, scale=(0.6, 1.0)), transforms.RandomRotation(30), transforms.RandomHorizontalFlip(p=0.5), transforms.RandomVerticalFlip(p=0.3), transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ])

第二组推荐配置加入了更激进的像素级扰动,适合模型在第一组配置下出现了过拟合趋势后使用:

train_transform = transforms.Compose([ transforms.RandomResizedCrop(size=224, scale=(0.5, 1.0)), transforms.RandomRotation(45), transforms.RandomAffine(degrees=0, translate=(0.1, 0.1)), transforms.RandomHorizontalFlip(p=0.5), transforms.RandomVerticalFlip(p=0.5), transforms.ColorJitter(brightness=0.4, contrast=0.4, saturation=0.4), transforms.GaussianBlur(kernel_size=3, sigma=(0.1, 1.0)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ])

第二组的RandomAffine平移操作在微生物图像上比较实用,因为目标在视野中不总是在中心位置;GaussianBlur模拟显微成像时可能出现的失焦模糊,能提升模型对成像质量波动的抗性。但注意GaussianBlur不能加在测试集的transform里,测试集只需要Resize、CenterCrop和Normalize。

我的经验是,先跑第一组配置做基线,观察测试集准确率;如果训练集准确率和测试集准确率的差距大于15%,说明过拟合严重,切到第二组配置重跑。需要说明的是,数据增强策略是否有效,最终要看在固定150张测试集上跑了多组实验后的对比结论,如果两组配置都有各自的明显优势,可以考虑在推理阶段做Test Time Augmentation(TTA),推理时把原始图片、水平翻转、垂直翻转各预测一次,取三个结果的平均概率作为最终输出,通常能再提升一到两个百分点的准确率。

从那次被微生物数据集折腾的经历以后,我每次拿到划分好的小数据集都会强制做一遍"先看目录+字典、跑一次基线、画一张混淆矩阵"这三步,然后再开始调参。这次的数据集胜在结构干净:630张训练图、150张测试图、类别字典完整、目录结构打包好,整个复现过程在普通笔记本上就能跑完。希望这份拆解能帮你在微生物图像分类上少走几步弯路。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询