如何可视化PyTorch中Vision Transformer各层输出及注意力区域
ViT自注意力可视化解决方案
核心问题拆解
你遇到的问题本质是对ViT数据流逻辑不清晰:
conv_proj输出的[1,768,14,14]是图像卷积投影后的全局特征图,未转换为序列格式,因此可视化的是整体特征而非分块注意力- 编码器输出的
[1,197,768]包含cls token+196个图像块的序列特征,直接可视化无意义,需提取注意力权重而非特征本身 - 遍历子模块传图时的维度错误,是因为跳过了ViT的关键步骤:将特征图展平并添加cls token转换为序列格式
正确实现步骤
1. 加载模型并提取注意力权重
ViT的每个编码器层都包含自注意力模块,我们需要通过钩子直接提取注意力权重:
import torch import matplotlib.pyplot as plt from torchvision.models import vision_transformer as vits # 加载训练好的模型 device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = torch.load("vit_mnist_model.pth").to(device) model.eval() # 存储注意力权重的列表 attention_weights = [] def attention_hook(module, input, output): # ViT自注意力模块的输出是(注意力输出, 注意力权重),取第二个元素 attention_weights.append(output[1].detach().cpu()) # 给最后一个编码器层的自注意力模块注册钩子 for name, module in model.named_modules(): if name == "encoder.layers.encoder_layer_11.self_attention": module.register_forward_hook(attention_hook)
2. 预处理输入图像
MNIST是单通道图,需转换为ViT要求的3通道输入,同时调整尺寸到224x224:
from torchvision import transforms from PIL import Image transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.Grayscale(num_output_channels=3), # 单通道转3通道 transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) # 加载测试图像(替换为你的MNIST图像路径) img = Image.open("mnist_test.png") input_tensor = transform(img).unsqueeze(0).to(device) # 转换为[1,3,224,224]格式
3. 前向传播并处理注意力权重
with torch.no_grad(): model(input_tensor) # 取最后一个编码器层的注意力权重,形状为[1, 12, 197, 197] # 12是ViT-B/16的注意力头数,197=1个cls token + 14*14个图像块 attn_weights = attention_weights[0][0] # 取第一个样本的注意力权重,形状[12,197,197] # 计算所有注意力头的平均,提取cls token对图像块的注意力 cls_attn = attn_weights[:, 0, 1:].mean(dim=0) # 忽略cls token自身,只看对图像块的注意力 cls_attn = cls_attn.reshape(14,14) # 转换为14x14的图像块注意力图
4. 可视化注意力热力图
plt.figure(figsize=(10,5)) # 绘制原图 plt.subplot(1,2,1) plt.imshow(img, cmap="gray") plt.title("Original MNIST Image") plt.axis("off") # 绘制注意力热力图 plt.subplot(1,2,2) heatmap = plt.imshow(cls_attn, cmap="viridis", interpolation="bilinear") plt.title("CLS Token Attention Heatmap") plt.axis("off") plt.colorbar(heatmap, fraction=0.046, pad=0.04) plt.tight_layout() plt.show()
关键问题解释
- 遍历子模块报错原因:
conv_proj输出的[B,C,H,W]格式,必须经过ViT内部的展平、添加cls token步骤,转换为[B, seq_len, C]序列格式后才能传入编码器层,直接传递会触发维度检查错误。 - conv_proj输出可视化全图原因:它是图像的卷积投影特征,未分割为独立图像块,若要得到分块特征需手动将
[1,768,14,14]展平为[1,196,768],但这只是特征而非注意力权重。 - 编码器输出的序列特征怎么用:
[1,197,768]是cls token和图像块的特征向量,无法直接可视化,必须提取注意力模块输出的权重矩阵才能反映模型的关注重点。
内容的提问来源于stack exchange,提问作者Mohsin Ali
相关产品推荐
相关产品推荐

