简介:这是一份面向医学图像处理与深度学习初学者的眼部疾病OCT图像分类项目源码。项目基于PyTorch 1.6实现,内置ResNet18/34/50与VGG16/19五种经典网络,在测试集上准确率可达90%以上;同时附有3D-ResNet实验记录,帮助读者理解参数量与数据规模对效果的影响。压缩包仅20KB,包含9个文件,以7个Python脚本和2个Markdown文档为主:Python脚本覆盖数据预处理、训练/验证集划分、模型构建、指标计算与训练分类等完整流程;Markdown文档提供运行说明与环境依赖清单,便于快速复现。项目还给出了matplotlib、seaborn、torchvision、sklearn等常用依赖的安装指引。目前已有411人学习参考,适合希望快速上手图像分类任务、需要可运行基线代码的开发者用于课程设计、毕业设计或论文复现。
1. 眼科图片分类选 pytorch+ResNet50,不是因为"它最先进"
把眼底彩照直接丢给 VGG 或自己搭的 CNN 训练,多半会撞上两个现象:一是几千张样本训几十轮还在震荡,二是验证集准确率虚高,换一台设备拍图就崩。这个基于 pytorch 的 ResNet50 眼部疾病图片分类方案,解决的核心问题就是"少样本 + 强相似性":不同眼疾的病灶区域往往只占整张图的几个像素块,类别差异小,背景噪声大。ResNet50 靠残差结构把梯度传得够深,又带着 ImageNet 上预训练好的纹理和边缘提取能力,微调成本比从零训练低一个量级。对 IT 从业者来说,这套东西的价值不只在一个医学模型,而是你接其他细粒度图像分类任务时可以直接平移的工程模板:数据组织、迁移学习、类别不平衡处理、训练与推理的 PyTorch 标准写法,全在项目源码里串了一遍。
2. 搭建 pytorch 基础环境并整理眼部图片数据集
2.1 用 Anaconda 建独立环境:pytorch 安装的 GPU 版与 CPU 版取舍
先说环境隔离。眼部疾病分类的训练集通常不大,但如果机器上有 NVIDIA 显卡,pytorch 训练 ResNet50 的收益非常明显:一张 12GB 显存的卡可以轻松跑 batch size 32 的 224×224 图,而纯 CPU 上同样的迭代次数大概要慢 10 倍。常见做法是先建一个干净的 conda 环境,把 pytorch 基础框架装进去,避免和系统里其他项目互相污染依赖。
# 创建 Python 3.9 环境并激活 conda create -n eye_cls python=3.9 -y conda activate eye_cls # 安装 pytorch + torchvision(CUDA 11.8 版本,按自己驱动版本选择) pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 # 验证 GPU 是否可用 python -c "import torch; print(torch.cuda.is_available(), torch.cuda.get_device_name(0))"torch.cuda.is_available()返回True才能走 GPU 训练,返回False时后面所有.cuda()调用都会闪退。这里有个容易翻车的细节:如果你之前用conda install pytorch装过,可能拿到的是 CPU 版;而pip install torch默认在 Linux 下会装带 CUDA 的版本,在 Windows 下默认不带。最稳妥的判断方式是看torch.version.cuda值。
2.2 数据集目录约定:按类名分文件夹
项目源码里通常使用 TorchVision 的ImageFolder接口,它对目录结构有硬性要求。假设我们有五类眼部疾病:normal(正常)、cataract(白内障)、glaucoma(青光眼)、diabetic_retinopathy(糖尿病视网膜病变)、myopia(近视病变),目录应该这样组织:
data/ train/ normal/ # 001.jpg ... cataract/ glaucoma/ diabetic_retinopathy/ myopia/ val/ normal/ cataract/ ...为什么用ImageFolder而不是自己写读取逻辑?因为它会在你调用DataLoader时自动生成 class 到 index 的映射,保证训练和验证阶段类别顺序一致。我自己写项目时会多做一个动作:建一个class_to_idx.json存下这个映射,避免后期推理阶段手工去猜类别编号。
2.3 transform 与归一化:对眼底图的均值和标准差不能随手写
图像预处理的写法直接决定 ResNet50 能不能"接得住"预训练权重。torchvision官方预训练模型的输入约束是:图像 resize 到 224×224,像素值除以 255 后用 ImageNet 数据集的mean=[0.485,0.456,0.406]和std=[0.229,0.224,0.225]做标准化。很多人会想"眼底图是灰绿调的,我是不是该自己算均值",在迁移学习场景下这是错误直觉——保留原预训练的统计量才能让网络前几层的卷积核继续正常工作。
from torchvision import transforms train_transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomHorizontalFlip(p=0.5), transforms.RandomRotation(15), transforms.ColorJitter(brightness=0.2, contrast=0.2), 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]) ])注意验证集这里没有做任何随机增强,目的是让每个 epoch 的评估结果可复现。RandomHorizontalFlip和RandomRotation属于空间增强,对眼底图特别合适——眼球的朝向不改变病变特征;ColorJitter模拟不同拍摄设备的光照差,但幅度别太大,否则会把血管纹理也破坏掉。
3. ResNet50 网络结构的迁移学习改造:加载预训练权重
3.1 从 torchvision 加载 ResNet50 的两代接口差异
ResNet50 的网络结构可以拆成四个 Stage,每个 Stage 由若干 Bottleneck 残差块组成。在 pytorch 里加载它不需要自己复刻这个结构,torchvision.models已经封装好了。需要注意接口已经发生过破坏性变更:旧代码里的pretrained=True参数在较新的 torchvision(0.13 之后)会直接报错,现在统一用weights参数。
import torch import torch.nn as nn from torchvision import models, transforms # 新版写法:显式指定预训练权重 weights = models.ResNet50_Weights.IMAGENET1K_V1 model = models.resnet50(weights=weights) # 查看最后两层结构,确认要替换的位置 print(model.fc) print(model.avgpool)IMAGENET1K_V1是官方在 ImageNet-1K 上训练得到的权重,avgpool会把 7×7 的特征图压缩成 2048 维向量,最后的fc层输出 1000 类。我们的眼部疾病分类任务只需要 5 类,所以要做的就是把fc层换掉,特征提取部分完全复用。
3.2 替换全连接层:五分类输出的正确接法
# 替换最后的全连接层为 5 分类 num_classes = 5 model.fc = nn.Sequential( nn.Dropout(p=0.5), nn.Linear(2048, 512), nn.ReLU(inplace=True), nn.Dropout(p=0.3), nn.Linear(512, num_classes) ) # 将模型送入 GPU device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = model.to(device)中间加一层 512 维的隐层是工程上比较稳的做法:2048 维直接压到 5 类容易让特征表达过于激进,Dropout 则降低全连接层过拟合的风险。inplace=True节省显存,但如果你之后要用钩子(hook)检查激活值,建议改成inplace=False,否则拿到的前向结果有可能是脏数据。
3.3 冻结 BatchNorm 与部分 Stage:显存不够或样本少时的通用技巧
眼部疾病数据如果不做额外扩增,可能只有两三千张训练图,全部层参与微调容易训飞。常见策略是冻结前几个 Stage,只更新后面的特征层和分类头。关键点:BatchNorm 层的均值和方差是根据当前 batch 统计的,如果你的 batch size 小(比如 8),BN 的统计量会很不稳,这时必须冻结所有 BN 层的参数。
# 冻结策略:前两个 stage 完全冻结 + 所有 BN 层冻结 running stats for name, param in model.named_parameters(): if name.startswith("layer1") or name.startswith("layer2"): param.requires_grad = False # 所有 BatchNorm 层固定 running_mean / running_var,只用学习到的 scale/shift for module in model.modules(): if isinstance(module, nn.BatchNorm2d): module.eval() module.weight.requires_grad = False module.bias.requires_grad = False| 冻结范围 | 显存占用 | 适合场景 | 备注 |
|---|---|---|---|
| 不冻结 | 最高 | 数据量 1 万+,且和 ImageNet 风格差异大 | 需要较小学习率 |
| 冻结 layer1~layer2 | 中等 | 数据量 3000~10000 | 兼顾特征复用与任务适配 |
| 全部冻结只训分类头 | 最低 | 数据量 <1000 | 本质是特征提取器 + 逻辑回归 |
| 额外冻结 BN | 不变 | batch size 极小 | 对稳定性的提升最明显 |
在训练循环里必须把模型切到model.train()状态,否则 BN 层即使requires_grad=False也会因为 eval 模式而用历史统计量,导致训练和验证行为不一致。
4. 训练主循环:损失函数选择、优化器参数与过拟合控制
4.1 眼科疾病数据集的类别不平衡与 loss 设计
眼科疾病数据天然是不平衡的——正常样本往往最多,早期病变样本少。直接最小化交叉熵会让模型偏向多数类。项目里常见做法有两个方向:一是给损失函数加类别权重,二是用WeightedRandomSampler在采样层面做平衡。我在工程上更推荐先做加权损失,因为它不改变每个 epoch 的数据分布,调试时更容易定位问题。
from torch.nn import CrossEntropyLoss # 按训练集各类别样本数的反比设置权重 class_counts = torch.tensor([3500, 1200, 800, 600, 400]).float() class_weights = class_counts.sum() / class_counts class_weights = class_weights.to(device) criterion = CrossEntropyLoss(weight=class_weights)CrossEntropyLoss的weight参数会在每个样本的 loss 上乘以对应类别的权重系数。比如normal类权重是 0.6,样本多的类贡献被压低;myopia类权重是 5.7,一个稀疏类别的错误预测会产生更大的梯度。这种做法的副作用是训练初期整体 loss 会偏高,学习率要相应调小一点。
4.2 两层训练循环:训练方差与验证评估
pytorch 的标准训练循环写起来不长,但里面有几个隐藏点:optimizer.zero_grad()必须放在前向之前;scaler.scale(loss).backward()是 AMP 混合精度训练的固定写法;验证阶段要用torch.no_grad()并且别忘了把模型切到eval()。
import torch from torch.cuda.amp import GradScaler, autocast def train_one_epoch(model, loader, criterion, optimizer, scaler): model.train() total_loss, correct, total = 0.0, 0, 0 for images, labels in loader: images, labels = images.to(device), labels.to(device) optimizer.zero_grad() # 混合精度前向传播 with autocast(): outputs = model(images) loss = criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() total_loss += loss.item() * images.size(0) correct += (outputs.argmax(1) == labels).sum().item() total += labels.size(0) return total_loss / total, correct / totalautocast只在 GPU 上生效,CPU 训练时会直接 pass,所以兼容性不用太担心。scaler的作用是防止混合精度下的梯度下溢,它把 loss 乘以一个放大因子再回传,更新参数前再缩小。每轮迭代后scaler.update()会动态调整放大因子。
4.3 训练超参数速查表:照着抄能稳定的起点
| 参数 | 推荐值 | 说明 |
|---|---|---|
| optimizer | AdamW | 比 Adam 多一个权重衰减修正,泛化更好 |
| base_lr | 1e-4 | 迁移学习统一用 1e-4 起步,预训练权重不需要大的更新步 |
| weight_decay | 1e-4 | 过大抑制 BN 外的参数,过小起不到正则作用 |
| batch_size | 32 | 显存不够降到 16,并同步把 lr 降到 5e-5 |
| epochs | 30 | 前 10 轮看趋势,后面 20 轮精调 |
| lr_scheduler | CosineAnnealingLR | 周期 30,eta_min=1e-6 |
| warmup_epochs | 2 | 从小 lr 线性升到 base_lr,稳住 BN 的统计量 |
CosineAnnealing 配合预训练模型的效果比 StepLR 平顺,因为它后期学习率无限趋近于 0,避免在损失面底部震荡。warmup 在 batch size 较大时尤其必要,因为前几个 batch 的梯度方向很不稳定,直接上大 lr 很容易破坏预训练权重。
4.4 过拟合判断与 checkpoint 保存
训练到第 5 轮左右,基本就能判断趋势了:训练 loss 稳步下降、验证 loss 开始回升,这是过拟合的典型信号。眼部疾病这种类内差异大的任务,验证集准确率通常在 60%~85% 之间徘徊,因为不同设备的成像色差会让模型"假装识别出病灶,实际在认色调"。
import os best_acc = 0.0 for epoch in range(30): train_loss, train_acc = train_one_epoch(...) # 省略 DataLoader 参数 val_loss, val_acc = validate(model, val_loader, criterion) if val_acc > best_acc: best_acc = val_acc torch.save({ "epoch": epoch, "model_state_dict": model.state_dict(), "class_to_idx": class_to_idx, }, "checkpoints/best_resnet50_eye.pth") print(f"Epoch {epoch}, val_acc improved -> {val_acc:.4f}")保存class_to_idx很容易被忽略,但推理时要恢复类别顺序,没有它就得靠猜或通过路径反推。state_dict只存参数不存模型结构,所以加载时你需要重新实例化一个结构相同的模型再load_state_dict。
5. 评估、推理脚本与一个被低估的验证技巧
5.1 验证集的指标不能只看 acc,要逐类别看 precision 和 recall
五分类任务里,准确率 85% 看着还行,但如果diabetic_retinopathy这种致盲率高的病 recall 只有 40%,模型的临床价值几乎为零。评估要从 sklearn 里拿classification_report和混淆矩阵,对着逐类指标找问题。
from sklearn.metrics import classification_report, confusion_matrix import numpy as np all_pred, all_label = [], [] model.eval() with torch.no_grad(): for images, labels in val_loader: images = images.to(device) outputs = model(images) all_pred.extend(outputs.argmax(1).cpu().numpy()) all_label.extend(labels.numpy()) print(classification_report(all_label, all_pred, target_names=["normal", "cataract", "glaucoma", "dr", "myopia"])) print(confusion_matrix(all_label, all_pred))混淆矩阵里对角线是预测正确的数量。如果某两类互相混淆严重,比如glaucoma和myopia交叉多,问题通常出在数据标注质量或者这两类病灶在视觉上确实高度重合,需要考虑对这两类做专门的二分类模型。
5.2 推理脚本:加载 checkpoint 并对单张眼底图预测
训练完的模型要能在生产环境跑单张图片。推理管道的输入处理必须和验证集完全一致:同样的Resize、同样的ToTensor、同样的Normalize。下面是完整可运行的推理代码。
import torch from PIL import Image from torchvision import transforms, models def predict_image(image_path, model, device, class_to_idx): # 推理时只做标准缩放和归一化,不做随机增强 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") image = transform(image).unsqueeze(0).to(device) model.eval() with torch.no_grad(): output = model(image) probs = torch.softmax(output, dim=1).cpu().numpy()[0] idx_to_cls = {v: k for k, v in class_to_idx.items()} sorted_idx = probs.argsort()[::-1] return [(idx_to_cls[i], probs[i]) for i in sorted_idx] # 使用方式 model = models.resnet50(weights=None) model.fc = ... # 和训练时完全一致的结构 checkpoint = torch.load("checkpoints/best_resnet50_eye.pth", map_location="cpu") model.load_state_dict(checkpoint["model_state_dict"]) model.to(device) result = predict_image("test/glaucoma_001.jpg", model, device, checkpoint["class_to_idx"]) print(result) # [('glaucoma', 0.91), ('myopia', 0.06), ...]map_location="cpu"是为了让在 GPU 上训练的模型也能在无 GPU 的机器上加载,实测中这一步能避免大量 "out of memory" 的部署事故。推理时的Dropout层不用手动关,model.eval()会自动将其切换成恒等映射。
5.3 一个最有价值的验证技巧:用 CAM 检查模型到底在看哪里
准确率达标不等于模型找到了病灶区域。眼科图像中,模型很可能学到的是"这张图偏暗所以是白内障"这种错误的捷径特征。pytorch 里可以对最后一个卷积层的输出做加权求和,生成类激活映射(CAM),然后把热力图叠加到原始图像上,直接看模型在预测某类疾病时关注了哪些区域。
# 注册钩子获取最后一个卷积层的输出 feature_map = None def hook_fn(module, input, output): global feature_map feature_map = output.detach() handle = model.layer4.register_forward_hook(hook_fn) # 前向传播获得 logits,取目标类别 output = model(image.unsqueeze(0)) target_class = output.argmax(1).item() # 取全连接层权重,对特征图加权求和 fc_weights = model.fc[2].weight[target_class] # 注意你的 fc 结构 cam = torch.matmul(fc_weights, feature_map.flatten(2)).reshape(224, 7, 7) cam = torch.relu(cam).unsqueeze(0).unsqueeze(0) cam = torch.nn.functional.interpolate(cam, size=(224, 224), mode="bilinear")model.layer4的输出分辨率是 7×7,正好对应 ResNet50 最后一层特征图的大小。如果 CAM 的热点集中在血管、视盘等结构而非渗出物或出血点,说明模型的决策依据不充分,需要回到数据层面清洗标注或增加这类病灶的样本。
若模型表现正常,你还可以把推理脚本包成一个 HTTP 接口,配合 FastAPI 提供/predict端点,输入眼底图像返回五类疾病的置信度,让眼科筛查工具真正被业务方调用起来。
本文还有配套的精品资源,点击获取