• 正文
  • 相关推荐
申请入驻 产业图谱

深度学习实战-基于Resnet50的青光眼疾病图像识别模型

08/07 11:00
176
加入交流群
扫码加入
获取工程师必备礼包
参与热点资讯讨论

 

1.项目背景

在现代眼科临床诊断与公共卫生防盲治盲体系中,青光眼作为全球首位不可逆性致盲眼病,其早期的隐匿性与突发性对传统筛查模式提出了严峻的挑战。由于该病变在发病初期症状极不明显,大量患者在确诊时已造成了不可逆的视神经萎缩与视野缺损。传统的临床筛查高度依赖资深眼科医生对眼底彩照中的视盘、视杯形态及杯盘比进行人工判读。然而,这种依靠主观经验的诊断模式不仅极度消耗专家资源,且在基层医疗或大面积人群普查中面临复盖率低、误诊率高的工程瓶颈。随着医学影像技术的数字化转型与智能医疗的深入推进,如何利用计算机视觉技术提取精细的解剖学异质性特征,构建自动、高效且稳健的眼底病变识别系统,已成为辅助临床诊断、优化医疗资源配置的紧迫技术课题。

本项目立足于医学影像分析与迁移学习的交叉前沿,展开了基于深度学习架构的青光眼全自动辅助筛查研究。实验围绕眼底彩照中视网膜血管走势及视神经乳头的微观形变规律,深度部署了具备卓越残差拟合能力的 ResNet50 神经网络模型。通过重构预训练骨干网络的末端分类决策头,模型成功实现了通用视觉先验向眼科垂直病理诊断空间的深度迁移。在整个技术闭环中,方案不仅打通了从医学图像流式动态增强、硬件加速自适应映射到多维定量评估的端到端管线,更通过对独立测试集的深度透视,科学量化了模型在保障高准确率的同时控制医疗漏诊率的实战效能,为开发下沉至社区和基层医院的便携式智能眼底筛查设备提供了可靠的算法支撑与工程技术参照。

2.数据集介绍

本实验数据集来源于Kaggle的青光眼数据集,Hillel Yaffe 青光眼数据集 (HYGD) 是一个经过精心整理的、黄金标准的带标注眼底图像集合,旨在用于青光眼的自动检测和分类。该数据集提供高质量的图像,对于训练和评估眼科诊断领域的机器学习和深度学习模型至关重要。
青光眼是全球导致不可逆性失明的主要原因之一。早期发现对于有效治疗和保护视力至关重要。本数据集旨在通过提供可靠的、带注释的资源,弥合临床眼科学与人工智能之间的差距,从而促进开发强大的自动化诊断工具。

数据集内容
图像:包含 747 张高分辨率眼底图像(图像目录)。

注释:包含一个 Labels.csv 文件,为每张图像提供结构化的元数据。

文档:包含 README.md 和 LICENSE.txt,其中包含技术规范和使用权信息。

完整性:提供了一个 SHA256SUMS.txt 文件,以确保在下载和处理过程中数据的完整性。

潜在应用案例
疾病检测:训练二元或多类分类器来识别青光眼眼底图像与健康眼底图像。

计算机视觉研究:开发视盘和视杯分割算法。

医疗人工智能开发:实现和评估用于医学图像分析的深度学习架构(例如,CNN、Vision Transformer)。

3.技术工具

Python版本:3.9

代码编辑器:jupyter notebook

4.实验过程

4.1导入数据

在临床眼科诊断与防盲治盲工作中,青光眼作为全球首位不可逆性致盲眼病,其早期的隐匿性使得临床筛查面临巨大挑战。传统的眼底彩照人工判读不仅极度消耗专家资源,且受限于主观经验,在基层医疗机构中难以实现大面积覆盖。本项目深入临床医学影像与计算机视觉的交叉领域,利用 PyTorch 深度学习框架,构建了一套基于 ResNet50 深度残差网络的眼底图像青光眼全自动识别方案。实验流程第一步聚焦于底层环境的初始化与病理标签数据的结构化解析。通过导入核心的张量计算、医学图像处理及模型评估组件,并利用 Pandas 对临床标注的 CSV 数据集进行深度清洗与字段剔除,不仅实现了计算硬件(GPU/CPU)的自适应分配,也为后续构建稳健的眼底图流式动态加载管道奠定了规范的数据形态。

