1. 项目概述:当深度学习遇上临床诊断,我们如何“看见”AI的决策?

在医疗影像AI领域,尤其是像青光眼诊断这样的关键任务中,模型的高准确率只是起点,而非终点。医生和临床专家真正关心的是:这个“黑箱”模型是基于什么做出判断的?它的诊断依据是否与人类专家的临床知识(如视盘形态、杯盘比、神经纤维层缺损等)相一致?如果模型只是“蒙对了”结果,但其决策逻辑与医学共识相悖,那么它在临床上的应用价值将大打折扣,甚至带来风险。

这就是“深度学习可解释性”的核心挑战。我们做的这个项目,正是聚焦于使用 类激活映射 方法,来验证一个用于青光眼筛查的深度学习模型,其内部决策是否与眼科医生的临床知识实现了“对齐”。简单来说,我们不仅要模型告诉我们“这张眼底照片有青光眼”,还要它用高亮区域“指”出来:“看,我判断的依据是这里的视盘凹陷变深、盘沿变窄了”。然后,我们将模型“指”出的区域,与资深眼科医生标注的关键病变区域进行量化对比,从而评估模型决策的临床合理性。

这不仅仅是技术验证,更是一种建立临床信任的桥梁。对于放射科医生、眼科医生而言,一个能提供合理解释的AI,更像是一位可以讨论病例的“同事”,而非一个无法沟通的“算命机器”。我们的工作,就是为这位“AI同事”做一次深入的“业务能力考核”,确保它的“诊断思路”是靠谱的。

2. 核心思路与技术选型:为什么是CAM?

2.1 可解释性方法的“全家福”与我们的选择

深度学习可解释性方法大致可分为两类: 事后解释方法 内置可解释模型 。事后解释方法是在训练好的模型上施加分析,如LIME、SHAP以及我们使用的CAM系列;内置可解释模型则试图在模型结构设计中融入可解释性,如注意力机制。

我们选择 基于梯度的类激活映射 系列方法作为核心工具,主要基于以下几点考量:

  1. 与卷积神经网络天然契合 :我们使用的青光眼诊断模型基于CNN(如ResNet、DenseNet)。CAM系列方法通过分析最后一个卷积层的特征图与最终分类权重的关系来生成热力图,这与CNN的层次化特征提取逻辑完全匹配,解释生成过程直观。
  2. 定位能力与可视化直观性 :CAM生成的热力图能清晰、直观地高亮出模型做出分类决策时所依赖的图像区域。对于眼底彩照,这直接对应了病变可能的位置(如视盘),非常便于医生进行视觉比对和定性评估。
  3. 计算效率与实现便捷性 :相较于LIME需要扰动大量输入样本,或SHAP基于博弈论的计算,标准CAM及其变种(Grad-CAM, Grad-CAM++)的计算相对高效,只需一次前向传播和反向梯度计算,易于集成到现有诊断流程中。
  4. 临床验证的适配性 :热力图输出的是一张与原始图像空间对应的显著性图,我们可以很容易地将其与医生手工标注的“金标准”区域(如视盘分割掩膜、杯盘区域)进行空间上的量化比较(如重叠度计算),这为客观的“知识对齐”验证提供了可能。

注意 :CAM方法并非万能。它主要适用于CNN,且解释的是“模型认为哪里重要”,而非“为什么这个区域重要”。它无法提供因果推理。但对于“定位诊断依据”这一临床核心关切,它目前是最直接、有效的工具之一。

2.2 项目技术栈与流程设计

