GradCAM与SHAP实战:让模型决策不再黑盒
发布日期: 2026/08/19 阅读总量: 1

被质检主管当场问住的下午三点

2024年3月,我把训练好的ResNet50缺陷分类模型接到产线MES系统上,准确率97.2%,误判率比原来用人工视觉的老方案低了一半。我正等着拿季度奖,结果质检主管老张拎着一块PCB板走进办公室:「小王,这块板子Model说是有划痕,但老师傅们看了半天都没找到划痕在哪儿。你这AI到底凭什么这么判?」

我打开推理脚本,输出一个概率值。

「0.967,置信度很高。」我说。

老张盯着屏幕上孤零零的数字:「所以呢?它为什么觉得是划痕?特征是什么?」

我答不上来。那一刻我意识到,光给一个准确率,生产线上没人会信你的模型。

后面我花了三周,把GradCAM和SHAP两条路都走了一遍,整理出这篇可以直接抄走的实战笔记。如果你也负责给业务方或客户解释模型行为,这篇文章能让你少掉几根头发。

一、问题界定:你要回答哪一类「为什么」

「为什么模型这么判」其实有至少三种不同的问法:

  • 空间维度的为什么:图像里的哪个区域让模型做出了判断?→ 用GradCAM这类基于梯度的热力图
  • 特征维度的为什么:是哪个特征、以什么方向、贡献了多少分数?→ 用SHAP这类基于博弈论的归因方法
  • 行为维度的为什么:哪些训练样本决定了这个判断?→ 用Influence Function或训练数据归因,不在本文范围

老张问的是第一类:缺陷在图像里的什么位置。但如果模型做的判断是「这张信用卡申请该不该批」,业务方问的会是第二类:「是因为收入低还是负债高?」我把两类场景的关键知识放在一篇文章里讲清楚,避免你重复踩坑。

二、方案概览:GradCAM和SHAP各自能干什么

这两个方法都被叫作「模型可解释性」,但底层原理和使用场景完全不同。选错工具,轻则答非所问,重则得出错误结论。

维度GradCAMSHAP
适用模型CNN(图像分类、目标检测)表格模型、树模型、深度学习均可
输入类型图像Tensor表格数据向量
输出形式与输入同尺寸的热力图特征对应的贡献值(可正可负)
原理梯度加权特征图Shapley值(合作博弈论)
计算成本一次反向传播,几乎免费指数级复杂度,靠近似采样
代码复杂度约60行(手写hook)调库一行出结果
最大的坑对错误类别的梯度会被软最大化掩盖树模型和深度模型的解释差异极大

总结一句话:输入是图像,用GradCAM;输入是表格,用SHAP。但如果你要工程化落地,光记住这句话不够,往下看。

三、手写GradCAM:从PyTorch hook到热力图

先说GradCAM的核心思想。

一句话版本:把最终分类得分对最后一个卷积层的特征图求梯度,把梯度做全局平均池化得到每个通道的权重,再对特征图做加权求和,最后过一个ReLU得到热力图。

如果你理解这句话,代码就很容易写了。不理解也没关系,先抄代码,跑通了再倒回来看公式。

版本说明:本文全部代码基于以下环境,跑不通先检查版本:Python 3.10.12 / PyTorch 2.2.0 / torchvision 0.17.0 / SHAP 0.44.0 / OpenCV 4.8.1 / numpy 1.24.3。GPU为单张NVIDIA GeForce RTX 3090。

3.1 完整可运行的GradCAM实现

"""
gradcam_demo.py
环境:Python 3.10.12 / PyTorch 2.2.0 / torchvision 0.17.0 / opencv-python 4.8.1
用法:python gradcam_demo.py --image ./defect_sample.jpg --class_idx 1
"""
import argparse
import cv2
import numpy as np
import torch
import torch.nn.functional as F
from torchvision import models, transforms
from PIL import Image

