ARTICLE DETAIL

资讯详情

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

Vision Transformer代码实战:用PyTorch从零实现ViT

Vision Transformer代码实战:用PyTorch从零实现ViT 第一次看ViT论文的时候不少人会忍不住倒吸一口凉气。明明是一张图片怎么就能像处理句子一样切成patch喂给Transformer虽然论文里写得清清楚楚但真到自己用PyTorch复现时Patch Embedding怎么写、Position Embedding怎么加、Class Token放在哪这些细节稍不留神就会翻车。这篇博文就是来填这个坑的。我会把Vision Transformer的PyTorch代码从里到外拆一遍关键位置加上图解思路配合完整可运行的代码让不熟悉Transformer结构的人也能把这套架构吃透并真正跑起来。1. 核心设计拆解ViT到底在模仿什么1.1 先理解ViT解决的核心问题ViTVision Transformer这个思路最颠覆的一点就是把图像彻底当成序列来处理。传统CNN靠卷积核逐层滑动天然具备局部性和平移等变性所以它能很快捕捉边缘、纹理这些局部特征。而ViT的做法是把一张图切成固定大小的patch每个patch展平后做一次线性映射变成token然后丢进标准的Transformer Encoder里面做全局自注意力建模。这个转变带来了几个关键影响。一是感受野问题CNN要看到全局信息必须依赖深层堆叠靠层数堆出足够大的感受野而Transformer的每一层用自注意力直接建立任意两个patch之间的联系第一层就已经是全局视野。二是归纳偏置问题CNN把“相邻像素大概率相关”这种偏置内置到网络结构里而ViT主动放弃了这个偏置完全靠数据来学习。这也是为什么ViT需要更大的数据集或者更强的数据增强才能训练好在ImageNet那种大规模数据集上效果才明显。从代码实现角度看ViT有一个特别友好的特点结构上它跟NLP领域的BERT几乎一模一样只是输入从词向量换成了图像patch的Embedding。所以你只要把Transformer Encoder那套吃透了ViT的代码其实就剩三块新东西Patch Embedding、Class Token、Position Embedding。这三块单独拎出来都不难难的是理解它们各自解决的问题和拼接时的维度变化。1.2 整体数据流向图解很多教程贴出ViT结构图时会画得很复杂但本质上数据流是这样的输入一张3通道、224x224的图片先切成一堆16x16的patch224/1614一共14x14196个patch。每个patch展平后是16x16x3768维的向量经过一个线性层映射成D维比如768维这样我们就得到了196个token。然后在序列最前面拼一个可学习的Class Token序列长度变成197。再给这197个token都加上一个可学习的位置编码保持维度不变。接着丢进L层Transformer Encoder。最后取序列第一个token也就是Class Token对应位置的输出过一层分类头得到类别概率。这里有个特别容易混淆的点最后一个输出到底取谁Transformer Encoder输出的序列长度还是197每个位置对应一个D维向量。分类只用第一个位置也就是当初拼上的Class Token经过N层编码之后的表示。后面代码里我会专门标出来这一行。整个数据流对应到PyTorch的Tensor形状变化就是下面这个流程输入(B, 3, 224, 224)Patch Embedding后(B, 196, 768)拼接Class Token后(B, 197, 768)加Position Embedding后(B, 197, 768)经过Transformer Encoder后(B, 197, 768)取第一个token并过分类头后(B, num_classes)后面所有代码都会围绕这个流程展开。2. 环境配置与关键依赖准备2.1 依赖安装与版本选择ViT的实现并不复杂核心依赖就是PyTorch和TorchVision。如果你没有现成的环境用Anaconda建一个干净的环境是最省事的做法。conda create -n vit_env python3.9 -y conda activate vit_env pip install torch torchvision pip install matplotlib tqdm这里有几个版本细节需要注意。Python版本建议3.9或更高PyTorch建议2.x版本因为2.x里有一些对Transformer比较友好的改进。如果你有NVIDIA显卡装CUDA版的PyTorch可以大幅加速训练如果只有CPU也没关系CIFAR-10这种小规模数据集用CPU跑几个epoch作为流程验证是完全可以接受的就是慢一点。还有个小提醒如果图片预处理需要做RandomResizedCrop、RandomHorizontalFlip之类的基础增强torchvision里全都有现成的不需要额外装albumentations这些第三方库。训练ViT时数据增强很重要但不必一开始就上太重的trick先把主体流程跑通更重要。2.2 数据准备CIFAR-10举例为了让大家能快速复现我用CIFAR-10来做演示。这个数据集有6万张32x32的彩色图片分10个类别下载方便、单张图很小训练起来压力不大。import torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision import datasets, transforms 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)), ]) trainset datasets.CIFAR10(root./data, trainTrue, downloadTrue, transformtransform_train) testset datasets.CIFAR10(root./data, trainFalse, downloadTrue, transformtransform_test) trainloader DataLoader(trainset, batch_size128, shuffleTrue, num_workers2) testloader DataLoader(testset, batch_size256, shuffleFalse, num_workers2)这里要注意一个关键尺寸问题CIFAR-10的图片是32x32如果要切成16x16的patch那整张图只有2x24个patch信息量太少Transformer几乎没法学出有效特征。所以对于CIFAR-10这种小图一般把patch设为4x4或者先用插值把图放大到64x64或224x224。我在后面的代码里采用patch_size4的方式这样能拿到8x864个token配合较小的模型规模在CIFAR-10上能跑出不错的效果。3. ViT核心代码逐行拆解3.1 Patch Embedding的两种实现方式与原理Patch Embedding的目标是把图像转成token序列。最直观的做法先用unfold把图像切成patch然后对每个patch做线性变换。不过在实际工程中更多是直接用nn.Conv2d一个卷积搞定这也是很多开源实现的写法。class PatchEmbed(nn.Module): def __init__(self, in_channels3, patch_size4, embed_dim192): super().__init__() self.patch_size patch_size # 用卷积实现patch切分 线性投影 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 assert H % self.patch_size 0 and W % self.patch_size 0, \ f输入尺寸 {H}x{W} 不能被 patch_size {self.patch_size} 整除 # 卷积后: (B, embed_dim, H/patch_size, W/patch_size) x self.proj(x) # 展平后两个维度: (B, embed_dim, num_patches) x x.flatten(2) # 转置成序列格式: (B, num_patches, embed_dim) x x.transpose(1, 2) return x用Conv2d实现的关键点在于卷积核大小等于patch大小步长也等于patch大小这样卷积输出特征图的每个点就对应原图的一个patch而且每个patch都经过了同一个卷积核的加权求和——这本质上就是一个线性映射。卷积输出的通道数设为embed_dim相当于把每个patch展平后的向量映射到了D维空间。为什么推荐用卷积而不是unfold主要是因为Conv2d在GPU上的实现经过了深度优化速度和显存利用效率都更好。而且卷积操作的表述非常简洁一行代码就完成了切patch和投影两件事。不过直接用unfold其实也不难的理解一下就行。维度变化是最容易绕晕的地方我把核心过程再捋一遍输入(B, 3, 32, 32)patch_size4Conv2d输出(B, 192, 8, 8)flatten(2)后是(B, 192, 64)transpose(1,2)后是(B, 64, 192)。这64个token每个都是192维。3.2 Class Token和Position Embedding的细节实现Patch Embedding之后95%的人会在Class Token和Position Embedding这里犯迷糊。先看位置编码。Transformer本身没有顺序概念而patch之间的相对位置对图像理解又至关重要所以必须把位置信息揉进输入里。ViT用的是可学习的Position Embedding也就是一个维度为(1, num_patches1, embed_dim)的参数训练时跟着网络一起更新。Class Token的来历更有意思。因为Transformer Encoder输出的每个位置都对应一个token的表示做分类时需要从这一堆token表示中汇聚出“整张图”的表示。最简单做法是对所有token做全局池化但ViT论文里选择在序列最前面放一个可学习的Class Token最后取这个token的输出当作整张图的特征。这个Class Token会在训练中学会“汇总”其他patch的信息。class ViT(nn.Module): def __init__(self, image_size32, patch_size4, num_classes10, embed_dim192, depth6, num_heads3, mlp_ratio4.0): super().__init__() self.patch_embed PatchEmbed(in_channels3, patch_sizepatch_size, embed_dimembed_dim) num_patches (image_size // patch_size) ** 2 # Class Token: 一个可学习的向量 self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim)) # Position Embedding: 序列长度是 num_patches 1因为要算上 Class Token self.pos_embed nn.Parameter(torch.zeros(1, num_patches 1, embed_dim)) self.pos_drop nn.Dropout(p0.1) self.blocks nn.ModuleList([ TransformerBlock(embed_dim, num_heads, mlp_ratio) 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, num_patches, embed_dim) # 把 Class Token 拼到序列前面 cls_tokens self.cls_token.expand(B, -1, -1) x torch.cat((cls_tokens, x), dim1) # (B, num_patches 1, embed_dim) # 加位置编码 x x self.pos_embed x self.pos_drop(x) for blk in self.blocks: x blk(x) # 取 Class Token 对应的输出 x self.norm(x) cls_out x[:, 0] return self.head(cls_out)这里有几个容易踩坑的细节我给你拆开说。第一cls_token的初始化我用了torch.zeros而不是随机初始化。实际训练中。两种方式差别不大因为后续的Dropout和Transformer层会迅速打破对称性所以zeros是完全可行的这也是很多官方实现的做法。第二cls_token.expand(B, -1, -1)只是扩展维度不会复制数据所以额外显存开销可以忽略不计。torch.cat之后序列长度从64变成65Position Embedding的维度必须对应改成65。如果你改了patch_size或者输入尺寸这个数字很容易对不上报错时会提醒你dimension mismatch很多人第一次写会在这里卡住。第三我把位置编码初始化为0。源码里经常用截断正态分布来初始化但在实践中位置编码的信号在学习初期会被输入的embedding信号淹没之后随着训练逐步调整。如果你的初始化方差过大反而可能拖慢收敛速度。3.3 自注意力机制的代码实现与计算过程现在到了整个ViT最核心的模块——Multi-Head Self-AttentionMSA。公式大家都见过Attention(Q, K, V) softmax(QK^T / sqrt(d_k)) V关键问题是Q、K、V在代码里怎么算多头又是什么意思class Attention(nn.Module): def __init__(self, embed_dim, num_heads, qkv_biasTrue): super().__init__() self.num_heads num_heads self.head_dim embed_dim // num_heads self.scale self.head_dim ** -0.5 # 一个全连接同时生成 Q、K、V self.qkv nn.Linear(embed_dim, embed_dim * 3, biasqkv_bias) self.proj nn.Linear(embed_dim, embed_dim) def forward(self, x): B, N, C x.shape # 生成 QKV并拆成三个张量 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] # 每个都是 (B, num_heads, N, head_dim) # 注意力分数: (B, num_heads, N, N) attn (q k.transpose(-2, -1)) * self.scale attn attn.softmax(dim-1) # 加权求和 x (attn v).transpose(1, 2).reshape(B, N, C) x self.proj(x) return x这段代码我建议你对着张量形状一行行看。self.qkv(x)输出的维度是(B, N, C*3)然后通过reshape把最后一维拆成3份对应Q、K、V。permute的作用是把维度顺序调整为(3, B, num_heads, N, head_dim)这样索引0、1、2分别就是Q、K、V。为什么要除以sqrt(head_dim)这是为了保证Q和K点积之后的结果方差维持在1左右避免softmax在输入较大时梯度消失。你也许注意到这里用的是head_dim而不是整个embed_dim因为每个头参与计算的是head_dim维的向量。多头注意力的本质是在同一个序列上并行运行多组不同的注意力。每个头有自己独立的Q、K、V投影这意味着不同的头可以关注不同的位置关系——有的头可能偏向关注相邻patch有的头可能偏向关注全局颜色分布有的头则专门抓边界纹理。8个头就有8种不同的“视角”。三维可视化一下这个过程对第h个头输入的token序列是(B, N, head_dim)Q和K做点积得到(B, N, N)的注意力矩阵第i行第j列表示第i个token对第j个token的关注度。然后softmax归一化确保一行加起来等于1。最后用这个权重矩阵去加权V得到每个token的新表示。这个新表示里第i个token的信息就是“整个序列所有token按注意力权重融合”的结果。3.4 MLP与残差连接的细节处理除了自注意力Transformer Block里还有MLP、LayerNorm和残差连接。这里有一个ViT跟原始Transformer不一样的地方值得专门讲一下ViT用的是Pre-LN结构也就是先做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 Attention(embed_dim, num_heads) self.norm2 nn.LayerNorm(embed_dim) self.mlp MLP(embed_dim, int(embed_dim * mlp_ratio), dropout) def forward(self, x): x x self.attn(self.norm1(x)) x x self.mlp(self.norm2(x)) return x残差连接的目的是让梯度可以顺畅地跨层传播。但Pre-LN和Post-LN有个重要区别Post-LN原始Transformer把LayerNorm放在残差相加之后训练深层Transformer时容易不稳定需要 careful的warmup策略Pre-LN把LayerNorm放在残差之前梯度能更直接地反传训练更稳定对warmup的需求也更低。ViT的实现统一采用Pre-LN好处是更稳坏处是表征能力略有一点损失——不过在实践里这个损失基本可以忽略。MLP部分相对简单就是一个两层的全连接网络中间跟一个GELU激活函数。为什么用GELU而不是ReLUGELU在负区间不硬截断而是平滑过渡训练更稳定这在Transformer类模型里已经成了事实标准。class MLP(nn.Module): def __init__(self, in_features, hidden_features, dropout0.1): super().__init__() self.fc1 nn.Linear(in_features, hidden_features) self.act nn.GELU() self.fc2 nn.Linear(hidden_features, in_features) self.drop nn.Dropout(dropout) def forward(self, x): x self.fc1(x) x self.act(x) x self.drop(x) x self.fc2(x) x self.drop(x) return xMLP的hidden维度一般是embed_dim的4倍。这个比例来自Transformer论文的实验结论4倍在计算量和表示能力之间取得了不错的平衡。再大效果提升不明显训练成本反而涨得很厉害。3.5 完整模型组装与参数量计算把前面所有模块拼起来就得到了我们完整的ViT模型。让我把每一层的维度变化和参数量算一下这样你在训练时对模型规模能有个底。def build_vit(image_size32, patch_size4, num_classes10): model ViT( image_sizeimage_size, patch_sizepatch_size, num_classesnum_classes, embed_dim192, depth6, num_heads3, mlp_ratio4.0, ) return model model build_vit() total_params sum(p.numel() for p in model.parameters()) trainable_params sum(p.numel() for p in model.parameters() if p.requires_grad) print(fTotal params: {total_params / 1e6:.2f}M) print(fTrainable params: {trainable_params / 1e6:.2f}M)在embed_dim192、depth6、patch_size4、输入32x32的配置下模型大概有几百万参数跟一个小型CNN差不多。但要注意参数量分布很不均匀6层Transformer占了绝大部分Patch Embedding和分类头的占比很小。如果你的输入图片更大比如224x224、patch_size16每个patch展平是768维embed_dim通常也设为768或更大的值参数量会直接涨到几千万甚至上亿。这就是为什么ViT模型动辄几百MB的原因——参数绝大部分都在Transformer Encoder里。4. 训练细节与超参数选择实战4.1 优化器、学习率与warmup策略ViT训练跟CNN有个很大不同它对优化器和学习率更敏感。如果用常规的SGD训练速度会比较感人。ViT官方实现和社区复现普遍推荐AdamW优化器并且配合warmup和cosine学习率衰减。import torch.optim as optim from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR, SequentialLR epochs 100 warmup_epochs 5 optimizer optim.AdamW(model.parameters(), lr1e-3, weight_decay0.05) warmup_scheduler LinearLR(optimizer, start_factor0.1, end_factor1.0, total_iterswarmup_epochs) cosine_scheduler CosineAnnealingLR(optimizer, T_maxepochs - warmup_epochs, eta_min1e-5) scheduler SequentialLR(optimizer, schedulers[warmup_scheduler, cosine_scheduler], milestones[warmup_epochs])warmup的作用很关键。Transformer在训练初期如果直接用较大学习率LayerNorm和注意力机制很容易产生震荡。先用小学习率跑几个epoch让网络对数据分布有基本认识再逐渐加大学习率进入正式训练阶段这是ViT能稳定收敛的重要前置条件。学习率本身的选择也值得斟酌。1e-3搭配batch_size128在我的配置下表现不错。如果你的batch_size翻倍到256学习率可以适当调到1.2e-3到1.5e-3前提是你用AdamW或类似的自适应优化器——这种线性缩放经验在Transformer训练里比CNN更实用。weight_decay设置为0.05这是ViT官方默认值对控制过拟合有帮助。4.2 训练循环完整代码训练循环本身不复杂关键是记得把模型切到train和eval两种模式并且保证每个epoch都验证一下测试集准确率。下面我给出一个完整的训练代码框架。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) 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 images.size(0) return total_loss / total, 100.0 * correct / total torch.no_grad() def evaluate(model, loader, criterion, device): model.eval() total_loss, correct, total 0.0, 0, 0 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) _, preds outputs.max(1) correct preds.eq(labels).sum().item() total images.size(0) return total_loss / total, 100.0 * correct / total device torch.device(cuda if torch.cuda.is_available() else cpu) model build_vit().to(device) criterion nn.CrossEntropyLoss() best_acc 0.0 for epoch in range(epochs): train_loss, train_acc train_one_epoch(model, trainloader, optimizer, criterion, device) test_loss, test_acc evaluate(model, testloader, criterion, device) scheduler.step() if test_acc best_acc: best_acc test_acc torch.save(model.state_dict(), best_vit_cifar10.pth) if (epoch 1) % 10 0: print(fEpoch [{epoch1}/{epochs}] fTrain Loss: {train_loss:.4f} Train Acc: {train_acc:.2f}% fTest Loss: {test_loss:.4f} Test Acc: {test_acc:.2f}% fBest Acc: {best_acc:.2f}%)这个训练循环里面的每个细节都有讲究。model.train()和model.eval()切换的是Dropout和LayerNorm的行为——Dropout在训练时随机丢弃、在评估时保持全量LayerNorm在训练时用当前batch的统计数据在评估时用运行时的统计均值。不切换的话评估结果会忽高忽低尤其是小的batch size时更明显。zero_grad()在每个batch开始前清零梯度。这一步漏了的话梯度会跨batch累积loss直接崩掉。torch.no_grad()告诉PyTorch不要构建计算图推理时省大量显存和内存。4.3 数据增强与正则化经验在CIFAR-10这种相对小的数据集上训练ViT数据增强和正则化是能否work的关键。ViT没有CNN那种先验偏置如果不做增强非常容易过拟合。我常用的一套增强组合如下transform_train transforms.Compose([ transforms.RandomCrop(32, padding4), transforms.RandomHorizontalFlip(), transforms.RandomApply([transforms.ColorJitter(0.4, 0.4, 0.4, 0.1)], p0.8), transforms.RandomGrayscale(p0.2), transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)), ])这套增强的核心思想是让模型不能仅仅依赖颜色或简单的位置先验而是必须学到更鲁棒的形状和纹理特征。RandomCrop和RandomHorizontalFlip是基础ColorJitter可以扰动亮度、对比度、饱和度RandomGrayscale则变相增强模型对颜色缺失的鲁棒性。除了这些Mixup和CutMix这类增强策略对ViT也有明显的正则化效果。它们本质上是把两张图的输入和标签都做线性插值强迫模型学习更平滑的决策边界。如果你训练集不大、又追求更高精度这两招值得一试在timm库里都有现成实现可以直接调用。5. 常见问题与踩坑实录5.1 训练不收敛或收敛过慢的排查思路我自己的经验是跑ViT遇到的绝大多数问题都集中在以下几类这里整理成一个排查清单问题现象可能原因排查/解决方案Loss不下降学习率过大或过小先用不同学习率做小规模实验搭配warmup训练震荡严重学习率太大 / 没做warmup降低初始学习率增加warmup轮数验证集准确率很低模型太浅 / patch_size过大适当增加depth减小patch_size过拟合训练好、验证差数据集太小 / 增强不足加数据增强增大weight_decay和Dropout显存不足OOMbatch_size过大 / 序列过长减小batch_size减小输入分辨率或用梯度累积准确率提升极慢位置编码未生效检查pos_embed是否正确加到每个token上训练和测试时结果差异大Dropout/LayerNorm模式切换错误确认train/eval模式是否正确调用这里我想特别强调两个容易困扰新手的点。第一个是关于Class Token的取法。有些实现会在最后对所有token做全局均值池化跟取Class Token两个结果差别不算大但如果你代码里取了x[:, 0]Class Token却在Transformer之后忘了做LayerNorm或者错把整个序列都过了一层nn.Linear分类效果就会明显下滑。Classifier放在最后输入维度必须是embed_dim不是num_patches 1这个维度对错了直接报错。第二个是输入图片尺寸必须能被patch_size整除。比如输入是224x224patch_size是16刚好整除。但如果你临时换了数据集比如图像是128x192patch_size还是16那就会出问题。代码里写了assert但实际使用中建议根据数据集灵活调整patch_size或者用插值先缩放图片。5.2 推理时输出特征可视化的简单实践虽然这次是代码解析但理解模型学到的模式很能帮助调参。简单来说可以取Transformer某个中间层比如第一层的注意力权重然后用热力图形式可视化。def visualize_attention(model, image, layer_idx0, head_idx0): model.eval() x image.unsqueeze(0) # 拿到指定层 block model.blocks[layer_idx] def hook_fn(module, input, output): global saved_attn # input[0] 是 norm1 之后的 x saved_attn module.attn.attn.detach() # 注册 forward hook 来获取注意力矩阵 handle block.attn.register_forward_hook(hook_fn) with torch.no_grad(): pred model(x) handle.remove() # saved_attn 的维度是 (B, num_heads, N, N) attn_map saved_attn[0, head_idx, 0, 1:].reshape(num_patches_side, num_patches_side) return attn_map当然这个实现依赖对Attention模块内部属性名的修改如果你根据自己的代码结构调整需要灵活适配一下但思路是一样的注册forward hook在forward过程中把注意力矩阵捞出来然后画热力图。通过热力图你能直观看到浅层attention往往关注局部结构相邻patch相关性高深层attention则可能学会跨距离的语义关联。如果发现浅层注意力看起来完全是均匀分布每个位置都一样那大概率是模型没有学到有效的位置信息需要检查Position Embedding是否加了、学习率是否合适。5.3 显存不足与训练速度优化如果你在更大的分辨率或更大的模型上训练ViT显存问题会非常突出。几个常用的减负手段分享给大家。第一个是梯度累积。如果理想batch_size是128但显存只够放32那就把batch_size设为32每4个step做一次参数更新效果基本等价。accumulation_steps 4 optimizer.zero_grad() for step, (images, labels) in enumerate(trainloader): images, labels images.to(device), labels.to(device) outputs model(images) loss criterion(outputs, labels) / accumulation_steps loss.backward() if (step 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad()第二个是使用混合精度训练。PyTorch自带torch.cuda.amp只需改动几行训练速度通常能提升一倍不止显存占用也明显下降。在支持Tensor Core的显卡上尤其明显。新手刚开始可能不想碰这块但等你在更大的数据集上跑ViT时这是绕不开的利器。第三个是多卡并行。在单机多卡环境下用torch.nn.DataParallel先做个最简单版本也能有不小提升。不过要注意PyTorch 2.x里DataParallel的调度开销比DistributedDataParallel大不少如果卡数多或者模型大建议直接上DDP。6. 从零到一复现ViT的经验总结我个人在实际操作中最大的体会是ViT的代码门槛不在Transformer本身而在于把图像转换成序列这个思维转变。只要把Patch Embedding、Class Token、Position Embedding这三块想明白后面的Attention和MLP都是标准组件照着NLP里成熟的实现搬过来就行。最后再分享一个小技巧。在CIFAR-10上调试ViT时建议先不急着上完整的数据增强和6层Transformer。先跑一个mini版本depth2、embed_dim96、只做RandomCrop和HorizontalFlip看看能不能在当前配置下过拟合训练集。如果连训练集都过拟合不了说明是代码逻辑有问题如果能过拟合但验证集很差说明是数据增强和正则化不够。这种“先求过拟合再谈泛化”的调试顺序能帮你快速定位问题出在哪一层。如果你后面打算在ImageNet这种大规模数据上使用ViT我建议直接基于timm库做二次开发里面的ViT实现经过充分验证、支持各种变体DeiT、Swin Transformer等比自己从头写稳妥得多。但在此之前非常推荐先按这篇文章的思路手写一遍——只有当你能把每一步的维度变化和每个模块的输入输出都烂熟于心时调参和改结构才能真正做到心里有数。
返回列表