热力图:网络到底学了个啥?

1 热力图作用

将传统机器学习的方法与深度学习做对比,深度学习方法的显著特点就是不用再人工设计特征,特征提取的过程由网络通过训练自动完成。然而深度学习的方法不仅非常依赖于海量的标记数据,并且至今其至今都有一个令人诟病的缺点:

那就是可解释性不足。

早些年的深度学习的文章可能只要有创新并且能够在公开数据集涨点,就能发出非常不错的论文,比如CVPR这种顶级的论文。

但是现如今发论文也是越来越卷,并且一味刷SOTA来发论文并不可靠,现在的论文审稿人评价一篇论文的好坏也不再局限于精度,更多的是看模型是否轻量,并且还要求你讲清楚你的模型为什么针对这个任务能够work。

而热力图就是模型可解释性的一大应用。它能够清晰的展现出模型究竟是依靠哪些特征来对样本做出诸如分类、检测和分割等任务,从而证明模型并不是瞎猫碰上死耗子,在胡乱猜测。

因此如果大家是在校学生,准备发论文,正在发愁工作量不充足,或者模型的可解释性不强,不知道怎么讲好论文的故事,不妨从热力图下手,来针对模型学习到的特征,一探究竟。

从上面图就可以看出来,在进行图像分类的时候,热力图的焦点集中于最关键的部分。比如在识别狗的时候,热力图在狗的部分,暖色调非常深,在识别猫的时候,热力图在猫的部分色调非常深。从中也可以窥探出为啥叫热力图,越关键的部分,图像的部分越"热"。

2 LeNet识别Mnist

下面我给出一个LeNet识别Mnist数据集,并且生成热力图的示范代码。

在将LeNet在Mnist上训练,并且生成对应的模型权重文件之后,我们在该权重文件基础上,对测试图像进行识别,并且在识别之后生成热力图,来观察模型究竟是集中于哪些图像特征来得到最终的识别结果的,从而使得我们的方法更加具有说服力。

具体步骤如下:

  • 捕获特征图

    #注册钩子来捕获特征图
    feature_maps = []
    def hook_fn(module, input, output):
    feature_maps.append(output.detach())
    #注册到最后一个卷积层
    hook = net.conv2.register_forward_hook(hook_fn)

其中钩子机制 :其实是使用 PyTorch 的 register_forward_hook 方法在模型的特定层(这里是最后一个卷积层 conv2 )注册一个回调函数,当模型进行前向传播时,钩子函数会捕获该层的输出特征图并存储起来。

并且我们选择最后一个卷积层 ,是因为最后一个卷积层的特征图包含了最抽象、最具判别力的特征,能更好地反映模型的关注点。

  • 特征图处理

    #获取最后一个卷积层的特征图
    last_conv_features = feature_maps[-1]
    #计算特征图的平均值作为热力图
    heatmap = torch.mean(last_conv_features, dim=1).squeeze()

主要做通道平均 :对捕获的特征图沿通道维度取平均值,得到单通道的热力图。每个通道的特征图对应不同的视觉模式,因此通道平均值能综合反映所有通道的激活情况。

  • 归一化热力图

    #将热力图调整到与输入图片相同的尺寸(28x28)
    heatmap = torch.nn.functional.interpolate(
    heatmap.unsqueeze(0).unsqueeze(0),
    size=(28, 28),
    mode='bilinear',
    align_corners=False
    )
    heatmap = heatmap.squeeze().cpu().numpy()
    heatmap = (heatmap - np.min(heatmap)) / (np.max(heatmap) - np.min(heatmap))

具体使用方法:

复制代码
python test.py --image 4.png --heatmap

比如我们测试的图片是:

最终得到结果:

可以看模型的确是将注意力集中于数字的特征上。

以上方法的核心原理主要在于卷积层对输入图像在特征提取时,针对激活值高的部分会形成可视化映射,从蓝到红表示激活值的高低,再将热力图叠加在原始图像上。

那么为什么这种方法有效呢?

  • 由卷积神经网络的性质决定 :卷积层可以通过局部感受野和权值共享,能够捕获图像中的局部特征。

  • 采用最后一层卷积层 :最后一个卷积层的特征图包含了最抽象、最具判别力的特征。

  • 通道平均的作用 :不同通道学习不同的特征,平均操作可以综合所有通道的信息,得到的信息更加全面。

3 其他热力图绘制方法

此外其实还有其他热力图的绘制方法,包括GradCAM、HiResCAM、ScoreCAM、GradCAMPlusPlus、AblationCAM、XGradCAM、LayerCAM、FullGrad、EigenCAM、ShapleyCAM 和 FinerCAM等。大家可以参考一下链接:

https://gitcode.com/gh_mirrors/py/pytorch-grad-cam

此代码库给出的例子非常全面,包括支持分类、目标检测、语义分割、嵌入相似度等。

如视觉方面:

语义检测和分割:

3D医学语义分割:

大家可以下载代码可以试试,代码的运行方法很简单:

复制代码
python cam.py --image-path ./dog.png --method gradcam --output-dir ./

亲测可用,大家可以试一试。

完整代码大家可以关注gzh:阿龙AI日记,回复**项目代码,**找到压缩包:LeNet_mnist.zip和pytyorch-grad-cam-master.zip。

相关推荐
hans汉斯17 分钟前
采煤工作面隐患目标检测中YOLO11与Faster R-CNN的对比研究
人工智能·目标检测·计算机视觉·cnn·信息与通信·信号处理
bulingg37 分钟前
bert输入长度有限,如何处理超长文本?
人工智能·深度学习·bert
菜冻鱼2 小时前
Python-pytorch-数据加载
开发语言·人工智能·pytorch·python·深度学习·机器学习
machnerrn3 小时前
智慧交通系列(一)-十字路口车辆闯红灯检测告警抓拍系统(附含数据+源码+模型)
人工智能·python·深度学习
阿瑞斯官方账号3 小时前
多台 CIS 拼接检测现场售后实录:条纹、暗带、拼接错位怎么处理
数码相机·计算机视觉·视觉检测
小O的算法实验室4 小时前
IEEE TCDS,基于改进双神经网络三维未知环境多机器人协同区域覆盖搜索
人工智能·神经网络·机器人
大数据点灯人4 小时前
【大模型】深度解答:OOV 与 Tokenizer 词汇表共用问题
人工智能·深度学习·ai·大模型·transformer
vivo互联网技术4 小时前
Octopus:基于无历史数据的梯度正交化的学习框架|CVPR 2026
深度学习·计算机视觉·llm
牧羊人.3335 小时前
计算机视觉基础 第 9 章|实战:银行卡号识别
图像处理·人工智能·opencv·计算机视觉·图搜索算法
eagle_Annie6 小时前
Python/torch/深度学习——Miniforge+uv环境安装
python·深度学习·uv