You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何可视化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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.19 00:05:24