基于PyTorch实现ViT训练CIFAR10:从零搭建与避坑指南
2026/9/24 21:15:39 网站建设 项目流程

简介:基于Vision Transformer(ViT)的CIFAR10图像分类训练与验证Python源码,面向人工智能、计算机、自动化等专业在校生及毕业设计、课程设计场景,帮助读者快速搭建图像分类模型并进行训练与验证,也可在此代码基础上修改实现其他分类任务。压缩包共2个文件,包含1个Python脚本和1个txt说明文件,整体大小仅2KB,结构非常精简,脚本承载数据加载、模型构建、训练与验证等完整流程,txt文件则提供配套说明或目录结构,便于按图索骥。目前已有525人学习/下载,代码经过运行测试,能够稳定完成CIFAR10数据集的分类实验,适合作为课程作业、毕设项目或初期立项演示的参考实现。通过Vit.py可以掌握基于Transformer架构的图像分类核心步骤,结合txt说明与代码注释进一步理解细节;资源整体小巧但功能完整,入门者可以借此熟悉ViT模型在标准数据集上的应用,进阶者也能快速替换或扩展模块,节省开发时间。

1. 用ViT训练CIFAR10:一张32×32小图里藏着的最关键选择

把 Vision Transformer(ViT)用在 CIFAR10 分类数据集的训练和验证上,很多人以为就是把经典代码换个数据集,结果一跑就翻车。原因很直接:CIFAR10 的图像只有 32×32,而 ViT 是为 224×224 这种大图设计的,模型结构里 patch 尺寸、位置编码长度、训练超参全都要跟着改。这篇文章要讲的,就是怎么用 Python 和 PyTorch 从零搭一个能在小分辨率图像上正常收敛的 ViT 训练验证源码,把 patch 怎么切、位置编码怎么做、训练循环怎么写、验证时看什么指标一次说清。

如果你正想把手里的 CNN 分类器换成 ViT 做对比实验,或者准备在 CIFAR10 上跑通流程后迁移到自己的小数据集,这篇提供的就是一条可以直接照着走的路线。下面从模型结构开始,逐步落到训练循环、验证评估和排错,每一段代码都可以直接保存运行。

2. 从patch切分到骨架代码:手写一个适配32×32小图的ViT模型

2.1 为什么 CIFAR10 会暴露 ViT 的短处

ViT 的设计目标是替代 CNN 做视觉特征建模,核心思想是把图像切成固定大小的 patch,然后当作 token 序列送进 Transformer encoder。在 ImageNet 上,标准做法是把 224×224 的图像切成 16×16 的 patch,得到 14×14=196 个 token,这个序列长度对自注意力来说是合理的计算量。

但 CIFAR10 的图像只有 32×32,如果沿用 16×16 的 patch,整张图只剩 2×2=4 个 token。四个 token 做自注意力,模型基本学不到任何空间关系,这就是很多人把 ImageNet 上的 ViT 代码直接搬到 CIFAR10 后 loss 不降的根本原因。所以要适配小分辨率图像,第一个动作是把 patch 改小,常见选择是用 4×4 的 patch,这样序列长度变成 8×8=64 个 token,虽然比 196 短,但已经足够让注意力机制工作起来。

另一个问题是 embedding 维度。ImageNet 上常用的 ViT-Base 是 768 维、12 层、12 头,这个规模放到只有 6 万张图的 CIFAR10 上会严重过拟合。我一般会把 embed_dim 压到 192 或 256,depth 控制在 8 到 10 层,num_heads 用 6 或 8。这里的直觉是:数据量小的时候,模型容量要跟着降,否则验证集准确率会长期停滞在 80% 左右上不去。

2.2 基于PyTorch实现ViT:一份最小可运行模型代码

下面这份代码是简化后的 ViT 实现,保留了 patch embedding、CLS token、位置编码、Transformer encoder 和分类头这五个核心部分,没有用任何第三方 Transformer 库,方便改参数和理解结构。把它存成vit_model.py即可。

