ViT模型逐层重构后输出维度不匹配,求原因及解决方案
问题:重构ViT模型后输出张量维度不一致
问题描述
尝试通过逐层复制timm库中ViT模型的组件来重构模型,结果新模型输出形状为[4,196,10],但原模型输出为[4,10]。
原模型代码
import torch import torch.nn as nn import timm class ViTImageClassifier(nn.Module): def __init__(self, num_classes): super(ViTImageClassifier, self).__init__() self.backbone = timm.create_model('vit_base_patch16_224', pretrained=True) self.backbone.head = nn.Linear(self.backbone.head.in_features, num_classes) def forward(self, x): x = self.backbone(x) return x model = ViTImageClassifier(num_classes=10) input = torch.randn(4, 3, 224, 224) output = model(input) print("原模型输出形状:", output.shape) # 输出: torch.Size([4, 10])
重构代码(存在问题)
new_model=[] new_model.append(model.backbone.patch_embed) new_model=nn.Sequential(*new_model) new_model.append(model.backbone.pos_drop) new_model.append(model.backbone.patch_drop) new_model.append(model.backbone.norm_pre) for i in range(12): new_model.append(model.backbone.blocks[i]) new_model.append(model.backbone.norm) new_model.append(model.backbone.fc_norm) new_model.append(model.backbone.head_drop) new_model.append(model.backbone.head) out=new_model(input) print("重构模型输出形状:", out.shape) # 输出: torch.Size([4, 196, 10])
原因分析
问题出在两个核心点:
- 缺少cls_token的添加与提取:ViT模型的核心设计是使用一个可学习的
cls_token作为分类特征的载体。原模型在patch embedding后会自动添加这个token,与196个patch特征拼接成[batch, 197, embed_dim]的序列(1个cls_token + 196个patch),最后分类时只取cls_token对应的特征(序列第0个元素)。而你的重构代码既没有添加cls_token,也没有提取该特征,导致head层直接处理196个patch的特征,输出维度变为[4,196,10]。 - 遗漏position embedding的加法:原模型的forward流程中,会将patch embedding结果与position embedding相加后再送入dropout层,这一步在你的重构代码中被跳过,不仅会导致维度问题,还会破坏特征的位置信息。
解决方案
方案1:自定义重构模型类(推荐)
通过自定义nn.Module完整复现原模型的forward逻辑,确保包含cls_token和position embedding的处理:
class ReconstructedViT(nn.Module): def __init__(self, original_backbone): super().__init__() # 复制原backbone的所有组件 self.patch_embed = original_backbone.patch_embed self.pos_drop = original_backbone.pos_drop self.patch_drop = original_backbone.patch_drop self.norm_pre = original_backbone.norm_pre self.blocks = nn.Sequential(*original_backbone.blocks) self.norm = original_backbone.norm self.fc_norm = original_backbone.fc_norm self.head_drop = original_backbone.head_drop self.head = original_backbone.head def forward(self, x): # 1. Patch embedding x = self.patch_embed(x) # shape: [4, 196, 768] # 2. 添加cls_token并拼接 cls_token = self.patch_embed.cls_token.expand(x.shape[0], -1, -1) # shape: [4, 1, 768] x = torch.cat((cls_token, x), dim=1) # shape: [4, 197, 768] # 3. 添加position embedding并送入dropout x = self.pos_drop(x + self.patch_embed.pos_embed) # 4. 后续层处理 x = self.patch_drop(x) x = self.norm_pre(x) x = self.blocks(x) x = self.norm(x) # 5. 提取cls_token特征 x = x[:, 0] # shape: [4, 768] # 6. 分类头处理 x = self.fc_norm(x) x = self.head_drop(x) x = self.head(x) return x # 初始化重构模型 new_model = ReconstructedViT(model.backbone) out = new_model(input) print("重构模型输出形状:", out.shape) # 输出: torch.Size([4, 10])
方案2:手动调整forward流程
如果不想自定义模型类,可以在现有代码基础上手动补全缺失的步骤:
# 手动执行forward流程 x = input # 1. Patch embedding + 添加cls_token + position embedding x = model.backbone.patch_embed(x) cls_token = model.backbone.patch_embed.cls_token.expand(x.shape[0], -1, -1) x = torch.cat((cls_token, x), dim=1) x = model.backbone.pos_drop(x + model.backbone.patch_embed.pos_embed) # 2. 依次通过后续层 x = model.backbone.patch_drop(x) x = model.backbone.norm_pre(x) for block in model.backbone.blocks: x = block(x) x = model.backbone.norm(x) # 3. 提取cls_token特征 x = x[:, 0] # 4. 分类头处理 x = model.backbone.fc_norm(x) x = model.backbone.head_drop(x) x = model.backbone.head(x) print("调整后输出形状:", x.shape) # 输出: torch.Size([4, 10])
关键注意点
- ViT的
cls_token是分类任务的核心特征,必须在patch embedding后添加,并在最后提取该token的输出。 - Position embedding负责编码图像的空间位置信息,必须与patch embedding相加后再送入后续层。
内容的提问来源于stack exchange,提问作者PikovO
相关产品推荐
相关产品推荐

