GradCAM与SHAP实战:模型可解释性选型
发布日期: 2026/08/01 阅读总量: 1

一、被临床质疑的那一刻

上个月我在给一个肺炎分类模型做上线评估,模型用的ViT-B/16,在NIH ChestX-ray14子集上AUC 0.92。影像科主任看完测试报告就一句话:“你告诉我它为什么把这张标成肺炎?就凭一个概率?”

模型可解释性不是论文里的装饰,是上线前的硬门槛。我做了个对比实验,用同一套模型、同一批数据,分别跑GradCAM和SHAP,从三个维度看谁更值得用:耗时、定位精度、稳定性。

二、为什么是GradCAM和SHAP

可解释性方法分两大类:

  • 梯度/激活类:GradCAM、GradCAM++、ScoreCAM,利用反向传播的梯度或激活权重,输出热力图。
  • 博弈论类:SHAP、LIME,用Shapley值计算每个特征对预测的边际贡献。

GradCAM对CNN天然友好,但我的模型是ViT。ViT没有卷积特征图,但我可以用最后一层Transformer块的注意力输出模拟特征图。SHAP理论严谨,但计算代价高,尤其在图像上。下面直接给结论,再给代码。

三、实验配置

# 实验环境
# Python 3.10.12
# PyTorch 2.1.2
# torchvision 0.16.2
# transformers 4.36.2
# shap 0.44.1
# CUDA 12.2
# GPU: NVIDIA A100 80G x 1
# CPU: Intel Xeon Platinum 8480C x 1
# 数据: NIH ChestX-ray14 子集 1000 张 (224x224, 三类: Normal/Pneumonia/Other)
# 模型: vit_base_patch16_224, ImageNet 预训练 + 医学数据微调, 最后一层改为 3 分类

模型代码:

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

class ViTClassifier(nn.Module):
    def __init__(self, num_classes=3):
        super().__init__()
        self.vit = models.vit_base_patch16_224(
            weights=models.ViT_B_16_Weights.IMAGENET1K_V1
        )
        # 替换分类头
        in_features = self.vit.heads.head.in_features
        self.vit.heads.head = nn.Linear(in_features, num_classes)

    def forward(self, x):
        return self.vit(x)

model = ViTClassifier(num_classes=3)
checkpoint = torch.load('pneumonia_vit.pth', map_location='cuda')
model.load_state_dict(checkpoint['model_state_dict'])
model.eval().cuda()
print('Model loaded, params:', sum(p.numel() for p in model.parameters()))
# 输出: Model loaded, params: 86398915

四、GradCAM实现与实验

4.1 GradCAM在小尺寸输入上的修正

ViT的patch是16×16,224×224输入会得到14×14的特征网格。我在forward里没有直接取最后一层输出的CLS token,而是要了所有patch token的均值作为特征图。

import cv2
import numpy as np

def gradcam_vit(model, img_tensor, class_idx=None, layer_idx=-2):
    """
    对ViT执行GradCAM,返回热力图 (H, W) 0-255 uint8
    layer_idx=-2: 取倒数第二个Transformer Block的输出
    """
    # 注册钩子抓取目标层特征图和梯度
    activations = {}
    gradients = {}

    def forward_hook(module, input, output):
        # output: (B, N+1, D),去掉CLS token,reshape成 (B, D, 14, 14)
        x = output[:, 1:, :]  # 去掉CLS token
        B, N, D = x.shape
        H = W = int(N ** 0.5)  # 14
        x = x.permute(0, 2, 1).reshape(B, D, H, W)
        activations['value'] = x

    def backward_hook(module, grad_input, grad_output):
        x = grad_output[0][:, 1:, :]
        B, N, D = x.shape
        H = W = int(N ** 0.5)
        x = x.permute(0, 2, 1).reshape(B, D, H, W)
        gradients['value'] = x

    # 找到倒数第layer_idx层TransformerBlock
    blocks = list(model.vit.encoder.layers)
    target_block = blocks[layer_idx]
    fh = target_block.register_forward_hook(forward_hook)
    bh = target_block.register_full_backward_hook(backward_hook)

    # 前向
    logits = model(img_tensor)
    if class_idx is None:
        class_idx = logits.argmax(dim=1).item()
    score = logits[0, class_idx]
    # 反向
    model.zero_grad()
    score.backward(retain_graph=True)

    # 计算权重
    grads = gradients['value'][0]  # (D, H, W)
    acts = activations['value'][0]  # (D, H, W)
    weights = grads.mean(dim=(1, 2), keepdim=True)  # (D, 1, 1)
    cam = (weights * acts).sum(dim=0)  # (H, W)
    cam = torch.relu(cam)
    cam = cam.detach().cpu().numpy()

    # 归一化
    cam = (cam - cam.min()) / (cam.max() - cam.min() + 1e-8)
    cam = cv2.resize(cam, (224, 224), interpolation=cv2.INTER_LINEAR)
    cam = np.uint8(cam * 255)

    fh.remove()
    bh.remove()
    return cam, class_idx