# vit_model.py import torch import torch.nn as nn class PatchEmbed(nn.Module): """把图像切成 patch 并映射到 embed_dim 维向量。 这里直接用 Conv2d 实现,kernel_size=patch_size, stride=patch_size, 等价于先切 patch 再做线性投影,计算效率更高。 """ def __init__(self, in_channels=3, embed_dim=192, patch_size=4): super().__init__() self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=patch_size) def forward(self, x): B, C, H, W = x.shape x = self.proj(x) # (B, embed_dim, H/p, W/p) x = x.flatten(2).transpose(1, 2) # (B, num_patches, embed_dim) return x class TransformerBlock(nn.Module): """标准 Transformer encoder 层:LayerNorm -> MHA -> 残差 -> LayerNorm -> MLP -> 残差。 预归一化(Pre-LN)是 ViT 训练稳定的关键,和 Post-LN 相比对小数据集更友好。 """ def __init__(self, embed_dim, num_heads, mlp_ratio=4.0, dropout=0.1): super().__init__() self.norm1 = nn.LayerNorm(embed_dim) self.attn = nn.MultiheadAttention( embed_dim, num_heads, dropout=dropout, batch_first=True, average_attn_weights=False ) self.norm2 = nn.LayerNorm(embed_dim) hidden_dim = int(embed_dim * mlp_ratio) self.mlp = nn.Sequential( nn.Linear(embed_dim, hidden_dim), nn.GELU(), nn.Dropout(dropout), nn.Linear(hidden_dim, embed_dim), nn.Dropout(dropout), ) # 保存最近一次前向的注意力权重,后面可视化会用到 self.attn_w = None def forward(self, x): x_norm = self.norm1(x) attn_out, attn_w = self.attn(x_norm, x_norm, x_norm) self.attn_w = attn_w.detach() x = x + attn_out x = x + self.mlp(self.norm2(x)) return x class ViT(nn.Module): def __init__(self, img_size=32, patch_size=4, in_channels=3, num_classes=10, embed_dim=192, depth=9, num_heads=6, dropout=0.1): super().__init__() self.patch_size = patch_size num_patches = (img_size // patch_size) ** 2 self.patch_embed = PatchEmbed(in_channels, embed_dim, patch_size) # CLS token 和位置编码都是可学习参数 self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim)) self.pos_drop = nn.Dropout(dropout) self.blocks = nn.Sequential(*[ TransformerBlock(embed_dim, num_heads, dropout=dropout) for _ in range(depth) ]) self.norm = nn.LayerNorm(embed_dim) self.head = nn.Linear(embed_dim, num_classes) self.apply(self._init_params) def _init_params(self, m): if isinstance(m, nn.Linear): nn.init.trunc_normal_(m.weight, std=0.02) if m.bias is not None: nn.init.zeros_(m.bias) elif isinstance(m, nn.LayerNorm): nn.init.ones_(m.weight) nn.init.zeros_(m.bias) def forward(self, x): B = x.shape[0] x = self.patch_embed(x) # (B, num_patches, embed_dim) cls_tokens = self.cls_token.expand(B, -1, -1) x = torch.cat([cls_tokens, x], dim=1) # (B, num_patches+1, embed_dim) x = x + self.pos_embed x = self.pos_drop(x) x = self.blocks(x) x = self.norm(x) x = x[:, 0] # 取 CLS token return self.head(x)

代码里值得注意的几个参数。embed_dim=192、depth=9、num_heads=6 是根据 CIFAR10 数据量调过的配置,如果你用 GPU 训练时显存有余,可以先把 depth 调到 12 看看验证集准确率是否继续上升;如果提升不足 0.5 个点,说明模型已经饱和,不要再加层数。dropout=0.1 分布在注意力层和 MLP 层,这是 ViT 在小数据集上防止过拟合的默认值,不要轻易改成 0。average_attn_weights=False是为了保留每个 head 的注意力权重,训练阶段多占一点显存,但后面可视化直接调用,不用再改模型。

2.3 CIFAR10数据准备:torchvision下载、增强策略与Dataloader写法

CIFAR10 数据集本身不用手工整理目录,torchvision 会自动下载。第一次运行时会从服务器拉取压缩包,网络不稳的话容易失败,我实际碰到的情况是下载到一半报ConnectionResetError,解决办法是用浏览器或下载工具手动下载 cifar-10-python.tar.gz,放到./data目录下,torchvision 检测到文件存在就不会重复下载。

数据增强对小图数据集影响很大。CIFAR10 上我用的策略是 RandomCrop(32, padding=4) 加 RandomHorizontalFlip,这两个是基础操作,Cutout(随机遮挡一块 8×8 区域)能再带来 0.5 到 1 个点的提升,但要注意遮挡块不能太大,32×32 的图上遮 16×16 基本就把主体盖没了。

