资讯动态

ViT图像分类实战:用Transformer做猫狗分类的完整指南

发布时间:2026/9/8 12:00:40 来源:尧图企业网站定制
简介面向Transformer初学者与计算机视觉入门者项目以“猫狗大战”数据集为载体完整演示了如何用Vision TransformerViT完成图像分类任务。资源共2000个文件其中1998张jpg图片构成可直接使用的猫狗分类数据集2个py脚本分别负责模型训练与推理整体压缩包218.41MB目录结构清晰下载后即可开始动手实践。通过项目实战读者能深入理解ViT的Patch划分、位置编码、自注意力等核心机制同时掌握一套通用的图像分类流程从数据加载、模型构建到训练评估代码均有清晰注释与逻辑分层便于逐行调试和学习。更重要的是项目具备良好的泛化设计只需修改数据集路径和类别数目即可将同一套模型迁移到花卉、交通标志等任意图像分类场景真正实现“一次学会、多处复用”。目前已有1448人学习或下载适合不同基础的学习者循序渐进地迈入视觉Transformer的大门。 做图像分类这几年我大部分时间都在跟CNN打交道。ResNet、EfficientNet、MobileNet换着用性能好、生态成熟、部署也方便。直到我拿Kaggle的猫狗大战Dogs vs Cats数据集用Vision TransformerViT完整跑通了一个图像分类项目才真正体会到“图像也可以像文本一样被当作序列来处理”到底是什么意思。这篇文章不念论文从一个可复现的猫狗分类项目出发把ViT的工程细节、位置编码选择、核心代码实现、训练调参和各类坑位一次性拆清楚。适合刚接触Transformer、想真正跑一个ViT图像分类任务的人也适合那些已经在用CNN做分类但还没搞清楚ViT到底该怎么落地的朋友。1. 为什么用ViT做猫狗分类1.1 “把图片变成一句话”ViT的核心思路ViT的做法一句话总结就是把图像切成固定大小的patch每个patch当作一个token所有token拼起来送进Transformer Encoder。拿猫狗分类来说一张224x224的图切成16x16的patch就是14x14196个patch展平后变成196个token每个token是一个16x16x3768维的向量。这个操作在ViT里叫Patch Embedding。听起来平平无奇但它的核心思想是不预设图像的局部先验也就是不提前告诉模型“相邻像素关系更紧密”所有patch之间的关系完全靠self-attention自己学。这和CNN靠卷积核天然拥有局部感受野的思路是两个完全不同的方向。ViT没有卷积里的平移等变性和局部归纳偏置所以如果数据量不够它反而比CNN难训练。猫狗分类这种中等规模数据集正好能逼你去思考“数据量、预训练、正则化”这三个ViT绕不开的问题特别适合用来做ViT的上手项目。1.2 位置编码选型1D可学习、2D还是相对位置编码ViT里有个特别容易踩坑的点就是位置编码。图像patch本身是二维排列的所以很多人一开始就想当然用二维位置编码或者纠结要不要用相对位置编码。实际上ViT原论文已经做了实验对比结论是在ImageNet这种规模的数据集上1D可学习位置编码、2D可学习位置编码以及相对位置编码最终精度差距很小1D可学习位置编码在工程实现上最简单、参数也最少。原因也好理解。位置编码的作用只是让attention能区分不同token的位置关系并不需要编码完整的二维坐标信息。Transformer Encoder是置换不变的也就是打乱token顺序结果完全一样所以必须靠位置编码把顺序感注入进去。1D可学习位置编码的做法是初始化一个(num_patches1, hidden_dim)的矩阵然后在forward里加到patch embedding上。那个加1是因为额外多了一个全局的CLS token也要有对应位置编码。我在项目里试过把1D可学习位置编码换成2D版本精度几乎没有差异但代码复杂度和调试难度明显上升。所以结论很明确默认就用1D可学习位置编码别在这个地方给自己找事。2. 数据准备与实验环境2.1 环境依赖与硬件说明我这次用的环境是Python 3.10 PyTorch 2.1训练在一张RTX 309024GB显存上完成。如果显存不够后面我会讲怎么把batch size调到很小时照样训练普通8GB显存的卡也能跑需要付出的只是时间。核心依赖就这几个pip install torch torchvision timm einops matplotlib tqdmtimm不是必须的但它提供了大量现成的ViT预训练权重和通用的训练工具非常适合快速验证想法。einops用来重排张量维度做patch reshape时非常爽比手动view清晰得多。2.2 猫狗数据集组织与预处理细节Kaggle的Dogs vs Cats数据集训练集有25000张图猫狗各半。我建议把数据集重新组织成以下结构PyTorch的ImageFolder直接能读data/ ├── train/ │ ├── cat/ │ └── dog/ └── val/ ├── cat/ └── dog/数据划分上我习惯把官方训练集重新按8:2划分比如训练集20000张、验证集5000张。验证集不要动训练完直接用它评估。预处理这一步非常关键ViT对数据归一化的敏感性比CNN更高。我使用的配置是from torchvision import transforms train_transform transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.8, 1.0)), transforms.RandomHorizontalFlip(), transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.2), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) val_transform transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ])这里的mean和std直接用了ImageNet的统计值因为后面要用ImageNet预训练权重。如果是从零训练理论上应该重新算数据集的mean和std但我实测下来用ImageNet统计值影响也不大可以直接沿用。Resize到224是ViT-base和ViT-small的默认输入尺寸patch size 16时刚好切成14x14如果size不是16的整数倍会报错。这是新手最常犯的错误切记。3. 从零实现一个简化版ViT这一节手写一个能跑通的简化版ViT。虽然直接用timm加载现成模型更省事但手推一遍能帮你看清patch embedding、位置编码、self-attention这些模块到底在干什么。理解了这段代码后面调参才有方向感。3.1 Patch Embedding把图片切成词ViT的Patch Embedding通常用nn.Conv2d实现比手动unfold再reshape更高效、写法也更简洁。把输入(B, 3, 224, 224)通过一个conv核大小和步长都是16输出通道是768直接得到(B, 768, 14, 14)然后再flatten成(B, 196, 768)import torch import torch.nn as nn class PatchEmbed(nn.Module): def __init__(self, in_channels3, patch_size16, embed_dim768): super().__init__() self.proj nn.Conv2d(in_channels, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): # x: (B, 3, 224, 224) - (B, 768, 14, 14) - (B, 196, 768) x self.proj(x) x x.flatten(2).transpose(1, 2) return x这里用卷积实现的好处是PyTorch的conv对底层做了优化还能自动处理padding和stride。你要是用F.unfold手动切patch再逐行拼接代码又长又慢没有必要。3.2 Transformer Encoder与分类头Transformer Encoder的核心是Multi-Head Self-Attention。它的逻辑一句话概括每个token都去跟所有token做相似度计算然后按相似度加权聚合信息。在猫狗分类里它可能学到“耳朵”这个patch跟“头部”附近几个patch的关联更强。实现上我直接用一个nn.MultiheadAttention简化代码但要注意batch_firstTrue否则张量维度顺序很容易搞混class EncoderBlock(nn.Module): def __init__(self, dim, num_heads, mlp_ratio4.0, dropout0.1): super().__init__() self.norm1 nn.LayerNorm(dim) self.attn nn.MultiheadAttention(dim, num_heads, dropoutdropout, batch_firstTrue) self.norm2 nn.LayerNorm(dim) self.mlp nn.Sequential( nn.Linear(dim, int(dim * mlp_ratio)), nn.GELU(), nn.Dropout(dropout), nn.Linear(int(dim * mlp_ratio), 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整个模型的拼装流程是先加CLS token再加位置编码然后过N个EncoderBlock最后取CLS token对应的输出接一个线性分类头class ViT(nn.Module): def __init__(self, img_size224, patch_size16, num_classes2, embed_dim768, depth12, num_heads12, dropout0.1): super().__init__() self.patch_embed PatchEmbed(in_channels3, patch_sizepatch_size, embed_dimembed_dim) num_patches (img_size // patch_size) ** 2 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(*[ EncoderBlock(embed_dim, num_heads, dropoutdropout) for _ in range(depth) ]) self.norm nn.LayerNorm(embed_dim) self.head nn.Linear(embed_dim, num_classes) def forward(self, x): B x.shape[0] x self.patch_embed(x) # (B, 196, 768) cls_tokens self.cls_token.expand(B, -1, -1) x torch.cat([cls_tokens, x], dim1) # (B, 197, 768) 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)CLS token最初是BERT里用来做分类的原理大概是既然没有一个固定的patch代表整张图就额外放一个可学习的token进去它的最终表示融合了全局信息。ViT沿用了这个做法。你也可以不用CLS token改为对所有patch输出做全局平均池化效果差不多但CLS token的写法更贴近原论文方便跟后续的attention可视化工具对接。4. 训练策略与调参实战4.1 用预训练还是从零训练这是最现实的问题。ViT从零训练需要海量数据像JFT-300M那种3亿张图。猫狗数据集25000张图从零训练一个ViT-base效果大概率打不过ResNet18。所以我强烈建议用ImageNet预训练权重做迁移学习。具体做法有两种一是用timm直接加载预训练模型简单高效另一种是用上面的手写结构然后手动加载官方权重但要注意位置编码和分类头的维度匹配。我现在更倾向用timm先把流程跑通手写版本保留来学习原理import timm # 用ImageNet预训练权重只改分类头 model timm.create_model(vit_base_patch16_224, pretrainedTrue, num_classes2)timm会自动处理好patch embedding、位置编码、分类头的维度pretrainedTrue会下载ImageNet权重。如果你想看一个轻量版本可以换成vit_tiny_patch16_224或vit_small_patch16_224GPU负载小很多猫狗分类这种二分类任务用tiny也已经足够。注意修改num_classes2时原模型的分类头会被替换成随机初始化的新头。前几轮训练时分类头的梯度会远大于其他层容易把预训练特征弄乱所以我会在transformer encoder部分加一点额外的weight decay甚至把backbone冻结前几层。4.2 超参数配置与训练流程我这次实践采用的超参数如下不是一个绝对标准但经过多轮对比后稳定性和效果都不错可以直接作为起点参数数值说明optimizerAdamW比Adam更稳对大模型更友好base_lr1e-4迁移学习时不宜太大weight_decay0.05ViT要比CNN更大的weight decaybatch_size6424GB显存可跑小显存配合AMP降到16~32warmup_epochs2学习率先线性升再cosine降total_epochs10迁移学习10轮就能收敛到不错效果label_smoothing0.1对防过拟合有效ampTrue混合精度训练显存和速度双优化训练循环我习惯封装成下面这种简洁风格关键是不要忘了开启混合精度和梯度缩放scaler torch.cuda.amp.GradScaler() for epoch in range(epochs): model.train() for images, labels in tqdm(train_loader): images, labels images.cuda(), labels.cuda() optimizer.zero_grad() with torch.amp.autocast(cuda): outputs model(images) loss criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()学习率采用warmup cosine schedule。用timm.scheduler可以一行搞定from timm.scheduler import CosineLRScheduler scheduler CosineLRScheduler( optimizer, t_initialepochs, warmup_twarmup_epochs, warmup_lr_init1e-6, lr_min1e-6, )4.3 训练曲线怎么读我实测的预训练ViT-base跑10个epochvalidation accuracy大概能到98%以上。前2个epoch是warmup阶段loss下降很快然后进入平稳上升期。这里最值得关注的是train loss和val loss之间的gap如果val loss开始反弹而train loss还在降说明过拟合了得提前停掉或加大正则。如果是自己写的这种从零训练的ViT效果会难看得多同样10个epoch可能只有85%左右而且train loss下降得也慢。这个差距不是代码写错了而是数据量不够模型学不到足够的通用表征。遇到这种情况别慌迁移学习是正常的选项。5. 常见问题排查与避坑清单5.1 训练震荡不收敛症状loss反复横跳甚至越训练越高。我会按顺序排查三件事。第一确认有没有加位置编码很多人手写模型时漏了这一层attention在没有位置信息的情况下基本乱学。第二确认数据归一化有没有做ViT对输入分布很敏感忘了Normalize会导致训练极不稳定。第三把学习率降10倍看看Transformer对学习率的容忍度比CNN低很多从1e-3开始大概率直接发散。5.2 显存不足与训练过慢24GB显存全开batch_size 64如果只有8GB显存也别慌。把batch_size调到16甚至8同时开启混合精度。AMP在ViT上的收益比CNN还明显因为Transformer里面大量矩阵乘法和softmax计算用fp16能省掉一大半显存。如果batch_size调到4还不够就加梯度累积把梯度累积步数设为4等价于batch_size 16的效果accum_steps 4 scaler.scale(loss).backward() if (step 1) % accum_steps 0: scaler.step(optimizer) scaler.update() optimizer.zero_grad()训练速度真的跟不上就换vit_tiny_patch16_224参数少一个数量级猫狗二分类的精度不会有太大损失。5.3 验证集过拟合严重过拟合的最直接特征是train accuracy接近100%val accuracy却停在92%左右不动了。除了提前停止我一般按优先级做三件事把随机增强开满比如RandomResizedCrop的scale下限降到0.5加上RandomErasing把dropout从0.1调到0.2把weight_decay从0.05加到0.1。注意weight_decay太大也可能让loss收敛变慢要观察train loss是否还正常下降。5.4 迁移学习后精读还不如CNN如果你用预训练ViT微调效果还不如MobileNet先检查分类头的初始学习率。新初始化的分类头需要更大学习率但全网络同一个学习率时会让分类头学得太猛反而破坏backbone。我习惯让分类头的学习率乘以10backbone用较小的base_lr效果通常会有1~2个百分点的提升。另一个可能原因是你微调了太久导致模型把ImageNet上的通用特征都覆盖掉了。10轮内解决问题不要傻乎乎跑50轮。6. 我个人踩过的坑与后续扩展方向这次实践下来最深的感受是ViT的代码实现并不难难的是它把“数据规模”和“正则化”这两个问题的权重放大了。CNN因为天生有局部先验小数据也能硬扛ViT没有这个先验所以要么用预训练权重补数据要么用更强的正则硬控。我自己在写项目时踩过一个很蠢的坑从torchvision下载预训练权重然后直接赋值给手写模型结果因为cls_token和pos_embed的维度对不上跑了半天损失都不降。后来统一用timm的接口省心太多。这套猫狗分类流程稍加改动就能直接迁移到其他二分类或细粒度分类任务。我后来拿它做过森林图像分类类别从2类改成10类数据增强和训练流程几乎原封不动效果也相当能打。如果你想进一步深挖可以研究attention rollout把模型关注的patch可视化出来看看它到底是在看猫的脸还是在看背景。那一刻你会觉得Transformer做视觉这件事越看越有意思。本文还有配套的精品资源点击获取

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价