做模型解释的时候,我最常被问到的一个问题就是:“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_attn=True,就能拿到每层每个 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', pretrained=True, num_classes=1000) 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_attn=True)注意,forward_features不会返回注意力,output_attn=True必须通过model.forward来触发。不同版本的 timm 处理方式也有区别,最稳妥的方式是直接看模型 forward 源码,或者用 hook 把 attention 矩阵截出来。我自己喜欢写一个小的 wrapper:
def get_attention_map(model, x, layer_id=-1, head_id=None): 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), mode='bilinear', align_corners=False ).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', pretrained=True) 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, dim=1) cls_idx = torch.argmax(prob, dim=1).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(gradient=one_hot) # 计算权重:对梯度做空间平均 weights = grad_map.mean(dim=1, keepdim=True) # (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), mode='bilinear', align_corners=False).squeeze().numpy() # 叠加到原图 img_np = np.array(img.resize((224, 224))) / 255.0 plt.imshow(img_np) plt.imshow(cam_up, cmap='jet', alpha=0.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(dim=1)做的是空间平均,等于把每个 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_id=None) # 前面定义过 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_name='vit_base_patch16_224', device='cuda'): ... def explain(self, img_path, method='cam+attn'): ... def save_visualization(self, save_path): ...封装时要注意几个易错点:每个样本跑完一定要清理 hook;backward 时记得model.zero_grad();如果是批量输入,batch size 最好设为 1,否则解释的是整张 batch 的梯度,很难对应到具体图像。工具化以后,你还能把同一张图片在不同 checkpoint 下的注意力图并排对比,快速定位模型退化问题。
另外一个扩展方向是把可解释性与错误分析结合。给模型喂一批错分样本,分别生成 CAM 和注意力图,人工观察可以发现模型是否依赖于水印、遮挡物、局部纹理等非语义特征。这也是我目前最推荐的实际应用场景,比起单独秀一张漂亮的注意力图,更有工程价值。
做可解释性最容易被忽略的一点,是要保持对“解释结果”本身的批判态度。无论是 CAM 还是注意力图,都只是我们对模型内部机制的近似侧写,不是模型的“思维过程”。我在实际项目里踩过几次坑之后,现在的习惯是:任何解释结论都必须至少有两张不同样本的对照,并且能用简单的因果实验验证——比如把高亮区域遮挡掉,看模型预测是否明显下降;遮挡非高亮区域,预测是否基本不变。只有通过了这种 sanity check,我才敢把可解释性结果写进报告里。这套流程虽然朴素,但比堆叠各种花哨的可视化方法可靠得多。