尧图网络科技YAOTU DIGITAL 获取报价
获取报价
首页 / 资讯中心 / 文章详情

从零手写ViT:PyTorch实现Vision Transformer的核心细节与完整代码

发布时间:2026/9/30 1:27:26

资讯中心
01
ARTICLE

从零手写ViT:PyTorch实现Vision Transformer的核心细节与完整代码

从零手写ViT:PyTorch实现Vision Transformer的核心细节与完整代码
1. 从零手写ViT前的准备工作环境搭建与数据集选型先把话说在前头ViT的PyTorch实现真正难的不是注意力机制本身而是把图像拆成Token序列这件事的代码表达。很多人看了Transformer的文本实现再看ViT还是一头雾水核心就在于怎么把一张224×224的图变成12×768的矩阵这个思维转换没打通后面全是死胡同。这篇文章我直接带你走一遍完整的ViT代码实现链路从Patch Embedding到Transformer Encoder再到分类头每一段代码都会配上图解逻辑和踩坑记录。建议你一边读一边把代码敲进编辑器里跑光看不练三个月也学不会。1.1 PyTorch环境配置别在这一步浪费太多时间动手之前先把环境搞定。我用的是PyTorch 2.x版本搭配CUDA 11.8实测下来稳定性和显存占用表现都不错。conda create -n vit python3.9 -y conda activate vit pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118 pip install numpy matplotlib einops tqdm tensorboard几个安装时的注意事项都是踩出来的经验einops库非常建议装。虽然ViT的标准实现不用它也能写但rearrange函数在做维度变换时可读性比一套permuteviewtranspose的组合强十倍调试体验完全不同。torch版本不要追新稳定版就行。PyTorch 2.x的torch.compile特性虽然能加速训练但初学阶段开启它反而会掩盖很多维度错误建议先关掉。Windows用户如果装GPU版本报错多半是CUDA版本没对上。先用nvidia-smi确认驱动支持的CUDA版本号再选择对应的torch wheel。1.2 数据集选择小数据集反而更适合学ViT初学者最容易犯的错误是一上来就搞ImageNet这个体量的数据结果训练一个epoch就要半天光调试维度问题就耗掉一周体验极差。我的建议是直接在torchvision.datasets里选CIFAR-10或者Tiny ImageNet来跑。CIFAR-10每张图32×32虽然和ViT原论文的224×224输入不一致但代码逻辑完全通用而且训练速度快一个数量级。等代码完全跑通、理解透彻了再无缝切到ImageNet也不迟。训练集50000张 32x32图片10个类别 测试集10000张 32x32图片这么小的图喂给ViT注意一个关键点Patch Size要相应调小。原论文用16×16的patch处理224×224是为了得到14×14的网格序列。32×32的图用4×4的patch同样能拿到8×864个patch的序列长度。序列长度适中注意力计算的开销也小非常适合在单卡上验证代码正确性。2. Patch Embedding与位置编码ViT把图像变成Token的核心逻辑ViT和CNN最大的思想分水岭在于它把一张图当作由固定大小图像块组成的序列来处理而不是用卷积核滑过整张图。Patch Embedding就是完成这个图像到序列转换的第一道工序。2.1 一个Conv2d搞定一切PatchEmbedding类的完整代码这是我最喜欢的部分。ViT原论文里把整张图切割成patch这一步用PyTorch代码实现只需一个卷积层。import torch import torch.nn as nn class PatchEmbedding(nn.Module): 将图像转换为Patch序列嵌入。 参数: img_size: 输入图像尺寸假设正方形 patch_size: 每个Patch的尺寸 in_chans: 输入图像的通道数 embed_dim: Embedding维度 def __init__(self, img_size224, patch_size16, in_chans3, embed_dim768): super().__init__() self.img_size img_size self.patch_size patch_size self.num_patches (img_size // patch_size) ** 2 self.embed_dim embed_dim # 核心一个stridepatch_size的卷积等效于切块展平线性投影 self.proj nn.Conv2d( in_chans, embed_dim, kernel_sizepatch_size, stridepatch_size ) def forward(self, x): # x: [B, C, 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] - 转置: [B, num_patches, embed_dim] x x.flatten(2).transpose(1, 2) return x用卷积来干这个活妙处在哪你可以脑补一个过程kernel_size16、stride16的卷积每次只在图像上做一次点乘滑动步长刚好等于patch大小这意味着每个卷积核的感受野恰好覆盖一个16×16的patch区域。卷积输出的每个空间位置的数值就是该patch所有像素值经过一组权重投影后的embedding向量。卷积天然就是滑动的patch提取器而且底层是高度优化的C实现效率远比自己写for循环切图再展平高得多。从维度上看输入[B, 3, 224, 224]经过卷积变成[B, 768, 14, 14]flatten后是[B, 768, 196]再转置成[B, 196, 768]。这196个token每个都对应原图一个16×16的区域。2.2 位置编码为什么不直接用绝对位置编码序列信息是Transformer的命脉。对于文本画面顺序天然重要对于图片patch的空间位置同样不可忽略。ViT的位置编码选择了一条有趣的路直接用可学习的参数矩阵。class PositionalEncoding(nn.Module): 可学习位置编码 def __init__(self, num_patches, embed_dim): super().__init__() # 多预留一个位置后面接class token用 self.pos_embed nn.Parameter(torch.zeros(1, num_patches 1, embed_dim)) nn.init.trunc_normal_(self.pos_embed, std0.02) def forward(self, x): return x self.pos_embed这里的num_patches 1是ViT的一个设计细节多出来的那个位置留给class token这个我们在第4节详细展开。你可能会问Transformer原论文用的是正弦余弦位置编码为什么ViT偏不这么干答案很直白可学习位置编码在ViT的实验里效果更好而且实现更简单。既然数据量够大、训练时间够长让网络自己学会第一行第一列的patch和第一行第二列的patch是什么位置关系比硬编码一个固定的三角函数更灵活。实际使用中有一点要注意位置编码和patch embedding的输出是相加而非拼接。所以维度必须完全一致写代码时不用检查模型就会报错。2.3 维度对齐最容易翻车的地方很多自己撸过ViT代码的人都会在positional encoding相加这一步翻车。常见的错误场景是忘了预留class token的一个位置导致num_patches和位置编码长度对不上输入是[B, H, W, C]的格式而模型期望[B, C, H, W]插值改变输入尺寸时位置编码没有做相应处理我的调试习惯是每写一个模块就在下面放一个打印维度的测试用例。# 验证PatchEmbedding patch_embed PatchEmbedding(img_size32, patch_size4, in_chans3, embed_dim64) x torch.randn(2, 3, 32, 32) x patch_embed(x) print(fPatchEmbedding输出: {x.shape}) # 期望: [2, 64, 64]正确输出应该是[2, 64, 64]Batch264个patch每个patch映射成64维向量。这一步对了再往下走省得后面每层都在排查来源不明的维度错。3. Transformer Encoder复刻注意力机制与MLP的PyTorch实现Patch变成token序列之后接下来的流程就和NLP里的Transformer Encoder完全一致了。但这部分依然有几个值得细说的点尤其是多头注意力的维度变换和LayerNorm的位置。3.1 多头自注意力从矩阵乘法视角理解Query/Key/Value多头注意力Multi-Head Self-Attention, MSA的代码不算复杂但每个维度都要想清楚为什么这么变。直接看代码class MultiHeadSelfAttention(nn.Module): def __init__(self, embed_dim, num_heads, dropout0.0): super().__init__() assert embed_dim % num_heads 0, embed_dim必须能被num_heads整除 self.embed_dim embed_dim self.num_heads num_heads self.head_dim embed_dim // num_heads # 一个全连接同时算Q、K、V比三个Linear更高效 self.qkv nn.Linear(embed_dim, embed_dim * 3) self.attn_drop nn.Dropout(dropout) self.proj nn.Linear(embed_dim, embed_dim) self.proj_drop nn.Dropout(dropout) def forward(self, x): B, N, C x.shape # Nnum_patches1 (含class token) # 生成Q, K, V qkv self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim) qkv qkv.permute(2, 0, 3, 1, 4) # [3, B, num_heads, N, head_dim] q, k, v qkv.unbind(0) # 缩放点积注意力 attn (q k.transpose(-2, -1)) * (self.head_dim ** -0.5) attn attn.softmax(dim-1) attn self.attn_drop(attn) # 加权求和 x (attn v).transpose(1, 2).reshape(B, N, C) x self.proj(x) x self.proj_drop(x) return x, attn # 返回attn是为了后续可视化展开讲讲几个思维要点为什么用一个Linear同时算QKV而不是三个独立的Linear从参数数量上看两者完全等价但合并成一个Linear后矩阵乘法的效率更高代码也更紧凑。其内部逻辑就是一次矩阵乘法后按块拆开相当于三个权重矩阵拼在一起做了一次大矩阵运算。reshape和permute的顺序是灵魂。第一眼看到qkv self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim)很多人会想为什么不直接reshape成[B, N, 3*num_heads*head_dim]再切原因在于内存布局。先reshape成五维张量再用permute(2, 0, 3, 1, 4)把维度重排成[3, B, num_heads, N, head_dim]这样q, k, v的分离只需要一次无拷贝的view操作。顺序错了分离出来的张量数值就是乱的。注意力矩阵为什么要除以sqrt(head_dim)这是Transformer稳定训练的关键技巧。当head_dim变大时Q和K点积的结果方差会随之增大导致softmax输入进入饱和区梯度趋近于零。除以sqrt(head_dim)可以把方差归一化到1左右保证softmax的梯度能正常回传。很多初学的同学会漏掉这一步结果训练出来注意力分布全是one-hot的模型完全学不动。3.2 MLP与LayerNorm的搭配取舍Pre-LN还是Post-LNViT的Encoder Block结构为LayerNorm → MSA → 残差连接 → LayerNorm → MLP → 残差连接。这种先归一化再计算的模式属于Pre-LN结构和原版Transformer Post-LN先计算再归一化不同。class TransformerBlock(nn.Module): 一个完整的ViT Encoder Block def __init__(self, embed_dim, num_heads, mlp_ratio4.0, dropout0.1): super().__init__() self.norm1 nn.LayerNorm(embed_dim) self.attn MultiHeadSelfAttention(embed_dim, num_heads, dropout) self.norm2 nn.LayerNorm(embed_dim) hidden_dim int(embed_dim * mlp_ratio) self.mlp nn.Sequential( nn.Linear(embed_dim, hidden_dim), nn.GELU(), nn.Dropout(dropout), nn.Linear(hidden_dim, embed_dim), nn.Dropout(dropout) ) def forward(self, x): # Pre-LN先归一化再做注意力 x x self.attn(self.norm1(x))[0] # 残差连接 MLP x x self.mlp(self.norm2(x)) return x为什么要选Pre-LN两个原因第一训练稳定性更好。Post-LN在深层网络中容易遇到梯度消失或爆炸的问题Pre-LN让每个子层的输入都经过归一化能有效缓解深层Transformer的训练难度。第二对学习率更不敏感。实际训练中Pre-LN的Transformer可以承受更大的学习率波动这在大规模训练时非常好用。3.3 GELU激活函数ViT为什么不用ReLUViT的MLP块选择的是GELUGaussian Error Linear Unit而不是ReLU。GELU和ReLU的关系其实很近GELU可以理解为在ReLU的基础上加入了随机正则的效果GELU(x) x * Φ(x)其中Φ(x)是标准正态分布的累积分布函数。直观理解ReLU对负值一刀切为0GELU则保留少量的负值信息但给予接近0的梯度。这样的好处是激活函数整体更平滑不会像ReLU那样产生神经元死亡dead ReLU的问题。在PyTorch里直接nn.GELU()调用就行需要注意的是在推理阶段GELU是确定性函数训练时它没有随机性这和Dropout有本质区别。4. 分类头与完整模型组装从Encoder到分类输出的链路模型的主体骨架搭好了现在需要把整个ViT完整串起来。这里有几个独属于ViT的设计选择必须掰开揉碎了讲。4.1 Class Token为分类任务设计的可学习哨兵ViT在编码器输入序列中额外插入了一个可学习的class token位置在patch序列的最前面。它的网络视角是经过多层Transformer编码后这个token的最终表示就是这个图像的全局整合特征。为什么要这么做因为Transformer的注意力机制会把每个token的信息和其他所有token做交互。Class token不包含具体的图像patch信息它的初始表示完全可学习经过12层注意力交互后它会从所有patch token中检索并聚合全局信息。相比之下如果直接对所有patch token做平均池化既丢掉了空间结构信息又无法让网络自适应地决定哪些patch更重要。这个设计思路的实现代码非常轻量class ClassToken(nn.Module): 可学习的分类令牌 def __init__(self, embed_dim): super().__init__() self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim)) nn.init.trunc_normal_(self.cls_token, std0.02) def forward(self, x): # x: [B, num_patches, embed_dim] cls_tokens self.cls_token.expand(x.shape[0], -1, -1) x torch.cat([cls_tokens, x], dim1) return x注意第三行初始化用了trunc_normal_截断正态分布这是ViT原论文实现中的一个关键细节。为什么要截断普通的正态分布可能产生远离均值的大值在网络早期导致attention分布过于尖锐。截断到±2倍标准差可以防止初始位置编码过大造成的训练不稳定性。4.2 完整ViT模型组装一个能跑的forward现在把所有模块拼起来class ViT(nn.Module): Vision Transformer完整模型 参数: img_size: 输入图像尺寸 patch_size: Patch大小 in_chans: 输入通道数RGB为3 num_classes: 分类类别数 embed_dim: Embedding维度 depth: Transformer Encoder层数 num_heads: 多头注意力头数 mlp_ratio: MLP隐藏层维度相对于embed_dim的倍数 def __init__( self, img_size224, patch_size16, in_chans3, num_classes1000, embed_dim768, depth12, num_heads12, mlp_ratio4.0, dropout0.1, attn_dropout0.0 ): super().__init__() # Patch Embedding self.patch_embed PatchEmbedding( img_sizeimg_size, patch_sizepatch_size, in_chansin_chans, embed_dimembed_dim ) num_patches self.patch_embed.num_patches # Class Token self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim)) nn.init.trunc_normal_(self.cls_token, std0.02) # 位置编码 self.pos_embed nn.Parameter(torch.zeros(1, num_patches 1, embed_dim)) nn.init.trunc_normal_(self.pos_embed, std0.02) self.pos_drop nn.Dropout(dropout) # Transformer Encoder堆叠 self.blocks nn.Sequential(*[ TransformerBlock( embed_dimembed_dim, num_headsnum_heads, mlp_ratiomlp_ratio, dropoutdropout ) for _ in range(depth) ]) # 分类头 self.norm nn.LayerNorm(embed_dim) self.head nn.Linear(embed_dim, num_classes) # 初始化权重 self.apply(self._init_weights) def _init_weights(self, module): if isinstance(module, nn.Linear): nn.init.trunc_normal_(module.weight, std0.02) if module.bias is not None: nn.init.zeros_(module.bias) elif isinstance(module, nn.LayerNorm): nn.init.ones_(module.weight) nn.init.zeros_(module.bias) elif isinstance(module, nn.Conv2d): nn.init.kaiming_normal_(module.weight, modefan_out) def forward(self, x): B x.shape[0] # Patch Embedding x self.patch_embed(x) # [B, num_patches, embed_dim] # Concatenate class token cls_tokens self.cls_token.expand(B, -1, -1) x torch.cat([cls_tokens, x], dim1) # [B, num_patches1, embed_dim] # Add position embedding x x self.pos_embed x self.pos_drop(x) # 经过所有Transformer Encoder层 x self.blocks(x) # 取class token的最终表示 x self.norm(x) cls_final x[:, 0] # 只取第一个tokenclass token # 分类输出 x self.head(cls_final) return xforward的具体流程图像经过Patch Embedding变成token序列shape为[B, num_patches, embed_dim]在序列第0位置插入class token变为[B, num_patches1, embed_dim]加上位置编码token间引入空间位置关系送入depth层Transformer Encoder每层都做全局注意力交互取最后一层输出的class token向量序列的第一个元素过LayerNorm和分类头这里有个值得品味的细节最终分类是只取class token的那一个向量而不是对所有token的向量做汇聚。这意味着整个模型的信息流可以理解为class token像一个信息收集员在每层注意力计算中它不断从各个patch中抽取关键特征最终携带整张图的语义信息去映射到分类空间。4.3 模型实例化与参数量验证模型组装完成后实例化并检查参数量是验证代码正确性的第一关def count_parameters(model): 统计可训练参数量 return sum(p.numel() for p in model.parameters() if p.requires_grad) # 小配置适配CIFAR-10 vit_small ViT( img_size32, patch_size4, in_chans3, num_classes10, embed_dim256, depth6, num_heads8, mlp_ratio4.0, dropout0.1 ) print(f参数量: {count_parameters(vit_small) / 1e6:.2f}M) # 验证前向传播 x torch.randn(2, 3, 32, 32) logits vit_small(x) print(f输出logits shape: {logits.shape})我跑这小模型时输出的参数量大约是2.9Mlogits是[2, 10]符合预期。你可以对比一下ViT-Base/16的参数是86MViT-Large/16是307M。从2.9M到86M本质就是embed_dim从256涨到768、depth从6涨到12、patch_size从4涨到16共同作用的结果。这里有个反直觉的经验小配置的ViT在CIFAR-10上未必打得过ResNet18。ViT没有CNN的归纳偏置局部性和平移等变性需要足够多的数据或者较强的数据增强才能发挥优势。所以拿小数据集练手时如果模型指标不好不代表代码错了很可能是训练策略的问题。关于这一点第5节会详细展开。5. 训练、调试与可视化让ViT真正跑起来的实战验证模型结构写完了但代码只有跑起来并得到合理结果才算真正掌握。这一节分享训练过程中必备的配套代码、常见错误和可视化技巧。5.1 训练循环与优化器选型ViT的训练循环和CNN的区别不大但有几个配置值得单独指出。完整的训练框架分为数据加载、模型前向、损失计算、反向传播、参数更新几个步骤下面是从项目中抽出来的核心逻辑import torch.optim as optim from torch.optim.lr_scheduler import CosineAnnealingLR from torch.utils.data import DataLoader from torchvision import datasets, transforms # 数据增强ViT在小数据集上非常依赖增强策略 transform_train transforms.Compose([ transforms.RandomCrop(32, padding4), transforms.RandomHorizontalFlip(), transforms.ColorJitter(0.4, 0.4, 0.4), 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(./data, trainTrue, downloadTrue, transformtransform_train) test_dataset datasets.CIFAR10(./data, trainFalse, downloadTrue, transformtransform_test) train_loader DataLoader(train_dataset, batch_size128, shuffleTrue, num_workers4) test_loader DataLoader(test_dataset, batch_size256, shuffleFalse, num_workers4) # 优化器ViT对优化器比较敏感AdamW是标配 model vit_small optimizer optim.AdamW(model.parameters(), lr1e-3, weight_decay0.05) scheduler CosineAnnealingLR(optimizer, T_max100, eta_min1e-5) # 训练一个epoch def train_one_epoch(model, loader, optimizer, criterion, device): model.train() total_loss, correct, total 0, 0, 0 for images, labels in loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() logits model(images) loss criterion(logits, labels) loss.backward() optimizer.step() total_loss loss.item() * images.size(0) preds logits.argmax(dim1) correct preds.eq(labels).sum().item() total images.size(0) return total_loss / total, correct / total def evaluate(model, loader, criterion, device): model.eval() total_loss, correct, total 0, 0, 0 with torch.no_grad(): for images, labels in loader: images, labels images.to(device), labels.to(device) logits model(images) loss criterion(logits, labels) total_loss loss.item() * images.size(0) preds logits.argmax(dim1) correct preds.eq(labels).sum().item() total images.size(0) return total_loss / total, correct / total device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device) criterion nn.CrossEntropyLoss() for epoch in range(1, 101): train_loss, train_acc train_one_epoch(model, train_loader, optimizer, criterion, device) scheduler.step() if epoch % 10 0 or epoch 1: test_loss, test_acc evaluate(model, test_loader, criterion, device) print(fEpoch {epoch:3d} | Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f} | Test Acc: {test_acc:.4f})优化器选型上AdamW是ViT训练的事实标准。ViT原论文在对比实验中验证了和SGD相比AdamW在Transformer结构上的优势解耦了权重衰减大幅提高了Transformer在小数据集上的收敛速度。注意weight_decay设到0.05左右别用0.0001后者在Transformer上是明显无效的。学习率还有一个我在实际问题中反复踩过的坑batch size变大的时候学习率要相应调大。ViT对于这种批次大小-学习率联动的敏感度比CNN更高。5.2 维度不匹配的排查链路我见过太多人卡在维度错误上包括我自己初学的时候一个mat1 and mat2 shapes cannot be multiplied能查半小时。这里给出一条标准排查链路遇到问题照着走一遍基本都能解决第一步查输入图像的shape。打印x.shape确认是[B, C, H, W]而不是[B, H, W, C]。第二步打印每一层的输出维度。最快的方式是在forward里临时加print。看似笨拙实际上比用debugger逐层断点更快。第三步验证PatchEmbedding。确认num_patches (img_size // patch_size) ** 2算对了。如果img_size224, patch_size16那num_patches196位置编码长度必须是197加class token。第四步核实QKV维度。embed_dim必须能被num_heads整除。这是非常容易忽略的点报错信息往往看起来和维度无关注意力计算之前reshape时就会挂。第五步单独把每个子模块从大模型里抽出来测试# 单独测试Attention模块 attn MultiHeadSelfAttention(embed_dim256, num_heads8) x torch.randn(2, 65, 256) # 65 num_patches 1 (class token) out, attn_weights attn(x) print(fAttention输出: {out.shape}, 注意力权重shape: {attn_weights.shape})记住一个原则先保证每个零件能工作再组装整机。整机报错时优先怀疑输入前一层某个维度恰好是1的情况其次怀疑embed_dim在MLP内部映射时没有对应上。5.3 注意力可视化验证模型到底在看哪里Transformer的注意力权重天然适合做可视化来验证模型是否像人类一样关注图像的关键区域。我在项目里常用的方式是提取最后一层所有头对class token的平均注意力图再插值到与原图相同的分辨率叠加显示。import matplotlib.pyplot as plt import torch.nn.functional as F def visualize_attention(model, image, device, patch_size4, img_size32): 可视化class token对图像各区域的注意力权重 model.eval() image_tensor image.unsqueeze(0).to(device) # 前向传播时获取注意力权重 model.blocks[0].attn.return_attention True # 需要在MultiHeadSelfAttention里设置一个标志位或者修改forward保存attn with torch.no_grad(): # 用修改后的forward获取注意力 x model.patch_embed(image_tensor) cls_tokens model.cls_token.expand(x.shape[0], -1, -1) x torch.cat([cls_tokens, x], dim1) x x model.pos_embed attention_maps [] for idx, block in enumerate(model.blocks): # 这里需要在forward中保存attn简化起见直接调用 assert hasattr(block.attn, last_attn), 请先在attention模块中保存last_attn x block(x) attn block.attn.last_attn # [B, num_heads, N, N] attention_maps.append(attn) # 取最后一层的cls attention在多头维度上求平均 last_attn attention_maps[-1][0] # 去掉batch维 cls_attn last_attn[:, 0, 1:].mean(dim0) # [num_patches]去掉class token自身 # 还原成网格 grid_size int(cls_attn.shape[0] ** 0.5) attn_map cls_attn.reshape(grid_size, grid_size) # 插值到原图分辨率 attn_map F.interpolate( attn_map.unsqueeze(0).unsqueeze(0), size(img_size, img_size), modebilinear, align_cornersFalse ).squeeze().cpu().numpy() # 反归一化并可视化 image_np image.cpu().numpy().transpose(1, 2, 0) mean np.array([0.4914, 0.4822, 0.4465]) std np.array([0.2470, 0.2435, 0.2616]) image_np image_np * std mean image_np np.clip(image_np, 0, 1) fig, (ax1, ax2) plt.subplots(1, 2, figsize(8, 4)) ax1.imshow(image_np) ax1.set_title(Original) ax1.axis(off) ax2.imshow(image_np) ax2.imshow(attn_map, cmapjet, alpha0.5) ax2.set_title(Attention Map) ax2.axis(off) plt.tight_layout() plt.savefig(attention_viz.png, dpi150)这段代码跑通后你会看到模型处理一张猫的图片时class token的注意力集中在猫的头部、耳朵和眼睛等判别性区域。如果注意力图完全是均匀分布的说明模型没有学到有效的特征判别通常是训练不充分或者数据增强太弱。提醒一下我在上面的代码里用到了block.attn.last_attn这个attr在原来的MultiHeadSelfAttention类里并不存在。实际使用时需要在class里加一行self.last_attn attn.detach() # 在forward里保存注意要detach避免显存累积5.4 训练中的常见意外情况与应对调试代码正确后训练过程也会不断出现新问题。我见过最多的是这几种Loss一开始就变成NaN。最常见的原因是学习率太大导致AdamW的更新步长过大。把学习率从1e-3降到3e-4或1e-4重试。另外如果用了torch.compile某些PyTorch版本下会出现算子融合后的精度问题也表现为NaN关掉compile即可确认。Loss下降缓慢测试准确率总是卡在某个值。先确认自己的数据增强是否够强CIFAR-10上的经验是ViT需要比ResNet更强的RandomCrop和HorizontalFlip再加上ColorJitter效果才明显。然后确认学习率schedule是否合理CosineAnnealing的T_max必须对齐总epoch数。显存溢出OOM。序列长度是ViT显存消耗的大头。CIFAR-10的32×32图patch_size4序列长度只有65显存压力还好。一旦切到224×224大图序列长度196×12层batch size稍大就会爆。应对方案减小batch size、开启梯度累积、降低图像分辨率作为热身或者把dropout适当加大。6. 更进一步从代码到论文复现的一些旁路经验把上面的代码跑通意味着你理解了ViT的核心。如果你想往深了走还有几个方向可以考虑这些都在我的实际项目中被验证过有明确的收益。6.1 ViT与CNN的对比没有免费的午餐很多人有个误解觉得ViT在视觉任务上全面碾压CNN。实际上在中小数据规模上ResNet的收敛速度和精度经常超越同样训练轮数的ViT——ViT的优势主要体现在超大数据集上或者有了强大的预训练权重做迁移学习之后。我在CIFAR-10上做过对比实验同等训练条件下ResNet18用50个epoch就能到92%的准确率而我手写的这个小ViT用100个epoch才勉强到91%。数据量少时CNN的归纳偏置优势非常明显。所以做工程选型时不要盲目上ViT先看你的数据量级和应用场景做研究学习时ViT的代码则无法绕开因为它代表了Transformer在视觉方向的基本范式。6.2 从ViT到Hybrid架构改进如果你想在这个基础上做改进最容易入手的方向是混合架构用CNN的浅层做stem把高分辨率的图像先下采样成特征图再切成patch喂给Transformer。这是Swin Transformer和很多高效ViT变体走过的路。实现上也只需要改PatchEmbedding层class CNNStem(nn.Module): 用卷积stem代替直接切patch先用CNN下采样降低序列长度 def __init__(self, in_chans3, embed_dim768): super().__init__() self.conv1 nn.Conv2d(in_chans, 64, kernel_size7, stride2, padding3) self.bn1 nn.BatchNorm2d(64) self.conv2 nn.Conv2d(64, embed_dim, kernel_size3, stride2, padding1) self.bn2 nn.BatchNorm2d(embed_dim) self.act nn.GELU() def forward(self, x): x self.act(self.bn1(self.conv1(x))) x self.act(self.bn2(self.conv2(x))) return x用CNN stem的好处是可以用较小的图像分辨率换取较高的特征抽象层级同时序列长度大幅缩短。比如输入224×224的图经过stride2的两个卷积后变成56×56×768如果再用patch_size4的切块序列长度只有196。这样比直接在224×224上切16×16的patch序列长度没有区别但CNN已经完成了初步特征抽象ViT负责全局关系建模两者优势互补。最后说一个我在这个项目里最想强调的建议写完代码后一定要手动推一遍每个tensor的shape变化。打开一张32×32的CIFAR图用笔在纸上写下从[1, 3, 32, 32]到最后[1, 10]的每一步shape变化。这个过程看似笨拙却是建立对Transformer架构直觉最有效的方法没有之一。我教你给团队同学讲原理、做Code Review时第一步先看他能不能完整推导shape能推出来代码基本已经有八成把握了。
02
RELATED NEWS

相关资讯

更多网站建设与数字化升级内容

03
WHY YAOTU

想打造同款高转化官网?

懂行业、懂生意,从建站到增长一站式陪跑

◈

场景化定制

不做模板站,围绕你的业务场景量身设计,小众不撞款。

◐

营销型架构

以转化目标组织内容与路径,让官网真正带来询盘。

▲

全周期服务

设计、开发、运营、运维一体,上线只是开始。

免费获取你的建站方案

留下需求,专属顾问 24 小时内为你输出方案建议。