news 2026/9/15 6:13:11

PyTorch+ResNet50实现眼部疾病分类:数据管道与训练调优实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PyTorch+ResNet50实现眼部疾病分类:数据管道与训练调优实战

简介:这是一份基于PyTorch与ResNet50的眼部疾病图片分类完整工程,主要面向正在准备课程设计、期末大作业的计算机相关专业学生,以及希望掌握深度学习图像分类实战流程的学习者。资源压缩包内共10个文件,含7个Python脚本,分别承担数据集划分与预处理、自定义Dataset、ResNet50网络构建、训练调度与多指标评估等任务,另附2个Markdown文档和gitignore配置;整体仅22KB,结构紧凑、便于直接阅读。该资源已有291人学习下载,源码经过严格调试与导师审核,下载解压后即可运行。配合文档说明,读者可以完整走通从图像数据准备、模型训练到准确率/召回率等指标评估的全流程,并能基于这份代码快速迁移至其他医学图像或通用图像分类课题,减少重复搭建工作,集中精力理解算法设计与调优。

1. 这个“眼部疾病分类”作业,难点根本不在模型上

用 PyTorch 和 ResNet50 做眼部疾病图片分类,是计算机视觉课程里出现频率极高的“高分大作业”题目。很多人拿到这类项目的第一反应是去找源码、跑通训练,但实际上真正决定成绩和运行效果的部分,往往不是那几行model = resnet50(),而是数据怎么组织、类别不均衡怎么处理、训练参数怎么设置,以及最后的推理脚本能不能拿一张新图给出可靠结论。

这类任务的典型场景是:数据集包含正常、白内障、糖尿病视网膜病变、青光眼等多个类别的眼底或裂隙灯图像,每个类别的样本量差异可能很大,图像尺寸、亮度、拍摄设备也不统一。和花卉分类这类相对“干净”的任务不同,医学图像分类的难点在于类间相似度高、噪声大、样本数量少,ResNet50 的预训练权重能缓解一部分问题,但完全依赖默认参数跑出来的结果往往只有 70% 出头的准确率,而经过合理调整后可以稳定提升到 90% 以上。

本文会从数据管道、模型改造、训练脚本、调优技巧到最终的验证推理,完整走一遍这个项目最可靠的实现路径。适合正在做课程大作业、准备毕设或者想用现成方案快速跑通眼病分类的人阅读,每一步都会给出可以直接抄的代码和参数解释。

2. 先把数据管道做对:目录编排、transform 与 dataset 类

2.1 数据目录应该怎么组织最省事

做图像分类项目,第一步不是写模型,而是把数据目录整理成 PyTorch 的ImageFolder可以直接读取的结构。这是最不容易出错的做法,也方便后续在DataLoader里直接按类别索引。

常见的目录结构是这样组织的:

data/ ├── train/ │ ├── normal/ │ ├── cataract/ │ ├── diabetic_retinopathy/ │ └── glaucoma/ ├── val/ │ ├── normal/ │ ├── cataract/ │ ├── diabetic_retinopathy/ │ └── glaucoma/ └── test/ ├── normal/ ├── cataract/ ├── diabetic_retinopathy/ └── glaucoma/

如果原始数据集不是这种结构,比如只有一个文件夹加一张 CSV 标注表,那就自己写一个整理脚本,把每张图片复制到对应类别的子目录里。ImageFolder会按照目录名的字典序生成类别索引,所以normal会排在cataract前面还是后面,取决于字母顺序而不是目录书写顺序,这一点在后续看混淆矩阵时要格外注意。

2.1.1 按病人划分数据集,防止“数据泄漏”

医学图像数据集里经常出现同一个病人的多张图片被同时分进训练集和验证集的情况,这会导致验证指标虚高。正确的做法是按病人 ID 划分而不是按图片划分。具体实现方式是:先读取标注文件,按病人 ID 做分组,再使用train_test_splitgroup参数或者直接手动切分,把同一病人的所有图片放进同一个集合。

2.2 眼科图像的 transform 配置与参数说明