class GradCAM:
    """手写GradCAM,不依赖第三方库,方便魔改"""
    def __init__(self, model, target_layer):
        self.model = model.eval()
        self.target_layer = target_layer
        self.gradients = None
        self.activations = None
        # 注册hook抓取目标层的激活值和梯度
        target_layer.register_forward_hook(self._forward_hook)
        target_layer.register_full_backward_hook(self._backward_hook)

    def _forward_hook(self, module, input, output):
        self.activations = output.detach()

    def _backward_hook(self, module, grad_input, grad_output):
        self.gradients = grad_output[0].detach()

    def generate(self, input_tensor, class_idx=None):
        """生成热力图"""
        output = self.model(input_tensor)
        if class_idx is None:
            class_idx = torch.argmax(output, dim=1).item()
        # 选定类别的预测分数
        score = output[0, class_idx]
        self.model.zero_grad()
        score.backward()

        # 梯度全局平均池化 -> 每个通道一个权重
        weights = torch.mean(self.gradients, dim=(2, 3), keepdim=True)  # [1, C, 1, 1]
        # 特征图加权求和
        cam = torch.sum(weights * self.activations, dim=1, keepdim=True)  # [1, 1, H, W]
        cam = F.relu(cam)
        # 归一化到0~1
        cam = cam.squeeze().cpu().numpy()
        cam = (cam - cam.min()) / (cam.max() - cam.min() + 1e-8)
        return cam, class_idx


def load_image(path, size=224):
    """读图并预处理,返回tensor和原始BGR图"""
    img_bgr = cv2.imread(path)
    img_rgb = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB)
    img_pil = Image.fromarray(img_rgb)
    transform = transforms.Compose([
        transforms.Resize((size, size)),
        transforms.ToTensor(),
        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
    ])
    input_tensor = transform(img_pil).unsqueeze(0)
    return input_tensor, img_bgr


def overlay_heatmap(cam, img_bgr, alpha=0.5):
    """将热力图叠加到原图上,opencv颜色工具有时会报错,统一用colormap"""
    cam_resized = cv2.resize(cam, (img_bgr.shape[1], img_bgr.shape[0]))
    # 这里用2.0是最大强度,实际场景建议1.5,太亮会覆盖原图细节
    heatmap = np.uint8(255 * cam_resized)
    heatmap = cv2.applyColorMap(heatmap, cv2.COLORMAP_JET)
    overlay = cv2.addWeighted(heatmap, alpha, img_bgr, 1 - alpha, 0)
    return overlay


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--image", type=str, required=True, help="输入图像路径")
    parser.add_argument("--class_idx", type=int, default=None, help="目标类别索引,默认取模型预测结果")
    args = parser.parse_args()

    # 加载预训练模型,最后一层卷积是layer4.2
    model = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V2)
    target_layer = model.layer4[2]
    cam_engine = GradCAM(model, target_layer)

    input_tensor, img_bgr = load_image(args.image)
    cam, pred_idx = cam_engine.generate(input_tensor, args.class_idx)
    overlay = overlay_heatmap(cam, img_bgr)

    # 保存结果而不是弹窗,方便在服务器上跑
    cv2.imwrite("gradcam_result.jpg", overlay)
    print(f"预测类别索引: {pred_idx}, 热力图已保存为 gradcam_result.jpg")


if __name__ == "__main__":
    main()

运行这个脚本:

# 先用任意一张测试图跑通流程
python gradcam_demo.py --image ./defect_sample.jpg

# 指定看第1类(例如"划痕")的激活区域
python gradcam_demo.py --image ./defect_sample.jpg --class_idx 1

我拿产线上1024×1024的PCB板缺陷图跑了一遍,每张图从加载到生成热力图平均耗时47.3ms(不含模型加载),其中前向传播18ms、反向传播21ms、后处理8ms。相比单独推理的16ms,GradCAM只增加了约30ms的耗时,完全可以放到实时检测的pipeline里。

3.2 为什么是layer4而不是layer3

ResNet50有四个残差阶段,layer4输出的特征图分辨率最低(7×7),语义最强。很多人一上来就用layer4,然后发现热力图特别粗糙,像打码一样,就以为是方法不行。

实际上,层越深,定位越粗但语义越准;层越浅,定位越细但语义越杂。我用同样的缺陷图跑了layer3和layer4的对比:layer3的热力图能大致框出划痕区域但带了很多背景噪声,layer4的热力图边界比较模糊但定位中心准确。

建议工程上同时输出两张图:一张用layer4做主判断,一张用layer3做辅助参考。或者直接用下面要讲到的GradCAM++做细化。

四、SHAP实战:表格数据的特征归因

GradCAM解决了「缺陷在哪」,但另一类问题它解决不了:表格模型。比如信贷审批模型判断「拒绝」,业务方要的是「因为收入低还是负债率高」——图像上没有坐标这个概念。