# prepare_data.py import torch from torch.utils.data import DataLoader from torchvision import datasets, transforms # 训练集增强:先随机裁剪+水平翻转,再归一化 train_transform = transforms.Compose([ transforms.RandomCrop(32, padding=4), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(mean=(0.4914, 0.4822, 0.4465), std=(0.2023, 0.1994, 0.2010)), ]) # 验证集不做随机增强,只做张量化和归一化 val_transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize(mean=(0.4914, 0.4822, 0.4465), std=(0.2023, 0.1994, 0.2010)), ]) train_ds = datasets.CIFAR10("./data", train=True, download=True, transform=train_transform) val_ds = datasets.CIFAR10("./data", train=False, download=True, transform=val_transform) train_loader = DataLoader(train_ds, batch_size=128, shuffle=True, num_workers=4, pin_memory=True) val_loader = DataLoader(val_ds, batch_size=256, shuffle=False, num_workers=4, pin_memory=True)

归一化的 mean 和 std 是 CIFAR10 数据集官方统计好的,不要换成 ImageNet 那套数值,否则虽然也能收敛,但收敛速度会变慢。num_workers=4在 Windows 上有概率报 DataLoader worker 相关的错,如果遇到可以把num_workers改成 0;在 Linux 服务器上保持 4 或 8 能明显加快数据读取。pin_memory=True只在 GPU 训练时有意义,配合to(device, non_blocking=True)能减少一次数据拷贝,CPU 训练建议关掉。

环境方面只要保证 Python 3.8 以上、PyTorch 1.12 以上、torchvision 0.13 以上即可,IDE 用 VS Code 还是 PyCharm 无所谓,但项目路径里别出现中文,否则部分 CUDA 环境下会报奇怪的路径编码错误。

3. 训练循环与验证循环:超参、warmup、保存最优模型

3.1 先定训练超参:ViT不是CNN,学习率策略要单独配

ViT 和 ResNet 这类 CNN 在训练行为上有明显差异。CNN 用 SGD 加上大学习率也能跑,ViT 对优化器的敏感度更高,常见的稳定组合是 AdamW、3e-4 左右的基础学习率、5e-2 的 weight decay,以及一个 5 到 10 个 epoch 的 warmup。直接用 ResNet 的 0.1 学习率跑 ViT,大概率会发现 loss 在前几个 epoch 完全不下降,这是注意力机制在训练初期还没稳定时的典型表现。

参数建议值说明
batch_size128显存不足时降到 64,配合梯度累积
base_lr3e-4小 ViT 在 CIFAR10 上的安全起点
min_lr1e-5cosine 退火的下界
weight_decay5e-2ViT 标准配置,仅对非 bias 和 norm 参数生效
warmup_epochs5总 epoch 的 5% 左右
epochs100快速验证可以只跑 30 个 epoch
dropout0.1模型定义里已设置,训练时不用再改

weight decay 这里有个细节:直接把 5e-2 传给优化器,会让 LayerNorm 里的 bias 和 scale 也被衰减,影响不大但不够规范。常见做法是给正则化分组,只对权重矩阵做 weight decay。下面代码里我用param_groups实现了这个分组,这也是从 ViT 开源实现里沿用到现在的标准做法。

3.2 训练与验证循环的完整实现

把训练和验证写在一个脚本里,每跑完一个 epoch 在验证集上算一次 top-1 准确率,然后根据准确率决定是否保存模型权重。这里有一个容易忽略的点:验证时一定要调用model.eval()并包在torch.no_grad()里,否则 BatchNorm、Dropout 会继续按训练模式工作,验证结果会虚高。

