☰
PyTorch眼睛疾病分类数据集训练:验证集划分与类别不平衡实战指南
2026/10/6 8:33:40 网站建设 项目流程

简介:眼睛疾病分类数据集是一份可直接用于图像分类任务的中小型医学影像资源,包含白内障、青光眼、正常、视网膜疾病四个类别,适合临床筛查模型练手、课程实验或YOLOv5分类项目。数据按train和test目录整理,训练集481张、测试集120张,均为JPEG格式,配合JSON分类字典和Python可视化脚本,可快速完成数据划分查看与模型迭代。压缩包共604个文件,除601张图片外,还有1个字典文件、1个可视化脚本和1张示例图,整包约61MB,轻量易下载。目前已有413人学习下载,脚本支持随机抽取4张图展示并保存结果,无需改动即可运行,能帮助使用者快速核对数据质量和类别分布。

1. 拿到眼睛疾病分类数据集,先别急着训:训练集和验证集到底在做什么

接手一个医学图像的眼睛疾病分类数据集时,真正让人栽跟头的往往不是模型,而是文件夹里那两个split:训练集和验证集。很多人直接把train和val合并重训,或者反复拿验证集调参,最后精度漂亮得可疑,一上真实场景就露馅。要解决的问题很具体:这个分类数据集该按什么结构读取、类别分布怎么看、训练流程怎么写、验证集怎么用才不作弊。适合用PyTorch做医学图像分类的算法工程师、做毕设的学生,以及想从yolo那套自定义数据习惯切到分类任务的人。

2. 拆解眼睛疾病分类数据集:目录结构、标签格式与划分合理性检查

2.1 先看目录结构和标签格式,再决定用什么姿势读取

常见眼睛疾病分类数据集的组织方式通常是:train目录下按类别建子文件夹,val目录同样按类别建子文件夹,图片文件散落在各自的类别文件夹里。类别名即标签,文件夹名就是医生给的诊断结论。公开数据集里的类别体系大体围绕眼底镜图像展开:正常、糖尿病视网膜病变、青光眼、白内障、黄斑变性、高血压视网膜病变、近视等。这类图像通常由眼底相机采集,也有部分是医院病历系统里导出的彩色照片。

拿到数据后第一件事不是写训练脚本,而是确认两类元信息:图片扩展名是否统一,.jpg和.png混用非常常见;类别文件夹里有没有混入非图片文件,比如隐藏的desktop.ini或macOS的.DS_Store。最有效的检查方式是直接对每个split做一次文件统计,把类别和数量一次性打出来。

import os from collections import Counter data_root = 'eye_disease_dataset' for split in ['train', 'val']: split_path = os.path.join(data_root, split) if not os.path.isdir(split_path): print(f'{split} 目录不存在,先检查数据集路径') continue classes = [d for d in os.listdir(split_path) if os.path.isdir(os.path.join(split_path, d))] per_class = {} for cls in sorted(classes): per_class[cls] = len(os.listdir(os.path.join(split_path, cls))) total = sum(per_class.values()) print(f'[{split}] 共 {total} 张,{len(classes)} 个类别') for cls, n in sorted(per_class.items(), key=lambda x: -x[1]): print(f' {cls}: {n} ({n / total * 100:.2f}%)')

这段脚本输出每个类别在训练集和验证集的数量占比。眼睛疾病数据集的通病是类别不平衡:正常眼通常是数量最多的类,而糖尿病视网膜病变的早期样本可能只有正常眼的零头。如果某个类在验证集里只有个位数,对应的acc、precision都不可信,后面必须换成per-class指标。另一个判断依据是目录层级,torchvision的ImageFolder要求类别文件夹直接挂在split下,如果数据集是train/class/subfolder这种二次封装结构,或者用csv索引标签,就要写自定义Dataset,不能硬套现成工具。

2.2 验证集、测试集和训练集:标题只给了两个split时怎么补第三个

很多公开医学图像数据集只划分了train和val,没有test。原因通常是数据量少,官方想给使用者预留调参空间。但作为落地的人,必须自己补出一个test split,否则报出来的所有指标都可能被验证集“污染”。验证集用来做模型选择、超参调优和早停,测试集用来估计最终交给业务方时的真实性能。

常见做法是从训练集里再切一小块出来当测试集。假如训练集有8000张,按分层抽样切出10%约800张作为test,剩下的做train。切分时用sklearn的train_test_split,stratify按类别标签分层,保证每个类在test里的比例和train一致。

