ARTICLE DETAIL

资讯详情

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

从Patch Embedding到PyTorch实现:Vision Transformer图像分类实战解析

从Patch Embedding到PyTorch实现:Vision Transformer图像分类实战解析 不用急着去看论文也不用一上来就啃ViT源码先把思路捋清楚Transformer原本是给NLP设计的它处理的是“一串词”而图像是一堆像素怎么把这两件事接起来是理解Transformer做图像分类的核心起点。这篇就围绕“Transformer在图像分类上的应用以及PyTorch代码实现”把原理、代码、训练细节一次讲透。我会带你从Patch Embedding开始手写一个可运行的Vision Transformer再对比Swin Transformer的思路顺便把我踩过的坑也一并交代。这个内容适合谁看刚入门深度学习、已经会一点PyTorch但没接触过Transformer的读者或者看过Transformer论文但不知道怎么写代码的人也可以参考。读完这篇你能独立用PyTorch训练一个ViT模型在CIFAR-10上跑分类并且能理解每个模块存在的意义。1. Transformer做图像分类的整体设计与思路拆解1.1 为什么图像也能用Transformer先明确一个基本问题Transformer的优势在长程依赖也就是它能建模序列中距离很远的位置之间的关联。NLP里一个句子可能有几十个词词与词之间跨距离的语义联系对理解句子很重要。而图像也有类似特点比如一张图上左上角的物体和右下角的物体可能属于同一个目标或者存在上下文关系。CNN是靠卷积核一点点扩大感受野看得远需要堆很多层或者把卷积核做得很大Transformer天然没有这个限制它在第一层就能让每个位置“看到”全部位置这是它在图像任务上能跟CNN掰手腕的重要原因。但图像天然是二维的、由像素构成的矩阵不能直接把像素扔给Transformer原因很简单计算量爆炸。假设一张224x224的图像素总数是50176如果每个像素当做一个token那么注意力矩阵就是50176x50176大概25亿个元素这在训练时是没法算的。所以ViT的做法是先把图像切成小块这个过程叫做Patch Embedding。1.2 Patch Embedding的动机和含义图像切成小块之后每个小块就相当于NLP里的一个“词”。比如224x224x3的输入图切成16x16大小的patch那么一共得到224/1614也就是14x14196个patch。每个patch展开就是16x16x3768长度的一维向量。这个768就相当于Transformer里的token特征维度。你可以理解成一张图变成了一句话这句话有196个“词”每个词对应768维的特征。PyTorch里的实现方式一般不需要真的手动切块再逐个展平而是用卷积一步到位。用Conv2d(in_channels3, out_channels768, kernel_size16, stride16)就可以实现。卷积核大小和步长都等于patch尺寸那输出的每个位置就对应一个patch的线性映射结果。这个做法又简单又高效这也是代码实现中大家最常用的技巧。1.3 Position Embedding不加上会怎样Transformer本身没有顺序概念。你让注意力算某一个patch和另一个patch的关联时如果没有任何位置信息模型并不知道哪个patch在左上角、哪个在右下角。对图像来说位置信息非常重要一个小球的patch出现在图片上方还是下方语义完全不同。所以ViT里给每个token加上了一个可学习的Position Embedding。这个Position Embedding的维度是和token一样的加在patch embedding上让每个token带上位置信息。ViT在实验里比较过可学习和固定编码的效果两者差别不大可学习编码稍微灵活一点所以默认用可学习的。这里有一个容易忽略的细节位置编码参与训练但它的学习率一般不单独设置跟着整体模型一起更新就行。1.4 分类Token和CLS的设计思路BERT里用了一个[CLS] token来聚合整个句子的语义ViT也采用了类似的设计在序列最前面额外加一个可学习的class token。这个token没有对应任何具体的patch它只在训练过程中通过注意力机制去聚合其他patch的信息最终拿它的输出去做分类。为什么不直接拿所有patch输出的平均池化理论上也可以ViT论文里做过实验class token的效果和池化差不多。但从实现角度讲class token让模型架构更接近BERT方便借力NLP那边的成熟经验同时下游任务如果要扩展比如做检测、分割class token这个语义聚合的位置保留下来也更容易迁移。代码里就是在patch embedding生成的序列前面cat一个shape为(1, 1, 768)的向量具体维度处理在后面代码部分再细说。1.5 Transformer Encoder堆叠起来之后在做什么把patch embedding加上position embedding和class token之后这个序列会经过若干层Transformer Encoder。每一层包含两个核心子结构多头自注意力Multi-Head Self-Attention和前馈网络MLP然后在每个子结构外接残差连接和LayerNorm。自注意力的作用是让每个位置去和其他所有位置计算关联权重不断的堆叠能够让高层特征逐渐语义化。从图像分类的角度看底层可能关注局部纹理和边缘高层会逐渐关注像是“有没有羽毛”、“有没有轮子”这类语义概念。这也是Transformer在图像分类上能得到不错效果的性能来源——它给模型提供了全局建模能力和灵活的特征交互方式。2. PyTorch环境准备与数据集处理2.1 环境配置的几个关键点做这个项目前我默认你已经装好了PyTorch。如果还没装建议用Anaconda创建一个虚拟环境然后根据显卡驱动选择对应版本的PyTorch。GPU版本怎么确认CUDA版本最稳妥的办法是在终端跑一下nvidia-smi看右上角的CUDA Version然后到PyTorch官网选择合适的安装命令。CPU环境也能跑但建议把图像尺寸调小、patch调大不然训练速度确实会让人崩溃。我在实际尝试中CPU环境下用CIFAR-10训练ViT-Tiny一个epoch大概要十几分钟GPU只需要几十秒。如果只是学习原理可以先拿小数据集跑通理解了再上GPU。安装完成后用这个命令验证python -c import torch; print(torch.__version__, torch.cuda.is_available())输出中第二项如果是True说明CUDA环境正常。CPU环境输出False也不用担心代码层面会自动识别。2.2 为什么用CIFAR-10而不是ImageNetImageNet有1000类、128万张图直接训练ViT需要非常长的训练时间和多卡资源个人环境很难复现。CIFAR-10有10类、6万张32x32的彩色图单卡就能快速跑通整个流程特别适合做模型验证和代码学习。但注意一个细节CIFAR-10图像分辨率是32x32ViT原论文的patch大小是16x16切出来的patch是2x24个。数量实在太少Transformer很难建模有效关系。所以实际做CIFAR-10时通常把patch大小设为4或8这样序列长度变成64或16。在后面的代码示例中我会把patch size设为4序列长度64这样效果验证充分而且训练也稳定。2.3 数据增强策略ViT相比CNN更依赖数据增强。原因很简单Transformer参数量大但归纳偏置弱——不像CNN天生就带“局部连接”和“平移等变性”如果训练数据不足或者增强不够很容易过拟合。我实际用过的增强策略RandomCrop配合padding4对32x32的图像先四周补4像素再随机裁剪到32x32等价于做了随机平移。RandomHorizontalFlip随机水平翻转对CIFAR-10大多数类别有意义比如猫水平翻转之后还是猫。NormalizeCIFAR-10的均值均值是(0.4914, 0.4822, 0.4465)标准差是(0.2470, 0.2435, 0.2616)这些值不是随便想的是官方统计的整数据集全局统计量。CutMix或MixUp我自己试验下来ViT用MixUp能把准确率提高3-5个百分点。但这两个策略会增加代码复杂度初学先把前三个用起来。3. Vision Transformer核心代码实现3.1 Patch Embedding代码详解前面说用卷积实现这里直接给代码import torch import torch.nn as nn class PatchEmbed(nn.Module): def __init__(self, in_channels3, patch_size4, embed_dim256): super().__init__() self.patch_size patch_size self.proj nn.Conv2d(in_channels, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): # x: (B, 3, H, W) B, C, H, W x.shape x self.proj(x) # (B, embed_dim, H/patch, W/patch) x x.flatten(2) # (B, embed_dim, num_patches) x x.transpose(1, 2) # (B, num_patches, embed_dim) return xH和W不整除patch_size时会报错或者产生不符合预期的形状所以输入图像尺寸要提前设计好。32x32、patch_size4时得到8x864个patch。如果图像是224x224、patch_size16就是196个patch。flatten(2)是把空间维度压成一维transpose(1,2)则是把embed_dim挪到最后一维因为Transformer的输入通常是(B, seq_len, embed_dim)格式。3.2 多头自注意力的实现要点多头自注意力是整个Transformer的核心但实现起来其实不太难class MultiHeadSelfAttention(nn.Module): def __init__(self, dim, num_heads8, dropout0.0): super().__init__() self.num_heads num_heads self.head_dim dim // num_heads self.scale self.head_dim ** -0.5 self.qkv nn.Linear(dim, dim * 3) self.proj nn.Linear(dim, dim) self.dropout nn.Dropout(dropout) def forward(self, x): B, N, C x.shape qkv self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim) qkv qkv.permute(2, 0, 3, 1, 4) q, k, v qkv[0], qkv[1], qkv[2] attn (q k.transpose(-2, -1)) * self.scale attn attn.softmax(dim-1) attn self.dropout(attn) x (attn v).transpose(1, 2).reshape(B, N, C) x self.proj(x) return x几个要点拆解一下。qkv从一次Linear里得到效率比分开三个Linear更高参数量一样但算子更集中。reshape和permute完成多头拆分key里每行是一个位置的向量query的shape是(B, heads, N, head_dim)key的转置是(B, heads, head_dim, N)矩阵乘出来就是(B, heads, N, N)的注意力分数矩阵。scale操作是防止点积结果过大导致softmax进入饱和区梯度消失。注意力矩阵每一行是所有位置对某个位置的相关性权重softmax让这些权重之和为1。用注意力权重对value做加权求和得到当前层输出。对新手来说最难理解的通常是reshape之后的维度变换我建议你实际print一下每个步骤的shape调通以后就豁然开朗了。3.3 Transformer Encoder层一个Encoder层是把自注意力、MLP、LayerNorm和残差连接组装起来class TransformerEncoderLayer(nn.Module): def __init__(self, dim, num_heads, mlp_ratio4.0, dropout0.0): super().__init__() self.norm1 nn.LayerNorm(dim) self.attn MultiHeadSelfAttention(dim, num_heads, dropout) 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)) x x self.mlp(self.norm2(x)) return x注意这里用的是Pre-LN先norm再做注意力而不是原始Transformer论文里的Post-LN。实际训练中Pre-LN更加稳定梯度传播更容易尤其是对深层Transformer来说能缓解训练不收敛的问题。这也是后面很多ViT变体都修改过的细节我可吃过不少亏奉劝你直接把Pre-LN用起来。MLP里激活函数用GELU而不是ReLU。GELU可以理解成ReLU的平滑版本在负区间不是直接置零而是有一个很小但非零的梯度Transformer这种深层网络用起来更稳。3.4 完整ViT模型组装把上面的模块拼起来就是一个完整可训练的ViTclass VisionTransformer(nn.Module): def __init__(self, in_channels3, image_size32, patch_size4, embed_dim256, num_heads8, num_layers6, num_classes10, dropout0.1): super().__init__() num_patches (image_size // patch_size) ** 2 self.patch_embed PatchEmbed(in_channels, patch_size, embed_dim) 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.ModuleList([ TransformerEncoderLayer(embed_dim, num_heads, dropoutdropout) for _ in range(num_layers) ]) 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, N, D) cls_tokens self.cls_token.expand(B, -1, -1) x torch.cat([cls_tokens, x], dim1) # (B, N1, D) x x self.pos_embed x self.pos_drop(x) for block in self.blocks: x block(x) x self.norm(x) cls_token_final x[:, 0] logits self.head(cls_token_final) return logits中间的num_patches计算可以更严谨一些但实际使用中patch_size能整除就行。embed_dim选256、num_layers6、num_heads8这种配置对应的是一个Tiny级别的ViT在CIFAR-10上大概有500万参数训练速度快效果也比简单CNN好。3.5 训练循环代码模型写好了训练代码反而简单下面是我常用的一个训练循环框架def train_one_epoch(model, train_loader, optimizer, criterion, device): model.train() total_loss, correct, total 0.0, 0, 0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() total_loss loss.item() * images.size(0) _, preds outputs.max(1) correct preds.eq(labels).sum().item() total labels.size(0) return total_loss / total, correct / total验证函数类似区别就是torch.no_grad()包住推断过程并且不更新梯度。优化器我推荐AdamW而不是Adam因为AdamW把权重衰减和自适应学习率解耦对Transformer这类大参数模型更友好。初始学习率按经验取1e-3到5e-4之间配合CosineAnnealingLR让学习率按周期衰减效果会比固定学习率好很多。3.6 完整训练脚本示意为了方便直接复制运行我把整个脚本串联起来展示import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader from torchvision import datasets, transforms from torch.optim.lr_scheduler import CosineAnnealingLR device torch.device(cuda if torch.cuda.is_available() else cpu) transform_train transforms.Compose([ transforms.RandomCrop(32, padding4), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)), ]) transform_test transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)), ]) train_dataset datasets.CIFAR10(root./data, trainTrue, downloadTrue, transformtransform_train) test_dataset datasets.CIFAR10(root./data, trainFalse, downloadTrue, transformtransform_test) train_loader DataLoader(train_dataset, batch_size128, shuffleTrue, num_workers2) test_loader DataLoader(test_dataset, batch_size128, shuffleFalse, num_workers2) model VisionTransformer().to(device) criterion nn.CrossEntropyLoss() optimizer optim.AdamW(model.parameters(), lr5e-4, weight_decay0.05) scheduler CosineAnnealingLR(optimizer, T_max100, eta_min1e-5) for epoch in range(100): train_loss, train_acc train_one_epoch(model, train_loader, optimizer, criterion, device) val_loss, val_acc evaluate(model, test_loader, criterion, device) scheduler.step() print(fEpoch {epoch1:03d} | Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f} | Val Acc: {val_acc:.4f})这个脚本直接跑在CIFAR-10上大概能到85%以上的准确率。如果想再高可以加MixUp增强、把embed_dim增大或者训练更久。4. 从ViT到Swin Transformer更强形态的实现思路4.1 ViT的短板和Swin的突破口ViT虽然效果不错但有两个比较明显的问题。一是图像切成固定大小patch后多尺度信息丢了这在目标检测、分割这类像素级任务里影响很大二是自注意力的计算复杂度是图像尺寸的平方高分辨率图像直接算不动。Swin Transformer给出的方案是层级化设计加窗口注意力。所谓层级化就是像CNN一样不断把token数量减少而维度增长形成类似特征金字塔的结构第一层是4x4 patch之后通过Patch Merging让分辨率减半、通道数翻倍。窗口注意力则是把图像按固定窗口切块注意力只在窗口内计算窗口之间再通过Shifted Window机制交换信息。这样计算复杂度和图像尺寸成正比而非平方所以叫Swin。4.2 Window Attention的实现特点Swin的窗口注意力在实现上有个很精细的地方为了避免不同窗口的token互相干扰PyTorch里用了一个巧妙的办法把图像reshape成(num_windows, window_size*window_size, embed_dim)再算注意力。计算完成后再reshape回原来的空间排布。这个操作在代码里就是view和permute的组合。你如果去阅读Swin的源码会发现这部分转来转去很容易把人绕晕。我建议拆开调试时打印每个步骤的shape关键点在于先分窗口再算注意力最后还原。理解了这一点Swin的代码就基本看懂一大半了。4.3 Swin的图像分类头设计Swin在分类任务里不像ViT那样有cls_token而是对最后一层输出的所有token做一个LayerNorm之后接全局平均池化再经过全连接层输出类别。从实验效果看Swin在ImageNet上比同规模的ViT能高1-2个点而且收敛更快。如果你有分类任务且数据量不大Swin往往比ViT更好调。4.4 ViT和Swin在PyTorch中的选择建议选型这件事不能只看榜单我建议按场景来数据量中等、图像分辨率固定、主要做分类ViT足够代码简单易改我这篇里的代码够用。需要多尺度信息或者要作为下游检测分割骨架Swin更合适。计算资源有限考虑Swin变体或者尝试把ViT的层数减小。训练数据很小只有几千张两个效果都很难好建议用预训练权重做迁移学习或者干脆用CNN。5. 训练过程中最容易踩的坑与排查技巧5.1 损失不下降的常见原因我自己调试Transformer踩过最多的坑就是模型压根训不动。如果loss一直在初始值附近徘徊首先检查学习率是不是太大或太小。Transformer对学习率比CNN更敏感一般1e-4到1e-3是合理区间我经常看到有人拿CNN的0.01直接套到ViT上效果惨不忍睹。另外检查一下数据归一化是否正确。CIFAR-10如果不做Normalize输入像素范围是0到1注意力计算时点积结果可能偏大softmax太锐利导致梯度不稳。我见过初学者只做ToTensor不做Normalize然后在训练时损失震荡非常厉害。还有一个容易忽略的是Position Embedding的初始化。前文代码里我用torch.zeros初始化这其实是参考了ViT官方实现从零开始学是没问题的。但如果你用的是更大的embed_dim有时候也能用小的随机值不然训练初期所有位置都一样注意力区分不开位置前几轮收敛偏慢。5.2 内存不足OOM处理显存溢出是训练Transformer时的家常便饭。最常见的原因是batch_size设太大或者模型embed_dim设太大。CIFAR-10上用Tiny级别的ViTbatch_size128在8G显存上是可以跑的但如果你想尝试embed_dim768、patch_size8序列长度变成16显存压力反而小如果patch_size4、embed_dim768序列长度64显存会明显增加。实在OOM的时候先把batch_size减半再用梯度累积模拟大batch。代码里加一句scaler来用混合精度训练也能省一半显存但需要import torch.cuda.amp这部分初学者可以先不搞等流程跑通再优化。5.3 过拟合的典型表现和应对方法训练准确率很高但验证准确率上不去是过拟合的典型信号。ViT参数多、归纳偏置弱在CIFAR-10这种6万张图上很容易过拟合。我实际测试过不加任何正则训练100个epoch训练集准确率能到99%以上但验证集只有75%左右差距非常明显。应对方法优先级数据增强 Dropout Weight Decay 减小模型。数据增强里MixUp效果尤其明显建议尽早用起来。Dropout建议加在MLP里和Position Embedding后面注意力矩阵里也可以加一点。Weight Decay用AdamW默认的0.05就行。5.4 复现一致性和随机种子设置Transformer训练受随机性影响比较大。同样的代码不同种子可能带来1%左右的准确率波动。所以做实验对比时一定要固定种子才能公平比较。我在训练脚本开头固定三处torch.manual_seed、torch.cuda.manual_seed_all和numpy.random.seed同时在DataLoader里设置generator。5.5 从头训练还是用预训练权重PyTorch官方和timm库都提供了大量预训练ViT权重。如果数据集不是CIFAR这种小图而是比较接近ImageNet的真实场景我强烈建议直接用timm里加载预训练模型做微调能省下大量训练时间和算力。但对于CIFAR-10这种32x32的图因为预训练权重通常适配224x224输入直接迁移效果反而一般这种情况下从头训练一个小ViT更靠谱。6. 效果对比与实操心得6.1 在CIFAR-10上的基准效果参考我在这篇代码配置下做过多次实验给一组参考数据ViT-Tinyembed_dim256, depth6, heads8patch4AdamWCosineAnnealingLR训练100个epochMixUp增强开启时验证集准确率大概在86%到88%之间。如果不加MixUp大概会掉到83%左右。作为对比一个经典的ResNet18在同样的数据增强下能做到约92%。ViT在CIFAR-10这种小数据集上并没有优势参数量大、归纳偏置弱是主要原因。这也提醒你不要神化Transformer。它的优势在大规模数据和长程依赖场景小图分类任务里CNN仍然是性价比很高的选择。不过如果你把patch尺寸调大并在更大的数据集上训练ViT的优势就会逐步体现出来。6.2 参数量与计算量的对比ViT-Tiny的参数量大概在500万到700万之间ResNet18参数量约1100万按道理ViT更小但实际训练时间却更长。原因在于自注意力的计算量和序列长度的平方成正比尤其是本文里patch_size4时序列长度达到64这个计算开销已经不低了。如果用原论文16x16的patch在224x224图上序列长度196计算量会更大。如果你希望在计算量上进一步优化可以考虑减少heads数量、降低depth或者用更大的patch_size。做实验时可以先跑一个depth4的小模型验证代码正确性再逐步放大。6.3 实际使用中觉得顺手的小技巧最后分享几个我在实战中总结的小技巧。第一个是断点续训训练周期长建议每个epoch结束都保存一次模型保存内容包括模型权重、优化器状态、scheduler状态和当前epoch数这样断了能接着跑。第二个是日志记录把每个epoch的损失和准确率写入CSV训练结束后直接画曲线比只看终端输出直观得多。第三个是测试阶段开启torch.no_grad()和model.eval()别忘了调用我见过有人漏掉eval模式导致BatchNorm和Dropout状态不对验证结果虚高。6.4 模型可视化验证训练完之后建议做一次注意力可视化把模型最后一层的cls_token对其他patch的注意力权重画成热力图叠加在原图上。这一步很多教程会跳过但对我理解模型帮助特别大。你会看到模型确实学到了任务的语义特征——比如在识别猫的时候注意力大部分集中在猫头、耳朵和轮廓位置。可视化可以用torchvision的transforms把Tensor转成PIL图然后用matplotlib的imshow显示叠加结果。这个验证过程能帮你检查模型是否真的在学图像特征而不是在瞎猜。7. 常见问题速查与补充建议问题现象可能原因解决办法损失不下降学习率不合适调到1e-4到1e-3区间损失震荡剧烈数据没归一化加上Normalize验证集准确率低但训练集高过拟合加MixUp数据增强显存溢出batch_size过大减小batch或开梯度累积注意力矩阵数值异常忘了除以sqrt(d)补上scale缩放patch尺寸不整除图像尺寸和patch不匹配修改patch或resize输入这个话题还有不少可以扩展的方向比如把ViT迁移到目标检测、语义分割或者尝试DeiT、CvT这类后续改进结构。如果你按照这篇代码把CIFAR-10跑通再去看timm库的实现会发现很多API设计和我这里讲的思路是一致的阅读源码也会轻松很多。
返回列表