ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

PyTorch+ResNet50实现眼部疾病图像分类实战

PyTorch+ResNet50实现眼部疾病图像分类实战 简介这是一份面向医学图像处理与深度学习初学者的眼部疾病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. 眼科图片分类选 pytorchResNet50不是因为它最先进把眼底彩照直接丢给 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 python3.9 -y conda activate eye_cls # 安装 pytorch torchvisionCUDA 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(p0.5), transforms.RandomRotation(15), transforms.ColorJitter(brightness0.2, contrast0.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已经封装好了。需要注意接口已经发生过破坏性变更旧代码里的pretrainedTrue参数在较新的 torchvision0.13 之后会直接报错现在统一用weights参数。import torch import torch.nn as nn from torchvision import models, transforms # 新版写法显式指定预训练权重 weights models.ResNet50_Weights.IMAGENET1K_V1 model models.resnet50(weightsweights) # 查看最后两层结构确认要替换的位置 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(p0.5), nn.Linear(2048, 512), nn.ReLU(inplaceTrue), nn.Dropout(p0.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 则降低全连接层过拟合的风险。inplaceTrue节省显存但如果你之后要用钩子hook检查激活值建议改成inplaceFalse否则拿到的前向结果有可能是脏数据。3.3 冻结 BatchNorm 与部分 Stage显存不够或样本少时的通用技巧眼部疾病数据如果不做额外扩增可能只有两三千张训练图全部层参与微调容易训飞。常见策略是冻结前几个 Stage只更新后面的特征层和分类头。关键点BatchNorm 层的均值和方差是根据当前 batch 统计的如果你的 batch size 小比如 8BN 的统计量会很不稳这时必须冻结所有 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_gradFalse也会因为 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(weightclass_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 训练超参数速查表照着抄能稳定的起点参数推荐值说明optimizerAdamW比 Adam 多一个权重衰减修正泛化更好base_lr1e-4迁移学习统一用 1e-4 起步预训练权重不需要大的更新步weight_decay1e-4过大抑制 BN 外的参数过小起不到正则作用batch_size32显存不够降到 16并同步把 lr 降到 5e-5epochs30前 10 轮看趋势后面 20 轮精调lr_schedulerCosineAnnealingLR周期 30eta_min1e-6warmup_epochs2从小 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(fEpoch {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, dim1).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(weightsNone) model.fc ... # 和训练时完全一致的结构 checkpoint torch.load(checkpoints/best_resnet50_eye.pth, map_locationcpu) 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_locationcpu是为了让在 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), modebilinear)model.layer4的输出分辨率是 7×7正好对应 ResNet50 最后一层特征图的大小。如果 CAM 的热点集中在血管、视盘等结构而非渗出物或出血点说明模型的决策依据不充分需要回到数据层面清洗标注或增加这类病灶的样本。若模型表现正常你还可以把推理脚本包成一个 HTTP 接口配合 FastAPI 提供/predict端点输入眼底图像返回五类疾病的置信度让眼科筛查工具真正被业务方调用起来。本文还有配套的精品资源点击获取
返回列表