import shutil from pathlib import Path from sklearn.model_selection import train_test_split train_root = Path('eye_disease_dataset/train') test_root = Path('eye_disease_dataset/test') test_root.mkdir(exist_ok=True) for cls_dir in train_root.iterdir(): if not cls_dir.is_dir(): continue imgs = list(cls_dir.glob('*')) _, test_imgs = train_test_split(imgs, test_size=0.1, random_state=42) dest = test_root / cls_dir.name dest.mkdir(parents=True, exist_ok=True) for img in test_imgs: shutil.copy(str(img), str(dest / img.name))

用random_state固定随机种子保证切分可复现;用copy而不是move,防止改主意后数据被搬走。有一个细节容易被忽略:如果这张数据来自同一病人的多角度拍摄,这个脚本是不安全的,先按病人分组再切,具体做法在2.3和避坑章展开。测试集切出来之后只碰一次,不要在它上面反复调参,否则测试集就变成了第二个验证集,失去了终极验证的意义。

2.3 划分合理性检查:同一个人可能出现在两个集合里吗

眼睛疾病分类数据集大多来自医院采集,同一个病人可能有两眼甚至多张不同时间拍摄的眼底图。如果切分按文件随机打散,同一病人的多张图很可能同时出现在训练集和验证集。模型会把“这个病人的视盘形态”记下来,而不是学习“这类疾病的通用特征”,验证集acc会虚高。

检查方式:先找数据集自带的元数据csv或DICOM头,看有没有patient_id字段。没有元数据时,部分数据集文件名会带patient前缀。如果两者都没有,可以用感知哈希做近似重复图片检测:

from PIL import Image import numpy as np def phash(path, size=16): img = Image.open(path).convert('L').resize((size, size)) pixels = np.array(img, dtype=np.float32) avg = pixels.mean() return ''.join('1' if p > avg else '0' for p in pixels.flatten()) def hamming(a, b): return sum(c1 != c2 for c1, c2 in zip(a, b))

phash把图像缩小成16x16的灰度指纹,汉明距离小于等于4的两张图基本可以认定是重复或近似重复。但这个方法只能找出“拷贝/裁剪”级别的重复,同一个病人两只眼的外观差异明显,phash查不出来。最稳妥还是靠patient_id切分:先按病人分组,再在所有病人上做train/val/test的分层切分。切完再回头看一眼train和val的类别分布,确认两个集合里的病人集合没有交集。

3. 用 PyTorch 跑通眼睛疾病分类的最小训练流程(ResNet 路线)

3.1 数据读取:ImageFolder 的两个注意点

torchvision的ImageFolder天然适配第2章的目录结构,不用写任何自定义Dataset。训练集和验证集分别挂不同的transform,训练集做随机增强,验证集只做尺寸统一和标准化:

import torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision import datasets, transforms, models train_tf = transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomResizedCrop(224, scale=(0.7, 1.0)), transforms.RandomHorizontalFlip(p=0.5), transforms.ColorJitter(brightness=0.15, contrast=0.15), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) val_tf = 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_ds = datasets.ImageFolder('eye_disease_dataset/train', transform=train_tf) val_ds = datasets.ImageFolder('eye_disease_dataset/val', transform=val_tf) train_loader = DataLoader(train_ds, batch_size=32, shuffle=True, num_workers=4, pin_memory=True) val_loader = DataLoader(val_ds, batch_size=32, shuffle=False, num_workers=4, pin_memory=True)

验证集不应用RandomResizedCrop和Flip,验证要的是确定性的结果;训练集用RandomResizedCrop能模拟眼底相机拍摄角度和视场范围的差异。Normalize沿用ImageNet的mean/std,对眼底图这种红色调为主的图像其实够用。如果发现图像分布差异很大,可以从数据集中采样几千张算出自己的mean/std替换,但多数场景没必要。

注意:Windows上num_workers设为0最稳,Linux下再按CPU核数往上加。先跑通流程,再优化加载速度。

3.2 训练脚本骨架:损失函数、优化器与验证时机

用预训练ResNet50作为backbone是多数眼睛疾病分类项目入门标配。参数少、权重好找、微调稳定。替换最后一层全连接,损失函数先用最朴素的CrossEntropyLoss,验证集每个epoch都算一次acc,保存val acc最高的checkpoint:

model = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V1) num_classes = len(train_ds.classes) model.fc = nn.Linear(model.fc.in_features, num_classes) model = model.cuda() criterion = nn.CrossEntropyLoss() optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-4) scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='max', factor=0.5, patience=3) def validate(model, loader): model.eval() correct = 0 total = 0 all_preds, all_labels = [], [] with torch.no_grad(): for images, labels in loader: images, labels = images.cuda(), labels.cuda() outputs = model(images) _, predicted = torch.max(outputs, 1) correct += (predicted == labels).sum().item() total += labels.size(0) all_preds.extend(predicted.cpu().tolist()) all_labels.extend(labels.cpu().tolist()) return correct / total, all_preds, all_labels best_acc = 0 no_improve = 0 early_stop_patience = 5 for epoch in range(30): model.train() for images, labels in train_loader: images, labels = images.cuda(), labels.cuda() optimizer.zero_grad() outputs = model(images) loss = criterion(outputs, labels) loss.backward() optimizer.step() acc, preds, labels = validate(model, val_loader) scheduler.step(acc) print(f'epoch {epoch+1}: val_acc={acc:.4f} lr={optimizer.param_groups[0]["lr"]:.2e}') if acc > best_acc: best_acc = acc no_improve = 0 torch.save(model.state_dict(), 'best_eye_cls.pth') else: no_improve += 1 if no_improve >= early_stop_patience: print(f'{early_stop_patience} 个epoch无提升,早停') break

lr=1e-4是微调全模型的安全起点;如果只解冻最后一层fc训练可以用1e-3,但全模型微调降到1e-4更稳。weight_decay=1e-4抑制医学图像上容易出现的过拟合。batch_size=32搭配ResNet50在24G以下显存基本舒适,显存受限改16时学习率相应减半。早停条件用“val acc连续epoch无提升”而不是“val loss无下降”,医学图像噪声大,loss和acc并非总是同步。

一个常见的翻车点:val acc已经连续5个epoch没涨,但因没保存最优模型,交付的是最后一个epoch的权重,性能大幅回退。上面把早停和模型保存写在一起,训完直接加载best_eye_cls.pth才是真正能用的模型。

3.3 关键参数设置:图像尺寸、batch size 与学习率的搭配

眼睛疾病分类数据集里的图像尺寸通常不统一。眼底相机常见2048x1536、1600x1200,也有手机翻拍的病历图。Resize到256再中心裁剪到224是ImageNet时代的标准做法。想追求速度可以缩到192或160,但代价是视网膜小血管、微动脉瘤这类细节可能被模糊掉,建议先用224跑通再压缩。

batch和学习率的搭配遵循线性缩放原则:batch=32配lr=1e-4,batch=16则lr减半到5e-5,batch=64可以尝试2e-4。下表只适用于单卡小batch场景,多卡时不这么算。

batch size学习率起点典型场景
165e-5显存受限的旧卡
321e-4最常见配置
642e-412G以上显存

DataLoader里还有个容易被忽略的参数drop_last。医学图像数据集样本数经常不是batch_size的整数倍,最后一个batch可能只有几张图,BN层的统计会不稳定。训练时建议设置drop_last=True,验证时保持drop_last=False以便统计所有样本。

4. 验证集评估与精度调优:眼睛疾病分类的 3 个必调参数

4.1 用混淆矩阵看模型到底错在哪一类

总acc对医学图像分类并不够。类别不平衡严重时,正常眼占大头,acc会被正常类拉高,模型把所有病变都判成正常也能到60%以上。要在验证集上算per-class的recall和混淆矩阵:

from sklearn.metrics import classification_report, confusion_matrix import seaborn as sns import matplotlib.pyplot as plt classes = val_ds.classes # 这两个列表来自validate()函数返回的preds和labels report = classification_report(labels, preds, target_names=classes, digits=3) print(report) cm = confusion_matrix(labels, preds) cm_norm = cm.astype('float') / cm.sum(axis=1, keepdims=True) plt.figure(figsize=(10, 8)) sns.heatmap(cm_norm, annot=True, fmt='.2f', cmap='Blues', xticklabels=classes, yticklabels=classes) plt.xlabel('Predicted') plt.ylabel('True') plt.tight_layout() plt.savefig('eye_confusion_matrix.png', dpi=150)

混淆矩阵呈现的是“真实类别vs预测类别”。在医学图像场景里,关注重点不是对角线多高,而是哪些非对角线值得警惕。糖尿病视网膜病变和黄斑变性早期都表现为黄斑区异常,模型容易把两者搞混;如果模型把青光眼判成正常,这种错误在临床上属于漏诊。看矩阵时先圈出“正常眼那行”的false negative,因为病患漏诊比误诊更危险。

