ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

ViT可解释性:注意力图与CAM如何互补

2026/9/15 13:07:24 拓冰建站 浏览量
ViT可解释性:注意力图与CAM如何互补 做模型解释的时候我最常被问到的一个问题就是“ViT 不是自带注意力图吗为什么还要折腾 CAM” 这个问题背后其实是很多同学对 Vision Transformer 可解释性工具链的误解。注意力图能告诉我们模型“看了哪里”但往往说不清“为什么看这里”CAM 能回答类别层面的“哪些像素支撑了这个判断”但用在 ViT 上又需要绕几个弯。把两者放在一起用互相印证才是 ViT 可解释性的完整姿势。这篇博文就围绕 CAM 和注意力图这两条主线讲清楚它们各自的原理、在 ViT 里的适用场景、具体怎么跑通、踩过哪些坑。不管你是刚入手 ViT 的新手还是已经在做模型可视化、故障分析、论文配图的老手希望能帮你省下几天的摸索时间。1. 为什么 ViT 需要专门聊可解释性1.1 CNN 时代怎么解释CAM / Grad-CAM 的基础想搞懂 ViT 里的 CAM得先把 CNN 时代那套逻辑回忆一下。CAM 的原始做法是把全局平均池化层之前的特征图拿出来跟全连接层的权重做加权求和得到一张和输入图像尺寸差不多的热力图反映模型分类时重点关注的区域。后来 Grad-CAM 把这事推广了不再要求特征图必须接 GAP可以直接用类别得分对特征图的梯度作为权重再对特征图做加权求和过个 ReLU 就能得到解释图。这里最核心的思想是“用梯度找特征图中对类别决策贡献最大的通道”。CNN 的特征图天然带有空间结构每一层都像是一堆“局部模式探测器”的堆叠所以把通道加权求和回投到输入空间解释起来非常直观。但这类方法有个通病你得到的解释分辨率受限于最后一层卷积特征图的分辨率通常也就是 7x7 或 14x14需要上采样才能覆盖原图边缘会比较糊。1.2 Transformer 来了老方法为什么不一定好用Vision Transformer 一出来原来的解释方法直接面临三个麻烦。第一特征图的概念变了。ViT 的核心操作是 self-attention图像被打成 patch 之后映射成 token特征图不再是 CNN 那种“通道 x 高 x 宽”的四维张量而是“序列长度 x 特征维度”的矩阵。你没法直接套用“通道加权求和”这个操作因为空间位置变成了 sequence position通道维也跟 token 混在一起。第二梯度传播路径变得更绕。CNN 的梯度可以从类别得分一路传到最后一个卷积层的特征图路径短且清晰。ViT 里从分类头到中间层要经过多层 transformer block每层都有多头注意力、MLP、LayerNorm梯度要穿越一堆非线性叠加直接拿梯度去解释中间注意力信号会散得很厉害。第三也是最关键的一点ViT 自带的注意力图让很多人产生了误判。注意力图确实能可视化出模型在 token 之间分配的权重但它表示的是 token 之间的相关性不是类别决策的“证据”。换句话说注意力图回答的是“模型内部信息怎么流动”CAM 回答的是“哪个区域直接支撑了当前分类结果”。两者需要配合不能互相替代。2. 注意力图ViT 自带的解释窗口2.1 注意力机制到底在可视化什么先不急着写代码把注意力机制本身弄清楚。ViT 的每一层 transformer block 都有多头自注意力每个头都会计算一个 attention matrix形状是(num_tokens, num_tokens)其中num_tokens num_patches 1多出来的那个是 class token。A[i][j]表示第 j 个 token 在计算第 i 个 token 的表示时被分配了多少权重。所有行加起来等于 1经过 softmax。当你想可视化“模型关注了图像的哪些区域”最直接的做法是取某个特定 head 的 attention map把[CLS]token 那一行拿出来去掉 class token 对应的那一列再 reshape 回原图的 patch 网格形状上采样到原图尺寸叠在输入图像上显示。这就是大多数人看到的第一张 ViT 注意力热图。但这里有个很容易忽略的细节单层单头的注意力图通常很 noisy而且不同 head 关注模式差别极大。有的 head 关注背景有的 head 关注目标中心有的 head 甚至呈现网格状规律这是因为某些 head 学到了位置编码带来的周期性模式。所以实际项目里我们很少直接拿单个 head 当解释结果而是把多层、多头的注意力矩阵做个聚合。2.2 如何提取并可视化 ViT 的注意力图用 timm 加载一个预训练 ViT 非常方便但需要注意不同代码库对 attention 的输出格式有差异。我常用的是timm里的VisionTransformer在 forward 时设置output_attnTrue就能拿到每层每个 head 的注意力矩阵。import torch import timm from PIL import Image import torchvision.transforms as T import matplotlib.pyplot as plt import numpy as np model timm.create_model(vit_base_patch16_224, pretrainedTrue, num_classes1000) model.eval() img Image.open(cat.png).convert(RGB) transform T.Compose([ T.Resize((224, 224)), T.ToTensor(), T.Normalize(mean[0.5, 0.5, 0.5], std[0.5, 0.5, 0.5]) ]) x transform(img).unsqueeze(0) with torch.no_grad(): output model.forward_features(x) # 只拿特征不拿分类结果 # timm 对 forward_features 的输出通常是 (B, N, C)不直接返回注意力 # 如果要用 attention 输出需要走 model.forward(x, output_attnTrue)注意forward_features不会返回注意力output_attnTrue必须通过model.forward来触发。不同版本的 timm 处理方式也有区别最稳妥的方式是直接看模型 forward 源码或者用 hook 把 attention 矩阵截出来。我自己喜欢写一个小的 wrapperdef get_attention_map(model, x, layer_id-1, head_idNone): attn_list [] def hook_fn(module, input, output): # 在 Attention 模块后面挂 hook把 attn 矩阵存下来 attn_list.append(output[1]) # 有些版本的 output 是个 tuple第一项是特征第二项是 attn handle model.blocks[layer_id].attn.register_forward_hook(hook_fn) with torch.no_grad(): model(x) handle.remove() attn attn_list[0] # (B, num_heads, N, N) cls_attn attn[0, head_id, 0, 1:] if head_id is not None else attn[0, :, 0, 1:].mean(0) patch_size model.patch_embed.patch_size[0] grid_size int(math.sqrt(cls_attn.shape[0])) cls_attn cls_attn.reshape(grid_size, grid_size) return torch.nn.functional.interpolate( cls_attn.unsqueeze(0).unsqueeze(0), size(224, 224), modebilinear, align_cornersFalse ).squeeze()这里把最后一个 attention block 的[CLS]行拿出来有需要就平均多个 head。实际跑下来你会发现最后一层的平均注意力图往往能粗略框出前景目标但边界很柔和会出现大片高亮区域覆盖背景。如果你希望解释更聚焦可以试试从不同层把注意力乘起来这个思路下面细说。2.3 注意力图的使用注意事项注意力图最大的优点是不需要额外的标签或梯度计算推理时顺手就能拿到。但用的时候要清楚它有几个先天局限。第一注意力权重不等于因果贡献。softmax 的归一化会让所有权重加起来等于 1意味着即使某个 token 完全不重要也会分到一点权重。你看到的“高亮区域”只能说明模型在计算[CLS]表示时对这个 token 的依赖度高不能证明它就是类别决策的直接原因。我曾经用一张只有猫头、背景全是纯色的图片做测试[CLS]注意力高亮区居然有一部分落在背景边缘原因是位置编码和边缘 token 形成的固定关联跟猫本身无关。第二跨层传播会造成注意力分散。ViT 每层的注意力矩阵都在变化直接可视化最后一层只能看到最后一步的信息汇聚。如果你想看模型“一层一层怎么把注意力收拢到目标上”可以拿相邻层或隔层的注意力矩阵做乘积这样能模拟从浅层到深层的信息传递路径。常见做法是把每层[CLS]行注意力累乘起来结果会更锐利但也更容易丢失某些 head 的多样性。第三别只看单张图下结论。注意力图对输入图片的扰动非常敏感换个尺寸或者加一点噪声高亮区域可能明显变化。建议至少用多张同类别图片跑一遍观察注意力分布的共性再下“模型关注了哪里”这种结论。3. 把 CAM 移植到 ViT 上方法与实操3.1 ViT-CAM 的基本思路从注意力到类激活有了注意力图还不够我们还需要“类别激活”级别的解释。CNN 里的 CAM 是把最后一个卷积层的特征图按类别权重加权求和ViT 里最接近“卷积特征图”的东西其实是最后一层所有 token 的特征向量。每个 token 对应原图的一个 patch这些 token 的特征向量包含了模型在当前层对各个 patch 的语义理解。如果我们能得到每个 token 对某个类别的重要性权重就能对这些特征向量做加权求和再 reshape 成热力图。问题是怎么得到权重。一个很自然的想法是直接拿分类头的权重ViT 分类时[CLS]token 会过一个 Linear 层得到 logits这个 Linear 层的权重就是“每个特征维度对类别的贡献”。但[CLS]token 的特征向量是全局信息不是每个 patch 的局部特征直接作用到 patch token 上并不严格。所以 ViT-CAM 各种变体其实都在回答同一个问题如何把从[CLS]学到的全局类别信息映射回每个 patch token 上。常用的解决思路有几种直接把[CLS]对应的分类权重当作补丁 token 的权重简单但粗糙效果一般。用 Grad-CAM 的做法计算类别得分对 patch token 特征的梯度再对梯度做空间维度的平均得到权重这个更接近原始 Grad-CAM。结合注意力图用注意力权重作为 mask乘到特征上再加权求和相当于把 CAM 和注意力图融合起来。我自己测试下来纯 Grad-CAM 的做法在 ViT 上效果没有在 CNN 上那么惊艳原因是梯度要穿过 attention 层传播路径长容易出现梯度饱和。反而是“注意力图做 mask 梯度做权重”的组合更稳因为注意力图提供了空间先验梯度负责修正类别方向。3.2 一个可落地的 ViT-CAM 实现示例我给出一个基于 hook 的最小实现基于timm的 ViT核心思路是在最后一个 block 的输出层前挂 hook获取 patch token 的特征然后在 backward 时获取类别 logits 对特征的梯度最后把梯度和特征做加权求和得到 CAM。import torch import timm import numpy as np import torch.nn.functional as F from torchvision import transforms from PIL import Image import matplotlib.pyplot as plt model timm.create_model(vit_base_patch16_224, pretrainedTrue) model.eval() # 我们需要拿到最后一个 block 的输入特征作为“特征图” # 通过在 block 输出之前插入 hook 实现 feat_map None grad_map None def forward_hook(module, input, output): global feat_map # output 是 (B, N, C)取 patch tokens去掉 class token feat_map output[:, 1:, :] def backward_hook(module, grad_input, grad_output): global grad_map grad_map grad_output[0][:, 1:, :] last_block model.blocks[-1] fwd_handle last_block.register_forward_hook(forward_hook) bwd_handle last_block.register_full_backward_hook(backward_hook) img Image.open(dog.png).convert(RGB) transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5]) ]) x transform(img).unsqueeze(0) output model(x) prob F.softmax(output, dim1) cls_idx torch.argmax(prob, dim1).item() print(predicted class:, cls_idx, prob:, prob[0, cls_idx].item()) model.zero_grad() one_hot torch.zeros_like(output) one_hot[0, cls_idx] 1 output.backward(gradientone_hot) # 计算权重对梯度做空间平均 weights grad_map.mean(dim1, keepdimTrue) # (B, 1, C) # 加权求和权重点乘特征再在通道维求和 cam F.relu((weights * feat_map).sum(dim-1)) # (B, N_patch) # reshape 成网格 patch_size model.patch_embed.patch_size[0] grid_size int(cam.shape[1] ** 0.5) cam cam[0].reshape(grid_size, grid_size).detach().numpy() cam np.maximum(cam, 0) cam (cam - cam.min()) / (cam.max() - cam.min()) cam_up F.interpolate(torch.from_numpy(cam).unsqueeze(0).unsqueeze(0), size(224, 224), modebilinear, align_cornersFalse).squeeze().numpy() # 叠加到原图 img_np np.array(img.resize((224, 224))) / 255.0 plt.imshow(img_np) plt.imshow(cam_up, cmapjet, alpha0.5) plt.axis(off) plt.show() fwd_handle.remove() bwd_handle.remove()这段代码有几个细节值得展开说。第一hook 挂的位置。我挂在了最后一个 block 的整体前后拿到的是 block 输入和输出。实际用的时候如果你只想解释某个特定层可以换成任意层但要确认该层输出仍然保留 patch token 的序列结构。第二backward hook 返回的grad_output是该层输出的梯度[:, 1:, :]去掉[CLS]就是为了让梯度与 patch token 对齐。第三weights grad_map.mean(dim1)做的是空间平均等于把每个 patch 位置的梯度求平均得到一个通道维的向量这个理解与 Grad-CAM 的全局平均池化一致。如果你想要更锐利的效果可以先对梯度取绝对值再做空间平均但对噪声会更敏感。3.3 关键参数与效果调优不同任务的 CAM 表现差异很大调优时有几个关键参数值得优先尝试。用哪个类别做解释通常用 argmax 类别能反映模型的主要决策。但在做错误分析时更建议把某个目标类别的 score 单独拿出来做 backward观察模型对“这个特定类别”的响应区域这时候注意要先把 logits 里的其他类别抑制掉。选择哪一层的特征一般来说越靠后的层语义越强但空间分辨率越低。ViT 的 patch token 空间分辨率是固定的patch16 就是 14x14所以不存在 CNN 那样高层分辨率骤降的问题反而适合直接用最后一层。如果你用的是 Deit 或 Swin情况略有不同Swin 有窗口注意力空间变化更复杂这一层特性以后单独写。是否叠加注意力 mask纯 CAM 容易出现散点状高亮我习惯把 CAM 结果乘上最后一层[CLS]行注意力图这样能过滤掉一些和类别无关的区域。注意要先对 CAM 做归一化再与注意力图相乘否则尺度差异会把结果带偏。# 接上面的 cam_up 结果 with torch.no_grad(): attn get_attention_map(model, x, layer_id-1, head_idNone) # 前面定义过 attn_np (attn - attn.min()) / (attn.max() - attn.min()) cam_atten cam_up * attn_np.numpy()用这个融合图时背景噪声明显更少热力图的高亮区更集中在实际目标上。4. 常见问题与排查技巧实录4.1 注意力图看起来“雾蒙蒙”怎么办这是最常见的吐槽尤其在你直接用最后一层[CLS]行注意力做可视化时。原因包括softmax 导致权重分布平缓、多 head 平均后互相抵消、以及位置编码带来的平滑先验。我的处理经验有三条按优先级排序把注意力矩阵做平方或指数放大。比如将cls_attn先减去最小值再取幂这样能人为拉大高权重区域和低权重区域的差异。但要注意这仅仅是可视化增强不是模型本身的解释权重。改用 Rollout 方法把各层注意力矩阵连乘起来。具体做法是先把每层注意力加上单位矩阵然后按层累乘最后取[CLS]行。这样能模拟信息从输入层到[CLS]的流动路径锐利度明显提高。做 head 筛选。计算每个 head 的注意力图与 ground truth mask如果有的话的 IoU选最优的 head。没有 ground truth 时可以观察 head 之间的方差选方差大的 head通常更聚焦。Rollout 有个实现要点由于 ViT 有残差连接每层注意力矩阵要加上单位矩阵再归一化不然连乘后会丢失最初的 token 信息。另外深层网络的注意力矩阵连乘容易过平滑所以一般最多乘到第 6~8 层就停。4.2 CAM 结果和注意力图哪个更可信这个问题我在团队内部讨论过很多次结论是不要把两者当成竞争关系要把它们当成互相验证的两个信号。如果 CAM 高亮区域和注意力图高亮区域重叠度高那说明模型的决策基础很明确可以放心用。如果两者差异很大优先怀疑 CAM 的梯度是否存在饱和问题或者注意力图是否被位置编码干扰。如果 CAM 在背景上有明显高亮但注意力图集中在前景说明模型可能是靠背景辅助决策这在很多数据集里其实是合理的背景和目标强相关。这时候不要急着删掉背景“解释出来的证据”和“我们希望模型学习的证据”是两码事。我个人的习惯是先用注意力图做粗筛剔除明显基于背景的坏样本再用 CAM 做细粒度分析看类别区分能力最后用注意力图与 CAM 的交集作为最终解释区域。这套组合拳已经在我处理过的不少视觉任务里表现出比单一方法更稳定的效果。4.3 后续扩展把可解释性做成可视化工具如果你不想每次都在 notebook 里手动跑流程可以把这些逻辑封装成一个类输入一张图片输出叠加了注意力图和 CAM 的可视化结果。接口设计可以参考下面这样class ViTExplainer: def __init__(self, model_namevit_base_patch16_224, devicecuda): ... def explain(self, img_path, methodcamattn): ... def save_visualization(self, save_path): ...封装时要注意几个易错点每个样本跑完一定要清理 hookbackward 时记得model.zero_grad()如果是批量输入batch size 最好设为 1否则解释的是整张 batch 的梯度很难对应到具体图像。工具化以后你还能把同一张图片在不同 checkpoint 下的注意力图并排对比快速定位模型退化问题。另外一个扩展方向是把可解释性与错误分析结合。给模型喂一批错分样本分别生成 CAM 和注意力图人工观察可以发现模型是否依赖于水印、遮挡物、局部纹理等非语义特征。这也是我目前最推荐的实际应用场景比起单独秀一张漂亮的注意力图更有工程价值。做可解释性最容易被忽略的一点是要保持对“解释结果”本身的批判态度。无论是 CAM 还是注意力图都只是我们对模型内部机制的近似侧写不是模型的“思维过程”。我在实际项目里踩过几次坑之后现在的习惯是任何解释结论都必须至少有两张不同样本的对照并且能用简单的因果实验验证——比如把高亮区域遮挡掉看模型预测是否明显下降遮挡非高亮区域预测是否基本不变。只有通过了这种 sanity check我才敢把可解释性结果写进报告里。这套流程虽然朴素但比堆叠各种花哨的可视化方法可靠得多。