整个验证流程是一个清晰的闭环:

  1. 数据基础 :收集带有两级标注的眼底彩照数据集。
    • 一级标注 :图像级标签(青光眼/非青光眼)。
    • 二级标注(金标准) :像素级标注,由至少两名资深眼科医生独立标注并协商一致得到,标注区域包括视盘、视杯、视网膜神经纤维层缺损区等关键解剖与病变结构。
  2. 模型训练 :使用图像级标签训练一个二分类(青光眼/正常)的CNN模型(如DenseNet-121)。不引入任何像素级标注信息,确保模型是纯粹从图像级监督中学习。
  3. 解释生成 :对测试集中的每一张图像,使用Grad-CAM++(我们选择了它,因其能更好地处理多个实例和更精细的定位)生成对应于“青光眼”类别的热力图。
  4. 知识对齐验证 :这是核心分析步骤。我们将模型生成的热力图与医生的像素级标注进行多维度对比:
    • 视觉定性对比 :将热力图叠加在原始图像上,邀请临床医生评估模型关注的区域是否与临床关注的解剖/病变区域相符。
    • 定量空间对齐分析
      • 将热力图通过阈值化(如取前20%的显著区域)转换为二值化显著性区域。
      • 计算该区域与医生标注的视盘区域的重叠度指标,如 Dice系数 交并比
      • 更精细地,可以计算热力图在视杯区域内的平均激活值是否显著高于视盘其他区域或背景,这能验证模型是否真的关注了“杯盘比”这个核心指标。
  5. 结果分析与迭代 :根据定量指标和医生反馈,评估模型决策的临床合理性。如果对齐度低,可能需要反思数据质量、模型结构或训练策略,并迭代优化。

这个流程的核心思想是: 用医生标注的“知识地图”作为尺子,去度量模型热力图这把“决策尺子”的刻度是否准确。

3. 实操详解:从数据到验证的全链路实现

3.1 数据准备与预处理的关键细节

数据的质量直接决定了验证的可信度。我们使用的是公开数据集与内部数据结合的方式。

  • 数据集 :主要使用了 REFUGE 挑战赛的部分数据,并补充了少量与医院合作收集的脱敏数据。所有数据均获得了相应的伦理许可。
  • 预处理标准化流程
    1. 分辨率统一 :将所有图像缩放到固定尺寸(如1024x1024),避免尺寸差异影响CNN特征提取。
    2. 颜色归一化 :采用 CLAHE 对图像进行对比度受限的自适应直方图均衡化,以减轻不同拍摄设备、光照条件造成的颜色和对比度差异,使模型更关注结构信息而非颜色偏差。
    3. 医生标注处理 :将多位医生的标注进行融合。对于视盘/视杯标注,我们采用 STAPLE 算法生成共识标注。对于存在分歧的区域,在数据分析阶段会特别注明,这本身也是探究模型与不同医生认知差异的切入点。
    4. 数据划分 :严格按照患者ID划分训练集、验证集和测试集,确保同一个患者的图像不会出现在不同集合中,防止数据泄露导致验证结果虚高。

实操心得 :与临床医生共同定义和审核标注标准至关重要。例如,“视盘边界”的精确界定可能存在细微差别。我们为此组织了一次标注培训,并使用一个小的校准集让所有参与医生进行试标注,讨论分歧直至达成明确协议。这个前期沟通成本不能省,它是后续一切定量分析的基础。

3.2 模型训练与Grad-CAM++集成

我们选择 DenseNet-121 作为主干网络,因为在有限的医疗数据上,它的特征复用特性有助于缓解过拟合。

import torch
import torch.nn as nn
import torchvision.models as models

class GlaucomaDenseNet(nn.Module):
    def __init__(self, num_classes=2):
        super(GlaucomaDenseNet, self).__init__()
        # 加载预训练的DenseNet-121,并移除最后的全连接层
        backbone = models.densenet121(pretrained=True)
        self.features = backbone.features
        # 获取分类器前的特征维度
        num_features = backbone.classifier.in_features
        # 自定义分类器
        self.classifier = nn.Linear(num_features, num_classes)
        # 注册钩子,用于获取最后一个卷积层的输出和梯度
        self.gradients = None
        self.activations = None

    # 前向传播钩子,保存激活值
    def activations_hook(self, grad):
        self.gradients = grad

    def forward(self, x):
        x = self.features(x)
        # 在需要计算CAM时,注册钩子
        if x.requires_grad:
            h = x.register_hook(self.activations_hook)
        self.activations = x  # 保存最后一个卷积层的输出
        x = torch.nn.functional.adaptive_avg_pool2d(x, (1, 1))
        x = torch.flatten(x, 1)
        x = self.classifier(x)
        return x

    # 方法:获取用于生成CAM的激活值和梯度
    def get_activations_gradient(self):
        return self.gradients

    def get_activations(self, x):
        return self.activations

训练过程中,我们使用了带权重的交叉熵损失,以应对数据中可能存在的类别不平衡(青光眼样本通常少于正常样本)。优化器使用 AdamW ,并配合 CosineAnnealingLR 学习率调度器。