import osimport numpy as npimport pandas as pdimport matplotlib.pyplot as pltfrom PIL import Imageimport torchimport torch.nn as nnfrom torch.utils.data import Dataset, DataLoaderimport torchvisionfrom torchvision import transformsfrom sklearn.model_selection import train_test_splitfrom sklearn.metrics import (    accuracy_score,    classification_report,    confusion_matrix)from torchvision.models import resnet50# --- 1. 计算硬件加速自适应配置 ---# 优先检测当前环境中是否存在可用的英伟达 GPU 加速芯片,否则回退至 CPU 计算device = torch.device("cuda" if torch.cuda.is_available() else "cpu")# --- 2. 物理数据集标注文件读取 ---csv_path = "/kaggle/input/datasets/nilesh2042/hillel-yaffe-glaucoma-dataset-hygd-images/hillel-yaffe-glaucoma-dataset-hygd-a-gold-standard-annotated-fundus-dataset-for-glaucoma-detection-1.1.0/Labels.csv"df = pd.read_csv(csv_path)# --- 3. 结构化表格清洗与冗余字段剔除 ---# 去除数据集中未命名的空白杂质列,确保 DataFrame 数据对齐df = df.drop(columns=["Unnamed: 4"])# 预览前五行数据,初步核验图像路径、患者ID及疾病标签的映射关系df.head()

4.2数据可视化

在医学图像分类任务中,深入了解病变样本的底层形态与色彩分布是设计鲁棒特征流水线的重要前提。青光眼的临床诊断通常依赖于眼底彩照中视盘、视杯的解剖学变异(如杯盘比扩大、视网膜神经纤维层缺损等),因此在将影像输入残差网络前,必须进行路径的映射重组与物理抽检。本阶段首先利用 Lambda 表达式实现图像文件名与物理硬盘存储路径的流式组装,并利用随机采样逻辑调取多幅眼底图进行网格化渲染,帮助我们从宏观上审视光照分布、对比度波动以及病灶特征的清晰度。

import random# --- 1. 物理图像存储根目录配置 ---IMAGE_DIR = "/kaggle/input/datasets/nilesh2042/hillel-yaffe-glaucoma-dataset-hygd-images/hillel-yaffe-glaucoma-dataset-hygd-a-gold-standard-annotated-fundus-dataset-for-glaucoma-detection-1.1.0/Images"image_files = os.listdir(IMAGE_DIR)# --- 2. 动态路径合成与 DataFrame 字段扩展 ---# 利用匿名函数将根目录路径与文件名列进行字符串拼接,生成完整的物理文件访问索引df["image_path"] = df["Image Name"].apply(    lambda x: os.path.join(IMAGE_DIR, x))# --- 3. 2x3 临床眼底图矩阵网络画布组装 ---fig, axes = plt.subplots(2, 3, figsize=(12, 8))# 扁平化轴阵列,循环抽取样本执行定量回显for ax in axes.flatten():    # 随机挑选当前图像池中的一个眼底样本文件名    img_name = random.choice(image_files)    img_path = os.path.join(IMAGE_DIR, img_name)    # 依托 PIL 引擎加载物理影像并渲染至对应的子图    img = Image.open(img_path)    ax.imshow(img)    ax.set_title(img_name) # 将子图标题设置为当前文件名以便于追溯病理编号    ax.axis("off")         # 移除传统物理坐标轴,聚焦于眼底解剖结构本身plt.tight_layout()plt.show()

通过这一阶段的映射与可视化,数据流水线的完备性得到了初步验证。将合成的物理路径存储于 DataFrame 中,避免了后续训练时频繁执行磁盘扫描带来的算力空耗;而随机采样的 2 x 3 眼底图像网格则揭示了真实医学影像中的常见技术挑战,如不同批次间的光照异质性、眼底视网膜血管分布的错综复杂等。这些直观的视觉反馈明确地告诉我们,在下一阶段的特征工程中,必须引入标准化的中心裁剪与对比度归一化,才能协助 ResNet50 绕开环境环境噪声,将有限的感受野精准聚焦于视神经乳头的微观病变区。

