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

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来适配输入维度,两种选择:

  1. 把卷积层换成全连接层:
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
  1. 把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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 21:20:56