ARTICLE DETAIL

资讯详情

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

ViT图像分类实战:PyTorch实现与训练调优指南

ViT图像分类实战:PyTorch实现与训练调优指南 简介VITVision Transformer图像分类项目将Transformer架构首次引入计算机视觉领域通过把图像切分为Patch序列进行建模以纯注意力机制替代传统卷积网络面向深度学习与CV方向学习者演示完整的图像分类流程是理解近年视觉模型演进的核心实践。资源包共2000个文件以大量jpg图像数据作为训练主体同时配有Python源码含模型定义、训练与预测脚本、xml标注、json配置、名称映射文件以及训练好的pth权重压缩包约539.35MB整体目录结构清晰便于按模块学习与二次开发。代码完整可一键运行自带数据集与预训练模型分类精度高达99%以上实测可直接跑通也可迁移到花卉、车型等自定义分类场景适配课程设计、算法研究或工程快速原型验证。目前已有12407人学习下载既能帮助初学者理解Transformer在视觉领域的应用也可作为进阶开发者的代码参考。1. 从 CNN 到 ViT图像分类任务里的 Transformer 之路ViTVision Transformer把一整张图切成 16×16 的 patch拉平成 token 序列后直接送进 Transformer 编码器这种思路在 2020 年刚出现时被很多人看成是 CNN 之外的奇技淫巧。但在超大规模数据集上预训练后ViT 的图像分类精度超过了同量级 ResNet并且迁移到 ImageNet 时表现出更好的可扩展性如今它已经是任何图像分类算法对比里绕不开的基线模型。接下来的内容会把 patch embedding、位置编码、分类头逐个拆开给一份能在 CIFAR-10 上跑通的最小 PyTorch 实现再把训练时最容易影响精度的学习率、数据增强和预训练权重用法讲透。适合已经掌握 PyTorch 基础、想快速进入 vision transformer 极简入门并理解每个参数在做什么的工程师。2. ViT 图像分类的原理拆解Patch、位置编码与分类头2.1 为什么要先做图像分块Patch Embedding 的两种写法一张 224×224×3 的图片如果直接把每个像素当作一个 token序列长度就是 150528自注意力的计算复杂度是 O(N²)这个规模任何 GPU 都扛不住。ViT 的做法是把图片切成 P×P 的小块P 通常取 16 或 14每个小块称为一个 patch。224×224 按 16×16 切得到 196 个 patch每个 patch 展平后的向量长度是 16×16×3768。这个长度仍然偏大所以还需要一个可学习的线性投影把 768 维压缩到指定的 Transformer 维度 D比如 768 或 1024。这一步叫 patch embedding本质是把每个 patch 映射到一个语义向量。在 PyTorch 里patch embedding 最简洁的实现是用一个卷积层代替手工切图和线性层。卷积核大小和步长都设为 patch_size输入通道数等于图片通道数输出通道数等于 D。因为卷积天然是按局部窗口滑动的它会把每个 P×P 窗口内的像素展开后乘上权重矩阵效果等价于对每个 patch 做线性变换。代码只有两行import torch.nn as nn # 用卷积实现 patch embeddingkernel_size 和 stride 都等于 patch_size self.patch_embed nn.Conv2d( in_channels3, out_channelsdim, kernel_sizepatch_size, stridepatch_size )forward 里拿到卷积输出后形状为 (B, dim, H/patch_size, W/patch_size)需要把它变成 (B, num_patches, dim) 才能交给 Transformer。写法是x.flatten(2).transpose(1, 2)。flatten(2) 把最后两个维度合并成 num_patchestranspose 再把通道维挪到序列维之后。这里的 dim 就是 Transformer 的隐藏维度也是后续所有 token 的统一表示长度CIFAR-10 这种小数据集上取 128 或 256 就够ImageNet 级别一般要 768 以上。提示patch_size 直接决定序列长度。224×224 用 patch16 时序列长度为 196用 patch8 时变成 784计算量约为前者的 16 倍。小数据集上 patch4 或 8 能保留更多局部细节但训练会明显变慢。2.2 位置编码与 CLS TokenTransformer 感知顺序的两种机制Transformer 的多头注意力本身是排列不变的。把 patch 序列随机打乱如果不加位置信息模型输出完全一样这显然不符合图像的空间结构。ViT 的做法是给每个位置初始化一个可学习的向量加到对应 patch token 上。这个位置编码矩阵形状是 (1, num_patches, dim)初始化一般用均值为 0、标准差为 0.02 的截断正态分布。num_patches 需要在创建模型时就固定所以 ViT 要求输入图片尺寸固定。想支持任意尺寸需要做插值或改为相对位置编码后面会提到。序列长度之外还有一个关键设计CLS token。BERT 在句子开头加一个[CLS]ViT 沿用同样的思路在 patch token 序列最前面拼接一个可学习的向量让它在整个编码过程中与其他 patch 做自注意力。编码结束后只取第一个位置的输出作为全局图像特征接一个线性分类头。为什么不用全局平均池化原论文对比过CLS 与 GAP 精度接近但 CLS 是预训练和微调之间更统一的接口。实践中有的工作直接去掉 CLS用所有 token 的均值池化也能得到相当的结果但 CLS 写法更通用。两个参数在__init__里这样初始化self.cls_token nn.Parameter(torch.randn(1, 1, dim) * 0.02) self.pos_embed nn.Parameter(torch.randn(1, num_patches 1, dim) * 0.02)这里 num_patches1 是因为多了一个 CLS token。注意 cls_token 的初始值不要设成零否则前几步所有 CLS token 完全相同梯度更新后可能长期停留在一个对称状态。乘 0.02 是为了让初始嵌入幅度远小于 patch embedding避免一开始就压过输入特征。2.3 Transformer 编码器与分类头归纳偏置去了哪里有了 token 序列接下来就是标准的 Transformer Encoder。每个 block 做四件事LayerNorm、多头自注意力、带 GELU 的 MLP、残差连接。用公式表示第 l 层的输出 z_l 为z_l MSA(LN(z_{l-1})) z_{l-1} z_l MLP(LN(z_l)) z_l其中 MSA 把 D 维向量分成 h 个头每个头独立计算 Q、K、V 的注意力权重然后拼接回 D 维。MLP 一般是先升维到 4D再降回 D中间用 GELU 激活。与 CNN 相比ViT 几乎没有内置归纳偏置。下面的表格整理了二者的关键差别这也是理解 ViT 训练难点的核心维度CNNResNet 为代表ViT局部性卷积核天然只看邻域第一层分块后全局自注意力平移等变性卷积权重共享平移后特征基本一致依赖数据学习无显式保证参数效率小数据下非常高效需要大数据或强正则才能发挥感受野浅层局部深层逐步扩大第一层就能看到整张图扩展性加大模型收益递减较快模型和数据同步增大时提升明显分类头在最后一个 block 后。先对 CLS 位置做一次 LayerNorm再通过线性层映射到类别数。如果类别数是 1000head 的形状就是 (D, 1000)。这一层初始化时一般用标准差 0.02 的小随机数微调时可以直接替换掉只保留前面编码器的预训练权重。到这一步ViT 图像分类的前向流程已经完整后面就是怎么把它写得能训练。3. 用 PyTorch 从零实现 ViT 图像分类的最小可运行代码3.1 搭建一个 6 层 MiniViT 模型这里给出一个不依赖任何第三方库的 ViT 实现总参数量在 2M 左右单张普通 GPU 甚至 CPU 都能在 CIFAR-10 上完成训练。模型输入尺寸设为 32×32patch_size 取 4这样序列长度为 64加上 CLS 一共 65 个 token。dim 取 128depth 取 6heads 取 8。关键组件都已写完import torch import torch.nn as nn class PatchEmbed(nn.Module): def __init__(self, in_ch3, patch4, dim128): super().__init__() self.proj nn.Conv2d(in_ch, dim, kernel_sizepatch, stridepatch) def forward(self, x): # x: (B, 3, 32, 32) - (B, dim, 8, 8) - (B, 64, dim) x self.proj(x) return x.flatten(2).transpose(1, 2) class Block(nn.Module): def __init__(self, dim128, heads8): super().__init__() self.norm1 nn.LayerNorm(dim) self.attn nn.MultiheadAttention(dim, heads, batch_firstTrue) self.norm2 nn.LayerNorm(dim) self.mlp nn.Sequential( nn.Linear(dim, dim * 4), nn.GELU(), nn.Linear(dim * 4, dim), ) def forward(self, x): # MultiheadAttention 返回 (attn_output, attn_weights) x x self.attn(self.norm1(x), self.norm1(x), self.norm1(x))[0] x x self.mlp(self.norm2(x)) return x class VisionTransformer(nn.Module): def __init__(self, img_size32, patch_size4, dim128, depth6, heads8, num_classes10): super().__init__() self.patch_embed PatchEmbed(3, patch_size, dim) num_patches (img_size // patch_size) ** 2 self.cls_token nn.Parameter(torch.randn(1, 1, dim) * 0.02) self.pos_embed nn.Parameter(torch.randn(1, num_patches 1, dim) * 0.02) self.blocks nn.ModuleList([Block(dim, heads) for _ in range(depth)]) self.norm nn.LayerNorm(dim) self.head nn.Linear(dim, num_classes) def forward(self, x): B x.shape[0] tokens self.patch_embed(x) # (B, N, dim) cls self.cls_token.expand(B, -1, -1) # (B, 1, dim) x torch.cat([cls, tokens], dim1) # (B, N1, dim) x x self.pos_embed for block in self.blocks: x block(x) # 取 CLS token 输出经过 LayerNorm 后接分类头 out self.norm(x[:, 0]) return self.head(out)代码里的nn.MultiheadAttention是 PyTorch 官方实现batch_firstTrue让输入输出形状都是 (B, seq_len, dim)和 torchvision 的 transformer 保持同一种习惯。Block 中先对输入做 LayerNorm 再进注意力这是 Pre-LN 结构目的是让梯度在深层传播时更稳定也是 ViT 最终从早期 Post-LN 架构修正过来的原因。MLP 中间层是 4 倍 dim这是一个从 BERT 沿袭下来的经验值增大到 4 倍以上收益有限但参数量和计算量却线性上升。3.2 加载 CIFAR-10 并执行训练CIFAR-10 有 60000 张 32×32 彩色图片类别是飞机、汽车、鸟等 10 类。它比 ImageNet 小很多但足以验证一个 ViT 实现是否正确。数据加载用 torchvision训练集加随机水平翻转和归一化。注意不能对测试集做随机翻转只能做归一化。这里用一个典型的数据管道from torch.utils.data import DataLoader from torchvision import datasets, transforms transform_train transforms.Compose([ transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)), ]) train_ds datasets.CIFAR10(./data, trainTrue, downloadTrue, transformtransform_train) train_dl DataLoader(train_ds, batch_size64, shuffleTrue, num_workers2)归一化使用的均值和标准差是 CIFAR-10 数据集的全局统计量。如果加载自己的数据集需要自己计算不能直接套 ImageNet 的数值否则图片整体偏暗或偏亮会让 ViT 的花卉图像分类或森林图像分类任务在前期损失下降变慢。训练循环直接写成最直观的形式没有额外封装device torch.device(cuda if torch.cuda.is_available() else cpu) model VisionTransformer(img_size32, patch_size4, dim128, depth6, heads8, num_classes10).to(device) opt torch.optim.AdamW(model.parameters(), lr3e-4, weight_decay5e-2) criterion nn.CrossEntropyLoss() for epoch in range(50): model.train() total_loss, correct, total 0.0, 0, 0 for x, y in train_dl: x, y x.to(device), y.to(device) opt.zero_grad() logits model(x) loss criterion(logits, y) loss.backward() opt.step() # 累计统计 total_loss loss.item() * x.size(0) correct (logits.argmax(1) y).sum().item() total y.size(0) acc correct / total print(fepoch {epoch 1}: loss{total_loss / total:.3f}, acc{acc:.3f})这里用了 AdamW学习率 3e-4weight_decay 0.05。ViT 的 attention 权重对学习率非常敏感用 SGD 或普通 Adam 时容易不收敛准确率会围绕 10% 上下跳动。AdamW 的 weight_decay 不要直接理解为 L2 正则它不会把梯度里的 weight_decay 项再乘一次学习率效果更干净。训练 50 轮后这个 MiniViT 在 CIFAR-10 测试集上大约能到 72%76% 的准确率比同规模的 ResNet-18 低 35 个百分点这是正常现象原因在 2.3 的表格里已经说明小型 ViT 缺少归纳偏置数据量又不够大。想继续提升精度直接用下一章的调整手段。3.3 评估代码与过拟合判断训练完成后做测试集评估需要把模型切成 eval 模式并整体包在torch.no_grad()里否则 Dropout 或 LayerNorm 的统计数据会出现偏差准确率会比真实水平高。下面的代码统计每个类别的正确数和总正确率model.eval() correct, total 0, 0 with torch.no_grad(): for x, y in test_dl: x, y x.to(device), y.to(device) pred model(x).argmax(1) correct (pred y).sum().item() total y.size(0) print(ftest acc: {correct / total:.3f})如果发现训练集准确率很高、测试集准确率低说明过拟合。ViT 在小数据集上的过拟合比 CNN 来得更快这时要优先检查 patch_size 是否过大比如 32×32 的图用 patch16 会把一张图压成 4 个 token信息保留太少。patch 越小模型保留的空间细节越多但序列变长后训练时间也会上升常用范围见下表数据集图片尺寸常见 patch_size序列长度CIFAR-1032×324 或 864 或 16Fashion-MNIST28×28449ImageNet224×22416196自定义高分辨率图224×224 以上16 或 32196 或 49小数据集上不要盲目照搬 ImageNet 的 patch16。比如 32×32 的图用 patch16N 只有 4位置编码几乎失去意义改用 patch4 之后模型才能区分局部纹理。序列长度对计算复杂度是平方关系所以表格里 64 与 16 的差距不是 4 倍而是 16 倍。4. ViT 图像分类训练参数与调优从数据增强到微调4.1 学习率是怎么决定 ViT 训练成败的ViT 对学习率的敏感程度远高于 ResNet。ResNet 用 0.1 的学习率和动量为 0.9 的 SGD 就能训练ViT 则普遍推荐 AdamW 配合 3e-4 以下的基础学习率。原论文里给出的规律是模型越大batch size 越大学习率可以适当提高但 Small/ViT-Base 级别超过 1e-3 后Loss 很容易冲高回不来表现为前几个 iteration 就出现 NaN 或准确率一直在 10%。除了基础学习率warmup 几乎是 ViT 训练的标配。原因有两个一是随机初始化的线性分类头在早期会产生很大的梯度噪声二是 Adam 在一阶矩估计还不稳定时过大的更新步长会破坏预训练或刚刚初始化的位置编码。常见做法是先线性从小到大上升再按余弦曲线衰减。PyTorch 1.13 之后可以用SequentialLR把两个调度器串起来from torch.optim.lr_scheduler import LinearLR, CosineAnnealingLR, SequentialLR warmup LinearLR(opt, start_factor0.01, end_factor1.0, total_iters10) cosine CosineAnnealingLR(opt, T_max90) scheduler SequentialLR(opt, schedulers[warmup, cosine], milestones[10])上面这段的含义是前 10 个 epoch学习率从 0.01 倍线性上升到目标值第 10 个 epoch 结束后切换到余弦调度在剩余 90 个 epoch 内从峰值缓慢衰减到接近 0。milestones 参数必须与 warmup 的 total_iters 保持一致它标记第一个调度器结束的位置。如果不做 warmupViT 在早期很容易损失震荡尤其是数据集本身分布明显不平衡时。4.2 数据增强与正则化弥补归纳偏置缺失ViT 没有 CNN 的平移等变性数据增强对最终精度的影响甚至比模型结构还大。中级规模数据集上仅增加一个 RandomResizedCrop就经常带来 3 到 5 个点的提升。推荐的组合按强度排列如下RandomResizedCrop、RandomHorizontalFlip、ColorJitter、RandAugment。前两个保留语义结构ColorJitter 提升颜色鲁棒性RandAugment 通过随机组合旋转、缩放、对比度等操作达到更强的正则效果。torchvision 的 v2 版本可以直接用from torchvision.transforms import v2 train_transform v2.Compose([ v2.RandomResizedCrop(size(32, 32), scale(0.8, 1.0)), v2.RandomHorizontalFlip(p0.5), v2.ColorJitter(brightness0.3, contrast0.3, saturation0.3), v2.ToImageTensor(), v2.ConvertImageDtype(torch.float32), v2.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)), ])RandomResizedCrop 的 scale 参数控制在原图上按多少比例裁剪。0.8 到 1.0 是防过拟合和保真度之间的中间值如果你的任务是森林图像分类训练样本里同一棵树在不同角度下尺度差异很大可以把下界调到 0.5增强模型对尺度变化的适应。但要注意裁剪过多会让类别特征消失比如一根树枝的局部照片根本没体现整棵树模型只能乱猜。常规的 CNN 花卉图像分类任务还会加一个随机旋转因为花卉不存在方向污染问题医学图像或带有明显朝向的数据集则不要加。正则化方面DropPath 比 Dropout 更适合 ViT。Dropout 作用在 token 的特征向量上把整个通道置零DropPath 则是在训练时随机跳过某个残差块效果更接近对网络深度的集成。timm 库中的 ViT 默认使用 0.1 到 0.2 的 drop_path_rate而自建的 MiniViT 没加这个模块所以小数据集上更容易出现过拟合。手动加的代价很小在 Block 的 forward 里用drop_path(x)替换x ...即可推理时该函数恒等返回。label smoothing 也是类别数较多时稳定训练的常用技巧交叉熵的 target 从 0/1 变成 0.1 和 0.9 的软标签模型不会过度自信。4.3 从预训练权重微调ViT 和 CNN 最大的不同干净地从零训练一个小型 ViT 在 CIFAR-10 上只有约 75% 准确率而加载在 ImageNet 上预训练过的 ViT-Base/16 做迁移即使只微调 20 个 epoch也能达到 96% 以上。这就是 ViT 在真实项目中的正确使用方式除非是在做视觉研究否则永远优先考虑预训练。使用 timm 库时微调代码很短import timm model timm.create_model( vit_base_patch16_224.augreg_in1k, pretrainedTrue, num_classes10, )这个模型期望输入是 224×224×3所以数据加载时要把图片 resize 到 224。如果显存不够可以换vit_small_patch16_224或vit_tiny_patch16_224前者参数量约 22M后者约 5M。加载后整个模型的感知能力来自预训练只替换最后的分类头因此最初几个 epoch 的分类头随机权重会产生较大梯度需要把基础学习率降到 3e-53e-4并用只替换输出维度的方式重初始化。微调时需要区分不同层的学习率。常见做法是分类头学习率乘以 10其他层保持小学习率position embedding 也建议使用较小的学习率因为它保存的是预训练数据中的绝对位置语义变化过快会破坏 patch 之间的空间关系。如果预训练输入是 224而你的数据是 384位置编码需要通过插值放缩在 timm 中传img_size384会自动处理自实现时则需要先F.interpolate再反序列化这段逻辑建议用现成库而不是自己写。超参数从零训练CIFAR-10预训练微调自定义数据集基础学习率3e-43e-51e-4分类头学习率同基础基础学习率 × 10weight_decay5e-21e-41e-3warmup epochs1035总 epochs50~20020~50这张表是针对 ViT 图像分类最常见的两组设置。从零训练为了对抗小数据过拟合weight_decay 会拉得很大微调时预训练特征已经很强weight_decay 过大反而会压制模型表达所以降到 1e-4 量级。如果显存吃紧batch size 只能用 16那学习率也应该等比下降到原来的 1/4线性缩放规则在这里依然适用。5. 用注意力权重验证 ViT 图像分类模型是否学到了目标区域训练或微调完成后不能只看准确率。ViT 的一个巨大优势在于注意力权重天然就是可视化的解释渠道最后一个 block 的 CLS token 对所有 patch 的注意力分布可以告诉我们模型在分类时主要盯着图像的哪个区域。如果注意力集中在背景而不是目标对象说明数据集存在偏差或者训练轮数不够。下面这段代码修改了 Block 的 forward让它在推理时同时返回 attention 矩阵class Block(nn.Module): def forward(self, x, need_weightsFalse): norm_x self.norm1(x) if need_weights: attn_out, attn_weights self.attn(norm_x, norm_x, norm_x) x x attn_out return x, attn_weights x x self.attn(norm_x, norm_x, norm_x)[0] x x self.mlp(self.norm2(x)) return x单张图片的注意力可视化这样取model.eval() with torch.no_grad(): x, _ test_ds[1] x x.unsqueeze(0) tokens model.patch_embed(x) cls model.cls_token.expand(x.size(0), -1, -1) tokens torch.cat([cls, tokens], dim1) tokens tokens model.pos_embed attn_sum None for block in model.blocks: tokens, attn block(tokens, need_weightsTrue) # attn: (B, heads, N, N)这里取 CLS 对所有 patch 的权重 if attn_sum is None: attn_sum attn.mean(dim1)[0, 0, 1:] else: attn_sum (attn_sum attn.mean(dim1)[0, 0, 1:]) / 2 p int(attn_sum.numel() ** 0.5) attn_map attn_sum.reshape(p, p)这里有一个细节没有继续把 CLS 输出传给最后的分类头但注意力是从 CLS 出发的已经能反映模型决策依据。求得 attn_map 后把它用F.interpolate放大到原图 32×32 大小再用 matplotlib 的imshow(attn_map, cmapjet, alpha0.6)叠加。观察时重点看注意力质量是否覆盖目标的大致轮廓不同 patch 的权重是否平滑有没有出现某个 patch 权重接近 1、其余全为 0 的极端情况。后者通常意味着位置编码分辨率不足或模型只学到 shortcut可以尝试增大 patch 数量或加深网络。如果想看更细的信号可以不要对多头取平均而是单独挑第 3 或第 5 个头不同头的注意力模式可能分别对应边缘、颜色区域和整体目标。综合考虑多个头的均值更容易获得稳定可解释的图。实践中还有一个反向用途把注意力图作为数据增强强度的参考。如果一个类别的注意力总是四散漂移优先增加该类别的旋转和缩放增强而不是一味调学习率。这就是可视化比准确率数字更进一步的原因。本文还有配套的精品资源点击获取
返回列表