# train_val.py import torch import torch.nn as nn import torch.optim as optim from vit_model import ViT device = "cuda" if torch.cuda.is_available() else "cpu" model = ViT(img_size=32, patch_size=4, num_classes=10, embed_dim=192, depth=9, num_heads=6).to(device) # 分组 weight decay:norm 层和 bias 不衰减 decay_params = [p for p in model.parameters() if p.requires_grad and p.ndim >= 2] no_decay_params = [p for p in model.parameters() if p.requires_grad and p.ndim < 2] optimizer = optim.AdamW([ {"params": decay_params, "weight_decay": 5e-2}, {"params": no_decay_params, "weight_decay": 0.0}, ], lr=3e-4) criterion = nn.CrossEntropyLoss() # warmup 结束后使用 cosine 退火到 min_lr from torch.optim.lr_scheduler import CosineAnnealingLR scheduler = CosineAnnealingLR(optimizer, T_max=95, eta_min=1e-5) best_acc = 0.0 epochs = 100 warmup_epochs = 5 for epoch in range(epochs): model.train() total_loss, total_correct, total_num = 0.0, 0, 0 for images, labels in train_loader: images, labels = images.to(device), labels.to(device) outputs = model(images) loss = criterion(outputs, labels) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() total_loss += loss.item() * images.size(0) total_correct += (outputs.argmax(dim=1) == labels).sum().item() total_num += images.size(0) # warmup 阶段手动调低学习率 if epoch + 1 < warmup_epochs: lr = 3e-4 * (epoch + 2) / warmup_epochs for g in optimizer.param_groups: g["lr"] = lr else: scheduler.step() # ---------- 验证 ---------- model.eval() val_correct, val_total = 0, 0 with torch.no_grad(): for images, labels in val_loader: images, labels = images.to(device), labels.to(device) outputs = model(images) val_correct += (outputs.argmax(dim=1) == labels).sum().item() val_total += labels.size(0) val_acc = val_correct / val_total print(f"epoch {epoch+1:3d} | train_loss {total_loss/total_num:.4f} " f"| train_acc {total_correct/total_num:.4f} | val_acc {val_acc:.4f}") # 保存验证集上最优的模型 if val_acc > best_acc: best_acc = val_acc torch.save(model.state_dict(), "best_vit_cifar10.pth")

训练循环里有两个关键操作。一是clip_grad_norm_(max_norm=1.0),这行把梯度模长限制到 1,能有效避免训练中期偶发的 loss 突变;ViT 的梯度分布比 CNN 更容易出现极端值,不加这个,跑着跑着可能突然出现一次loss=nan。二是 warmup 的写法,我手动判断epoch + 1 < warmup_epochs来调整学习率,前 5 个 epoch 线性上升到 3e-4,之后才进入 cosine 退火阶段。这里我让 warmup 阶段不调用scheduler.step(),这样 CosineAnnealingLR 的 T_max 设置为 95,正好覆盖剩余 epoch。

3.3 保存模型、断点续训与训练日志

训练过程看train_lossval_acc两个指标就够了。train_loss 在前 10 个 epoch 从 1.8 左右降到 1.0 以下属于正常节奏,如果 20 个 epoch 后还在 1.5 以上徘徊,需要回头检查数据增强是否过强、学习率是否过高。val_acc 会呈现出阶梯式上升,也就是连续几个 epoch 停顿,然后突然涨一个点,这是 Transformer 训练的常见现象,不是 bug,不用着急调参数。

保存模型时只存state_dict,不存整个 model 对象,这样换环境加载时不用等 torchvision 重新实例化模型。如果想支持断点续训,可以附加保存 optimizer 状态和当前 epoch。日志方面,我习惯直接用print输出,因为 CIFAR10 训练一轮用不了几分钟,写完文件再画图反而浪费时间;做对比实验时再用 TensorBoard 记录 loss 曲线。

4. 训练ViT避坑指南:5个常见问题与排查方法

4.1 训练了20个epoch,loss完全不动

现象是 loss 从 2.3 开始缓慢下降,到第 10 个 epoch 还在 2.2 附近,几乎看不出变化。最常见的原因是学习率太大,ViT 在小数据集上的有效学习率区间很窄,3e-4 能跑,3e-3 就会让注意力权重在初始化附近震荡,梯度方向互相抵消。另一个隐蔽原因是位置编码和 CLS token 的初始化问题,代码里我用trunc_normal_(std=0.02)初始化,如果漏掉这一步或者用zeros_,模型前几层会一直输出相近的特征,梯度信号混乱。

解决办法是先用很小的学习率 1e-4 跑 5 个 epoch 确认 loss 在下降,再逐步调大。如果 1e-4 下能降,但 3e-4 下明显变差,说明当前这套模型结构对学习率太敏感,可以检查是否漏了 LayerNorm,或者 dropout 是否设置成了 0。

4.2 改patch_size或输入分辨率后位置编码维度报错

现象是训练正常跑通了,想试试 8×8 的 patch,改完ViT(img_size=32, patch_size=8)之后报size mismatch for pos_embed。原因很简单:位置编码的第一个维度num_patches+1是在__init__里算好的,patch_size 从 4 改成 8 后,序列长度从 65 变成 17,和已生成的位置编码参数维度直接冲突。

解决方法是不要复用旧的权重文件,改了 patch_size 或 img_size 之后,要把模型和权重都重新初始化。另外一个常见坑是在验证阶段直接传入不同分辨率的图,比如训练用 32×32,测试时某张图是 64×64,模型前向时 patch_embed 能跑,但位置编码对不上。CIFAR10 场景里验证集分辨率必须和训练集完全一致。