4.3特征工程

在深度学习临床医学影像任务中,特征工程的核心价值在于构建严谨的受试者数据隔离流水线,并通过针对性的光学扰动增强模型面对多机构影像时的泛化边界。本阶段首先实现临床病理文本标签向离散整数张量的数字化映射,接着采用多阶分层抽样技术,将患者数据集严密划分为训练集、验证集与独立测试集。随后,针对眼底照相机光源不均及视网膜形态的异质性,定制了双轨制的 torchvision.transforms 预处理算子,并继承自 PyTorch 的 Base Dataset 构建了流式数据异步加载管道,实现了物理眼底图像向张量矩阵的无缝转化与高效批分发。

# =========================================================# 第一部分:临床标签数字化、多阶分层划分与数据增强配置# =========================================================# --- 1. 临床医学文本标签向离散整数标签的映射映射 ---label_map = {    "GON-": 0,   # Healthy (健康眼底彩照)    "GON+": 1    # Glaucoma (确诊青光眼病理影像)}df["target"] = df["Label"].map(label_map)# --- 2. 引入分层抽样机制执行多阶数据集切分 ---# stratify=df["target"] 确保训练集和测试集中的健康与青光眼样本比例与原始大盘严格保持一致train_df, test_df = train_test_split(df, test_size=0.2, random_state=42, stratify=df["target"])train_df, val_df = train_test_split(train_df, test_size=0.2, random_state=42, stratify=train_df["target"])# --- 3. 训练集数据增强与标准化管道配置 ---train_transform = transforms.Compose([    transforms.Resize((224,224)),       # 将眼底图统一缩放至 ResNet50 的标准输入尺寸    transforms.RandomHorizontalFlip(),  # 模拟左/右眼底镜像的空间对称性    transforms.RandomRotation(10),      # 模拟由于患者头部轻微倾斜带来的解剖位置偏移    transforms.ColorJitter(        brightness=0.2,        contrast=0.2    ),                                  # 模拟不同品牌眼底相机(如Topcon或Zeiss)的光源照度与对比度波动    transforms.ToTensor(),              # 将 HWC 格式的 PIL 图像转化为 CHW 的 PyTorch 张量    transforms.Normalize(        mean=[0.485,0.456,0.406],        std=[0.229,0.224,0.225]    )                                   # 基于 ImageNet 数据集的均值与标准差执行全通道归一化])# --- 4. 验证集与测试集标准化管道配置 ---# 注意:测试阶段严禁使用任何随机翻转与色彩扰动,以保障评估的客观客观性test_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]    )])# 打印各子集样本规模,核验数据资产分布print("Train:", len(train_df))print("Validation:", len(val_df))print("Test:", len(test_df))# =========================================================# 第二部分:自定义 PyTorch 数据集定义与高并发 DataLoader 实例化# =========================================================# --- 5. 继承 Dataset 类,自定义眼底图像流式加载器 ---class GlaucomaDataset(Dataset):    def __init__(self, dataframe, transform=None):        # 必须重置索引,确保通过 loc 检索连续逻辑序号时不会引发 KeyError 报错        self.df = dataframe.reset_index(drop=True)        self.transform = transform    def __len__(self):        return len(self.df)    def __getitem__(self, idx):        # 动态索引当前批次的物理图片文件存储路径        image_path = self.df.loc[idx, "image_path"]        # 转换为三通道 RGB 模式,消除部分特殊医学格式产生的 alpha 通道杂质        image = Image.open(image_path).convert("RGB")        # 调取临床真值标签        label = self.df.loc[idx, "target"]        # 流式触发张量处理与标准化增强        if self.transform:            image = self.transform(image)        return image, torch.tensor(label, dtype=torch.long)# --- 6. 实例化三个独立的医学数据集对象 ---train_dataset = GlaucomaDataset(train_df, transform=train_transform)val_dataset = GlaucomaDataset(val_df, transform=test_transform)test_dataset = GlaucomaDataset(test_df, transform=test_transform)# --- 7. 构建多线程异步流式数据分发加载器 ---# 配合 16 的小批次大小,num_workers=2 开启多进程加速数据矩阵解码train_loader = DataLoader(train_dataset, batch_size=16, shuffle=True, num_workers=2)val_loader = DataLoader(val_dataset, batch_size=16, shuffle=False, num_workers=2)test_loader = DataLoader(test_dataset, batch_size=16, shuffle=False, num_workers=2)# --- 8. 抽取首批数据执行管线深度打通验证 ---images, labels = next(iter(train_loader))print(images.shape)  # 预期输出: torch.Size([16, 3, 224, 224])print(labels.shape)  # 预期输出: torch.Size([16])print(labels[:10])   # 打印当前批次前 10 个样本的离散疾病真值标签