现实里还有一个常见现象:模型对验证集里“背景亮度过高”的样本常常成片判错。这类样本往往在混淆矩阵某一列扎堆,先别急着加数据,回去看那一类图像是不是存在设备差异。

4.2 类别不平衡:损失函数替换与样本权重

用CrossEntropyLoss时,class weight是最直接的平衡手段。先统计训练集的每类样本数,再算权重,注意归一化:

import os from collections import Counter import torch split_path = 'eye_disease_dataset/train' class_counts = Counter() for cls in sorted(os.listdir(split_path)): cls_path = os.path.join(split_path, cls) if os.path.isdir(cls_path): class_counts[cls] = len(os.listdir(cls_path)) counts_tensor = torch.tensor([class_counts[c] for c in sorted(class_counts)]) weights = 1.0 / counts_tensor.float() weights = weights / weights.mean() # 归一化,让权重均值保持在1附近 criterion = nn.CrossEntropyLoss(weight=weights.cuda())

直接用1/count会出现极端类权重过大的问题,比如某类只有80张,权重会变成正常眼的几十倍,训练反而震荡。除以均值把正常类压回1附近,正常眼权重小于1,稀有类权重大于1,但不会离谱。

如果加了class weight后召回率还是上不去,可以换Focal Loss。它自动降低易分样本的loss贡献,让模型把注意力放在难分的病变类别上:

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_loss = nn.functional.cross_entropy(logits, targets, weight=self.alpha, reduction='none') pt = torch.exp(-ce_loss) focal_loss = (1 - pt) ** self.gamma * ce_loss return focal_loss.mean()

gamma=2.0是常用起点。对眼睛疾病分类,我建议先从class weight入手,因为它只改一个参数,跑两个epoch就能看出趋势;focal loss要调gamma,gamma太大模型会过度聚焦难分样本,出现验证集acc原地抖动。见过有人把alpha和class weight混着用,结果正常眼权重被压到0.1以下,模型开始大量误报,没必要叠这么多。

4.3 学习率策略:从视频动作分类实战里常用的余弦退火说起

很多做视频动作分类(比如跑UCF101这类基准)的团队,长训练时几乎默认用余弦退火。这不是眼睛疾病分类里的新东西,但确实好用。第3.2节用的ReduceLROnPlateau是验证集驱动的,适合训练中期;如果数据集不大,余弦退火的确定性调度往往更稳:

scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=30, eta_min=1e-6) for epoch in range(30): # 现有训练循环 scheduler.step()

T_max设为总epoch数,eta_min设为初始学习率的百分之一,从1e-4退到1e-6足够。用它替代ReduceLROnPlateau时要注意:余弦退火是“从当前值一路往下”,没有回头涨的机会,所以初始lr宁可偏低。有人把初始lr设成1e-3跑余弦退火,前几个epoch loss直接炸穿。

在眼睛疾病分类这类中小规模医学图像数据集上,我的使用顺序是:先用ReduceLROnPlateau跑20个epoch看baseline,如果尾部loss震荡明显,再换余弦退火重跑一次,验证集acc通常能再上1到2个点。先有baseline再调调度器,比一上来就堆各种trick更省时间。

5. 眼睛疾病分类数据集落地避坑:5 条真实的血泪经验

下面这几条都是从实际训练过程中踩出来的,按“现象、原因、解决”的顺序写,遇到类似问题可以直接对号入座。

5.1 验证集acc漂亮得可疑:训练集里混进了验证集图像

现象:训练到一半验证集acc飙到99%,但换一批新采集的图,acc直接掉到60%。

原因:数据清洗不彻底。很多公开数据集的train和val是从原始资料里分出来的,原始文件里有重复截图、图像拷贝,同一个病例的不同版本文档被误放进了两个集合。

解决:用2.3节的phash全库跑一遍去重,汉明距离小于等于4的图像对确认后只保留一份。更稳的是在训练前用文件名或元数据查一下“同一病人文件是否被分到两个split”。

5.2 灰度图与RGB通道不一致

现象:训练到中途dataloader报错“Expected a 3-channel input”,或者loss变成nan。

原因:眼底相机输出一般是彩色JPEG,但医院导出的历史数据里存在灰度PNG和带透明通道的图。ImageFolder遇到灰度图时ToTensor会把它变成单通道,和Normalize的三通道统计不匹配。

解决:在transform之前统一转RGB:

def load_as_rgb(path): img = Image.open(path) if img.mode != 'RGB': img = img.convert('RGB') return img

