ConvNeXt Small添加自定义块后加载预训练权重报错求助
解决ConvNeXt添加自定义层后加载预训练权重的缺失键错误
错误原因
你给ConvNeXt新增了custom_block模块,但官方提供的预训练权重里完全没有这个模块的参数,直接调用model.load_state_dict(checkpoint["model"])时,PyTorch会严格校验所有参数键是否匹配,找不到自定义层的键就会抛出缺失键错误。
另外你代码最后重复赋值model.custom_block = CustomBlock()属于冗余操作,会覆盖模型初始化时创建的自定义块,完全没必要,直接删掉即可。
解决方案
方法1:加载时忽略缺失的键
最简单的方式是给load_state_dict加上strict=False参数,让PyTorch只加载权重中存在的键,忽略自定义层的缺失键,自定义层的参数会用你初始化时的配置(比如trunc_normal_和常数偏置),之后训练过程中再更新这些参数。
修改convnext_small函数中的加载代码:
@register_model def convnext_small(pretrained=False, in_22k=False, **kwargs): model = ConvNeXt(depths=[3, 3, 27, 3], dims=[96, 192, 384, 768], **kwargs) if pretrained: url = model_urls['convnext_small_22k'] if in_22k else model_urls['convnext_small_1k'] checkpoint = torch.hub.load_state_dict_from_url(url=url, map_location="cpu") # 添加strict=False忽略缺失键 model.load_state_dict(checkpoint["model"], strict=False) return model # 初始化模型,去掉冗余的custom_block重新赋值 model = convnext_small(pretrained=True, in_22k=False) print(model)
方法2:手动过滤预训练权重(更严谨)
如果想更精准地加载原模型的参数,避免意外加载不匹配的键,可以手动过滤预训练权重,只保留当前模型中存在的键:
@register_model def convnext_small(pretrained=False, in_22k=False, **kwargs): model = ConvNeXt(depths=[3, 3, 27, 3], dims=[96, 192, 384, 768], **kwargs) if pretrained: url = model_urls['convnext_small_22k'] if in_22k else model_urls['convnext_small_1k'] checkpoint = torch.hub.load_state_dict_from_url(url=url, map_location="cpu") # 过滤出当前模型存在的参数键 pretrained_dict = checkpoint["model"] model_dict = model.state_dict() filtered_dict = {k: v for k, v in pretrained_dict.items() if k in model_dict} # 更新模型参数 model_dict.update(filtered_dict) model.load_state_dict(model_dict) return model model = convnext_small(pretrained=True, in_22k=False) print(model)
额外注意:自定义块的维度不匹配问题
你的CustomBlock还存在一个隐藏错误:原模型forward_features中,x = self.norm(x.mean([-2, -1]))输出的是2维张量(batch_size, 768),但CustomBlock里的conv1是nn.Conv2d层,它需要4维输入(batch_size, channels, height, width),运行时会直接报错。
你需要修改CustomBlock来适配输入维度,两种选择:
- 把卷积层换成全连接层:
class CustomBlock(nn.Module): def __init__(self): super(CustomBlock, self).__init__() self.linear1 = nn.Linear(768, 512) # 替换Conv2d为Linear self.linear2 = nn.Linear(512, 512) self.identity = nn.Identity() self.multihead_attention = nn.MultiheadAttention(embed_dim=512, num_heads=4) self.linear = nn.Linear(512, 512) self.dropout = nn.Dropout(0.2) def forward(self, x): # MultiheadAttention需要输入形状为(seq_len, batch_size, embed_dim),所以先转置 x = self.linear1(x) x = self.linear2(x) x = self.identity(x) x = x.unsqueeze(0) # 添加seq_len维度,变成(1, batch_size, 512) x = self.multihead_attention(x, x, x)[0] x = x.squeeze(0) # 去掉seq_len维度 x = self.linear(x) x = self.dropout(x) return x
- 把2维张量转成4维,适配卷积层:
class CustomBlock(nn.Module): def __init__(self): super(CustomBlock, self).__init__() self.conv1 = nn.Conv2d(768, 512, kernel_size=1, padding=0) self.conv2 = nn.Conv2d(512, 512, kernel_size=1) self.identity = nn.Identity() self.multihead_attention = nn.MultiheadAttention(embed_dim=512, num_heads=4) self.linear = nn.Linear(512, 512) self.dropout = nn.Dropout(0.2) def forward(self, x): # 把2维张量转成4维:(batch_size, 768) -> (batch_size, 768, 1, 1) x = x.unsqueeze(-1).unsqueeze(-1) x = self.conv1(x) x = self.conv2(x) x = self.identity(x) # 转成MultiheadAttention需要的形状:(seq_len, batch_size, embed_dim) x = x.flatten(2).transpose(0, 1) # (batch_size, 512, 1,1) -> (1, batch_size,512) x = self.multihead_attention(x, x, x)[0] x = x.transpose(0, 1).squeeze(-1) # 转回(batch_size,512) x = self.linear(x) x = self.dropout(x) return x
内容的提问来源于stack exchange,提问作者user2158736
相关产品推荐
相关产品推荐