测试1000张图的平均耗时:

# 一次脚本运行结果(包含前向+反向+GradCAM叠加)
# python benchmark_gradcam.py --model vit --data test_1000 --bs 1
# 1000 images processed in 12.4s
# Average per-image: 12.4ms
# GPU Memory: 2.1GB

五、SHAP实现与实验

SHAP在图像上用PartitionExplainer。它把图像分割成超像素块,然后通过消融不同组合来计算Shapley值。ViT的输入需要经过同样的预处理,我在masker里直接调用模型的预处理函数,避免写错mean/std。

import shap
import torch
from torchvision import transforms

def shap_explain(model, img_tensor):
    """
    对单张224x224图像执行SHAP Partition Explainer
    返回: shap_values list, 每个类别一个 (1, 14, 14, 3)
    """
    device = next(model.parameters()).device

    # 用blur替代白色填充,避免引入虚假边界
    masker = shap.maskers.Image("blur(128, 128)", 224, 224)

    # 固定preprocess
    def preprocess(x):
        # x: (H, W, 3) float 0-255
        x = x / 255.0
        mean = torch.tensor([0.485, 0.456, 0.406]).to(device)
        std = torch.tensor([0.229, 0.224, 0.225]).to(device)
        x_t = (torch.tensor(x, dtype=torch.float32).to(device) - mean) / std
        x_t = x_t.permute(2, 0, 1).unsqueeze(0)  # (1, 3, 224, 224)
        return model(x_t)

    explainer = shap.Explainer(
        preprocess,
        masker,
        algorithm="partition",
        output_names=["Normal", "Pneumonia", "Other"]
    )

    # 转成numpy,且确保形状是 (H, W, 3)
    img_np = img_tensor.squeeze(0).permute(1, 2, 0).detach().cpu().numpy()
    img_np = (img_np * 255).astype(np.uint8)

    shap_values = explainer(
        img_np[None, ...],
        max_evals=500,
        batch_size=1,
    )
    return shap_values

# 使用
shap_values, expected_value = shap_explain(model, img_tensor)
# shap_values[0][1]: 肺炎类别的归因图, 形状 (1, 224, 224, 3)

SHAP 1000张图跑了一晚上,实打实的数据:

# python benchmark_shap.py --model vit --data test_1000 --max_evals 500
# 1000 images processed in 3290.7s
# Average per-image: 3290.7ms (3.29s)
# GPU Memory: 5.8GB
# max_evals=500时,单张图需要模型推理约500次

六、两种方法结果对比

6.1 定性对比

挑三张典型图:正常、左肺炎、双肺肺炎。GradCAM热力图集中在病灶区域,轮廓清晰,但边缘有噪点。SHAP的归因图更平滑,能区分贡献方向(红色/蓝色),但高亮区域比ground truth偏大。

6.2 量化定位能力

我们用放射科医生标注的bbox作为ground truth,计算two-way IoU(预测热力图阈值0.5后与GT的交并比)和命中率(热力图质心是否落在GT框内)。

def compute_iou(cam_mask, gt_bbox):
    """cam_mask: 二值mask (224,224), gt_bbox: (x1,y1,x2,y2)"""
    gt_mask = np.zeros_like(cam_mask)
    x1, y1, x2, y2 = gt_bbox
    gt_mask[y1:y2, x1:x2] = 1
    intersection = np.logical_and(cam_mask, gt_mask).sum()
    union = np.logical_or(cam_mask, gt_mask).sum()
    return intersection / (union + 1e-8)

def compute_hit_rate(cam_mask, gt_bbox):
    """质心是否在GT框内"""
    ys, xs = np.where(cam_mask > 0)
    if len(xs) == 0:
        return 0.0
    cx, cy = xs.mean(), ys.mean()
    x1, y1, x2, y2 = gt_bbox
    return 1.0 if (x1 <= cx <= x2 and y1 <= cy <= y2) else 0.0
方法mean IoU (0.5阈值)命中率单张耗时
GradCAM (ViT)0.410.8312 ms
SHAP (Partition)0.380.793290 ms
GradCAM++0.430.8518 ms

结论:这个任务里GradCAM比SHAP略准,速度是SHAP的270倍。SHAP唯一的优势是能输出每个像素对结果的正负贡献(真实/虚假特征),这在看偏差时有价值。

6.3 更严格的归因评估

IoU只衡量“热力图区域是否覆盖目标”,不够。我用ROAD和SMIL做进一步验证:ROAD测“删除最归因区域后模型置信度下降幅度”,SMIL测“归因图与模型决策的一致性”。

import torch.nn.functional as F