数据管道的全面实例化与自检成功完成。控制台回显的训练集、验证集与测试集数量展示了清晰的受试者分配逻辑。通过 next(iter(train_loader)) 的边界测试,我们成功获取到了第一组形如 [16, 3, 224, 224] 的四维批次图像张量。这意味着 CPU 已经成功在内存中完成了图像的解码、缩放、灰度变换与多通道归一化,整个特征工程的输出形态与后续 ResNet50 的第一层三维空间感受野达成了完美的维度对齐,数据流可以安全地推向计算硬件。

4.4构建模型

在医学影像识别领域,利用在大规模通用视觉数据集上训练成熟的深度残差网络进行迁移学习,是解决医疗样本稀缺、加速模型收敛的经典策略。ResNet50 凭借其独特的跨层恒等映射机制,能够有效避免网络加深带来的梯度色散,其浅层和中层卷积组积攒了强大的边缘、纹理及高级空间结构感知能力。针对眼底彩照中视盘与视杯的微观解剖形变,我们无需从零训练庞大的特征提取器,而是通过保留其强大的骨干权重,专职重构其末端的全连接分类决策头,使其在保持原有高阶视觉感知力的同时,精准适配本次青光眼二分类任务的输出维度。

# --- 1. 导入具备高阶先验特征的预训练 ResNet50 骨架网络 ---# 选用官方深度优化的 IMAGENET1K_V2 权重版本,其相比 V1 具有更强的特征泛化与表示上限model = resnet50(weights="IMAGENET1K_V2")# --- 2. 动态解析原始分类头的输入特征维度 ---# 提取 ResNet50 最后一层全连接层(fc)的输入神经元通道数(预期为 2048 维)num_features = model.fc.in_features# --- 3. 重构全连接层分类决策头 ---# 将原先面向 ImageNet 千分类的输出拓扑,替换为面向“健康 / 青光眼”的全新线性映射层(输出维度为 2)model.fc = nn.Linear(num_features, 2)# --- 4. 将计算图与网络参数全面推向目标计算硬件 ---# 配合前期的硬件检测,将模型权重整体载入 GPU 显存或 CPU,为前向传播做准备model = model.to(device)

本环节成功完成了核心分类网络的拓扑重组与硬件常驻。通过调取代表通用视觉特征最高水平之一的 IMAGENET1K_V2 权重,模型在物理层面上已经具备了辨识微观色彩渐变与几何边界的“视力”。通过破坏原有的 model.fc 并重新实例化一个输入为 2048 维、输出为 2 维的 nn.Linear 算子,我们强迫网络在后续的微调(Fine-tuning)训练中,将所有注意力聚焦于如何将眼底图像的高维语义特征映射到临床诊断标签上,完成了从通用目标检测向专业医疗辅助诊断的技术跨越。

4.5训练模型

在深度学习模型的生命周期中,训练控制循环是决定模型能否将先验视觉能力转化为垂直领域诊断精度的关键阶段。青光眼识别作为一项严谨的医疗辅助诊断任务,其参数优化过程需要兼顾收敛的平稳性与泛化边界。本阶段通过配置标准的标准交叉熵损失函数(CrossEntropyLoss)来量化预测概率与临床真值之间的偏差,并引入 Adam 优化器以 0.0001 的精细步长对网络参数实施渐进式微调。通过显式切换模型的训练(train())与评估(eval())状态,我们在 10 个 Epoch 的双轨控制循环中交替执行前向传播、梯度反向传播以及无梯度的验证集性能盘点,将每一轮的损失值与准确率动态固化至历史监控容器中。