医学图像的预处理比自然图像更需要克制。过度的数据增强反而会让模型学不到有判别力的纹理细节,尤其是在血管、渗出物、视盘边缘这些关键结构上。

from torchvision import transforms train_transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomRotation(degrees=15), transforms.RandomAffine(degrees=0, translate=(0.05, 0.05)), transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.1, hue=0.05), transforms.RandomHorizontalFlip(p=0.5), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) val_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]) ])

注意这里没有使用RandomResizedCrop,因为眼科图像中病变区域可能分布在图像的各个位置,而随机裁剪会丢失边缘信息。使用Resize把图像统一缩放到 224×224,保证全部输入到 ResNet50 之前保持相同尺寸。训练集做了轻微旋转和仿射变换,模拟拍摄角度差异;ColorJitter控制亮度对比度的微小波动,但范围都不大,避免让模型学到错误的颜色关联。

参数取值为什么这么设
Resize(224, 224)匹配 ResNet50 标准输入尺寸
RandomRotation15°眼底相机拍摄角度偏差通常很小
RandomAffinetranslate=0.05模拟图像轻微位移,但不过度
ColorJitter亮度0.2 对比度0.2适配不同设备的光照差异
NormalizeImageNet 均值方差使用迁移学习时保持统计分布一致

2.3 自定义 Dataset 类的完整实现

虽然ImageFolder可以直接用,但实际做眼部疾病分类时往往需要加载额外的标注信息,比如疾病分级标签或者病人 ID,这种情况下自定义 Dataset 类更灵活。

