You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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])

原因分析

问题出在两个核心点:

  1. 缺少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]。
  2. 遗漏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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.16 06:17:00