把load_as_rgb放进自定义Dataset里。灰度图转RGB是复制通道,RGBA图则丢弃Alpha。不要指望现成的ImageFolder帮你处理这些。

5.3 验证集acc虚高的背后:没有按病人切分

现象:模型在验证集上对青光眼类acc达到98%,业务方拿新数据实测,准确率远低于预期。

原因:医院采集中同一个病人可能提供两只眼的图片,随机划分后同一病人的两只眼一个在train一个在val,模型学到的其实是病人特征。病人ID可能隐藏在文件名里,比如patient_001_left.png和patient_001_right.png,直接listdir根本看不出来。

解决:先抽出patient_id,按病人划分:

import pandas as pd from sklearn.model_selection import train_test_split df = pd.read_csv('image_patient_labels.csv') patients = df['patient_id'].unique() train_patients, val_patients = train_test_split(patients, test_size=0.2, random_state=42) train_df = df[df['patient_id'].isin(train_patients)] val_df = df[df['patient_id'].isin(val_patients)]

注意分层:如果病人总数少,还要按主诊断做stratify,否则可能出现某类病人只进val的情况。

5.4 从 yolo 自定义数据集的习惯迁移过来的误区

现象:做检测的人第一次拿到分类数据集时,习惯性去找label文件、找标注框txt,发现没有这些文件,不知道该怎么训练。

原因:分类数据集和yolov8、yolo26乃至deim这类目标检测框架的数据约定不同。检测用边界框加txt标签,分类数据集的标签全部隐含在文件夹名里,不需要再生成任何label文件。

解决:直接用ImageFolder按目录读。实在习惯用CSV,就自己写一个映射文件:

import csv from pathlib import Path rows = [] for split in ['train', 'val']: root = Path(f'eye_disease_dataset/{split}') for cls_dir in root.iterdir(): if not cls_dir.is_dir(): continue for img in cls_dir.glob('*'): rows.append([str(img), cls_dir.name]) with open('eye_cls_labels.csv', 'w', newline='') as f: writer = csv.writer(f) writer.writerow(['path', 'label']) writer.writerows(rows)

CSV的作用是方便后续做病人级切分和脏样本过滤,而不是替代目录结构。

5.5 早停判断标准别只盯val loss

现象:训练时val loss一直在降但验证集acc纹丝不动,多跑几个epoch后acc突然跳几个点;另一边val loss止跌,你以为可以停了,结果再跑两个epoch又涨一点。

原因:医学图像类别特征差异大,loss和acc并非同步变化。小类别在loss里的贡献占比低,loss下降反映的只是大类特征收敛,小类的acc没有变化。

解决:以“验证集acc连续patience个epoch无提升”作为早停条件,patience设5到8。如果用了class weight或focal loss,同时盯per-class recall的调和平均,不要只盯总acc,因为总acc会被正常眼主导。每个epoch都备份一次最优checkpoint,就算误停也有后悔药。

6. 把验证集利用到极致:错误分析是医学图像分类的最后一公里

训练结束不等于交付。从验证集里筛出预测错误的样本,一张一张看,才是真正提升模型价值的部分。做法是保存验证集的softmax输出,抽最底部的错误样本:

probs = torch.softmax(outputs, dim=1) max_probs, preds = torch.max(probs, dim=1) filter_mask = (preds != labels) | (max_probs < 0.6)

把mask筛出来的图像路径和预测结果写成csv,对照训练集里的人工复核清单再查一遍。这步常会发现“预测错误”其实是标注错误,比如早期白内障被标成正常眼。这类脏样本如果不清理,会一直污染指标。从val里剔除后重新评测,模型能力才是真实的。

我的习惯是每个epoch结束都保存val acc和混淆矩阵;训练完成后用验证集里置信度低于0.6的样本生成一份人工复核清单,而不是把模型输出当裁决者。眼睛疾病分类的落地价值不在acc多高,而在于辅助医生把漏诊率降下来,所以让模型学会说“我不确定”比强行输出一个错误类别更安全。置信度阈值在验证集上扫描一遍再定:0.6或0.7,看不同阈值下被标记为需复核的样本数量、以及复核样本里的误检率,选一个业务能接受的操作点。我在设备色彩偏移的新数据集上翻过车,教训是验证集只能证明模型在同类数据上有效,真要在新采集设备上跑,还得单独留一批数据做上线前验证。希望这些从数据集结构到验证集用法的经验,能真正帮到你。

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

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

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

立即咨询