SHAP的全称是SHapley Additive exPlanations,核心是Shapley值,来自合作博弈论。每个特征被看作一个「玩家」,模型预测结果就是「总收益」,Shapley值回答的是:每个玩家对总收益的边际贡献是多少。

4.1 直接调库的SHAP样例

SHAP库封装得很好,表格数据基本上五行代码出结果。但「五行出结果」不等于「结果是对的」,后文避坑段落专门讲这个。先给你一段能直接跑的代码:

"""
shap_demo.py
环境:Python 3.10.12 / XGBoost 2.0.3 / SHAP 0.44.0
数据:UCI German Credit(可替换成自己的CSV)
"""
import pandas as pd
import numpy as np
import xgboost as xgb
import shap
from sklearn.model_selection import train_test_split
from sklearn.preprocessing import LabelEncoder

# 1. 加载数据(如果你的特征不是数值型,先做编码)
df = pd.read_csv("german_credit.csv")
# 假设最后一列是标签(0/1)
X = df.iloc[:, :-1]
y = df.iloc[:, -1]

# 2. 对非数值列做标签编码
for col in X.select_dtypes(include=["object"]).columns:
    X[col] = LabelEncoder().fit_transform(X[col].astype(str))

# 3. 训练XGBoost模型
X_train, X_test, y_train, y_test = train_test_split(
    X, y, test_size=0.2, random_state=42, stratify=y
)
model = xgb.XGBClassifier(
    n_estimators=300,
    max_depth=6,
    learning_rate=0.05,
    subsample=0.8,
    colsample_bytree=0.8,
    eval_metric="logloss",
    tree_method="hist",   # 统一用hist,避免exact在Python 3.10下出问题
    random_state=42
)
model.fit(X_train, y_train)

# 4. 计算SHAP值
explainer = shap.TreeExplainer(model)
# 这里只对测试集前100行做解释,全量500行会慢很多
shap_values = explainer.shap_values(X_test[:100])

# 5. 两个最常用的可视化
# 5.1 summary_plot:全局特征重要性+方向
shap.summary_plot(shap_values, X_test[:100], show=False)
import matplotlib.pyplot as plt
plt.savefig("shap_summary.png", dpi=150, bbox_inches="tight")
plt.close()

# 5.2 force_plot:单个样本的预测分解
shap.initjs()
force_plot = shap.force_plot(
    explainer.expected_value,
    shap_values[0, :],
    X_test.iloc[0, :],
    matplotlib=True,
    show=False
)
plt.savefig("shap_force_sample0.png", dpi=150, bbox_inches="tight")
plt.close()

print("done:shap_summary.png 和 shap_force_sample0.png 已生成")

4.2 SHAP的两种重要输出

跑完上面代码,你手上有两个东西:

  • summary_plot:全局视角,横轴是SHAP值(正负表示推动/抑制预测),纵轴是特征,颜色代表特征值高低。一眼看出「哪个特征最重要、影响方向是什么」
  • force_plot:个体视角,把单个样本的基准预测值(expected_value)一步步分解到每个特征头上。信贷员拿去跟客户解释「为什么拒贷」特别好用

但注意样本量。我在真实的信贷数据集(2万条)上对测试集5000条全量计算SHAP值,耗时如下表:

样本数TreeExplainer耗时内存占用
1000.84s38MB
10008.2s312MB
500041.5s1.6GB

看到没,基本是线性增长,但内存翻得很快。SHAP值本身是个shape为[样本数, 特征数]的矩阵,特征多、样本多的话内存是主要瓶颈。我自己在16GB内存的服务器上跑到8000条就接近上限了。

如果你只想对单条样本做解释,我建议用shap.Explanation对象做增量计算,而不是一次性算全量。但SHAP库目前没有真正的增量API,所以实践中更常见的做法是只对抽样的100-500条做解释。这个细节后面「避坑」段落还要讲。

五、实战案例复盘:工业PCB缺陷分类

回到老张的PCB板。我用GradCAM跑通后,最直接的价值是:让老师傅确认了模型看的地方是对的。

流程是这样:

# 1. 离线批量跑1000张带标注的缺陷图,生成热力图
python gradcam_demo.py --image ./defect_0001.jpg --class_idx 1   # 划痕
python gradcam_demo.py --image ./defect_0002.jpg --class_idx 2   # 脏污
...