import torch from torch.utils.data import Dataset from PIL import Image import os class EyeDiseaseDataset(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 = {cls: idx for idx, cls in enumerate(self.classes)} self.samples = [] for cls in self.classes: cls_dir = os.path.join(root_dir, cls) for fname in os.listdir(cls_dir): if fname.lower().endswith(('.jpg', '.jpeg', '.png')): self.samples.append((os.path.join(cls_dir, fname), 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

这段代码的逻辑是:初始化时遍历目录结构,把所有图片路径和标签构建成samples列表;每次取数据时打开图片、应用 transform、返回张量和标签。convert('RGB')这一步很有必要,因为部分医学图像是灰度图或带透明通道的 PNG,直接不处理会导致输入通道数不匹配,训练时报错或者在推理时出现维度异常。目录名通过sorted排序保证类别顺序稳定,否则每次运行脚本标签索引都可能变化。

3. ResNet50 与训练骨架:把模型参数和训练脚本一次写清楚

3.1 加载 ResNet50 预训练权重并替换分类头

PyTorch 的torchvision.models库提供了 ResNet50 的完整实现,使用weights=ResNet50_Weights.DEFAULT可以自动下载在 ImageNet 上预训练好的权重。这里有一个常见的错误是在跑大作业时习惯随便找一段别人训练脚本里resnet50(pretrained=True)的代码,但不知道这个权重到底解决了什么问题。

from torchvision import models import torch.nn as nn model = models.resnet50(weights=models.ResNet50_Weights.DEFAULT) num_features = model.fc.in_features num_classes = 4 model.fc = nn.Sequential( nn.Dropout(0.3), nn.Linear(num_features, num_classes) )

重要的是model.fc.in_features这一行,先把原始全连接层的输入维度取出来,再替换成新的分类头。ResNet50 的fc层原始输出是 1000 类,对应 ImageNet 的类别数,必须改成自己的类别数。这里加了一个Dropout(0.3),可以减少分类层过拟合的概率,代价是训练收敛稍慢一点,但对最终准确率有正向帮助。

3.1.1 ResNet50 网络结构示意图的解读方式

很多人在看 ResNet50 的结构图时只关注它“有 50 层”,但没有看到关键的分层结构。ResNet50 由五个阶段组成:输入经过 7×7 卷积和 3×3 最大池化后进入四个残差层,每个残差层包含不同数量的 Bottleneck 块,分别是 [3, 4, 6, 3]。最后一层全局平均池化后接全连接层。了解这个结构的意义在于:当你想做“冻结浅层只训练深层”的策略时,就知道该从哪个 stage 开始解冻。

3.2 训练脚本核心代码详解

有了模型和数据集,接下来就是把训练循环写完整。下面这段是一个可直接运行的训练代码核心部分,包含了优化器、损失函数、学习率调度和每个 epoch 的验证逻辑。

import torch import torch.nn as nn import torch.optim as optim from torch.optim import lr_scheduler device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model = model.to(device) criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=1e-4) scheduler = lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.5) best_acc = 0.0 num_epochs = 30 for epoch in range(num_epochs): model.train() running_loss = 0.0 correct = 0 total = 0 for inputs, labels in train_loader: inputs, labels = inputs.to(device), labels.to(device) optimizer.zero_grad() outputs = model(inputs) loss = criterion(outputs, labels) loss.backward() optimizer.step() running_loss += loss.item() * inputs.size(0) _, predicted = torch.max(outputs, 1) total += labels.size(0) correct += (predicted == labels).sum().item() train_loss = running_loss / total train_acc = correct / total scheduler.step() model.eval() val_correct = 0 val_total = 0 with torch.no_grad(): for inputs, labels in val_loader: inputs, labels = inputs.to(device), labels.to(device) outputs = model(inputs) _, predicted = torch.max(outputs, 1) val_total += labels.size(0) val_correct += (predicted == labels).sum().item() val_acc = val_correct / val_total if val_acc > best_acc: best_acc = val_acc torch.save(model.state_dict(), 'best_model.pth') print(f'Epoch {epoch+1}/{num_epochs}, ' f'Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.4f}, ' f'Val Acc: {val_acc:.4f}')

学习率设置为 1e-4 而不是常用的 1e-3,因为迁移学习时模型的大部分参数已经接近最优区域,学习率过大会把预训练权重破坏掉。优化器选择Adam,因为它对学习率的敏感度较低,在大作业场景下不需要精细调参也能得到可接受的结果。调度器每 5 个 epoch 把学习率减半,让训练后期用更小的步长做精细调整。验证集在每轮训练结束后评估,并且只在验证准确率创新高时保存模型权重,确保最终拿到的是验证集表现最好的模型。

3.3 关键参数组合参考

参数推荐值备注
输入尺寸224×224与预训练权重对齐
batch_size16 或 32显存不够时优先降到 16
初始学习率1e-4迁移学习的标准起点
优化器Adam不用手动调动量
损失函数CrossEntropyLoss自带 softmax
Epoch 数30~50看验证集是否提前饱和
学习率调度StepLR step=5 gamma=0.5每 5 轮减半

batch_size的选择需要根据 GPU 显存调整。ResNet50 在 224×224 输入下单张图片的前向计算大约占用 0.4GB 显存(batch=32 时约 12GB),如果显存不足就调低 batch size,同时可以把num_workers调大到 4 或 8 来加速数据加载,但 Windows 系统下num_workers大于 0 时需要把训练代码放到if __name__ == '__main__':中保护,否则会报多进程错误。

4. 类别不平衡、过拟合与冻结策略:三个高频问题逐个拆

4.1 类别不平衡问题的两种处理路径

眼部疾病数据集里,正常样本往往远多于病变样本,糖尿病视网膜病变的样本可能只有正常样本的三分之一。这种不均衡会导致模型偏向多数类,整体准确率看着不低,但少数类的召回率很低。验证这一点的最快方法是打印每个类别的分类报告。

处理这个问题有两种常见做法。第一种是加权损失函数,在CrossEntropyLoss中传入weight参数,让少数类样本的梯度贡献更大:

from sklearn.utils.class_weight import compute_class_weight import numpy as np all_labels = [] for _, labels in train_loader: all_labels.extend(labels.numpy()) class_weights = compute_class_weight( class_weight='balanced', classes=np.unique(all_labels), y=np.array(all_labels) ) class_weights = torch.tensor(class_weights, dtype=torch.float32).to(device) criterion = nn.CrossEntropyLoss(weight=class_weights)

第二种做法是使用WeightedRandomSampler对数据加载过程进行重采样,让每个 batch 中各个类别的样本比例更接近均匀。

from torch.utils.data import WeightedRandomSampler sample_counts = np.bincount(np.array(all_labels)) weights = 1.0 / sample_counts[all_labels] sampler = WeightedRandomSampler(weights, num_samples=len(all_labels), replacement=True) train_loader = torch.utils.data.DataLoader( train_dataset, batch_size=32, sampler=sampler )

两种方法可以单独使用,也可以叠加。实际项目里我一般先用加权损失函数,因为它改动最少,如果少数类的准确率仍然偏低,再叠加WeightedRandomSampler。需要注意的是,用采样器之后每个 epoch 看到的数据顺序变了,但由于replacement=True,同一个 batch 内可能出现重复样本,这会略微增加过拟合风险。

4.2 过拟合的判断指标与数据增强止损线

过拟合在医学图像分类中几乎必然出现,尤其是样本总数不多于几千张的情况。最明显的信号是:训练准确率持续上升达到 98% 以上,验证准确率却停滞不再增长,两者差距越来越大。此时优先调整数据增强的强度,而不是换模型。

增强操作初始值过拟合后调整建议
RandomRotation15°增大到 30°
RandomAffinetranslate=0.05增大到 0.1
ColorJitter亮度0.2提高到 0.3
RandomHorizontalFlip0.5保持不变
新增 RandomErasing不使用增加 p=0.3

还有一个值得注意的细节是,眼病图像中左右眼可能呈现对称性,水平翻转是安全的增强方式;但垂直翻转在部分眼底图像中会导致视盘位置异常,不建议开启。RandomErasing可以随机遮挡图像中的一块矩形区域,强迫模型学习非局部特征,在训练后期引入往往比一开始就使用效果更好。

4.3 冻结 ResNet50 浅层特征,只训练深层

如果你的 GPU 显存有限,或者数据量实在太小,冻结部分 ResNet50 参数是减少训练开销和过拟合的有效手段。常见做法是冻结前三个残差层,只训练第四个残差层和全连接层:

for name, param in model.named_parameters(): if 'layer4' in name or 'fc' in name: param.requires_grad = True else: param.requires_grad = False

只传入需要梯度的参数给优化器:

optimizer = optim.Adam(filter(lambda p: p.requires_grad, model.parameters()), lr=1e-4)

冻结方案适用于数据量很少的情况,例如每个类别只有 200~500 张图。如果数据量在每类上千张以上,全量微调的效果通常更好。从 stage 结构来看,layer1学到的是边缘和纹理等基础特征,这些特征在不同领域的图像中差异不大;layer4学到的是高层语义特征,与具体任务强相关。所以冻结浅层、训练深层是合理的折中。

5. 推理脚本里补上置信度与 Top-3,验收才算完成

训练完模型之后,最后的成果不是损失曲线图,而是一个能对单张图片给出预测结果和置信度的推理脚本。这对大作业答辩尤其重要,因为老师通常会在现场随机拿一张图让你测试。写一个可靠的predict.py

import torch from torchvision import transforms from PIL import Image import torch.nn.functional as F def predict_image(image_path, model, class_names, device): model.eval() 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]) ]) image = Image.open(image_path).convert('RGB') input_tensor = transform(image).unsqueeze(0).to(device) with torch.no_grad(): outputs = model(input_tensor) probabilities = F.softmax(outputs, dim=1) top3_prob, top3_idx = torch.topk(probabilities, 3) print(f'输入图片: {image_path}') for i in range(3): idx = top3_idx[0][i].item() print(f'Top-{i+1}: {class_names[idx]} - {top3_prob[0][i].item():.4f}') return class_names[top3_idx[0][0].item()]