Grad-CAM++的核心实现步骤

  1. 前向传播输入图像,获取模型对目标类别(青光眼)的预测分数。
  2. 对目标分数进行反向传播,计算最后一个卷积层特征图的梯度。
  3. 利用Grad-CAM++特有的权重计算方式,对梯度进行加权求和,得到每个特征通道的重要性权重。Grad-CAM++改进了标准Grad-CAM,能更好地处理图像中多个同类对象的情况,并更精确地定位。
  4. 计算最后一个卷积层激活值的加权和,然后通过 ReLU 过滤掉负贡献(因为负贡献通常意味着对当前类别有抑制),生成粗热力图。
  5. 将粗热力图进行上采样,使其尺寸与原始输入图像一致。
import numpy as np
import cv2

def generate_gradcampp(model, input_image, target_class_idx, device):
    model.eval()
    input_image = input_image.to(device).unsqueeze(0)
    input_image.requires_grad_()

    # 前向传播
    output = model(input_image)
    model.zero_grad()

    # 获取目标类别的分数
    score = output[:, target_class_idx]
    score.backward()

    # 获取梯度和激活值
    gradients = model.get_activations_gradient()
    activations = model.get_activations(input_image).detach()

    # Grad-CAM++ 权重计算
    gradients_pow_2 = gradients.pow(2)
    gradients_pow_3 = gradients_pow_2 * gradients
    # 全局和
    sum_activations = activations.sum(dim=(2,3), keepdim=True)
    eps = 1e-7
    aij = gradients_pow_2 / (2*gradients_pow_2 + sum_activations * gradients_pow_3 + eps)
    # 加权系数
    weights = (aij * gradients).sum(dim=(2,3), keepdim=True)

    # 生成CAM
    cam = (weights * activations).sum(dim=1, keepdim=True)
    cam = torch.nn.functional.relu(cam) # ReLU过滤
    cam = cam.squeeze().cpu().numpy()

    # 归一化并上采样
    cam = cv2.resize(cam, input_image.shape[2:][::-1]) # 调整到输入图像大小
    cam = (cam - cam.min()) / (cam.max() - cam.min() + eps) # 归一化到[0,1]
    return cam

3.3 知识对齐的定量化验证策略

这是将主观临床知识转化为客观指标的关键一步。

  1. 热力图后处理 :我们将归一化的热力图 cam 通过阈值(例如,激活值最高的前15%或20%的区域)二值化,得到模型的“决策区域” M_model
  2. 金标准区域 :使用医生的视盘标注掩膜 M_disc 作为基础区域。有时,我们更关注视杯区域 M_cup ,因为杯盘比是青光眼诊断的核心。
  3. 计算空间重叠指标
    • Dice相似系数 Dice = 2 * |M_model ∩ M_disc| / (|M_model| + |M_disc|) 。这个指标对区域内部填充的均匀性不敏感,更关注空间重叠程度,非常适合我们的场景。值越接近1,重叠越好。
    • 交并比 IoU = |M_model ∩ M_disc| / |M_model ∪ M_disc|
  4. 计算统计显著性 :我们不是简单看平均值。对于测试集所有样本,我们计算其 Dice 系数的分布。同时,我们构建一个 随机基线 :生成与模型热力图相同大小的随机噪声图,同样阈值化后计算与视盘区域的 Dice 。使用非参数检验(如Mann-Whitney U检验)来验证模型热力图与视盘区域的重叠是否显著优于随机噪声。
  5. 区域特异性分析 :我们定义了一个更精细的指标——“ 杯区聚焦比 ”。
    • 计算模型热力图在视杯区域 M_cup 内的平均激活值 Mean_cup
    • 计算热力图在整个视盘区域 M_disc 内的平均激活值 Mean_disc
    • 计算 Focus_Ratio = Mean_cup / Mean_disc
    • 如果模型真正学到了“杯盘比扩大”这个特征,那么 Focus_Ratio 应该显著大于1(即更关注杯区)。我们可以统计测试集上该比值大于1的样本比例。

通过这套组合指标,我们不仅能回答“模型关注的地方对不对”(Dice/IoU),还能初步回答“它关注的重点是否与临床重点一致”(杯区聚焦比)。