# 2. 用脚本把热力图和原图拼成对比图
python batch_visualize.py \
    --input_dir ./defect_images \
    --output_dir ./heatmap_output \
    --model_path ./resnet50_finetuned.pth
"""
batch_visualize.py
批量生成GradCAM热力图对比图,输出到指定目录
"""
import os
import cv2
import torch
from torchvision import models, transforms
from PIL import Image
import numpy as np
from gradcam_demo import GradCAM, load_image, overlay_heatmap

def main():
    input_dir = "./defect_images"
    output_dir = "./heatmap_output"
    os.makedirs(output_dir, exist_ok=True)

    # 加载微调后的模型(结构一样,权重不同)
    model = models.resnet50(weights=None)
    model.fc = torch.nn.Linear(2048, 4)  # 4类缺陷
    model.load_state_dict(torch.load("./resnet50_finetuned.pth"))
    model.eval()
    cam_engine = GradCAM(model, model.layer4[2])

    total_time = 0.0
    count = 0
    for fname in sorted(os.listdir(input_dir)):
        if not fname.lower().endswith((".jpg", ".png", ".jpeg")):
            continue
        path = os.path.join(input_dir, fname)
        input_tensor, img_bgr = load_image(path)
        cam, pred_idx = cam_engine.generate(input_tensor)
        overlay = overlay_heatmap(cam, img_bgr, alpha=0.5)
        out_path = os.path.join(output_dir, f"cam_{fname}")
        cv2.imwrite(out_path, overlay)
        # 统计耗时
        # 这里需要再计时一次,简化起见省略
        count += 1

    print(f"处理完成:{count}张,结果保存于{output_dir}")

if __name__ == "__main__":
    main()

第三步是我觉得最出彩的一步:让老师傅们在1000张热力图上做盲评,判断模型关注的区域是不是人类认为的缺陷区域。结果如下:

评估方式数量结论
模型热力图中心点在缺陷标注框内897/100089.7%定位准确
老师傅认为「模型看错地方」的图63/10006.3%疑似误判
老师傅说不清楚「模型看哪里」40/10004.0%无法判断

这个结果直接改变了老张的态度。从「你这AI不靠谱」变成「哦,它确实是看划痕附近的纹理异常,可以接受」。

顺便说一句,我之所以敢做盲评,是因为我先做了约束:只统计「热力图高亮区域的几何中心」是否落在标注框内。这个指标叫Pointing Game Accuracy,是GradCAM类方法论文里常用的定量评估指标。如果你要给客户汇报,别只拿几张好看的图说「你看挺准的」,要拿这个指标说话。

六、GradCAM与SHAP结合:多模态解释方案

一个容易被忽略的需求是:模型输入既有图像又有结构化特征。我在做一个设备预测性维护项目时,输入是振动信号的频谱图加上设备温度、转速等8个表格特征。业务方问的问题既有「哪个时间段的频谱异常」又有「是温度高还是转速异常导致的」。

这时候需要GradCAM和SHAP同时上。我的做法是:

{
  "解释方案": {
    "图像分支": {
      "方法": "GradCAM",
      "目标层": "cnn_backbone.layer4",
      "输出": "频谱热力图",
      "回答": "哪个频率区间异常"
    },
    "表格分支": {
      "方法": "SHAP TreeExplainer",
      "特征": ["温度", "转速", "负载", "湿度", "振动RMS", "轴承温度", "油压", "运行时长"],
      "输出": "force_plot和summary_plot",
      "回答": "哪个物理量贡献最大"
    },
    "融合层": {
      "策略": "分别解释,不强行融合",
      "原因": "融合层特征语义不可控,解释反而失真"
    }
  }
}

这个方案让我在客户现场扛住了一个多小时的技术质询,对方从「你们这模型就是个黑盒」变成「那我们能不能把规则引擎换成你们的模型」。方案文档里就两张图:一张热力图叠加原图,一张SHAP的summary plot。

七、效果数据:两种方案的量化对比

口说无凭,把耗时、一致性和落地效果都放出来:

# 压测:同一台机器(RTX 3090 / i9-12900K / 32GB RAM)
# 输入:批量PCB缺陷图1024x1024,1000张,batch_size=32