class_names必须是训练时同一个class_to_idx列表的逆映射,否则类别名和索引对不上,会得到完全错误的语义标签。一个实用的技巧是:训练完保存模型时,顺手把类别列表也存成 JSON 文件:

import json class_names = train_dataset.classes with open('class_names.json', 'w') as f: json.dump(class_names, f)

推理时加载类别列表,再搭配torch.load('best_model.pth', map_location=device)加载权重,就形成了一个完整闭环。

验证模型可靠性的方法值得多说一步:不要只记录整体准确率,要看每个类别单独的召回率和混淆矩阵。使用 scikit-learn 在验证集上计算这几项指标:

from sklearn.metrics import classification_report, confusion_matrix model.eval() all_preds = [] all_labels = [] with torch.no_grad(): for inputs, labels in val_loader: inputs, labels = inputs.to(device), labels.to(device) outputs = model(inputs) _, preds = torch.max(outputs, 1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) print(classification_report(all_labels, all_preds, target_names=class_names)) print(confusion_matrix(all_labels, all_preds))

classification_report会输出每个类别的精确率、召回率和 F1 分数。如果某个疾病类别的召回率明显低于其他类别,说明该疾病的大量样本被误分类为正常或其他疾病,此时优先检查该类别的样本数量和数据增强配置,而不是盲目增加训练轮数。混淆矩阵可以帮助定位具体的误分类方向——比如白内障和青光眼之间出现混淆,可以考虑增加这些类别之间的差异特征样本,或者为这两个类别单独增加损失权重。

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/15 6:12:49

