news 2026/9/13 13:03:47

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
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. 眼科图片分类选 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 的评估结果可复现。RandomHorizontalFlipRandomRotation属于空间增强,对眼底图特别合适——眼球的朝向不改变病变特征;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)

CrossEntropyLossweight参数会在每个样本的 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 / total

autocast只在 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-5
epochs30前 10 轮看趋势,后面 20 轮精调
lr_schedulerCosineAnnealingLR周期 30,eta_min=1e-6
warmup_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(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))

混淆矩阵里对角线是预测正确的数量。如果某两类互相混淆严重,比如glaucomamyopia交叉多,问题通常出在数据标注质量或者这两类病灶在视觉上确实高度重合,需要考虑对这两类做专门的二分类模型。

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端点,输入眼底图像返回五类疾病的置信度,让眼科筛查工具真正被业务方调用起来。

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

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

DB-GPT 接入 Ollama:本地运行开源模型的完整配置指南

DB-GPT 接入 Ollama&#xff1a;本地运行开源模型的完整配置指南 【免费下载链接】DB-GPT open-source agentic AI data assistant for the next generation of AI Data products. 项目地址: https://gitcode.com/GitHub_Trending/db/DB-GPT 导读 本文围绕 DB-GPT 项目…

作者头像 李华
网站建设 2026/9/13 13:01:50

Auracast与全双工对讲方案详解:LE Audio低延时高音质工程实践

最近圈子里聊得比较多的 Auracast 广播音频&#xff0c;泰凌这次给了一套能直接落地的方案&#xff0c;亮点是把 Auracast 广播接收/发射和高音质低延时的全双工对讲打包在一起。我拿到样品之后做了好几轮测试&#xff0c;今天把这套方案背后的技术思路、SDK 里的处理细节&…

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

Maven PKIX报错:证书信任链排查与cacerts修复指南

1. 破译报错&#xff1a;PKIX path building failed到底是谁在说话先别急着百度复制粘贴解决方案&#xff0c;我们花两分钟把这段报错真正看懂。绝大多数Maven用户在IDEA里导入项目时看到这行红字&#xff0c;第一反应是“Maven崩了”“镜像挂了”“IDEA坏了”&#xff0c;其实…

作者头像 李华
网站建设 2026/9/13 12:58:07

LabVIEW比较运算符详解与应用实践

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/13 12:57:07

研究生论文写作必备:9大AI工具评测与使用策略

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华