# GradCAM(含前向+反向+后处理)
python benchmark_gradcam.py --num_images 1000 --batch_size 32
# 结果:平均单张耗时 47.3ms,吞吐 21.1 fps,GradCAM额外的反向传播占21ms

# SHAP(XGBoost,500个测试样本,20个特征)
python benchmark_shap.py --num_samples 500 --num_features 20
# 结果:TreeExplainer总耗时 3.95s,平均单样本 7.9ms,内存峰值 512MB

解释性的一致性量化上,我用了一个指标:针对模型错误预测的样本,看GradCAM热力图是否还能指向正确特征。在1000张测试图上,我统计了两个数字:

  • 分类正确样本的热力图定位准确率:93.1%
  • 分类错误样本的热力图定位准确率:61.7%

也就是说,模型判错的时候,它看的地方往往也是错的。这个结论反过来验证了GradCAM确实是模型决策的真实依据,而不是生成了一张「看起来合理」的图。

八、避坑指南(血泪总结)

以下每一条都是我自己踩过的坑,花了一周时间填平。你直接拿走。

坑1:GradCAM在batch_size>1时梯度会累积

如果一次输入多张图,然后对其中一张图的score做backward,其他图的梯度会累积到同一个特征图上,导致热力图错乱。最初我的pipeline是batch推理,出了好几张完全错位的热力图,排查了三个小时。解决方式有两个:

# 方法一:保证batch_size=1(最省事)
input_tensor = input_tensor.unsqueeze(0)  # [1, 3, H, W]

# 方法二:batch模式下只对目标样本的梯度做mask
# (不推荐,实现麻烦且容易出边界问题)

坑2:SHAP对LightGBM的feature_importance和SHAP重要性结论不一致

LightGBM自带的feature_importance有两种:split(按分裂次数)和gain(按信息增益)。我在一个客户数据集上发现,模型自带的gain重要性排名第一的特征是「收入」,但SHAP重要性排名第一的是「负债率」。客户当场质疑我的解释结论。

原因在于:gain是把所有分裂点的增益加起来,不考虑方向;SHAP是边际贡献,考虑的是对最终预测的净影响。两个指标回答的问题不同,没有对错。正确姿势是,汇报时两个指标同时呈现,说清楚各自含义,避免被挑战。

坑3:TabNet等深度表格模型不能直接用TreeExplainer

SHAP的TreeExplainer只能用于树模型(XGBoost/LightGBM/CatBoost)。如果你用的是TabNet或MLP,得改用DeepExplainer或GradientExplainer,而这两个的速度慢一个数量级。我之前图省事直接拿TreeExplainer去解释TabNet,跑完结果数值错得离谱但程序不报错。后来查了文档才发现问题。

坑4:GradCAM热力图下采样太严重

ResNet50的layer4输出是7×7,直接上采样回原图后,热力图边缘非常粗糙。我在PCB小缺陷(只有30×30像素)上几乎看不出定位效果。解法是:

# 在hook里手动插值放大特征图到输入尺寸,再算加权
cam = F.interpolate(cam, size=(input_h, input_w), mode="bilinear", align_corners=False)

插值到原图尺寸后再加权,定位边缘细腻很多。但这会让热力图有棋盘格效应,后续可以用高斯模糊平滑一下。实际我用kernel size=7的GaussianBlur效果最好。

坑5:SHAP的expected_value在不同版本间会变

SHAP 0.42.0从0.41.0升级后,TreeExplainer的expected_value从标量变成了数组(多类模型)。升级后我force_plot的基线偏移了0.15,生成的所有解释结果全变了。后来排查发现是shap包的接口变更,不是模型变了。建议项目里锁定shap版本,或者每次在requirements.txt里写明版本号。

九、最终建议:先想清楚你要回答哪种「为什么」

最后把思路整理一遍,免得你学了一堆细节,到现场不知道怎么选:

  • 输入是图像,回答「模型看的是图像哪个区域」→ GradCAM,简单、快、效果好
  • 输入是表格,回答「哪个特征贡献最大」→ SHAP TreeExplainer,树模型首选
  • 输入是图像+表格的多模态 → 分别用两种方法各自解释,不强行融合
  • 千万别上来就两个都学,先看你的业务方到底在问什么

GradCAM和SHAP不是竞争关系,它们是解决不同问题的两把钥匙。你的职责是判断门锁类型,然后选对钥匙。希望这篇笔记能帮你少走我走过的弯路。