# --- 1. 损失函数与参数优化器配置 ---criterion = nn.CrossEntropyLoss()optimizer = torch.optim.Adam(model.parameters(), lr=1e-4) # 使用保守的学习率平滑微调预训练权重epochs = 10# --- 2. 初始化核心监控数据指标容器 ---train_losses = []val_losses = []train_accuracies = []val_accuracies = []# --- 3. 端到端双轨控制训练循环 ---for epoch in range(epochs):    # =========================================================    # 训练阶段 (Training Phase)    # =========================================================    model.train() # 显式激活训练模式,允许 Dropout 与 Batch Normalization 动态更新    running_loss = 0    correct = 0    total = 0    for images, labels in train_loader:        # 将输入矩阵与标签动态载入指定的物理计算设备        images = images.to(device)        labels = labels.to(device)        # 梯度清零,防止历史批次的残余梯度对当前权重产生干扰        optimizer.zero_grad()        # 前向传播:计算眼底图像在当前网络下的置信度得分        outputs = model(images)        loss = criterion(outputs, labels)        # 反向传播:自动计算各层神经元的偏导数        loss.backward()        # 参数更新:利用 Adam 算子沿着梯度负方向更迭权重        optimizer.step()        # 统计当前批次的累加损失值与正确频数        running_loss += loss.item()        _, preds = torch.max(outputs, 1)        total += labels.size(0)        correct += (preds == labels).sum().item()    # 计算当前轮次的平均训练损失与准确率    train_loss = running_loss / len(train_loader)    train_acc = correct / total    train_losses.append(train_loss)    train_accuracies.append(train_acc)    # =========================================================    # 验证阶段 (Validation Phase)    # =========================================================    model.eval() # 显式激活评估模式,锁死 BN 统计量与神经元随机失活机制    val_running_loss = 0    val_correct = 0    val_total = 0    # 锁定内存计算图,关闭梯度流计算,节约训练期间的显存与算力    with torch.no_grad():        for images, labels in val_loader:            images = images.to(device)            labels = labels.to(device)            outputs = model(images)            loss = criterion(outputs, labels)            val_running_loss += loss.item()            _, preds = torch.max(outputs, 1)            val_total += labels.size(0)            val_correct += (preds == labels).sum().item()    # 计算当前轮次的平均验证损失与准确率    val_loss = val_running_loss / len(val_loader)    val_acc = val_correct / val_total    val_losses.append(val_loss)    val_accuracies.append(val_acc)    # 实时回显当前周期的宏观性能指标    print(        f"Epoch [{epoch+1}/{epochs}] "        f"Train Loss: {train_loss:.4f} "        f"Train Acc: {train_acc:.4f} "        f"Val Loss: {val_loss:.4f} "        f"Val Acc: {val_acc:.4f}"    )

4.6模型评估

训练与验证损失曲线复盘

本部分利用 Matplotlib 算子将训练阶段固化下来的 train_losses 与 val_losses 历史日志进行双线同框渲染。通过观察两条损失随 Epoch 演进的收敛轨迹,我们可以直观判断残差网络在寻找局部最优解时的稳定度。

# --- 1. 渲染训练集与验证集损失函数对比曲线 ---plt.figure(figsize=(10,5))plt.plot(train_losses, label="Train Loss")plt.plot(val_losses, label="Validation Loss")plt.xlabel("Epoch")plt.ylabel("Loss")plt.title("Training vs Validation Loss")plt.legend()plt.show()

准确率增量演进对比

本部分聚焦于模型判别效能的成长状态,将训练集与验证集的准确率轨迹(train_accuracies 与 val_accuracies)映射至同一坐标系中,用以复盘 ResNet50 在特征空间重组过程中的学习增益速度。

# --- 2. 渲染训练集与验证集分类准确率对比曲线 ---plt.figure(figsize=(10,5))plt.plot(train_accuracies, label="Train Accuracy")plt.plot(val_accuracies, label="Validation Accuracy")plt.xlabel("Epoch")plt.ylabel("Accuracy")plt.title("Training vs Validation Accuracy")plt.legend()plt.show()