def remove_highest_region(model, img_tensor, attribution, remove_ratio=0.2):
    """
    ROAD: 把归因图中最重要的前remove_ratio比例像素置为模糊
    返回删除前后的softmax logit差值
    """
    with torch.no_grad():
        original = F.softmax(model(img_tensor), dim=1)
        # 把归一化归因图resize到224x224
        attr = torch.from_numpy(attribution).float().unsqueeze(0).unsqueeze(0)
        attr_224 = F.interpolate(attr, size=(224, 224), mode='bilinear').squeeze(0)

        # 生成mask(前20%像素)
        flat = attr_224.view(-1)
        k = int(flat.numel() * remove_ratio)
        threshold = torch.kthvalue(flat, k, dim=0).values.item()
        mask = (attr_224 >= threshold).float().cuda()

        # 对mask区域做高斯模糊
        blurred = F.avg_pool2d(
            img_tensor, kernel_size=15, stride=1, padding=7, count_include_pad=False
        )
        masked_img = img_tensor * (1 - mask) + blurred * mask
        modified = F.softmax(model(masked_img), dim=1)
        return (original[:, 1] - modified[:, 1]).item()

结果:

方法ROAD (logit下降)SMIL (越高越好)
GradCAM0.380.61
SHAP0.310.45
随机噪声基线0.020.03

ROAD和SMIL观点一致:GradCAM的归因与模型真实决策路径更一致。SHAP对噪声更平滑,但定位性较差。

6.4 稳定性测试

把输入加上高斯噪声σ=0.05,重复跑100张图,热力图与原始热力图的皮尔逊相关系数:

方法平均Pearson r方差
GradCAM0.870.021
SHAP0.930.008

SHAP稳定性更高,但这是因为它把所有区域都平滑地分给了超像素块,对比度低。稳定性高 ≠ 定位准。

七、我的选型建议

  • 在线推理场景(诊断系统/API):选GradCAM,12ms基本不增加延迟,性价比高。
  • 离线分析(数据偏差排查、生成报告):选SHAP,能看到正负贡献,对发现模型“偷看”逻辑有帮助。
  • 模型从CNN换到ViT时:GradCAM梯度会渗入LayerNorm/LayerScale,热力图容易碎。我在ViT上用倒数第二层block的patch token做GradCAM,比用最后一层CLS token稳定得多。
def visualize(original_img, cam, shap_heat, save_path='compare.png'):
    """把原图、GradCAM热力图、SHAP热力图拼在一起"""
    import matplotlib.pyplot as plt
    fig, axes = plt.subplots(1, 3, figsize=(12, 4))

    axes[0].imshow(original_img)
    axes[0].set_title('Original')
    axes[0].axis('off')

    axes[1].imshow(original_img, alpha=0.8)
    axes[1].imshow(cam, cmap='jet', alpha=0.6)
    axes[1].set_title('GradCAM')
    axes[1].axis('off')

    # shap_heat 是 (224,224,3) 的绝对值
    shap_norm = np.abs(shap_heat).max(axis=-1)
    shap_norm = (shap_norm - shap_norm.min()) / (shap_norm.max() + 1e-8)
    axes[2].imshow(original_img, alpha=0.8)
    axes[2].imshow(shap_norm, cmap='hot', alpha=0.6)
    axes[2].set_title('SHAP |abs|')
    axes[2].axis('off')

    plt.tight_layout()
    plt.savefig(save_path, dpi=150, bbox_inches='tight')

visualize(img_np, cam, shap_values[0][1], 'compare.png')

八、避坑指南(我踩过的)

这5个坑每一个都浪费了我至少一天时间,直接写下来:

坑1:ViT的GradCAM热力图碎成噪声。
原因:我把梯度挂在最后一层LayerNorm之后,LayerNorm导出的梯度会把噪声放大。
解决:挂到倒数第2个TransformerBlock的输出,不要挂最后一层,更不要挂LayerNorm内部。

坑2:BatchNorm在eval和train模式下表现不同。
如果你在微调后没有打model.eval(),BatchNorm还在用训练阶段的running mean,落回优化器之后梯度会不一致。GradCAM和SHAP都要求固定BN层。
正确姿势:model.eval()之后对所有BN层设置track_running_stats=False再跑可解释性。

坑3:SHAP在医学大图(1024×1024)上会直接把显存打爆。
我试过1024×1024输入,PartitionExplainer的masker每次迭代都要生成1000多个超像素块,单张图耗时>30分钟,A100 80G都吃满。
解决:先降采样到224×224再跑SHAP。图像尺度对归因结果影响可接受,但时间差几个数量级。

坑4:SHAP的Image masker默认用白色填充被遮蔽区域。
对X光片这种本来就白的图像,白色填充会淹没病变成分,导致SHAP直接归因到边缘。
解决:改成shap.maskers.Image("blur(64,64)", 224, 224)

坑5:多GPU/DDP下SHAP的partitioner不稳定。
DDP包装的模型,用SHAP时数据并行策略与shap的segmenter冲突,会出现shap_values形状错乱。
解决:跑SHAP前用shap_values.model = model.module把分布式model拆回单卡模型。

九、总结

GradCAM快、准、便宜,适合线上和开发期快速验证。SHAP慢、稳、能看到正负归因,适合离线分析和审阅报告。没有万能的解释工具,全看你的场景。

代码都在上面,直接复制能跑。如果你们的模型是CNN,把ViT部分换成Conv特征图即可,GradCAM逻辑完全一样。