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

基于PyTorch与timm库修改ViT模型自定义分类头的问题

ViT自定义分类头:移除默认头部但保留类Token的问题

问题背景

我用PyTorch和timm库开发Vision Transformer(ViT)模型,目标是替换默认分类头,实现一个取所有Token均值并添加新分类层的自定义分类头。

原ViT模型末尾结构如下:

LayerNorm-247             [-1, 197, 768]           1,536
        Identity-248                  [-1, 768]               0
         Dropout-249                  [-1, 768]               0
          Linear-250                 [-1, 1000]         769,000
VisionTransformer-251                 [-1, 1000]               0

我编写了这段代码尝试移除最后几层:

class VisionTransformerWithoutHead(nn.Module):
    
    def __init__(self, model_name):
        super(VisionTransformerWithoutHead, self).__init__()

        # Load the ViT model
        vit_model = timm.create_model(model_name, pretrained=True)

        # Remove the final layers
        self.features = nn.Sequential(*list(vit_model.children())[:-1])

    def forward(self, x):
        # Forward pass through the modified model
        output = self.features(x)
        return output

但修改后模型末尾结构变成这样,Token数量从197降到196,类Token被移除了:

LayerNorm-247             [-1, 196, 768]           1,536
        Identity-248             [-1, 196, 768]               0
         Dropout-249             [-1, 196, 768]               0

原因分析

问题出在timm库ViT模型的逻辑实现方式:

  • ViT中类Token(cls token)的添加、位置嵌入的相加等操作,是写在VisionTransformer类的forward方法里的,并不是作为独立的子模块存在。
  • 你用list(vit_model.children())[:-1]提取的是模型的子模块集合,然后用nn.Sequential包装,但这个Sequential只会依次执行子模块的forward,不会触发原模型forward里添加类Token的逻辑。
  • 原模型的patch_embed子模块输出的是196个图像Patch对应的Token,类Token是在原模型forward中通过张量拼接手动添加的,你的自定义模型没有执行这一步,所以最终输出只有196个Token。

解决方案

推荐两种简洁的实现方式,都能保留类Token并替换分类头:

方式1:利用timm的num_classes=0参数移除默认头

timm的create_model函数支持num_classes=0参数,设置后模型会跳过默认分类头,直接输出包含类Token的所有特征张量(形状为[batch_size, 197, 768]),之后你可以直接添加自定义分类头:

import torch
import torch.nn as nn
import timm

# 加载移除默认分类头的ViT模型,输出包含类Token的所有特征
vit_backbone = timm.create_model("vit_base_patch16_224", pretrained=True, num_classes=0)

# 自定义分类头:取所有Token均值后加线性层
class CustomClassificationHead(nn.Module):
    def __init__(self, embed_dim, num_classes):
        super().__init__()
        self.norm = nn.LayerNorm(embed_dim)
        self.fc = nn.Linear(embed_dim, num_classes)
    
    def forward(self, x):
        # 对所有Token(含类Token)取均值
        x = x.mean(dim=1)
        x = self.norm(x)
        x = self.fc(x)
        return x

# 组合骨干网络和自定义头
model = nn.Sequential(
    vit_backbone,
    CustomClassificationHead(vit_backbone.embed_dim, num_classes=10)  # 替换为你的目标类别数
)

# 测试输出
x = torch.randn(2, 3, 224, 224)
output = model(x)
print(output.shape)  # 应为 [2, 10]

方式2:复用原模型的forward逻辑

如果你需要更精细的控制,可以手动复现原模型的特征提取逻辑,再添加自定义头:

import torch
import torch.nn as nn
import timm

class CustomViT(nn.Module):
    def __init__(self, model_name, num_classes):
        super().__init__()
        # 加载预训练ViT模型
        self.vit = timm.create_model(model_name, pretrained=True)
        # 可选:冻结骨干网络参数
        for param in self.vit.parameters():
            param.requires_grad = False
        # 自定义分类头
        self.custom_head = nn.Sequential(
            nn.LayerNorm(self.vit.embed_dim),
            nn.Linear(self.vit.embed_dim, num_classes)
        )
    
    def forward(self, x):
        # 复现原模型的特征提取逻辑(保留类Token)
        x = self.vit.patch_embed(x)
        # 添加类Token
        cls_token = self.vit.cls_token.expand(x.shape[0], -1, -1)
        x = torch.cat((cls_token, x), dim=1)
        # 位置嵌入与Dropout
        x = self.vit.pos_drop(x + self.vit.pos_embed)
        # Transformer编码器
        x = self.vit.blocks(x)
        # 最后的LayerNorm
        x = self.vit.norm(x)
        
        # 自定义逻辑:取所有Token均值
        x = x.mean(dim=1)
        # 过自定义分类头
        x = self.custom_head(x)
        return x

# 实例化模型
model = CustomViT("vit_base_patch16_224", num_classes=10)

# 测试输出
x = torch.randn(2, 3, 224, 224)
output = model(x)
print(output.shape)  # 应为 [2, 10]

内容的提问来源于stack exchange,提问作者Mohsin Ali

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 14:19:57