独立测试集流式推断与多指标精细报告

为了确保评估结果不带有任何信息泄露(Data Leakage)偏置,本部分全面切入完全独立的 test_loader 测试集管线。通过关闭梯度流、将模型强制置为 eval() 评估模式,流式收集未见样本的预测结果,并调用 Scikit-Learn 算子输出涵盖精确率(Precision)、召回率(Recall)与 F1-Score 的细粒度医学分类报告。

# --- 3. 独立测试集流式前向推断与评估 ---model.eval() # 冻结批归一化统计量,确立测试模式all_preds = []all_labels = []with torch.no_grad(): # 彻底锁定计算图,阻断梯度传播,降低测试能耗    for images, labels in test_loader:        # 矩阵流跨硬件载入        images = images.to(device)        labels = labels.to(device)        # 模型前向推理得到预测置信度        outputs = model(images)        _, preds = torch.max(outputs, 1) # 提取置信度最高的类别索引        # 将测试结果与临床真值同步回收至 CPU 显存空间中        all_preds.extend(preds.cpu().numpy())        all_labels.extend(labels.cpu().numpy())# --- 4. 打印针对健康眼底与青光眼病灶的细粒度分类报告 ---print(    classification_report(        all_labels,        all_preds,        target_names=[            "Healthy",            "Glaucoma"        ]    ))

独立测试集整盘综合跑分判定

本部分作为全盘定量的硬性标准,通过调用 accuracy_score 直接计算最优权重模型在整个独立测试集上的总正确比率,直观反映该医疗辅助诊断算法的宏观可靠度。

# --- 5. 计算并打印测试集全局总准确率 ---test_acc = accuracy_score(    all_labels,    all_preds)print("Test Accuracy:", test_acc)

混淆矩阵空间可视化

作为多分类性能盘点的压轴环节,本部分计算并绘制了结构化的二维混淆矩阵(Confusion Matrix)。通过结合 plt.imshow 将数值分布矩阵转化为直观的色块热图,可视化展示各样本类型的分类归宿。

# --- 6. 计算混淆矩阵并执行色块热图渲染 ---cm = confusion_matrix(    all_labels,    all_preds)print(cm) # 打印原始矩阵频数阵列plt.figure(figsize=(6,5))plt.imshow(cm) # 渲染混淆空间热度plt.colorbar() # 挂载热图色彩刻度尺plt.xlabel("Predicted")plt.ylabel("Actual")plt.title("Confusion Matrix")plt.show()

最后可以保存模型

torch.save(    model.state_dict(),    "glaucoma_resnet50.pth")print("Model Saved Successfully")

5.总结

本实验围绕全球首位不可逆性致盲眼病——青光眼的早期临床自动检测展开,依托黄金标准的 Hillel Yaffe 青光眼数据集(HYGD)成功实现了一套高精度的医学影像辅助诊断管线。实验利用数据集中 747 张高分辨率眼底彩照作为底层资产,在经历文本标签转换、受试者多阶分层切分以及光学异质性数据增强后,全流程打通了预训练 ResNet50 深度残差网络的微调训练。

定量评估结果表明,模型在未参与任何优化迭代的 150 例独立测试集大盘中展现出了卓越的泛化与判别边界,最终交出了全局总准确率 98.67% 的优秀成绩。深度透视其细粒度多指标分类报告,网络针对“Healthy(健康)”与“Glaucoma(青光眼)”的加权平均 F1-score 达到了 0.99,且双向混淆矩阵显示在 110 例真实的青光眼病理影像中,模型仅产生 1 例漏诊,对阳性病灶的召回率高达 99%

这一杰出的数据表现充分论证了深度残差网络在复用通用视觉先验的基础上,能够极其敏锐地捕捉到视盘、视杯等细微解剖结构的变异与形变。本实战不仅弥合了临床眼科学与计算机视觉之间的技术鸿沟,更为未来将高精度医学人工智能算法下沉部署至基层医疗机构、开发便携式全自动眼底筛查系统提供了极具数理背书与工程落地价值的规范化范例。

相关推荐