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

加载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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 13:15:24