4.3 验证集准确率在某个epoch突然掉4个点

现象是 val_acc 本来稳定上升,第 40 个 epoch 突然从 88% 掉到 84%,但 train_loss 还在下降,后面几个 epoch 又涨回来了。导致这种现象最常见的原因是验证集数据没有固定增强流程——验证集用了包含 RandomCrop 的 train_transform,每次验证都在不同随机裁剪下评估,结果天然不稳。另一种原因是学习率在 cosine 退火中期出现了一个陡坡,模型还在探索阶段,验证集上正好落在不稳定的权重快照上。

解决方法是验证代码里强制使用无随机增强的 transform,并且只保存验证集上历史最优的模型权重,最后提交模型时用best_vit_cifar10.pth,不要用最后一轮保存的权重。这也是我坚持在训练循环里单独维护best_acc变量而不是每一轮都覆盖权重的原因。

4.4 GPU显存溢出

现象是 batch_size=256 时直接 OOM,4090 显卡也顶不住。ViT 的计算开销集中在自注意力上,嵌入维度 192、序列长度 65 虽然不大,但 depth=9 的层数会让中间激活值累积起来,实际显存占用比同参数量的 CNN 高不少。另一个隐性开销是average_attn_weights=False保留的注意力权重矩阵,每层多存一份(batch, heads, seq, seq)的激活。

解决办法有两条路。一是把 batch_size 降到 64 或 32,同时用梯度累积补偿;二是开启自动混合精度(AMP),CIFAR10 这类任务上 float16 对精度影响很小,显存能省将近一半。如果 AMP 开启后 loss 出现 NaN,检查模型里是否有未适配 fp16 的自定义操作,最常见的是手动算的数值稳定性不过关。

4.5 训练集acc冲到99%,验证集停在80%

这是 ViT 在小数据集上最典型的过拟合症状。CIFAR10 只有 5 万张训练图,192 维、9 层的 ViT 已经有约 2000 万参数,容量远超过数据量能支撑的复杂度。此时优先检查三件事:数据增强是否只用了 RandomCrop 和 Flip,如果太弱,换成 2.3 节里的完整方案;dropout 是否被无意中设为 0;weight decay 是否因为分组配置错误没生效。

如果以上都正常但仍过拟合,下一步是减小模型容量,把 embed_dim 从 192 降到 128、depth 从 9 降到 6。这个改动会牺牲一点训练集准确率,但验证集通常能提升 1 到 2 个点。小数据集上不要盲目追求跟 ImageNet 一样大的模型,这是我从一开始就反复确认过的一条经验。

5. 验证不只是准确率:混淆矩阵、错例分析与泛化检查

5.1 验证循环里顺手收集预测结果,避免二次加载数据

第 3 章的验证循环只输出了一个 acc 数字,要做更细的分析,就得在验证过程中把每个样本的预测类别和真实类别收集起来。值得注意的坑是,DataLoader 的 shuffle 在验证集上要设为 False,不然两次验证跑出的样本顺序不一致,后续保存 confusion matrix 时和原始 label 对不上。

# evaluate.py import numpy as np from sklearn.metrics import confusion_matrix, classification_report import matplotlib.pyplot as plt all_preds, all_labels = [], [] model.eval() with torch.no_grad(): for images, labels in val_loader: images = images.to(device) outputs = model(images) preds = outputs.argmax(dim=1).cpu().numpy() all_preds.extend(preds) all_labels.extend(labels.numpy()) all_preds = np.array(all_preds) all_labels = np.array(all_labels) conf_mat = confusion_matrix(all_labels, all_preds) plt.figure(figsize=(8, 8)) plt.imshow(conf_mat, cmap="Blues") plt.colorbar() plt.xlabel("Predicted Label") plt.ylabel("True Label") plt.savefig("confusion_matrix.png", dpi=120)

confusion matrix 能直接揭示模型在哪两类之间频繁混淆。CIFAR10 中常见的混淆对是 cat 和 dog、bird 和 deer,这符合视觉直觉,因为这两组类别的外形和纹理确实接近。如果模型把 truck 和 automobile 混淆得很厉害,说明模型主要依赖整体轮廓特征,没有学到车斗和货箱这种细粒度差异。这一步比分对准确率更能判断数据集本身的可分性。

5.2 分类报告告诉你:准确率掩盖了哪些细节

