简介:本资源是一套基于Vision Transformer(ViT)架构实现的轻量级图像分类完整项目,专为计算机类专业本科生毕业设计、课程设计及AI入门实践打造。面向计科、人工智能、大数据等方向的学习者,提供从数据加载、ViT模型构建、训练调优到预测推理的一站式可运行方案,兼顾理论理解与工程落地。压缩包共15个文件,含6个核心Python源码(如vit_model.py、train.py、predict.py)、3个编译缓存文件、2份Markdown说明文档、1个类别索引JSON及若干运行日志目录,总大小仅31KB,结构简洁、依赖明确、开箱即用。已有264人下载学习,项目代码经实测稳定可靠,附带详细项目说明与使用提示,特别强调路径命名规范等易错点,并支持二次开发拓展,是初学者掌握ViT原理与PyTorch实战的高性价比入门范例。
1. 为什么用 ViT 做图像分类不是“炫技”,而是毕设落地的务实选择:3 分钟跑通、显存友好、代码干净可讲清楚
你手头有一份标注好的图像数据集(哪怕只有 200 张猫狗图),导师说“毕设要体现深度学习能力”,但你刚学完 CNN,YOLOv5 训练起来像在调参玄学,ResNet50 又怕显存炸、怕过拟合、怕答辩时被问“为什么不用更前沿的结构”。这时候,“python实现基于ViT的图像分类任务源码+数据集(可作毕设,运行简单).zip”不是标题党——它直指一个被低估的现实:ViT 在中小规模图像分类任务上,收敛快、调参少、结构透明、显存占用比同等精度的 CNN 更可控。我带过 7 届毕设,用 ViT 的学生答辩通过率最高,不是因为模型多神,而是因为:训练日志干净(loss 下降平滑不抖)、推理速度够用(单图 20ms 内)、模型结构能画出清晰流程图(patch embedding → transformer encoder → cls token → classifier)、代码不到 300 行且全是 PyTorch 原生 API(无黑匣子封装)。它不追求 ImageNet Top-1 85% 的极限精度,但能让你在 4GB 显存的笔记本上,3 小时内完成数据准备、训练、验证、导出 ONNX 全流程,并把每一步讲清楚——这才是毕设最需要的“可解释性”和“可复现性”。
2. 从零跑通 ViT 图像分类:不装额外库、不改环境、只靠 PyTorch 1.12+ 和 torchvision 0.13+
ViT 不是必须用 timm 或 transformers 库才能跑。毕设场景下,用 PyTorch 原生模块手写 ViT 主干,反而更利于理解、调试和答辩展示。本方案完全基于torch.nn和torchvision.transforms,不依赖任何第三方模型库,所有代码可直接粘贴进.py文件执行。核心逻辑分三块:Patch Embedding 模块(把图像切成块并线性映射)、Transformer Encoder 堆叠(标准 multi-head attention + MLP)、Classification Head(取 [CLS] token 后接全连接)。下面给出最小可运行版本,已适配 PyTorch 1.12~2.0(CUDA 11.3+ / CPU 均可)。
2.1 构建 ViT 模型:126 行纯 PyTorch 实现,无外部依赖
import torch import torch.nn as nn import torch.nn.functional as F class PatchEmbedding(nn.Module): def __init__(self, img_size=224, patch_size=16, in_chans=3, embed_dim=768): super().__init__() self.img_size = img_size self.patch_size = patch_size self.n_patches = (img_size // patch_size) ** 2 # 卷积代替展平切块,更稳定(避免 torch.unfold 的梯度问题) self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size) def forward(self, x): x = self.proj(x) # [B, embed_dim, H', W'] x = x.flatten(2) # [B, embed_dim, H'*W'] x = x.transpose(1, 2) # [B, H'*W', embed_dim] return x class Attention(nn.Module): def __init__(self, dim, n_heads=12, qkv_bias=True, attn_p=0., proj_p=0.): super().__init__() self.n_heads = n_heads self.dim = dim self.head_dim = dim // n_heads self.scale = self.head_dim ** -0.5 self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias) self.attn_drop = nn.Dropout(attn_p) self.proj = nn.Linear(dim, dim) self.proj_drop = nn.Dropout(proj_p) def forward(self, x): B, N, C = x.shape qkv = self.qkv(x).reshape(B, N, 3, self.n_heads, self.head_dim) qkv = qkv.permute(2, 0, 3, 1, 4) # [3, B, n_heads, N, head_dim] q, k, v = qkv[0], qkv[1], qkv[2] attn = (q @ k.transpose(-2, -1)) * self.scale attn = attn.softmax(dim=-1) attn = self.attn_drop(attn) x = (attn @ v).transpose(1, 2).reshape(B, N, C) x = self.proj(x) x = self.proj_drop(x) return x class MLP(nn.Module): def __init__(self, in_features, hidden_features=None, out_features=None, drop=0.): super().__init__() out_features = out_features or in_features hidden_features = hidden_features or in_features self.fc1 = nn.Linear(in_features, hidden_features) self.act = nn.GELU() self.fc2 = nn.Linear(hidden_features, out_features) self.drop = nn.Dropout(drop) def forward(self, x): x = self.fc1(x) x = self.act(x) x = self.drop(x) x = self.fc2(x) x = self.drop(x) return x class Block(nn.Module): def __init__(self, dim, n_heads, mlp_ratio=4., qkv_bias=True, p=0., attn_p=0.): super().__init__() self.norm1 = nn.LayerNorm(dim) self.attn = Attention(dim, n_heads, qkv_bias, attn_p, p) self.norm2 = nn.LayerNorm(dim) self.mlp = MLP(dim, int(dim * mlp_ratio)) def forward(self, x): x = x + self.attn(self.norm1(x)) x = x + self.mlp(self.norm2(x)) return x class VisionTransformer(nn.Module): def __init__(self, img_size=224, patch_size=16, in_chans=3, n_classes=1000, embed_dim=768, depth=12, n_heads=12, mlp_ratio=4., qkv_bias=True, p=0., attn_p=0.): super().__init__() self.patch_embed = PatchEmbedding(img_size, patch_size, in_chans, embed_dim) self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed = nn.Parameter( torch.zeros(1, 1 + self.patch_embed.n_patches, embed_dim) ) self.pos_drop = nn.Dropout(p) self.blocks = nn.Sequential(*[ Block(embed_dim, n_heads, mlp_ratio, qkv_bias, p, attn_p) for _ in range(depth) ]) self.norm = nn.LayerNorm(embed_dim) self.head = nn.Linear(embed_dim, n_classes) def forward(self, x): B = x.shape[0] x = self.patch_embed(x) cls_tokens = self.cls_token.expand(B, -1, -1) x = torch.cat((cls_tokens, x), dim=1) x = x + self.pos_embed x = self.pos_drop(x) x = self.blocks(x) x = self.norm(x) x = x[:, 0] # 取 [CLS] token x = self.head(x) return x提示:这段代码是 ViT 的最小可行实现,与原始论文结构一致。关键点在于
PatchEmbedding使用Conv2d而非unfold,避免某些 PyTorch 版本下unfold的梯度不稳定问题;Block中的 LayerNorm 位置严格按论文放在 attention 和 MLP 前(pre-norm),这是训练稳定的关键;cls_token和pos_embed均为可学习参数,初始化用torch.zeros即可,无需特殊初始化——ViT 对初始化鲁棒性远高于 CNN。
2.2 数据加载与增强:适配任意本地文件夹结构,支持小数据集过拟合验证
毕设数据集往往样本少(<1000 张/类),必须用强增强防过拟合,但又要保留语义不变性。以下Dataset类支持标准文件夹格式(./data/train/cat/xxx.jpg),自动识别类别,且增强策略针对 ViT 优化:不用 RandomResizedCrop(破坏 patch 结构),改用 Resize + CenterCrop 组合;ColorJitter 强度降低(ViT 对颜色扰动更敏感);增加 CutMix(对小数据集提升显著)。
from torch.utils.data import Dataset, DataLoader from torchvision import transforms from PIL import Image import os import random import numpy as np class ImageFolderDataset(Dataset): def __init__(self, root_dir, transform=None, is_train=True): self.root_dir = root_dir self.transform = transform self.is_train = is_train self.classes = sorted(os.listdir(root_dir)) self.class_to_idx = {cls: idx for idx, cls in enumerate(self.classes)} self.samples = [] for cls in self.classes: cls_path = os.path.join(root_dir, cls) if not os.path.isdir(cls_path): continue for img_name in os.listdir(cls_path): if img_name.lower().endswith(('.png', '.jpg', '.jpeg')): self.samples.append((os.path.join(cls_path, img_name), self.class_to_idx[cls])) def __len__(self): return len(self.samples) def __getitem__(self, idx): img_path, label = self.samples[idx] img = Image.open(img_path).convert('RGB') if self.transform: img = self.transform(img) return img, label # ViT 专用增强(训练集) train_transform = transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomHorizontalFlip(p=0.5), transforms.CenterCrop(224), transforms.ColorJitter(brightness=0.1, contrast=0.1, saturation=0.1, hue=0.05), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) # 验证集(无随机增强) val_transform = transforms.Compose([ transforms.Resize((256, 256)), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) # 加载数据(示例:data/train 和 data/val 目录) train_dataset = ImageFolderDataset('./data/train', transform=train_transform, is_train=True) val_dataset = ImageFolderDataset('./data/val', transform=val_transform, is_train=False) train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=2, pin_memory=True) val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=2, pin_memory=True)参数说明:
Resize((256,256))确保所有图像先统一到略大于 224 的尺寸,再CenterCrop(224)切出标准输入;ColorJitter参数比 CNN 场景降低 50%,因 ViT 的 attention 机制对局部颜色变化更敏感;Normalize使用 ImageNet 均值方差——即使你的数据集不是 ImageNet,也建议沿用,这是 ViT 预训练权重的归一化基准,迁移学习时效果更稳。
2.3 训练循环:带早停、学习率预热、梯度裁剪的毕设友好版
ViT 训练容易震荡,尤其小数据集上。本训练脚本内置三项关键保护:①Linear Warmup(前 10 个 epoch 学习率从 0 线性升到峰值,防 early collapse);②Gradient Clipping(norm=1.0,防 attention softmax 梯度爆炸);③Early Stopping(验证 loss 连续 5 epoch 不下降则终止,防过拟合)。全程无 wandb/tensorboard 依赖,日志输出到 console 和train.log文件。
import torch.optim as optim from torch.optim.lr_scheduler import LambdaLR import time def train_epoch(model, loader, criterion, optimizer, device, epoch): model.train() running_loss = 0.0 correct = 0 total = 0 for batch_idx, (data, target) in enumerate(loader): data, target = data.to(device), target.to(device) optimizer.zero_grad() output = model(data) loss = criterion(output, target) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) # 关键! optimizer.step() running_loss += loss.item() _, predicted = output.max(1) total += target.size(0) correct += predicted.eq(target).sum().item() acc = 100. * correct / total print(f'Epoch {epoch} | Train Loss: {running_loss/len(loader):.4f} | Acc: {acc:.2f}%') return running_loss / len(loader), acc def validate(model, loader, criterion, device): model.eval() val_loss = 0 correct = 0 total = 0 with torch.no_grad(): for data, target in loader: data, target = data.to(device), target.to(device) output = model(data) val_loss += criterion(output, target).item() _, predicted = output.max(1) total += target.size(0) correct += predicted.eq(target).sum().item() acc = 100. * correct / total print(f'Val Loss: {val_loss/len(loader):.4f} | Val Acc: {acc:.2f}%') return val_loss / len(loader), acc # 初始化 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model = VisionTransformer( img_size=224, patch_size=16, in_chans=3, n_classes=len(train_dataset.classes), # 自动适配你的类别数 embed_dim=768, depth=12, n_heads=12, mlp_ratio=4.0 ).to(device) criterion = nn.CrossEntropyLoss() optimizer = optim.AdamW(model.parameters(), lr=3e-4, weight_decay=0.05) # ViT 推荐用 AdamW # 学习率预热调度器(前 10 epoch 线性 warmup) def warmup_lr_lambda(epoch): if epoch < 10: return float(epoch + 1) / 10 else: return 1.0 scheduler = LambdaLR(optimizer, lr_lambda=warmup_lr_lambda) # 训练主循环 best_val_loss = float('inf') patience_counter = 0 log_file = open('train.log', 'w') for epoch in range(1, 51): # 最大 50 epoch print(f'\n=== Epoch {epoch} ===') train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, device, epoch) val_loss, val_acc = validate(model, val_loader, criterion, device) # 早停逻辑 if val_loss < best_val_loss: best_val_loss = val_loss torch.save(model.state_dict(), 'best_vit_model.pth') patience_counter = 0 print('Saved best model!') else: patience_counter += 1 if patience_counter >= 5: print(f'Early stopping at epoch {epoch}') break scheduler.step() log_file.write(f'{epoch},{train_loss:.4f},{train_acc:.2f},{val_loss:.4f},{val_acc:.2f}\n') log_file.close()为什么这样设参数?
lr=3e-4是 ViT Base 在中小数据集上的经验最优值(比 ResNet 常用的 1e-3 低一个数量级);weight_decay=0.05比 CNN 常用的 1e-4 高,因 ViT 的 attention 权重更易过拟合;clip_grad_norm_=1.0是 ViT 训练稳定性底线,不加此行,第 3~5 epoch 极易 loss nan;warmup 10 epoch防止 ViT 在初期因 attention softmax 梯度不稳定导致训练崩溃——这是 ViT 区别于 CNN 的核心行为差异。
3. 数据集准备与转换:从任意图片文件夹到 ViT 可训格式,含 3 种常见毕设场景处理方案
毕设数据集来源五花八门:手机拍的植物照片、爬虫下载的商品图、公开数据集裁剪版。ViT 对输入尺寸敏感(必须整除 patch_size),且要求类别目录结构清晰。本节提供三种典型场景的落地方案,全部用 Python 脚本一键完成,不依赖 labelImg 等 GUI 工具。
3.1 场景一:你只有杂乱图片(无分类文件夹),需自动聚类并人工校验
常见于“用手机拍了 200 张校园植物,但没分好类”。此时不能直接扔给 ViT,需先粗分。我们用CLIP 零样本特征 + KMeans 聚类快速生成初始标签,再人工修正——比纯手工打标快 5 倍,且聚类结果可直接用于 ViT 的弱监督预训练。
# cluster_images.py:用 CLIP 提取特征并聚类 import torch import clip from PIL import Image import numpy as np from sklearn.cluster import KMeans import os from tqdm import tqdm # 加载 CLIP 模型(CPU 可跑,无需 GPU) device = "cpu" model, preprocess = clip.load("ViT-B/32", device=device) # 读取所有图片路径 img_paths = [] for root, _, files in os.walk('./raw_photos'): for f in files: if f.lower().endswith(('.jpg', '.jpeg', '.png')): img_paths.append(os.path.join(root, f)) # 提取 CLIP 图像特征 image_features = [] for path in tqdm(img_paths, desc="Extracting CLIP features"): image = preprocess(Image.open(path)).unsqueeze(0).to(device) with torch.no_grad(): feat = model.encode_image(image).cpu().numpy() image_features.append(feat[0]) image_features = np.array(image_features) # shape: (N, 512) # KMeans 聚类(假设你预估有 5 类植物) kmeans = KMeans(n_clusters=5, random_state=42, n_init=10) labels = kmeans.fit_predict(image_features) # 创建聚类结果目录 os.makedirs('./clustered_data', exist_ok=True) for i in range(5): os.makedirs(f'./clustered_data/class_{i}', exist_ok=True) # 按聚类结果移动图片 for idx, path in enumerate(img_paths): class_id = labels[idx] fname = os.path.basename(path) dst = f'./clustered_data/class_{class_id}/{fname}' os.system(f'cp "{path}" "{dst}"') # Linux/Mac;Windows 用 shutil.copy print("Clustering done! Check ./clustered_data/") print("Now manually rename class_0/ to meaningful names like 'maple_leaf', 'oak_leaf'...")操作后动作:运行完脚本,你会得到
./clustered_data/class_0/,class_1/等文件夹,每个里面是 CLIP 认为相似的图片。此时打开每个文件夹,人工检查并重命名文件夹为真实类别名(如class_0→roses,class_1→tulips)。这步不可跳过,但只需 10 分钟——比从零打标 200 张图快得多。
3.2 场景二:你有公开数据集(如 Oxford-IIIT Pets),但格式是 .mat 或 .txt 标签,需转为标准文件夹
以 Oxford-IIIT Pets 为例,官方提供annotations.tar.gz,标签在.mat文件里。ViT 需要train/cat/xxx.jpg这种结构。以下脚本自动解压、解析、复制,全程命令行执行,无 GUI 依赖:
# convert_pets.py import scipy.io as sio import os import shutil from PIL import Image # 解压 annotations 并读取 mat_path = './annotations/list.mat' data = sio.loadmat(mat_path) classes = [name[0] for name in data['species'][0]] images = data['file_list'][0] labels = data['species_labels'][0] - 1 # MATLAB 索引从 1 开始,转为 0-based # 创建目标目录 os.makedirs('./pets_data/train', exist_ok=True) os.makedirs('./pets_data/val', exist_ok=True) # 按 8:2 划分训练/验证集(固定随机种子保证可复现) np.random.seed(42) indices = np.random.permutation(len(images)) train_idx = indices[:int(0.8*len(indices))] val_idx = indices[int(0.8*len(indices)):] # 复制图片并按标签建文件夹 for idx_set, split in [(train_idx, 'train'), (val_idx, 'val')]: for i in idx_set: img_name = images[i][0].strip() label = int(labels[i]) class_name = classes[label].replace(' ', '_') # 处理空格 src = f'./images/{img_name}' dst_dir = f'./pets_data/{split}/{class_name}' os.makedirs(dst_dir, exist_ok=True) dst = f'{dst_dir}/{img_name}' if os.path.exists(src): shutil.copy(src, dst) else: print(f"Warning: {src} not found") print("Oxford-IIIT Pets converted to folder structure!")关键细节:
classes[label].replace(' ', '_')处理类别名中的空格(如Persian cat→Persian_cat),避免 Linux 下路径错误;np.random.seed(42)保证每次划分一致,答辩时可复现;shutil.copy比os.system('cp')更跨平台,Windows 也能跑。
3.3 场景三:你只有单张大图(如卫星图、医学切片),需切割为 patch 并标注
森林图像分类、病理切片分析等毕设常遇此场景。ViT 本身不处理大图,需先切 patch。切 patch 不是简单 grid 切割,必须加 overlap 和 ignore 边缘噪声:
# tile_large_image.py from PIL import Image import numpy as np import os def tile_image(img_path, patch_size=224, overlap=32, ignore_edge=16): """ 将大图切成带重叠的 patch,边缘 ignore_edge 像素不参与切分 """ img = Image.open(img_path) w, h = img.size # 有效区域(去掉边缘) w_eff, h_eff = w - 2*ignore_edge, h - 2*ignore_edge # 计算起始坐标(居中 crop) left = ignore_edge + (w - w_eff) // 2 top = ignore_edge + (h - h_eff) // 2 right = left + w_eff bottom = top + h_eff img_cropped = img.crop((left, top, right, bottom)) patches = [] # 步长 = patch_size - overlap step = patch_size - overlap for i in range(0, h_eff - patch_size + 1, step): for j in range(0, w_eff - patch_size + 1, step): patch = img_cropped.crop((j, i, j+patch_size, i+patch_size)) patches.append(patch) return patches # 示例:切一张 forest.jpg patches = tile_image('./forest.jpg', patch_size=224, overlap=32) os.makedirs('./forest_patches', exist_ok=True) for i, p in enumerate(patches): p.save(f'./forest_patches/patch_{i:04d}.jpg') print(f"Generated {len(patches)} patches from forest.jpg")为什么 overlap=32?
ViT 的 patch 是局部感受野,单 patch 可能只含树干或树叶,无完整语义。overlap=32(约 14% 重叠)确保相邻 patch 共享上下文,提升分类鲁棒性;ignore_edge=16剔除扫描/拍摄引入的模糊边缘,避免 ViT 学到噪声模式。
4. 避坑:ViT 毕设训练中 5 个高频翻车点,现象→原因→解决全闭环
ViT 看似结构简洁,但训练行为与 CNN 有本质差异。以下 5 条是我带毕设时学生踩过的真坑,每条都附带print级别的快速验证方法,不靠猜。
4.1 现象:训练前 5 个 epoch loss 从 7.0 直线掉到 0.1,然后卡在 0.05 不动,验证 acc 停在 10%(随机水平)
原因:cls_token初始化为torch.zeros,但未加nn.init.trunc_normal_,导致 [CLS] token 初始向量全零,attention softmax 输出坍缩,模型只学到了 bias。
解决:在VisionTransformer.__init__()中self.cls_token初始化后加一行:
nn.init.trunc_normal_(self.cls_token, std=0.02)验证:训练前打印
model.cls_token的 norm,应为 ~0.02;若为 0,则确认修复。
4.2 现象:验证 loss 波动极大(0.3 → 1.2 → 0.4),acc 在 50% 上下抖动
原因:BatchNorm层混入 ViT 主干(ViT 用 LayerNorm,CNN 才用 BatchNorm)。常见于 copy-paste CNN 代码时误留nn.BatchNorm2d。
解决:全局搜索代码中BatchNorm,ViT 全链路必须只用LayerNorm。检查PatchEmbedding、Block、MLP内部,确认无BatchNorm。
验证:
print([name for name, m in model.named_modules() if isinstance(m, nn.BatchNorm2d)]),输出应为空列表。
4.3 现象:训练 loss 正常下降,但验证 acc 始终低于训练 acc 20% 以上,且越往后差距越大
原因:数据增强太强(如RandomRotation(30)),破坏了 patch 的空间连续性,ViT 的 attention 无法建模扭曲后的局部关系。
解决:删除所有RandomRotation、RandomAffine,仅保留RandomHorizontalFlip和ColorJitter(强度≤0.1)。ViT 对几何变换鲁棒性远低于 CNN,靠数据增强提升泛化效果有限。
验证:临时注释掉
train_transform中所有旋转/仿射,只留 flip + jitter,观察 val acc 是否收敛。
4.4 现象:RuntimeError: CUDA out of memory,即使 batch_size=8 也报错
原因:ViT 的 attention 计算复杂度为 O(N²),N 是 patch 数(224/16=14 → 14²=196)。当img_size=224时内存尚可,但若误设img_size=448,N=28 → N²=784,显存暴涨 4 倍。
解决:检查VisionTransformer初始化时img_size参数是否与transforms.Resize一致;ViT 毕设推荐固定用img_size=224,勿盲目增大。
验证:
print(model.patch_embed.n_patches),应为(224//16)**2 == 196;若为 784,则img_size设错。
4.5 现象:模型导出 ONNX 后推理结果全为 0,或类别概率全相同
原因:ONNX 导出时未设置training=False,导致 dropout 层在推理时仍生效,输出随机。
解决:导出前必须model.eval(),且torch.onnx.export的training参数设为torch.onnx.TrainingMode.PRESERVE或显式training=False:
model.eval() # 关键! dummy_input = torch.randn(1, 3, 224, 224) torch.onnx.export( model, dummy_input, "vit.onnx", input_names=['input'], output_names=['output'], dynamic_axes={'input': {0: 'batch'}, 'output': {0: 'batch'}}, opset_version=12, training=torch.onnx.TrainingMode.PRESERVE # 关键! )验证:用
onnxruntime加载 ONNX,输入全 0 tensor,输出不应全 0;若仍异常,检查model.eval()是否在 export 前调用。
5. 毕设答辩加分技巧:3 个让导师眼前一亮的 ViT 可视化与分析方法
答辩时,光说“我用了 ViT”不够,要证明你真正理解了它在你数据上的行为。以下三个技巧无需额外库,纯 PyTorch + matplotlib,5 分钟内可完成,却能让答辩分数拉满。
5.1 可视化 attention map:证明 ViT 真在看“关键区域”,而非胡猜
ViT 的 attention map 揭示模型关注点。我们提取最后一层 encoder 的 attention weights,反向映射到原图——这比 Grad-CAM 更符合 ViT 机理。关键点:只可视化 [CLS] token 对所有 patch 的 attention 权重,因 [CLS] 聚合全局信息。
import matplotlib.pyplot as plt import numpy as np def visualize_attention(model, img_tensor, save_path='attention_map.png'): """ img_tensor: [1, 3, 224, 224],已 normalize """ model.eval() with torch.no_grad(): # 获取中间 attention 输出(需修改 Block.forward 返回 attn) # 临时 monkey patch Block orig_forward = model.blocks[-1].forward attn_weights = [] def new_forward(x): x_norm = model.blocks[-1].norm1(x) attn_out = model.blocks[-1].attn(x_norm) # 保存最后一层的 attention weights B, N, C = x_norm.shape qkv = model.blocks[-1].attn.qkv(x_norm).reshape(B, N, 3, model.blocks[-1].attn.n_heads, model.blocks[-1].attn.head_dim) q, k, v = qkv[0], qkv[1], qkv[2] attn = (q @ k.transpose(-2, -1)) * model.blocks[-1].attn.scale attn = attn.softmax(dim=-1) # [B, n_heads, N, N] attn_weights.append(attn[0, :, 0, 1:].cpu().numpy()) # [n_heads, N-1],忽略 cls self-attention return attn_out model.blocks[-1].forward = new_forward _ = model(img_tensor.unsqueeze(0)) model.blocks[-1].forward = orig_forward # 聚合多头 attention avg_attn = np.mean(attn_weights[0], axis=0) # [N-1] # 映射到图像网格(14x14) h, w = 14, 14 attn_grid = avg_attn.reshape(h, w) # 可视化 plt.figure(figsize=(6, 6)) plt.imshow(attn_grid, cmap='hot', interpolation='nearest') plt.title('ViT Attention Map ([CLS] to Patches)') plt.axis('off') plt.savefig(save_path, bbox_inches='tight', dpi=300) plt.close() print(f"Attention map saved to {save_path}") # 使用示例(取验证集第一张图) img, _ = next(iter(val_loader)) visualize_attention(model, img[0]) # 输入单张图答辩话术:“您看这张猫图的 attention map,热点集中在猫脸和耳朵区域,证明 ViT 并非黑箱,它和人类一样,优先关注判别性部位——这验证了模型决策的合理性。”
5.2 分析 patch embedding 的 PCA 散点图:揭示数据内在结构是否适合 ViT
ViT 的 patch embedding 是后续 attention 的输入基础。用 PCA 将 768 维 embedding 降到 2D,看同类样本是否聚拢——若同类散开,说明 patch 切割或数据质量有问题。
from sklearn.decomposition import PCA import numpy as np def analyze_patch_embedding(model, dataloader, n_samples=200): model.eval() embeddings = [] labels = [] with torch.no_grad(): for data, target in dataloader: if len(embeddings) >= n_samples: break data = data[:min(32, n_samples-len(embeddings))].to(device) target = target[:min(32, n_samples-len(embeddings))] # 提取 patch embedding(不含 cls token) x = model.patch_embed(data) # [B, N, D] # 取第一个样本的 embedding(或平均) emb = x[0].cpu().numpy() # [ <p> <a href="https://download.csdn.net/download/Runnymmede/89484670" style="color:#ec7500;font-size:14px;"> 本文还有配套的精品资源,点击获取 </a> <img alt="menu-r.4af5f7ec.gif" src="https://csdnimg.cn/release/wenkucmsfe/public/img/menu-r.4af5f7ec.gif" style="width:16px;margin-left:4px;vertical-align:text-bottom;cursor:text;"> </p>