4. 结果分析与临床洞见:我们发现了什么?

经过对测试集上百张图像的分析,我们得到了一些有启发性的结果。

4.1 定量结果展示

我们用一个表格来清晰呈现主要定量指标:

评估指标 模型热力图 vs. 视盘区域 (均值±标准差) 随机噪声 vs. 视盘区域 (均值±标准差) P值 (Mann-Whitney U检验) 临床意义解读
Dice系数 0.62 ± 0.18 0.21 ± 0.09 < 0.001 模型决策区域与视盘解剖结构高度重合,显著优于随机猜测。
交并比 (IoU) 0.46 ± 0.17 0.12 ± 0.06 < 0.001 进一步确认了空间重叠的显著性。
杯区聚焦比 >1 的样本比例 78% (不适用) (不适用) 在大部分病例中,模型在视杯区域表现出更高的关注度,与“杯盘比”诊断依据相符。

从数据上看,模型的表现是令人鼓舞的。Dice系数达到0.62,意味着模型的热力图区域与医生标注的视盘区域有相当程度的重叠。更重要的是,这种重叠不是随机的,其显著性极强(P<0.001)。杯区聚焦比的结果进一步说明,模型并非均匀地关注整个视盘,而是有倾向性地聚焦于杯区——这个青光眼病理改变的核心部位。

4.2 典型案例分析:成功与“失败”的样本

案例一(成功对齐) : 一张典型的晚期青光眼眼底彩照,视杯明显扩大,盘沿变窄。模型生成的热力图清晰地、高强度地覆盖了整个视杯区域,并向颞侧盘沿延伸,与医生标注的病变区域几乎完美重合。临床医生反馈:“这个AI指出的地方,正是我第一眼就觉得有问题的地方。”

案例二(有趣的不一致) : 一张早期青光眼病例,杯盘比轻度增大,但更重要的特征是下方盘沿的“切迹”和伴随的视网膜神经纤维层缺损。模型的热力图主要聚焦于视杯中心,对下方的盘沿切迹区域激活较弱。然而,分类结果是正确的(青光眼)。我们与医生讨论后,得出一个关键洞见: 模型可能学习到了一种更“粗粒度”但有效的特征组合 。它可能将“整体杯盘比增大”与“视盘区域某些纹理/颜色分布模式”关联起来,而这些模式人类医生不一定能直观描述,但确实存在于数据中。这提示我们,模型的知识与人类知识是“对齐”而非“等同”。模型可能发现了人类视觉不易察觉的、统计意义上的关联特征。

案例三(模型“失误”带来的启发) : 一张高度近视眼的眼底彩照,伴有巨大的视盘和倾斜的视杯(近视弧)。模型将其误判为青光眼,热力图高亮区域集中在倾斜的颞侧。医生指出,这是典型的“近视性视盘改变”,容易与青光眼混淆。这个“失败”案例极具价值:它暴露了模型训练数据中可能缺乏足够多、标注清晰的高度近视非青光眼样本。同时,它也提示, 可解释性工具能帮助我们发现模型的“认知偏差”或数据集的“盲区” ,为下一步数据收集和模型优化提供了明确方向。

4.3 临床医生反馈与模型信任建立

我们将热力图可视化结果(原始图、热力图叠加图、医生标注对比图)制作成简单的评估界面,邀请5位未参与标注的眼科主治及以上医师进行盲评(不告知模型诊断结果)。评估维度包括:

  1. 热力图高亮区域是否与您认为的关键病变区域相符?(是/部分/否)
  2. 基于热力图,您对模型诊断该病例的信心是?(1-5分,5分最高)

收集到的反馈显示,超过85%的病例被医生认为热力图区域“相符”或“部分相符”。在热力图与医生判断高度一致的病例中,医生对模型诊断的信心评分平均达到4.2分。多位医生表示:“能看到它‘看’哪里,让我更愿意参考它的结果,尤其是在疑难病例上,可以作为一个交叉验证的视角。”

5. 踩坑实录与进阶思考

