ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

ViT源码复现:从Patch Embedding到位置编码的PyTorch实现

2026/9/26 16:43:41 拓冰建站 浏览量
ViT源码复现:从Patch Embedding到位置编码的PyTorch实现 很多读者对 ViT 的认知是从“把图片切成 16×16 的小块然后丢给 Transformer”这段动画开始的。动画看完好像什么都懂了patch、position embedding、class token每个名词都能念出来。但真正动手做源码复现的时候第一个问题就能把人卡住——序列长度到底是 196 还是 197pos_embed 要加在哪一层为什么模型跑起来 loss 一点不降我觉得好的结构理解不应该只停留在“图像”层面而要落到“张量”层面。ViT 的本质是把图片变成一组 token再用 Transformer 做全局建模。这件事听上去简单但涉及图像 Patch 化、特征投影、位置编码、多头自注意力、分类头五个环节任何一环的维度对不上代码都会报错或者训练不起来。本文不使用动画直接给你一个可以跑通的 PyTorch 版 ViT 源码从原理到实现一步步拆解带你彻底搞清楚 ViT 的全过程。1. 这篇文章真正要解决的问题网上讲 ViT 的教程并不少但你很可能遇到过下面这几种情况第一个情况是“按论文复现但跑不懂”。ViT 的核心代码只有几十行但每个子模块的输入输出维度环环相扣。patch_embed输出的形状是[B, N, D]可一旦插入了cls_token序列长度就从N变成N1如果你还按原来的N去初始化和添加位置编码模型直接报维度不匹配。第二个情况是“小数据集上完全训练不起来”。ViT 在学术界的成功建立在 JFT-300M 这种超大规模数据集上。你直接拿 CIFAR-10 去训练一个 ViT-Base效果往往还不如一个简单的 ResNet-18。于是很多初学者下了个结论ViT 不如 CNN。这个结论其实忽略了分辨率、patch 大小、数据增强和训练策略等因素。第三个情况是“不知道 ViT 对今天的大模型有什么意义”。ViT 不是一个只活在 2020 年论文里的模型。CLIP 的视觉编码器、DINO、SAM以及很多多模态大模型的视觉底座本质上都是 ViT 或类似像素成 token 的结构。不把 ViT 吃透后面看视觉大模型、多模态大模型都会很吃力。这篇文章要做的就是三件事第一把 ViT 的设计动机和每个模块的原理讲清楚第二用 PyTorch 从零写一个可运行的简化版 ViT 并跑通前向传播和训练第三把复现中最高频的报错和最容易被忽视的训练细节整理出来让你少走弯路。读完以后你可以自己动手改 patch 大小、改层数、换数据集也能理解为什么很多视觉大模型会把分辨率、patch 大小和位置编码插值这些东西当成关键配置。2. ViT 的核心原理一张图如何变成一串 Token2.1 为什么图像需要 Transformer在 ViT 之前图像领域是卷积神经网络CNN的天下。CNN 靠卷积核在局部窗口里提取特征通过不断堆叠卷积层来扩大感受野。但卷积有两个天然限制一是感受野扩张速度慢底层特征很难直接看到图像远处的内容二是它对“局部性”和“平移等变性”有很强的先验这在大规模数据下反而不一定是最优选择。Transformer 最早用在自然语言处理中核心机制是自注意力。自注意力让序列里任意两个位置的 token 直接做信息交互一步就能完成全局建模。于是研究者提出了一个很自然的想法图像能不能也像句子一样变成一串 token然后用 Transformer 来建模这就是 ViTVision Transformer的基本思路。ViT 最初由 Google 在 2020 年提出论文标题是An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale。这个标题已经把核心讲透了一张图像可以被看成 16×16 的“单词”序列。它带来的是一次结构上的简化——不再需要卷积和池化的复杂堆叠把图像直接映射成序列让纯 Transformer 来处理视觉任务。2.2 Patch Embedding图像如何变成 Token假设输入是一张 224×224×3 的 RGB 图片patch_size16。ViT 会把图片均匀切成互不重叠的小块每个块的大小是 16×16一共切出[ \frac{224}{16} \times \frac{224}{16} 196 ]个 patch。每个 patch 的形状是 16×16×3展平后是 768 维的向量。这一步很关键展平后的 768 维不能直接作为 token 使用因为它的维度太大而且每个维度的语义分散。ViT 的做法是通过一个可学习的线性投影把 768 维映射成embed_dim维通常 768 或 1024。这个线性投影本质上就是一张像素块到 token 特征的映射表。在代码实现里一个常用的等价做法是用卷积层实现。Conv2d(in_channels3, out_channelsembed_dim, kernel_sizepatch_size, stridepatch_size)正好等价于“切块 展平 线性投影”三个操作的融合。这也是我们在自己复现时最干净利落的实现方式。每个 patch 经过投影后得到一个向量叫做 token。224×224 的图会得到 196 个 token形状变为[B, 196, embed_dim]。到了这一步图像已经变成和 NLP 里句子序列类似的结构每个 token 对应图片的一个局部区域token 之间的序列关系就是空间位置关系。2.3 Position Embedding给 Token 加上位置信息Transformer 本身是一种置换等变结构。也就是说你把 token 序列的顺序打乱自注意力输出的集合是不会变的只是顺序跟着变。但图像是有空间结构的左上角的 patch 和右下角的 patch 显然不同因此必须在输入序列里显式加入位置信息。ViT 使用了一个简单而有效的方案可学习的一维位置嵌入Position Embedding。它初始化一个形状为[1, 1961, embed_dim]的矩阵加在 patch embedding 的输出上。这里的“1”是预留给了后面的cls_token我们下一节再解释。这里有一个初学者容易问的问题为什么不直接用卷积里那种二维位置编码ViT 论文做过对比实验可学习的一维位置编码和更复杂的二维位置编码效果几乎持平。原因在于 patch 本身已经包含了很强的局部结构信息一维位置编码的序号足够让模型推断出二维的相对位置关系。所以在源码复现时不需要特意设计 2D 位置编码直接用可学习的一维位置编码即可。2.4 Class Token分类任务的关键设计在 BERT 里输入序列开头有一个[CLS]token最后用它对应的输出去做分类。ViT 借鉴了这个设计在 patch 序列的最前面插入一个cls_token。最终序列长度是1961197。为什么要单独放一个类别 token而不是直接用全局平均池化的结果从直觉上说cls_token是一个“可学习的位置”它没有任何图像 patch 的语义完全靠注意力机制去聚合整张图片的信息。训练过程中它会慢慢学会“关注最有助于分类的 patch”充当一个全局信息收集器。论文中的实验也表明在同等条件下cls_token的效果要优于简单的全局平均池化。推理时我们只需要取 Transformer Encoder 输出中第一个位置对应cls_token的向量后面接一个 LayerNorm 和线性分类头就能得到分类 logits。2.5 Transformer Encoder全局自注意力经过 Patch Embedding、位置编码和cls_token拼接之后输入已经变成了一个形状为[B, 197, embed_dim]的序列。接下来这个序列会进入标准 Transformer Encoder。Encoder 的每个 block 由三部分组成LayerNorm、多头自注意力Multi-Head Self-Attention、MLP 模块。ViT 采用的是 Pre-LN 结构也就是先做 LayerNorm再进入注意力层最后用残差连接相加。MLP 部分通常使用两层线性层和 GELU 激活函数隐藏层维度通常设置为embed_dim的 4 倍。为什么 Pre-LN 很重要在原始 Transformer 中Post-LN 的深层网络在训练早期很容易出现梯度爆炸或者不收敛的问题往往需要复杂的学习率预热策略。而 Pre-LN 把 LayerNorm 放在残差分支入口可以有效稳定深层 Transformer 的训练过程。这也是 ViT 能在几百层甚至上千层深度上稳定训练的重要保证。多头注意力让模型可以从多个子空间去捕捉 patch 之间的关联。比如某些头关注颜色和纹理某些头关注物体的大致轮廓另一些头可能关注跨区域的上下文关系。这种全局建模能力正是 ViT 与 CNN 在结构上最本质的差异。2.6 ViT 与 CNN、混合模型的技术路线对比理解 ViT最好把它放到视觉主干网络的技术路线里看。下面这张表整理了三类路线的典型代表和核心差异路线典型代表核心思想适合场景局限性CNN 路线ResNet、RegNet、ConvNeXt局部卷积 逐步扩大感受野中小数据集、移动端、需要低延迟的场景长距离依赖建模弱纯 ViT 路线ViT、DeiT图像切块成 token全局自注意力大数据集、视觉预训练、多模态底座小数据容易欠拟合训练成本高混合路线Swin Transformer、PVT分层 token窗口局部注意力 全局建模分割、检测等密集预测任务结构更复杂实现难度高从材料来看纯 ViT 路线的价值更多在于它为视觉模型打开了“比卷积更通用”的可能性。Swin Transformer 又引入了 CNN 的分层思想来缓解 ViT 的计算量问题这也是视觉 Transformer 主流技术路线里很重要的一支。对初学者来说先把纯 ViT 的代码跑通再去理解 Swin 的分层窗口注意力会发现思路是相通的。3. 环境准备与前置条件在正式开始源码复现之前先把环境准备好。为了不引入额外复杂度这里直接使用 PyTorch不依赖timm这类现成实现库所有模块都自己写一遍。需要准备的环境如下Python 3.9 或更高版本。PyTorch 2.x任意稳定版本均可本文以官方 PyTorch 为准。torchvision用于加载 CIFAR-10 数据集和做图像预处理。一个带 GPU 的环境最好没有 GPU也可以把输入尺寸和模型缩小用 CPU 验证前向传播。安装命令参考pip install torch torchvision如果你有 CUDA 版本的 GPU建议到 PyTorch 官网根据机器的 CUDA 版本选择对应的安装命令。版本差异不会影响本文代码的逻辑因为核心 API 都是 PyTorch 长期稳定的接口。准备好以后新建一个项目文件夹结构如下vit-demo/ ├── vit_model.py # ViT 模型定义 ├── train_cifar.py # 训练与验证脚本 └── data/ # 数据集下载目录下面所有代码都围绕这两个 Python 文件展开。4. 用 PyTorch 从零实现一个可训练的 ViT4.1 定义 Patch EmbeddingPatchEmbedding是整个 ViT 结构的第一步。实现思路是用Conv2d完成 cut and projection。输入是[B, 3, H, W]经过一个kernel_sizepatch_size, stridepatch_size的卷积后得到[B, embed_dim, H/patch, W/patch]然后 flatten 成[B, embed_dim, N]最后转置成[B, N, embed_dim]。# 文件路径vit_model.py import math import torch import torch.nn as nn class PatchEmbedding(nn.Module): 把图像切分成 patch 并线性投影成 token 序列 def __init__(self, in_channels3, patch_size16, embed_dim768, img_size224): super().__init__() self.patch_size patch_size self.num_patches (img_size // patch_size) ** 2 self.proj nn.Conv2d( in_channelsin_channels, out_channelsembed_dim, kernel_sizepatch_size, stridepatch_size, ) def forward(self, x): # x: [B, C, H, W] x self.proj(x) # [B, embed_dim, H/p, W/p] x x.flatten(2) # [B, embed_dim, num_patches] x x.transpose(1, 2) # [B, num_patches, embed_dim] return x这里需要注意num_patches的计算依赖于H和W都能被patch_size整除。如果输入图像的宽高不是 patch_size 的整数倍卷积会报错或者产生不对称的 patch 数量后续的位置编码维度就会全部对不上。这也是输入 ViT 前必须做 resize 的原因。4.2 定义 Transformer Encoder BlockTransformerBlock是 ViT 中的基础组成单元。我们使用 Pre-LN 结构先对输入做 LayerNorm然后进入多头自注意力残差相加再接一次 LayerNorm进入 MLP残差相加。PyTorch 自带的nn.MultiheadAttention可以直接使用注意设置batch_firstTrue这样输入输出形状都是[B, seq_len, embed_dim]更容易理解。class TransformerBlock(nn.Module): Pre-LN Transformer Encoder Block def __init__(self, embed_dim768, num_heads12, mlp_ratio4.0, dropout0.1): super().__init__() self.norm1 nn.LayerNorm(embed_dim) self.attn nn.MultiheadAttention( embed_dim, num_heads, dropoutdropout, batch_firstTrue ) 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.Linear(hidden_dim, embed_dim), nn.Dropout(dropout), ) def forward(self, x): # x: [B, N, D] normed self.norm1(x) attn_out, _ self.attn(normed, normed, normed) x x attn_out x x self.mlp(self.norm2(x)) return x这里一个容易出错的细节是nn.MultiheadAttention的q、k、v三个参数都应该来自同一个normed否则就变成了先用未归一化的数据做注意力结构上和 Pre-LN 不一致。很多代码首次复现时直接传x也会训练起来但不够规范。4.3 组装 VisionTransformer现在把PatchEmbedding、cls_token、position embedding、若干TransformerBlock、LayerNorm 和分类头组装成完整的 ViT。class VisionTransformer(nn.Module): 简化版 Vision Transformer结构对齐原论文 def __init__( self, img_size224, patch_size16, in_channels3, embed_dim768, depth12, num_heads12, num_classes1000, dropout0.1, ): super().__init__() self.patch_embed PatchEmbedding( in_channelsin_channels, patch_sizepatch_size, embed_dimembed_dim, img_sizeimg_size, ) num_patches self.patch_embed.num_patches 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.blocks nn.Sequential(*[ TransformerBlock(embed_dim, num_heads, dropoutdropout) for _ in range(depth) ]) self.norm nn.LayerNorm(embed_dim) self.head nn.Linear(embed_dim, num_classes) # 参数初始化参考 DeiT 的初始化风格 nn.init.trunc_normal_(self.pos_embed, std0.02) nn.init.trunc_normal_(self.cls_token, std0.02) 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) def forward(self, x): B x.shape[0] x self.patch_embed(x) # [B, N, D] cls_token self.cls_token.expand(B, -1, -1) # [B, 1, D] x torch.cat([cls_token, x], dim1) # [B, N1, D] x x self.pos_embed x self.blocks(x) x self.norm(x) cls_out x[:, 0] # 取 cls_token 位置 logits self.head(cls_out) return logits核心逻辑在forward中看得最清楚先通过 Patch Embedding 得到[B, N, D]然后拼接cls_token变成[B, N1, D]加上位置编码后过 Encoder最后取第一个位置进行分类。4.4 前向传播验证模型定义好后先不急着训练。我们用一个很小的配置验证前向传播能否跑通。这里为了节约资源把embed_dim调成 192depth调成 6num_heads调成 6这就是一个“微型 ViT”。在 CPU 上也能直接运行。def test_forward(): device torch.device(cpu) model VisionTransformer( img_size224, patch_size16, embed_dim192, depth6, num_heads6, num_classes10, ) x torch.randn(2, 3, 224, 224) logits model(x) print(输入:, x.shape) print(输出:, logits.shape) if __name__ __main__: test_forward()预期输出输入: torch.Size([2, 3, 224, 224]) 输出: torch.Size([2, 10])如果你的输出形状也是[2, 10]说明 ViT 的整体结构已经正确。接下来只需要把数据接入这个模型就能开始训练。5. 适配小数据集将 ViT 用于 CIFAR-10很多初学者直接拿 ViT-Base 去跑 CIFAR-10结果发现训练很慢、效果很差。这并不意外。标准 ViT-Base 有 12 层、12 个头、embed_dim 768参数量约 86M在 CIFAR-10 这样的小数据上非常容易过拟合。更关键的是CIFAR-10 的分辨率只有 32×32如果 patch_size 还是 16那么每张图只有 4 个 patch序列长度太短注意力机制几乎没有发挥空间。所以在代码演示中我会把 ViT 调整成适合小数据集的微型版本patch_size4embed_dim128depth4num_heads4。这样 32×32 的图片会被切成 8×864 个 patch序列长度足够模型也能在几分钟内完成训练验证。def build_tiny_vit(img_size32, patch_size4, num_classes10): model VisionTransformer( img_sizeimg_size, patch_sizepatch_size, embed_dim128, depth4, num_heads4, num_classesnum_classes, ) return model数据加载部分使用 torchvision 的 CIFAR-10。为了提升 ViT 在小数据上的稳定性我们使用 RandomCrop 和 RandomHorizontalFlip 做基础数据增强并归一化到[-1, 1]附近。# 文件路径train_cifar.py import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader from torchvision import datasets, transforms from vit_model import VisionTransformer transform_train transforms.Compose([ transforms.RandomCrop(32, padding4), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(mean[0.4914, 0.4822, 0.4465], std[0.2470, 0.2435, 0.2616]), ]) transform_test transforms.Compose([ transforms.ToTensor(), transforms.Normalize(mean[0.4914, 0.4822, 0.4465], std[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_size64, shuffleTrue, num_workers2) test_loader DataLoader(test_dataset, batch_size64, shuffleFalse, num_workers2)注意CIFAR-10 的原始图片是 32×32 的 PIL 图像RandomCrop需要输入大于等于 32 才能正常工作所以先不额外 Resize。patch_size4恰好能整除 3264 个 patch 也让位置编码有足够的信息量。6. 训练与验证训练循环本身和普通 PyTorch 分类任务没有差别。ViT 的优化器通常使用 AdamW并且建议设置一个较小的 weight decay。这里为了演示直接给出一个最小可运行的训练脚本。def train_one_epoch(model, loader, optimizer, device): model.train() total_loss 0.0 correct 0 total 0 for images, labels in loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() logits model(images) loss nn.functional.cross_entropy(logits, labels) loss.backward() optimizer.step() total_loss loss.item() correct (logits.argmax(dim1) labels).sum().item() total labels.size(0) avg_loss total_loss / len(loader) acc correct / total return avg_loss, acc def evaluate(model, loader, device): model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in loader: images, labels images.to(device), labels.to(device) logits model(images) correct (logits.argmax(dim1) labels).sum().item() total labels.size(0) return correct / total def main(): device torch.device(cuda if torch.cuda.is_available() else cpu) print(使用设备:, device) model VisionTransformer( img_size32, patch_size4, embed_dim128, depth4, num_heads4, num_classes10, ).to(device) optimizer optim.AdamW(model.parameters(), lr1e-3, weight_decay5e-2) epochs 20 for epoch in range(epochs): train_loss, train_acc train_one_epoch(model, train_loader, optimizer, device) test_acc evaluate(model, test_loader, device) print(fEpoch {epoch 1:02d}/{epochs} | floss {train_loss:.4f} | ftrain acc {train_acc:.4f} | ftest acc {test_acc:.4f}) if __name__ __main__: main()运行方式python train_cifar.py如果你的机器有 GPU这个微型 ViT 在 20 个 epoch 内应该能看到训练 loss 持续下降测试准确率逐步提升。如果只有 CPU也建议先跑 1 到 2 个 epoch确认数据管道和模型都能正常工作。这里的重点不是冲击 SOTA而是跑通“数据 - 模型 - 损失 - 优化 - 验证”的完整闭环。需要说明的是以上训练脚本在数据集规模、模型容量和训练轮数上都不是为高精度设计的。要和大规模预训练的 ViT 对比不能用这种微型配置。它真正的价值在于让你能在几分钟内看到 ViT 的训练动态排除结构错误理解过拟合和欠拟合分别在什么现象中出现。7. 常见问题与排查思路源码复现过程中我见过的高频问题基本都集中在维度不匹配、训练不收敛、预训练权重加载失败这三类。下面整理成一张排查表问题现象可能原因排查方式解决方案模型报 size mismatch序列长度对不上输入图像宽高不是 patch_size 的整数倍打印x.shape和num_patches先 resize 到固定尺寸确保能被 patch_size 整除pos_embed 加载或初始化维度错误忘了cls_token会占用一个位置检查num_patches 1是否一致位置编码长度必须是num_patches 1训练 loss 不降学习率过高或过低数据未归一化尝试 1e-3 或 1e-4 的 AdamW检查输入范围小模型用小 lr大模型做大 lr配合 warmup小数据集严重过拟合模型容量大、数据增强不足看 train acc 和 test acc 差距缩小 embed_dim/depth增加数据增强或正则化CPU 训练特别慢序列长度长、模型层数多统计单 step 耗时降低 patch_size 的倒数效应减小输入或 batch加载官方预训练权重报错分类头类别数与预训练不一致打印 state_dict key 和 shape只加载 backbone 部分或者忽略head.weight这里特别想提醒一个容易被忽略的细节ViT 的位置编码是在“训练分辨率”下定义的。如果你先把模型在 224×224 上预训练后面又想用 384×384 输入做微调patch 数量会从 196 变成 576位置编码矩阵跟不上。这时候需要做位置编码插值通常是使用双线性插值将pos_embed从[1, 197, D]插值到[1, 577, D]。timm库中提供了现成实现但如果自己写模型这一步必须单独处理。8. 最佳实践与工程建议8.1 输入尺寸固定是第一原则ViT 不是一个天然支持任意尺寸输入的结构。CNN 可以用全卷积结构在不同分辨率下滑动扫描而 ViT 的 position embedding 和 patch 数量绑定在一起。工程上需要统一所有图片的分辨率或者使用一个支持分辨率变化的 wrapper 来处理位置编码插值。最简单的做法就是在 DataLoader 的 transform 里强制Resize到固定尺寸。8.2 优先使用 Pre-LN 结构很多初学者照着原始 Transformer 论文写 Post-LN结果训练深层 ViT 时频繁遇到发散的情况。代码实现里LayerNorm - Attention - 残差的 Pre-LN 结构才是 ViT 实际训练中更稳妥的选择。如果你要改造自己的 Transformer 模块建议从 Pre-LN 起步。8.3 小数据集慎重选择 ViT在数据量有限的场景ViT 不是万能方案。除非你有足够的预训练权重否则直接在小型业务数据集上训练标准 ViT效果大概率不如 ResNet 或 EfficientNet。一种折中方案是使用 DeiT 这类在 ImageNet-1k 上蒸馏出来的 ViT 预训练权重做微调另一种是使用 Swin Transformer 这类带局部先验的层级结构。8.4 位置编码插值与大分辨率微调在检测、分割等场景中输入分辨率往往比 224×224 大。使用预训练 ViT 时必须处理位置编码的尺寸变化。推荐做法是在加载预训练权重之后对pos_embed做插值同时可以延长 warmup 轮数让模型适应新的位置编码分布。8.5 用好注意力可视化调试模型ViT 有一个自带的可解释性优势我们可以直接读取cls_token对不同 patch 的注意力权重画出热力图观察模型关注到了哪些区域。如果热力图上模型关注点很分散或者完全忽略目标主体往往说明训练不充分、数据偏差或者位置编码处理不当。这种可视化调试成本很低值得在生产项目中引入。8.6 考虑混合精度与梯度累积ViT 相比同级别的 CNN训练显存开销大。实际项目中经常使用 AMP混合精度训练来降低显存占用并在 batch size 受限时用梯度累积来模拟更大 batch。这些策略不影响模型结构但对训练效率和稳定性提升明显。如果是生产环境的分布式训练还需要注意shuffle的随机种子设置保证数据流的可复现性。9. 总结与后续学习方向这篇文章从 ViT 的设计动机讲起把图像 patch 化、token 投影、位置编码、cls_token 和 Transformer Encoder 逐层拆开并把完整的 PyTorch 源码走了一遍。跑通了前向传播也跑通了 CIFAR-10 上的微型训练流程。比“看动画理解 ViT”更进一步的是你现在知道每个张量的形状是怎么变化的也知道维度不匹配通常发生在哪几个环节。如果接下来要深入建议按下面几条路线继续走第一看 ViT 的变体。Swin Transformer 把层级设计重新引入视觉 TransformerDeiT 用知识蒸馏解决了中小数据集上的训练困难这些模型都是在 ViT 基础上做结构性改进的典型代表。第二看 ViT 在自监督中的使用。MAE 通过掩码图像建模直接在大规模无标注图像上预训练DINO 则用自蒸馏方式让 ViT 学到富含语义的特征。理解 ViT 结构之后这些模型的代码会好读很多。第三看视觉和语言的统一。CLIP 使用 ViT 作为图像编码器把图像和文本拉到同一个向量空间之后的多模态大模型也大量使用类似思路。你会发现图像切块成 token 这个设计几乎是视觉语言融合的前提。最后给你一个实际项目的提醒选型时别只看模型结构还要看你的数据规模和算力预算。数据量不够就用预训练权重微调算力受限就用混合模型或直接选 CNN。真正把 ViT 吃透的人不是只会跑通一个模型而是知道在什么场景下该用它在什么场景下不该用它。希望你读完本文后能带着这份判断力去继续探索。