运行EfficientDet加载模型时出现state_dict参数shape不匹配错误如何解决?
错误原因定位
该报错是分类头权重形状不匹配导致的,核心问题如下:
- 配置修改顺序错误:你在实例化
EfficientDet模型之后才修改config的num_classes、image_size参数,模型初始化时用的仍是默认90类的配置,默认配置下分类头的输出通道为90类 × 9个锚框 = 810,和 checkpoint 中1类训练得到的9通道分类头权重形状完全不匹配。 - 部分版本的EfficientDet实现中,
HeadNet的num_outputs参数需要手动传入类别数×锚框数的总通道数,若直接传入num_classes=1也可能导致输出通道数不符合预期。
可行解决方法
方法1:调整配置修改顺序(优先推荐)
调整load_net函数的参数修改顺序,先改配置再实例化模型,从根源对齐形状:
def load_net(checkpoint_path): config = get_efficientdet_config('tf_efficientdet_d5') # 先修改配置参数,再实例化模型 config.num_classes = 1 config.image_size = 512 net = EfficientDet(config, pretrained_backbone=False) net.class_net = HeadNet(config, num_outputs=config.num_classes, norm_kwargs=dict(eps=.001, momentum=.01)) checkpoint = torch.load(checkpoint_path) net.load_state_dict(checkpoint['model_state_dict']) del checkpoint gc.collect() net = DetBenchEval(net, config) net.eval() return net.cuda()
方法2:加载权重时关闭严格匹配
如果调整顺序后仍有少量层不匹配(比如自定义层和checkpoint层命名差异),可以在load_state_dict时添加strict=False参数,自动跳过形状不匹配的层,仅加载匹配的骨干网络、特征融合网络权重:
net.load_state_dict(checkpoint['model_state_dict'], strict=False)
方法3:手动对齐分类头输出通道
如果你的HeadNet实现不会自动计算锚框倍数,手动传入总输出通道数即可:
# EfficientDet默认单锚点锚框数为9,总输出通道=类别数 × 锚框数 net.class_net = HeadNet(config, num_outputs=config.num_classes * 9, norm_kwargs=dict(eps=.001, momentum=.01))
内容的提问来源于stack exchange,提问作者Ni7san
相关产品推荐
相关产品推荐