5.1 实操中遇到的典型问题与解决方案

  1. 热力图模糊、定位不准

    • 问题 :早期使用标准CAM时,热力图经常显得弥散,像一团雾覆盖大片区域,无法精确定位到视杯。
    • 排查 :检查模型最后一个卷积层的空间分辨率。如果分辨率过低(如7x7),上采样后自然会模糊。同时,检查梯度是否出现“饱和”或“消失”,特别是在使用 ReLU 激活函数的网络中,可能导致梯度信息弱。
    • 解决
      • 升级方法 :从CAM切换到 Grad-CAM++ ,其对梯度的加权方式能产生更集中、更定位准确的热力图。
      • 调整网络 :考虑使用保留更高空间分辨率的网络结构(如移除部分下采样层),或采用特征金字塔结构。
      • 尝试Guided Grad-CAM :将Grad-CAM热力图与导向反向传播的像素级梯度结合,能生成更清晰、像素级的显著性图,但计算量稍大。
  2. 热力图与金标准区域存在系统性偏移

    • 问题 :计算出的Dice系数始终不高,可视化发现热力图区域整体偏向视盘的某一侧。
    • 排查 :检查数据预处理(特别是中心裁剪)是否破坏了空间一致性。检查医生标注与图像是否严格对齐。
    • 解决 :确保在生成热力图和计算指标时,使用的是与原始训练/测试图像经过 完全相同预处理 (包括裁剪、缩放)后的坐标空间。所有标注掩膜必须随图像进行相同的空间变换。
  3. 定量指标好,但医生认为“没抓到重点”

    • 问题 :Dice系数不错,但医生反馈热力图虽然覆盖了视盘,却均匀激活,没有突出杯盘比这个核心。
    • 排查 :Dice系数只衡量空间重叠,不衡量激活强度分布。模型可能只是学会了“找到视盘”,但未深入理解其内部结构。
    • 解决 :引入更细粒度的评估指标,如前述的“杯区聚焦比”。同时,在训练阶段可以尝试 弱监督定位 的辅助任务,或在损失函数中加入鼓励模型关注小区域的约束(需谨慎,可能影响主分类任务性能)。

5.2 对CAM方法局限性的再认识

通过这个项目,我们更深刻地认识到CAM作为工具的边界:

  • 相关性而非因果性 :热力图展示的是模型决策与图像区域的相关性,而非因果性。一个区域被高亮,可能是因为它确实重要,也可能是因为它与真正重要的区域高度共现。
  • 对模型结构的依赖 :CAM系列严重依赖卷积层和全局池化层。对于Transformer-based模型(如ViT)或包含大量全连接层的网络,需要其他解释方法(如注意力 rollout)。
  • “沉默的证据”问题 :CAM只显示对当前决策有 正面贡献 的区域。那些对决策有 强烈抑制 作用的区域(例如,一个非常健康的盘沿特征可能让模型倾向于判断为正常)无法显示。这可能导致解释不完整。
  • 阈值选择的任意性 :将连续的热力图转为二值区域用于定量比较时,阈值的选择会影响指标数值。需要报告不同阈值下的结果,或使用类似AUC的曲线下面积指标来评估。

5.3 项目延伸:走向更严谨的临床验证

本次工作是一个良好的起点,但距离严格的临床验证还有距离。未来的工作可以沿着以下几个方向深入:

  1. 多中心、前瞻性验证 :在当前回顾性数据集上验证后,需要在独立的多中心、前瞻性收集的数据集上进行测试,以评估其泛化能力和临床实效。
  2. 与更多诊断标准对齐 :除了视盘/视杯,尝试与更广泛的青光眼体征对齐,如视网膜神经纤维层缺损的分布、盘沿出血、血管屈膝等。这需要更精细的多标签标注。
  3. 开发交互式解释工具 :将热力图生成集成到临床诊断软件中,允许医生点击热力图不同区域,询问模型“为什么这部分重要?”(结合反事实解释或概念激活向量等更高级的方法)。
  4. 探索因果解释 :结合因果推断的方法,尝试构建更接近临床因果路径的解释模型,而不仅仅是相关性图谱。

这个项目的最终价值,不在于证明某个模型多优秀,而在于展示了一种 用临床知识检验AI决策逻辑的方法论 。它让AI的“黑箱”打开了一条缝,让光透进去,也让医生的经验照进来。在医疗AI落地的漫漫长路上,可解释性不是锦上添花,而是建立信任、确保安全、实现人机协同的基石。每一次“知识对齐”的验证,都是向更可靠、更负责任的医疗AI迈进的一步。

更多推荐