ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

PyTorch眼睛疾病分类数据集训练:验证集划分与类别不平衡实战指南

PyTorch眼睛疾病分类数据集训练:验证集划分与类别不平衡实战指南 简介眼睛疾病分类数据集是一份可直接用于图像分类任务的中小型医学影像资源包含白内障、青光眼、正常、视网膜疾病四个类别适合临床筛查模型练手、课程实验或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(), keylambda 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_splitstratify按类别标签分层保证每个类在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_okTrue) 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_size0.1, random_state42) dest test_root / cls_dir.name dest.mkdir(parentsTrue, exist_okTrue) 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, size16): img Image.open(path).convert(L).resize((size, size)) pixels np.array(img, dtypenp.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(p0.5), transforms.ColorJitter(brightness0.15, contrast0.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, transformtrain_tf) val_ds datasets.ImageFolder(eye_disease_dataset/val, transformval_tf) train_loader DataLoader(train_ds, batch_size32, shuffleTrue, num_workers4, pin_memoryTrue) val_loader DataLoader(val_ds, batch_size32, shuffleFalse, num_workers4, pin_memoryTrue)验证集不应用RandomResizedCrop和Flip验证要的是确定性的结果训练集用RandomResizedCrop能模拟眼底相机拍摄角度和视场范围的差异。Normalize沿用ImageNet的mean/std对眼底图这种红色调为主的图像其实够用。如果发现图像分布差异很大可以从数据集中采样几千张算出自己的mean/std替换但多数场景没必要。注意Windows上num_workers设为0最稳Linux下再按CPU核数往上加。先跑通流程再优化加载速度。3.2 训练脚本骨架损失函数、优化器与验证时机用预训练ResNet50作为backbone是多数眼睛疾病分类项目入门标配。参数少、权重好找、微调稳定。替换最后一层全连接损失函数先用最朴素的CrossEntropyLoss验证集每个epoch都算一次acc保存val acc最高的checkpointmodel models.resnet50(weightsmodels.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(), lr1e-4, weight_decay1e-4) scheduler torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, modemax, factor0.5, patience3) 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(fepoch {epoch1}: 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无提升早停) breaklr1e-4是微调全模型的安全起点如果只解冻最后一层fc训练可以用1e-3但全模型微调降到1e-4更稳。weight_decay1e-4抑制医学图像上容易出现的过拟合。batch_size32搭配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和学习率的搭配遵循线性缩放原则batch32配lr1e-4batch16则lr减半到5e-5batch64可以尝试2e-4。下表只适用于单卡小batch场景多卡时不这么算。batch size学习率起点典型场景165e-5显存受限的旧卡321e-4最常见配置642e-412G以上显存DataLoader里还有个容易被忽略的参数drop_last。医学图像数据集样本数经常不是batch_size的整数倍最后一个batch可能只有几张图BN层的统计会不稳定。训练时建议设置drop_lastTrue验证时保持drop_lastFalse以便统计所有样本。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_namesclasses, digits3) print(report) cm confusion_matrix(labels, preds) cm_norm cm.astype(float) / cm.sum(axis1, keepdimsTrue) plt.figure(figsize(10, 8)) sns.heatmap(cm_norm, annotTrue, fmt.2f, cmapBlues, xticklabelsclasses, yticklabelsclasses) plt.xlabel(Predicted) plt.ylabel(True) plt.tight_layout() plt.savefig(eye_confusion_matrix.png, dpi150)混淆矩阵呈现的是“真实类别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(weightweights.cuda())直接用1/count会出现极端类权重过大的问题比如某类只有80张权重会变成正常眼的几十倍训练反而震荡。除以均值把正常类压回1附近正常眼权重小于1稀有类权重大于1但不会离谱。如果加了class weight后召回率还是上不去可以换Focal Loss。它自动降低易分样本的loss贡献让模型把注意力放在难分的病变类别上class FocalLoss(nn.Module): def __init__(self, gamma2.0, alphaNone): super().__init__() self.gamma gamma self.alpha alpha def forward(self, logits, targets): ce_loss nn.functional.cross_entropy(logits, targets, weightself.alpha, reductionnone) pt torch.exp(-ce_loss) focal_loss (1 - pt) ** self.gamma * ce_loss return focal_loss.mean()gamma2.0是常用起点。对眼睛疾病分类我建议先从class weight入手因为它只改一个参数跑两个epoch就能看出趋势focal loss要调gammagamma太大模型会过度聚焦难分样本出现验证集acc原地抖动。见过有人把alpha和class weight混着用结果正常眼权重被压到0.1以下模型开始大量误报没必要叠这么多。4.3 学习率策略从视频动作分类实战里常用的余弦退火说起很多做视频动作分类比如跑UCF101这类基准的团队长训练时几乎默认用余弦退火。这不是眼睛疾病分类里的新东西但确实好用。第3.2节用的ReduceLROnPlateau是验证集驱动的适合训练中期如果数据集不大余弦退火的确定性调度往往更稳scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max30, eta_min1e-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之前统一转RGBdef 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_size0.2, random_state42) 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(feye_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, dim1) max_probs, preds torch.max(probs, dim1) filter_mask (preds ! labels) | (max_probs 0.6)把mask筛出来的图像路径和预测结果写成csv对照训练集里的人工复核清单再查一遍。这步常会发现“预测错误”其实是标注错误比如早期白内障被标成正常眼。这类脏样本如果不清理会一直污染指标。从val里剔除后重新评测模型能力才是真实的。我的习惯是每个epoch结束都保存val acc和混淆矩阵训练完成后用验证集里置信度低于0.6的样本生成一份人工复核清单而不是把模型输出当裁决者。眼睛疾病分类的落地价值不在acc多高而在于辅助医生把漏诊率降下来所以让模型学会说“我不确定”比强行输出一个错误类别更安全。置信度阈值在验证集上扫描一遍再定0.6或0.7看不同阈值下被标记为需复核的样本数量、以及复核样本里的误检率选一个业务能接受的操作点。我在设备色彩偏移的新数据集上翻过车教训是验证集只能证明模型在同类数据上有效真要在新采集设备上跑还得单独留一批数据做上线前验证。希望这些从数据集结构到验证集用法的经验能真正帮到你。本文还有配套的精品资源点击获取
返回列表