自定义smp DeepLabV3+编码器加载预训练权重报错求助
解决SMP自定义DeepLabV3+编码器预训练权重加载错误
问题场景
尝试为Segmentation Models Pytorch(SMP)的DeepLabV3+模型自定义编码器,选用torchvision.models.segmentation.deeplabv3_resnet50作为骨干网络测试,但执行model = smp.DeepLabV3Plus(encoder_name="deeplab_resnet50")实例化模型时,出现预训练权重加载错误,提示state_dict存在缺失键与意外键。
原测试代码
import torch import torch.nn as nn import segmentation_models_pytorch as smp import torchvision.models.segmentation as segmentation from torchvision.models.resnet import Bottleneck class resnet50_encoder(nn.Module, smp.encoders._base.EncoderMixin): def __init__(self, **kwargs): super().__init__() # Define your encoder module self.encoder = segmentation.deeplabv3_resnet50(pretrained=True) # Set the number of output channels for each feature tensor self._out_channels = [3, 64, 64, 128, 256, 512] # Set the depth (number of downsampling operations) self._depth = 5 # Set the default number of input channels (usually 3 for RGB images) self._in_channels = 3 def forward(self, x: torch.Tensor): # Get features from the encoder features = self.encoder(x)['out'] # Return features sorted in descending order of spatial resolution return [features[f] for f in ['0', '1', '2', '3', '4', '5']] smp.encoders.encoders["deeplab_resnet50"] = { "encoder": resnet50_encoder, "pretrained_settings": { "imagenet": { "mean": [0.485, 0.456, 0.406], "std": [0.229, 0.224, 0.225], "url": "https://download.pytorch.org/models/deeplabv3_resnet50_coco-cd0a2569.pth", "input_space": "RGB", "input_range": [0, 1], }, }, "params": { "out_channels": (3, 64, 256, 512, 1024, 2048), "block": Bottleneck, "layers": [3, 4, 6, 3], }, }
报错信息
Traceback (most recent call last): File "<string>", line 1, in <module> File "/da/aics/projects/ComPath/envs/sweenke4/compath_art_det_repl/lib/python3.9/site-packages/segmentation_models_pytorch/decoders/deeplabv3/model.py", line 146, in __init__ self.encoder = get_encoder( File "/da/aics/projects/ComPath/envs/sweenke4/compath_art_det_repl/lib/python3.9/site-packages/segmentation_models_pytorch/encoders/__init__.py", line 85, in get_encoder encoder.load_state_dict(model_zoo.load_url(settings["url"])) File "<string>", line 26, in load_state_dict File "/da/aics/projects/ComPath/envs/sweenke4/compath_art_det_repl/lib/python3.9/site-packages/torch/nn/modules/module.py", line 2189, in load_state_dict raise RuntimeError('Error(s) in loading state_dict for {}:\n\t{}'.format( RuntimeError: Error(s) in loading state_dict for resnet50_encoder: Missing key(s) in state_dict: "encoder.backbone.conv1.weight", "encoder.backbone.bn1.weight", "encoder.backbone.bn1.bias", "encoder.backbone.bn1.running_mean", "encoder.backbone.bn1.running_var", "encoder.backbone.layer1.0.conv1.weight", "encoder.backbone.layer1.0.bn1.weight", "encoder.backbone.layer1.0.bn1.bias", ... Unexpected key(s) in state_dict: "backbone.conv1.weight", "backbone.bn1.weight", "backbone.bn1.bias", "backbone.bn1.running_mean", "backbone.bn1.running_var", "backbone.bn1.num_batches_tracked", "backbone.layer1.0.conv1.weight", "backbone.layer1.0.bn1.weight", "backbone.layer1.0.bn1.bias", ...
错误原因分析
- 权重键路径不匹配:预训练权重的参数键为
backbone.xxx,但自定义编码器将整个deeplabv3模型存在self.encoder下,对应参数键应为encoder.backbone.xxx,直接加载会导致键完全不匹配。 - Forward逻辑错误:
deeplabv3_resnet50(x)返回的是包含out(最终分割输出)和aux(辅助输出)的字典,并非带0-5键的多层特征字典,原forward方法无法返回SMP编码器要求的各层级特征。 - 通道配置不一致:自定义编码器的
_out_channels与注册参数中的out_channels数值不符,会导致后续解码器适配失败。
修正方案
直接基于resnet50的backbone构建编码器,手动处理预训练权重的键映射,同时正确返回各层级特征:
修正后的代码
import torch import torch.nn as nn import segmentation_models_pytorch as smp from torchvision.models import resnet50 from torchvision.models.resnet import Bottleneck class ResNet50Encoder(nn.Module, smp.encoders._base.EncoderMixin): def __init__(self, pretrained=True, **kwargs): super().__init__() # 初始化resnet50 backbone,移除分类头 self.backbone = resnet50(pretrained=False) self.backbone.fc = nn.Identity() # SMP编码器要求的输出通道数,对应resnet50各阶段特征 self._out_channels = [3, 64, 256, 512, 1024, 2048] self._depth = 5 self._in_channels = 3 if pretrained: # 加载deeplabv3_resnet50的预训练权重 state_dict = torch.hub.load_state_dict_from_url( "https://download.pytorch.org/models/deeplabv3_resnet50_coco-cd0a2569.pth", progress=True ) # 提取backbone部分的权重,去掉"backbone."前缀以匹配当前模型键 backbone_state_dict = {k.replace("backbone.", ""): v for k, v in state_dict.items() if k.startswith("backbone.")} # 加载权重,忽略无关键(如原模型的分类头、分割头) self.backbone.load_state_dict(backbone_state_dict, strict=False) def forward(self, x: torch.Tensor): # 按顺序提取resnet各阶段的特征,返回从输入到最深层的特征列表 features = [] features.append(x) # 输入层,对应depth=0 x = self.backbone.conv1(x) x = self.backbone.bn1(x) x = self.backbone.relu(x) features.append(x) # 第一层卷积输出,depth=1 x = self.backbone.maxpool(x) x = self.backbone.layer1(x) features.append(x) # stage1输出,depth=2 x = self.backbone.layer2(x) features.append(x) # stage2输出,depth=3 x = self.backbone.layer3(x) features.append(x) # stage3输出,depth=4 x = self.backbone.layer4(x) features.append(x) # stage4输出,depth=5 return features # 注册自定义编码器到SMP编码器库 smp.encoders.encoders["deeplab_resnet50"] = { "encoder": ResNet50Encoder, "pretrained_settings": { "imagenet": { "mean": [0.485, 0.456, 0.406], "std": [0.229, 0.224, 0.225], "input_space": "RGB", "input_range": [0, 1], }, }, "params": { "out_channels": (3, 64, 256, 512, 1024, 2048), "block": Bottleneck, "layers": [3, 4, 6, 3], }, } # 测试实例化模型 if __name__ == "__main__": model = smp.DeepLabV3Plus(encoder_name="deeplab_resnet50", in_channels=3, classes=10) print("模型实例化成功")
核心修正点
- 骨干网络拆分:直接使用resnet50的backbone作为编码器,而非整个deeplabv3模型,符合SMP编码器需要输出多层特征的要求。
- 权重键映射:手动处理预训练权重的键,去掉
backbone.前缀,匹配自定义编码器的参数结构。 - Forward逻辑修正:按顺序提取resnet各阶段的特征,返回SMP要求的特征列表(从输入到最深层的顺序)。
- 通道配置统一:确保
_out_channels与注册参数中的out_channels完全一致,避免解码器适配错误。
内容的提问来源于stack exchange,提问作者Kevin Sweeney
相关产品推荐
相关产品推荐

