尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

ViT图像分类毕设实战:300行PyTorch代码跑通小数据集

ViT图像分类毕设实战:300行PyTorch代码跑通小数据集 简介本资源是一套基于Vision TransformerViT架构实现的轻量级图像分类完整项目专为计算机类专业本科生毕业设计、课程设计及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 张猫狗图导师说“毕设要体现深度学习能力”但你刚学完 CNNYOLOv5 训练起来像在调参玄学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.13ViT 不是必须用 timm 或 transformers 库才能跑。毕设场景下用 PyTorch 原生模块手写 ViT 主干反而更利于理解、调试和答辩展示。本方案完全基于torch.nn和torchvision.transforms不依赖任何第三方模型库所有代码可直接粘贴进.py文件执行。核心逻辑分三块Patch Embedding 模块把图像切成块并线性映射、Transformer Encoder 堆叠标准 multi-head attention MLP、Classification Head取 [CLS] token 后接全连接。下面给出最小可运行版本已适配 PyTorch 1.122.0CUDA 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_size224, patch_size16, in_chans3, embed_dim768): 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_sizepatch_size, stridepatch_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_heads12, qkv_biasTrue, attn_p0., proj_p0.): 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, biasqkv_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_featuresNone, out_featuresNone, drop0.): 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_ratio4., qkv_biasTrue, p0., attn_p0.): 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_size224, patch_size16, in_chans3, n_classes1000, embed_dim768, depth12, n_heads12, mlp_ratio4., qkv_biasTrue, p0., attn_p0.): 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), dim1) 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, transformNone, is_trainTrue): 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(p0.5), transforms.CenterCrop(224), transforms.ColorJitter(brightness0.1, contrast0.1, saturation0.1, hue0.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, transformtrain_transform, is_trainTrue) val_dataset ImageFolderDataset(./data/val, transformval_transform, is_trainFalse) train_loader DataLoader(train_dataset, batch_size32, shuffleTrue, num_workers2, pin_memoryTrue) val_loader DataLoader(val_dataset, batch_size32, shuffleFalse, num_workers2, pin_memoryTrue)参数说明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 Clippingnorm1.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_norm1.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(fEpoch {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(fVal 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_size224, patch_size16, in_chans3, n_classeslen(train_dataset.classes), # 自动适配你的类别数 embed_dim768, depth12, n_heads12, mlp_ratio4.0 ).to(device) criterion nn.CrossEntropyLoss() optimizer optim.AdamW(model.parameters(), lr3e-4, weight_decay0.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_lambdawarmup_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(fEarly 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()为什么这样设参数lr3e-4是 ViT Base 在中小数据集上的经验最优值比 ResNet 常用的 1e-3 低一个数量级weight_decay0.05比 CNN 常用的 1e-4 高因 ViT 的 attention 权重更易过拟合clip_grad_norm_1.0是 ViT 训练稳定性底线不加此行第 35 epoch 极易 loss nanwarmup 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, devicedevice) # 读取所有图片路径 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, descExtracting 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_clusters5, random_state42, n_init10) labels kmeans.fit_predict(image_features) # 创建聚类结果目录 os.makedirs(./clustered_data, exist_okTrue) for i in range(5): os.makedirs(f./clustered_data/class_{i}, exist_okTrue) # 按聚类结果移动图片 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(fcp {path} {dst}) # Linux/MacWindows 用 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→rosesclass_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_okTrue) os.makedirs(./pets_data/val, exist_okTrue) # 按 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_okTrue) dst f{dst_dir}/{img_name} if os.path.exists(src): shutil.copy(src, dst) else: print(fWarning: {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_size224, overlap32, ignore_edge16): 将大图切成带重叠的 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, jpatch_size, ipatch_size)) patches.append(patch) return patches # 示例切一张 forest.jpg patches tile_image(./forest.jpg, patch_size224, overlap32) os.makedirs(./forest_patches, exist_okTrue) for i, p in enumerate(patches): p.save(f./forest_patches/patch_{i:04d}.jpg) print(fGenerated {len(patches)} patches from forest.jpg)为什么 overlap32ViT 的 patch 是局部感受野单 patch 可能只含树干或树叶无完整语义。overlap32约 14% 重叠确保相邻 patch 共享上下文提升分类鲁棒性ignore_edge16剔除扫描/拍摄引入的模糊边缘避免 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, std0.02)验证训练前打印model.cls_token的 norm应为 ~0.02若为 0则确认修复。4.2 现象验证 loss 波动极大0.3 → 1.2 → 0.4acc 在 50% 上下抖动原因BatchNorm层混入 ViT 主干ViT 用 LayerNormCNN 才用 BatchNorm。常见于 copy-paste CNN 代码时误留nn.BatchNorm2d。解决全局搜索代码中BatchNormViT 全链路必须只用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_size8 也报错原因ViT 的 attention 计算复杂度为 O(N²)N 是 patch 数224/1614 → 14²196。当img_size224时内存尚可但若误设img_size448N28 → N²784显存暴涨 4 倍。解决检查VisionTransformer初始化时img_size参数是否与transforms.Resize一致ViT 毕设推荐固定用img_size224勿盲目增大。验证print(model.patch_embed.n_patches)应为(224//16)**2 196若为 784则img_size设错。4.5 现象模型导出 ONNX 后推理结果全为 0或类别概率全相同原因ONNX 导出时未设置trainingFalse导致 dropout 层在推理时仍生效输出随机。解决导出前必须model.eval()且torch.onnx.export的training参数设为torch.onnx.TrainingMode.PRESERVE或显式trainingFalsemodel.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_version12, trainingtorch.onnx.TrainingMode.PRESERVE # 关键 )验证用onnxruntime加载 ONNX输入全 0 tensor输出不应全 0若仍异常检查model.eval()是否在 export 前调用。5. 毕设答辩加分技巧3 个让导师眼前一亮的 ViT 可视化与分析方法答辩时光说“我用了 ViT”不够要证明你真正理解了它在你数据上的行为。以下三个技巧无需额外库纯 PyTorch matplotlib5 分钟内可完成却能让答辩分数拉满。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_pathattention_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], axis0) # [N-1] # 映射到图像网格14x14 h, w 14, 14 attn_grid avg_attn.reshape(h, w) # 可视化 plt.figure(figsize(6, 6)) plt.imshow(attn_grid, cmaphot, interpolationnearest) plt.title(ViT Attention Map ([CLS] to Patches)) plt.axis(off) plt.savefig(save_path, bbox_inchestight, dpi300) plt.close() print(fAttention map saved to {save_path}) # 使用示例取验证集第一张图 img, _ next(iter(val_loader)) visualize_attention(model, img[0]) # 输入单张图答辩话术“您看这张猫图的 attention map热点集中在猫脸和耳朵区域证明 ViT 并非黑箱它和人类一样优先关注判别性部位——这验证了模型决策的合理性。”5.2 分析 patch embedding 的 PCA 散点图揭示数据内在结构是否适合 ViTViT 的 patch embedding 是后续 attention 的输入基础。用 PCA 将 768 维 embedding 降到 2D看同类样本是否聚拢——若同类散开说明 patch 切割或数据质量有问题。from sklearn.decomposition import PCA import numpy as np def analyze_patch_embedding(model, dataloader, n_samples200): 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 hrefhttps://download.csdn.net/download/Runnymmede/89484670 stylecolor:#ec7500;font-size:14px; 本文还有配套的精品资源点击获取 /a img altmenu-r.4af5f7ec.gif srchttps://csdnimg.cn/release/wenkucmsfe/public/img/menu-r.4af5f7ec.gif stylewidth:16px;margin-left:4px;vertical-align:text-bottom;cursor:text; /p
返回列表