基于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
相关产品推荐
相关产品推荐

