PyTorch torchvision自0.12版起原生支持ViT,如vit_b_16可直接加载预训练模型;骨干由conv_proj(patch embedding)和encoder(Transformer)组成,输出(N,197,768),需截取patch tokens并reshape为(N,768,14,14)适配下游任务。

PyTorch 官方 torchvision 从 0.12 版本起已原生支持 Vision Transformer(ViT)骨干,无需手动拼装 patch embedding 或 Transformer encoder——直接调用预训练模型即可用于下游任务。
用 torchvision.models.vit_b_16 等函数快速加载标准 ViT 骨干
PyTorch 提供了多个预训练 ViT 变体,命名规则清晰:vit_b_16 表示 Base size、16×16 patch;vit_l_16 是 Large;vit_h_14 是 Huge(需 torchvision ≥ 0.15)。它们返回的是完整 nn.Module,但默认包含分类头。
若只需骨干(即去掉最后的 heads 分类层),推荐写法:
from torchvision.models import vit_b_16 <p>model = vit_b_16(weights="DEFAULT")</p><p><span>立即学习</span>“<a href="https://pan.quark.cn/s/00968c3c2c15" style="text-decoration: underline !important; color: blue; font-weight: bolder;" rel="nofollow" target="_blank">Python免费学习笔记(深入)</a>”;</p><h1>去掉分类头,只保留 backbone</h1><p>backbone = torch.nn.Sequential(*list(model.children())[:-1])</p><h1>注意:这会保留 norm 层,但丢掉 heads;更稳妥的方式是直接访问 model.encoder
不过更推荐直接使用 model.encoder + model.conv_proj 组合,因为 ViT 的“骨干”逻辑实际分布在:
-
model.conv_proj:等价于 patch embedding 的卷积实现(输入 3×224×224 → 输出 768×14×14) -
model.encoder:纯 Transformer encoder(含 LayerNorm、MSA、MLP) -
model.encoder.ln:最终的层归一化,输出 token features
手动构建 ViT 骨干时 patch embedding 的常见错误
很多人照着论文用 nn.Conv2d 或 nn.Unfold 实现 patch 切分,结果 shape 对不上或梯度中断。关键点在于:
SkillSub Pro - Python 题解与代码注释双功能技能功能概述SkillSub Pro - Python 题解与代码注释双功能技能是一项面向实际任务的技能,主要用于SkillSub Pro 是一个 Python 题解生成与代码注释的 双功能合体技能 ,专为学生、算法学习者和开发者设计;✅ 一个技能,两种用途 :;核心要点📝 题解模式 :输入题目/题号,自动生成完整 Python 题解(含详细注释、解题思路、复杂度分析);💬 注释模式 :输入 Python 代码,自动添加详细中。它将相关步骤、
- ViT 的 patch embedding 不是“先 unfold 再线性”,而是用
nn.Conv2d(in_c, embed_dim, kernel_size=patch_size, stride=patch_size)等效实现,它天然保证 spatial 维度对齐 - 若用
nn.Unfold,必须配合permute(0, 2, 1)和linear,且unfold输出 shape 是(N, C×P², L),容易漏掉维度重排 - 位置编码要加在 class token 拼接之后,不是 patch embedding 之后——即顺序是:
[cls_token; patches]→+ pos_embed
错误示例(导致 shape mismatch):
x = x.unfold(2, 16, 16).unfold(3, 16, 16) # 错:输出为 N×C×14×14×16×16,难处理 x = rearrange(x, 'b c h w p1 p2 -> b (h w) (c p1 p2)') # 易出错且无必要
ViT 骨干输出 shape 与下游适配的关键细节
标准 vit_b_16 输入 3×224×224,输出 features 是 (N, 197, 768):其中 197 = 1(class token)+ 14×14(patch tokens)。但多数检测/分割任务需要 spatial feature map(如 N×768×14×14),这时不能直接 squeeze class token。
正确做法:
- 取
model.encoder(x)输出(shape(N, 197, 768)) - 剔除第一个 token:
x_feat = x[:, 1:, :] - reshape:
x_feat = x_feat.permute(0, 2, 1).reshape(N, 768, 14, 14)
注意:model.encoder 不包含 final layer norm,而 model 整体会做。若需带 norm 的特征,应使用 model(x) 后截取,或显式调用 model.encoder.ln(model.encoder(x))。
兼容性与性能陷阱:为什么不用 torch.hub 加载 timm 的 ViT
虽然 timm 提供更多 ViT 变体(如 vit_base_patch16_224),但它和 torchvision 的接口不一致:
-
timm默认输出是(N, 768)(仅 class token),无 patch token;开启global_pool='token'才能控制 -
timm的 position embedding 是可学习参数,torchvision的是固定正弦编码(v0.15+ 已统一为可学习) - 混合精度训练时,
timm某些版本的DropPath在torch.compile下有 bug,而torchvision的Dropout更稳定
除非你明确需要 Deformable ViT 或 CrossViT 等非标结构,否则优先用 torchvision.models.vit_*——它的 backbone 抽取路径清晰、导出 ONNX 支持好,且和 detection/segmentation reference scripts 开箱兼容。
ViT 骨干最易被忽略的其实是输入预处理:它要求像素值归一化到 [0, 1] 后再按 ImageNet 统计值标准化(mean=[0.485,0.456,0.406], std=[0.229,0.224,0.225]),和 CNN 一致,但有人误用 -1~1 范围导致性能断崖下跌。

















