一、被临床质疑的那一刻
上个月我在给一个肺炎分类模型做上线评估,模型用的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.41 | 0.83 | 12 ms |
| SHAP (Partition) | 0.38 | 0.79 | 3290 ms |
| GradCAM++ | 0.43 | 0.85 | 18 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 (越高越好) |
|---|---|---|
| GradCAM | 0.38 | 0.61 |
| SHAP | 0.31 | 0.45 |
| 随机噪声基线 | 0.02 | 0.03 |
ROAD和SMIL观点一致:GradCAM的归因与模型真实决策路径更一致。SHAP对噪声更平滑,但定位性较差。
6.4 稳定性测试
把输入加上高斯噪声σ=0.05,重复跑100张图,热力图与原始热力图的皮尔逊相关系数:
| 方法 | 平均Pearson r | 方差 |
|---|---|---|
| GradCAM | 0.87 | 0.021 |
| SHAP | 0.93 | 0.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逻辑完全一样。