sklearn 的classification_report会给每一类输出 precision、recall、f1-score 三个指标。top-1 准确率是整体正确率,但某一类如果只有 70% 的 recall,可能这类样本在增强后变得太难识别。这个函数还会自动计算宏平均和加权平均,方便对比不同类别的平衡程度。

CIFAR10 的类别是均衡的,每类 1000 张验证图,所以 micro 和 macro 指标差别不大。如果你的自己的数据集类别不均衡,就一定要看每个类别的 recall,而不是只看整体 acc。这直接影响后续要不要针对难分类别做过采样或加重 loss 权重。

5.3 把错例图像画出来,验证效果提升最快

保存错例图片是最直观也最容易被跳过的环节。把验证集中预测错误的样本统一画在一张大图上,不要只看数字,你会立刻发现很多规律:有些错例是图像本身就模糊到人眼也无法分辨,这类模型答错是数据集标注噪声,不用管;更多情况是增强强度过猛,比如 RandomCrop 把主体的关键部位裁掉了。

# visualize_error.py import math import torchvision.utils as vutils wrong_idx = np.where(all_preds != all_labels)[0] wrong_images = torch.stack([val_ds[i][0] for i in wrong_idx[:16]]) grid = vutils.make_grid(wrong_images, nrow=4, normalize=True) plt.imshow(grid.permute(1, 2, 0)) plt.axis("off") plt.savefig("wrong_predictions.png", dpi=120)

这里的 normalizze=True 会把归一化后的张量反变换回 0-1 范围显示,避免图片发黑或发灰。如果发现错例集中出现在某几个类别,回到 5.2 的 classification_report 里对照 recall,就能定位是数据问题还是模型能力不足。

5.4 验证的最后一环:单独留一份原始未增强的样本

我在做验证时会把原始 CIFAR10 测试集额外存一份,不套任何增强,专门给训练好的模型做最终评估。这是因为训练过程中我们反复用同一份验证集选最优模型,存在轻微的选择偏差,真实上线时的表现往往比验证 acc 低 0.5 到 1 个点。如果你只需要一个数字证明模型有效,用增强后的验证集没问题;但如果要写报告或对比不同模型,务必用未增强的原始样本重新跑一次 final test,这个数字才是可信的。

6. 用attention map复盘模型学到了什么:一个10行的可视化技巧

验证准确率达标只说明模型能分清类别,但不说明它看对了目标。CIFAR10 上有一个容易被忽略的问题:背景和主体高度耦合,比如 airplane 的背景常为蓝天,ship 的背景常为水面,模型可能学到了背景特征而非物体本身。要验证这一点,可以用第 2 章代码里保存的注意力权重,把 CLS token 对每个 patch 的注意力画成热力图,看看模型在分类时到底在关注图像哪个区域。

# visualize_attention.py import numpy as np import matplotlib.pyplot as plt from torchvision import transforms def show_attention(model, image_tensor, layer_index=8, head_dim=0): model.eval() with torch.no_grad(): model(image_tensor.unsqueeze(0)) attn = model.blocks[layer_index].attn_w[0] # (heads, 65, 65) attn = attn.mean(dim=0)[0, 1:].reshape(8, 8) # 平均所有head,去掉CLS attn = np.kron(attn.cpu().numpy(), np.ones((4, 4))) # 放大到32x32 mean, std = (0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010) img = image_tensor.permute(1, 2, 0).numpy() img = img * np.array(std) + np.array(mean) # 反归一化 plt.imshow(img) plt.imshow(attn, cmap="jet", alpha=0.5) plt.savefig("attention_map.png", dpi=120)

使用步骤很简单:先加载best_vit_cifar10.pth权重,从 val_ds 里取一张图,直接调用上面的函数。如果热力图集中在目标物体的轮廓上,说明模型学到了可解释的区域特征;如果热力图分散在背景甚至四个角上,说明模型在靠背景作弊,这时要做的就是加强数据增强里的 RandomCrop 强度,或者用 Cutout 强制模型去关注局部主体。

我自己的习惯是每次训练完成都至少看三层注意力:第 2 层看 low-level 边缘特征、第 5 层看局部纹理、最后一层看分类决策依据。曾有模型 val_acc 到了 90%,但 attention map 显示它一直在看天空区域,那个 90% 在真实场景里根本靠不住。这个习惯帮我避免了好几次把数据集 bias 当成模型能力的误判。希望这条经验对你也有用,动手从第一张 attention map 开始吧。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询