加载MMDetection的Evaluator类时遇state_dict不匹配问题求助
MMDetection加载模型时state_dict不匹配的问题排查与修复
尝试运行他人开发的深度学习工具,加载封装MMDetection功能的Evaluator类时,反复出现警告:
mmdet - WARNING - The model and loaded state dict do not match exactly. unexpected key in source state_dict:
模型配置(自动下载预训练权重)
model = dict( type='FCOS', pretrained='open-mmlab://detectron/resnet101_caffe', backbone=dict( type='ResNet', depth=101, num_stages=4, out_indices=(0, 1, 2, 3), frozen_stages=1, norm_cfg=dict(type='BN', requires_grad=False), norm_eval=True, style='caffe'), neck=dict( type='FPN', in_channels=[256, 512, 1024, 2048], out_channels=256, start_level=1, add_extra_convs=True, extra_convs_on_inputs=False, num_outs=5, relu_before_extra_convs=True), bbox_head=dict( type='FCOSHead', num_classes=15, in_channels=256, stacked_convs=4, feat_channels=256, strides=[8, 16, 32, 64, 128], loss_cls=dict( type='FocalLoss', use_sigmoid=True, gamma=2.0, alpha=0.25, loss_weight=1.0), loss_bbox=dict(type='IoULoss', loss_weight=1.0), loss_centerness=dict( type='CrossEntropyLoss', use_sigmoid=True, loss_weight=1.0)))
Evaluator类核心代码(utils.py中)
from skimage.draw import rectangle_perimeter import skimage.io as io from skimage.transform import resize import numpy as np import skimage import pickle import torch from mmcv import Config, DictAction from mmdet.models import build_detector from mmcv.runner import load_checkpoint import mmcv from mmdet.datasets.pipelines import Compose # TO LOOK AT from mmcv.parallel import collate, scatter from mmdet.core import bbox2result from skimage import data, io, filters from matplotlib.pyplot import figure import os class_to_number = {"Yeast White": 0, "Budding White": 1, "Yeast Opaque": 2, "Budding Opaque":3,"Yeast Gray": 4, "Budding Gray": 5, "Shmoo":6,"Artifact": 7, "Unknown ": 8, "Pseudohyphae": 9, "Hyphae": 10, "H-junction": 11, "P-junction":12,"P-Start":13,"H-Start":14} number_to_class = {y:x for x,y in class_to_number.items()} class Evaluator(): def __init__(self,config,checkpoint_file): self.cfg = Config.fromfile(config) self.cfg["gpu-ids"] = 6 self.model = build_detector( self.cfg.model, train_cfg=self.cfg.train_cfg, test_cfg=self.cfg.test_cfg) checkpoint_dict = load_checkpoint(self.model,checkpoint_file) state_dict = checkpoint_dict["state_dict"] self.model.CLASSES = checkpoint_dict['meta']['CLASSES'] self.model.load_state_dict(state_dict) self.model.eval()
已尝试的排查步骤
- 核对并匹配了原作者使用的mmdet、mmcv、PyTorch版本
- 重新下载模型文件,确认文件未损坏
问题分析与修复方案
1. 定位具体不匹配的参数
警告未给出具体不匹配的key名称,先添加代码打印差异:
在model.load_state_dict(state_dict)前插入:
# 获取模型和checkpoint的参数集合 model_keys = set(self.model.state_dict().keys()) checkpoint_keys = set(state_dict.keys()) # 打印差异 print("Unexpected keys in checkpoint:", checkpoint_keys - model_keys) print("Missing keys in model:", model_keys - checkpoint_keys)
2. 修复核心重复加载问题
原代码中load_checkpoint已经完成了权重加载到模型的操作,后续又调用model.load_state_dict(state_dict)属于重复加载,这是导致警告的主要原因。修改Evaluator的初始化方法:
class Evaluator(): def __init__(self,config,checkpoint_file): self.cfg = Config.fromfile(config) self.cfg["gpu-ids"] = 6 self.model = build_detector( self.cfg.model, train_cfg=self.cfg.train_cfg, test_cfg=self.cfg.test_cfg) # load_checkpoint会自动完成权重加载,无需手动调用load_state_dict checkpoint_dict = load_checkpoint(self.model, checkpoint_file) self.model.CLASSES = checkpoint_dict['meta']['CLASSES'] self.model.eval()
3. 处理其他常见不匹配场景
场景1:checkpoint带module.前缀(多GPU训练保存)
如果打印的差异显示checkpoint的key都带module.前缀,而模型的key没有,可移除前缀:
checkpoint_dict = load_checkpoint(self.model, checkpoint_file, strict=False) state_dict = {k.replace('module.', ''): v for k, v in checkpoint_dict["state_dict"].items()} self.model.load_state_dict(state_dict, strict=False)
场景2:模型结构细微差异(如分类数、层结构)
如果确认模型结构和checkpoint有少量无关差异,可允许非严格加载:
checkpoint_dict = load_checkpoint(self.model, checkpoint_file, strict=False)
注意:此方法仅适合非核心参数不匹配,若核心结构差异会影响模型性能,需确保配置文件和训练时完全一致。
内容的提问来源于stack exchange,提问作者jjacob
相关产品推荐
相关产品推荐

