ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

ViT图像分类实战:从Patch Embedding到注意力可视化

ViT图像分类实战:从Patch Embedding到注意力可视化 简介这是一份基于Vision Transformer的图像分类项目实现随附可直接使用的数据集。项目面向计算机相关专业在校学生与初入深度学习领域的开发者适用于课程设计、期末课设、毕业设计或初期项目立项演示也可作为Transformer图像分类学习与二次开发的起点。压缩包共32个文件以12个Python脚本为主涵盖模型定义、数据读取、训练、预测与工具函数等完整流程并包含项目说明文档、文本说明与JSON配置整体仅66KB轻量易部署。目前已有325人学习浏览代码结构清晰、模块划分明确便于对照论文理解Vision Transformer在图像分类中的实现细节。包内已包含可直接运行的训练与预测脚本及配套数据测试运行通过基础薄弱者可按说明逐步跑通基础较好的读者可在此基础上修改模型结构或数据流程扩展至其他分类任务。1. 从课设出发认识 Vision Transformer 图像分类课程设计群里十个做图像分类的八个还在用 ResNet。导师一句“要用最新的图像分类模型”就把你推进了 vision transformer 的坑。这里要讲的是这类课设项目在 Python 里从零落地 ViT 的完整路径如何把数据集切成 patch、堆叠 transformer block、训练到能在验证集上交差。我不打算直接用timm一行加载模型因为课设答辩时老师会追问“位置编码是什么”。适合已经会 Python 基础、想深入 ViT 源码结构、又不想只做调包侠的同学。沿着这个思路走你会得到一份可复现的工程骨架。2. ViT 结构拆解用 CNN 的直觉理解 Patch Embedding 与注意力2.1 为什么课设项目选 ViT 而不是 CNN图像分类经典的做法是卷积神经网络。CNN 靠卷积核扫描局部区域再通过堆叠深度扩大感受野。Vision Transformer简称 ViT完全舍弃卷积把图片分割成固定大小的 patch序列化成 token然后用 Transformer 的编码器做全局注意力建模。看一张 224×224 的图patch16会得到 196 个 token每一个 token 都有机会直接和所有其他 token 计算关联这会带来两个直接变化第一模型一上来就有全局视野不需要像 CNN 那样靠深度堆叠第二自注意力的计算量随 patch 数量的平方增长所以 ViT 通常需要更多数据和更强正则才能超过 CNN。对课设而言ViT 的优势在于它足够新答辩时有话题可讲同时它的结构清晰Patch Embedding、位置编码、Transformer Encoder、分类头每一块都可以独立画出网络结构图。缺点是数据需求量高所以我们需要在数据准备和训练技巧上多花功夫。这里采用 PyTorch 实现因为 PyTorch 的动态计算图便于打印中间张量形状也方便把每个模块拆开逐步验证。搭配标准的 CIFAR-10 或花卉数据集一样能跑出让人满意的结果。2.2 Patch Embedding 与位置编码的 Python 实现ViT 里第一步是把图像切成 patch这一步一般用卷积层一次完成。卷积核大小等于 patch 大小步长也等于 patch 大小输出通道等于 embedding 维度。下面是最小实现import torch import torch.nn as nn class PatchEmbed(nn.Module): def __init__(self, img_size224, patch_size16, in_channels3, embed_dim192): super().__init__() self.n_patches (img_size // patch_size) ** 2 self.proj nn.Conv2d( in_channels, embed_dim, kernel_sizepatch_size, stridepatch_size ) def forward(self, x): # 输入 x: [B, 3, 224, 224] x self.proj(x) # [B, embed_dim, 14, 14] x x.flatten(2) # [B, embed_dim, 196] x x.transpose(1, 2) # [B, 196, embed_dim] return x代码里最关键的是nn.Conv2d的使用它把每个 patch 映射成一个向量同时保持了局部空间结构。flatten(2)保留 batch 和 channel 维度把最后两维压平transpose把 token 数放到中间维度方便后续输入 Transformer。如果你把patch_size改成 32token 数量会变成 49序列变短全局建模能力下降计算量明显降低这是课设里可以演示的敏感度实验。位置编码我通常使用可学习参数因为它在小数据集上比固定三角函数更灵活self.pos_embed nn.Parameter(torch.zeros(1, self.n_patches 1, embed_dim)) # 加 1 是因为要预留 class token 的位置这里的1对应后面提到的 class token。参数形状是 (1, 197, 192)训练时会在正向传播中加到 patch token 上。注意初始化最好用截断正态分布避免开始时位置信号压过 Patch Embedding 的输出。2.3 Transformer Encoder 与分类头怎么拼得到 patch token 序列后ViT 会在序列最前面拼接一个可学习的 class token它的初始向量在训练中逐渐聚合整个图像的信息。之后堆叠 L 个 Transformer Encoder Block每个 Block 由多头自注意力、MLP、LayerNorm 和残差连接组成。核心实现如下class TransformerBlock(nn.Module): def __init__(self, embed_dim, num_heads, mlp_ratio4.0, dropout0.1): super().__init__() self.norm1 nn.LayerNorm(embed_dim) self.attn nn.MultiheadAttention(embed_dim, num_heads, dropoutdropout, batch_firstTrue) self.norm2 nn.LayerNorm(embed_dim) self.mlp nn.Sequential( nn.Linear(embed_dim, int(embed_dim * mlp_ratio)), nn.GELU(), nn.Dropout(dropout), nn.Linear(int(embed_dim * mlp_ratio), embed_dim), nn.Dropout(dropout) ) def forward(self, x): # 第一个残差分支 x x self.attn(self.norm1(x), self.norm1(x), self.norm1(x))[0] # 第二个残差分支 x x self.mlp(self.norm2(x)) return x注意nn.MultiheadAttention在 PyTorch 2.x 里需要设置batch_firstTrue这样输入输出都是[B, seq_len, embed_dim]。我在实际实现时会在每个 Block 后打印输出形状确认序列长度没有变化只有embed_dim是稳定的。分类头很简单取出 class token 对应的向量过一层 LayerNorm 和 Linear 即可self.head nn.Linear(embed_dim, num_classes) # forward 返回 logits到这里整个 ViT 的主干就拼完了。把PatchEmbed、pos_embed、TransformerBlock列表和head组合成一个nn.Module就是一个完整的 ViT 模型。课设项目里如果要求画网络结构图这个三层的组合式结构直接拿来对应论文的 Figure 1。3. 数据准备用 Python 把数据集变成 ViT 能吃的张量3.1 数据集选择与标签处理课设中我一般建议先选小图数据集比如 CIFAR-10、CIFAR-100 或者 Oxford Flowers-102。CIFAR-10 只有 32×32如果直接输入 ViT 会用 patch4 或者先 resize 到 224否则 patch16 会导致 token 数太少。Oxford Flowers-102 是天然的高分辨率图像类别也适中适合展示 ViT 对细节的区分能力。要注意一点torchvision 里的ImageFolder要求数据按类别文件夹存放目录结构类似dataset/ train/ class_0/ img1.jpg class_1/ ... val/如果你的数据集是 CSV 带标签的需要自己写一个Dataset子类让__getitem__返回(image_tensor, label_index)。标签一定要从 0 开始连续编号否则CrossEntropyLoss会报类索引越界。3.2 用 torchvision 完成增广、打乱和批次加载ViT 对数据增广更敏感。常见做法是先用transforms.Resize((224, 224))统一尺寸然后接RandomHorizontalFlip和RandomCrop做随机裁剪相当于给模型制造更多训练样本。注意RandomCrop需要先pad一下或者改用RandomResizedCrop。下面是一份可直接抄的配置from torchvision import datasets, transforms from torch.utils.data import DataLoader, random_split train_transform transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomResizedCrop(224, scale(0.8, 1.0)), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) val_transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) dataset datasets.ImageFolder(rootdataset/train, transformtrain_transform) train_ds, val_ds random_split(dataset, [0.8, 0.2]) train_loader DataLoader(train_ds, batch_size64, shuffleTrue, num_workers4, pin_memoryTrue) val_loader DataLoader(val_ds, batch_size64, shuffleFalse, num_workers4, pin_memoryTrue)这里的增广参数很有讲究scale(0.8, 1.0)限制了随机裁剪的面积范围避免切到只剩一角Normalize用的均值标准差是 ImageNet 统计值ViT 预训练权重通常依赖这个标准化方式如果从零训练也可以改用数据集自身统计值。num_workers在 Windows 上建议设成 0否则可能进入死锁。注意如果你的课设机器显存只有 4G把batch_size降到 32同时把pin_memory设为False否则训练时会把显存打满。3.3 数据划分与类别分布检查random_split虽然方便但不能保证类别比例一致。课设建议用分层划分先读一遍所有样本的标签用train_test_split的stratify参数划分。数据集规模小的时候我还会统计每个类别的样本数打印出来看是否有不平衡问题。ViT 在数据不平衡时更容易偏向高频类别因为注意力机制会把更多权重放在常见特征上。另外如果原始图片的宽高比差异很大Resize加RandomResizedCrop会丢失一些内容。此时可以在Resize前补Pad或使用CenterCrop。这些细节在课设报告里可以作为“数据预处理优化”写进去。我还习惯打印一个批次的数据形状确认[B, 3, 224, 224]正确因为 ViT 对输入尺寸非常敏感一旦 patch 数对不上后面的 Transformer 维度就会全部错位。4. 训练与调参从零跑通 ViT 图像分类的 Pipeline4.1 损失函数、优化器与学习率调度策略ViT 的分类任务仍然使用交叉熵损失。这里要重点说的是优化器。实践中最稳的是 AdamW而不是 Adam。AdamW 把权重衰减从梯度更新中解耦对 Transformer 这类大参数模型能明显降低过拟合风险。我常用的配置criterion nn.CrossEntropyLoss() optimizer torch.optim.AdamW(model.parameters(), lr1e-4, weight_decay0.05) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_maxepochs)lr1e-4是 224 分辨率下从头训练 ViT 的常用起点如果迁移学习可以调到 1e-5。weight_decay0.05对卷积头来说可能偏大但对 attention 矩阵的抑制效果很好。学习率调度我倾向于余弦退火前几个 epoch 先用线性 warmup。没有 warmup 的话大学习率会让 LayerNorm 的统计量在早期剧烈抖动。4.2 训练循环与验证循环的代码模板下面是一个可直接改造的训练循环核心支持分类准确率输出和模型保存def train_one_epoch(model, loader, optimizer, criterion, device): model.train() total_loss, correct, total 0.0, 0, 0 for images, labels in loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) # [B, num_classes] loss criterion(outputs, labels) loss.backward() optimizer.step() total_loss loss.item() * images.size(0) correct (outputs.argmax(dim1) labels).sum().item() total images.size(0) return total_loss / total, correct / total def validate(model, loader, criterion, device): model.eval() total_loss, correct, total 0.0, 0, 0 with torch.no_grad(): for images, labels in loader: images, labels images.to(device), labels.to(device) outputs model(images) loss criterion(outputs, labels) total_loss loss.item() * images.size(0) correct (outputs.argmax(dim1) labels).sum().item() total images.size(0) return total_loss / total, correct / total这个模板里outputs.argmax(dim1)直接取 logits 中最大值索引作为预测类别。model.eval()和torch.no_grad()必须成对使用否则 Dropout 和 BatchNorm 会干扰验证结果。注意 ViT 中的 Dropout 在训练时是开着的验证时要自动关闭。4.3 超参数速查表与调参顺序超参数从头训练建议值迁移学习建议值备注batch_size64 或 12864受显存限制ViT 在大图上很吃显存lr1e-4 到 3e-41e-5 到 1e-4迁移时最好降低 lrweight_decay0.050.01 到 0.05常用 0.05epochs50 到 10020 到 30CIFAR-10 从零训练建议 100warmup_epochs52线性从 0 升到 lr调参顺序我的习惯是先固定 batch_size 和 patch_size用较小的 embed_dim 跑通一次完整训练确认 loss 有下降再逐步增大 embed_dim 或增加 Transformer Block 数最后才调 weight_decay 和 dropout。不要一开始就上 ViT-Base课设机器通常扛不住。mini 版本常用配置是 patch_size16embed_dim192num_heads3depth4参数量在 500 万以下CIFAR-10 上都能跑出 85% 左右的准确率。5. 评估与可视化准确率、混淆矩阵和注意力热力图5.1 评估指标与混淆矩阵绘制训练完只看 loss 曲线不够课设报告中必须给出混淆矩阵。它比准确率更直观地展示哪些类别容易混淆。用 sklearn 一行就能计算from sklearn.metrics import confusion_matrix, classification_report import numpy as np all_preds, all_labels [], [] model.eval() with torch.no_grad(): for images, labels in val_loader: images images.to(device) preds model(images).argmax(dim1).cpu().numpy() all_preds.extend(preds) all_labels.extend(labels.numpy()) cm confusion_matrix(all_labels, all_preds) print(classification_report(all_labels, all_preds, digits3))画热力图时用matplotlib的imshow横纵轴都写类别名。注意打印 classification_report 时每个类别的 f1-score 能直接暴露模型对哪个类不敏感。ViT 通常对纹理简单的类有优势对背景和物体颜色接近的类容易混。5.2 用 Rollout 可视化注意力权重如果要给答辩增加亮点可以可视化 ViT 的注意力热力图。常见的方法是 Attention Rollout把多头注意力的权重按层累乘得到一个从输入 token 到分类结果的注意力分数矩阵。实现思路是向前传播时缓存每层nn.MultiheadAttention的 attention weights然后逐层与恒等矩阵做矩阵乘法并求平均。def rollout_attention(model, image): # image: 预处理后的单张图片 [1, 3, 224, 224] model.eval() attn_weights [] hooks [] def hook_fn(module, input, output): # 取出 attention_probs形状 [B, num_heads, query_len, key_len] attn_weights.append(module.attention_probs.detach().cpu()) for layer in model.blocks: hooks.append(layer.attn.register_forward_hook(hook_fn)) with torch.no_grad(): model(image) for h in hooks: h.remove() # 对多头取平均 attn_maps torch.stack([w.mean(dim1) for w in attn_weights]) # [num_layers, query_len, key_len] # 加上恒等连接并递归乘 rollout torch.eye(attn_maps.size(1)) for a in attn_maps: rollout a rollout rollout (rollout - rollout.min()) / (rollout.max() - rollout.min()) return rollout[0, 1:] # 去掉 class token 所在行得到每个 patch 的注意力分数注意这里的 hook 需要注册在MultiheadAttention模块上PyTorch 2.x 中 attention 权重通常会缓存在module.attention_probs属性中。如果你的 PyTorch 版本不同可能需要修改 hook 内部逻辑。最终热力图可以用matplotlib缩放到原图尺寸叠加显示。5.3 从结果反推哪里需要调如果混淆矩阵中某两类频繁混淆首先检查这两类是不是包含相似的全局形状比如猫和狗。ViT 的全局注意力会放大这种相似性此时可以增加 Attention Dropout 或增大数据增广强度。如果热力图显示注意力集中在背景区域说明 class token 没有学到合理的聚合方式尝试在损失函数中加入对比正则或者使用更强的位置编码初始化。另一个常见现象是验证 loss 持续下降但准确率不再上涨这说明模型在过度自信可以降低 dropout 或者使用标签平滑。criterion nn.CrossEntropyLoss(label_smoothing0.1)标签平滑对 ViT 效果非常明显它能把硬标签变成软标签防止 class token 输出过于极端的概率分布。这个技巧在课设报告里可以单独作为一个小节写。6. 课设报告与答辩前必须会的 3 个进阶技巧6.1 用预训练权重做微调ViT 在 ImageNet-21k 上预训练的权重迁移到小数据集通常比从零训练高 10 个百分点以上。PyTorch 官方有torchvision.models.vit_b_16可以直接加载预训练权重。用weightstorchvision.models.ViT_B_16_Weights.IMAGENET1K_V1即可然后把分类头替换成自己的类别数。微调时把 stem 和前几个 block 冻结或使用较小的学习率只更新后端和分类头可以明显减少过拟合。import torchvision.models as models model models.vit_b_16(weightsmodels.ViT_B_16_Weights.IMAGENET1K_V1) model.heads.head torch.nn.Linear(model.heads.head.in_features, num_classes)代码里heads.head是分类层替换后只训练这一层的随机初始化参数其他层用 1e-5 学习率微调。6.2 导出模型为 TorchScript无论用什么框架交作业时不能只交.pth。常见做法是导出为 TorchScript 或者 ONNX便于写一个可执行的 demo。TorchScript 导出很简单scripted_model torch.jit.script(model, example_inputs[torch.rand(1, 3, 224, 224)]) scripted_model.save(vit_cifar10.pt)用torch.jit.script时要注意代码里不要有动态 if 依赖张量形状否则会报错。导出后可以用torch.jit.load验证输出是否一致。6.3 用梯度加权热力图解释分类依据Attention Rollout 是纯 forward 的注意力分布另一种更符合直觉的解释是用分类结果对最后一个 Transformer block 的输出做梯度加权得到每个 patch 对最终预测的贡献。这个技巧不需要改模型结构只需要在 backward 时保留梯度然后用torch.autograd.grad计算。放到课设报告里会让整个项目从“调包跑通”变成“可以解释的模型”。在答辩现场把热力图和分类错误案例放在同一页幻灯片里比任何文字说明都有效。本文还有配套的精品资源点击获取
返回列表