羽毛球目标检测数据集:2879张图三分类标注与YOLOv8训练实践

做体育视频分析这些年,羽毛球可能是最让我头疼的检测对象之一。场上的运动员还算好认,真正麻烦的是那个时速轻松突破三百公里的羽毛球——在画面里往往只有十几个像素,一眨眼就飞出视野。找来找去,公开的目标检测数据集大多集中在…

作者头像 李华
网站建设 2026/9/15 6:12:06

电子元器件缺陷检测实战:从YOLOv8到YOLO26的选型与融合大模型

我接手这个项目的时候,心里其实没底:一条电子元器件检测线,几千种物料,引脚弯曲、表面划痕、缺件漏焊,靠人眼盯久了必然疲劳,靠传统机器视觉的规则得写到手软,稍微换个光源就崩。最后我选择了YO…

作者头像 李华
网站建设 2026/9/15 6:12:02

Windows上nnU-Net实战:CUDA匹配、环境配置与显存调优

简介:nnUnet 是面向医学图像分割任务的主流深度学习框架,但在 Windows 下部署需自行处理大量编译与依赖问题。压缩包面向 Windows 用户提供了可直接运行的 nnUnet 编译版本,已提前完成路径修正、依赖适配等兼容性调整,适合希望绕过…

作者头像 李华
网站建设 2026/9/15 6:11:55

扶梯逆行检测实战:YOLOv8轻量定制与方向感知优化

简介:本资源是一套基于YOLOv8实现的商场扶梯逆行行为智能预警系统,面向计算机、人工智能、自动化等专业本科生及课程设计/毕业设计需求者,解决公共场所安全监管中实时异常行为识别与主动预警的实际问题。资源共8个文件,含3个核心P…

作者头像 李华
网站建设 2026/9/15 6:11:51

AI大模型技术解析与投资机会全指南

1. 项目概述作为一名在AI领域摸爬滚打多年的技术老兵,我经常被刚入行的程序员朋友问到同一个问题:"现在学AI大模型还有机会吗?"这个问题背后,其实隐藏着对技术趋势的迷茫和对职业发展的焦虑。今天,我就从一个…

作者头像 李华
网站建设 2026/9/15 6:11:48

OpenCV实战指南:从图像处理到DNN推理与实时视频流技术要点

简介:一套完整的OpenCV计算机视觉库源码与配套示例资料包,面向从事图像处理、目标检测、人脸识别等方向的开发者与研究人员。压缩包共包含7052个文件,大小约91.37MB;覆盖C、Python、Java等多种编程语言,包括cpp、hpp、…

作者头像 李华