使用Detectron2训练自定义数据集时遇Model Zoo配置文件找不到错误
解决Detectron2 Model Zoo配置文件找不到的问题
错误原因及修复步骤:
- 函数定义与参数不匹配
utils.py中的get_train_cfg被错误定义为带self的类方法(实际是普通函数),且参数列表缺少test_dataset_name,导致调用时参数数量不对。 - Model Zoo路径前缀冗余
使用model_zoo.get_config_file()和model_zoo.get_checkpoint_url()时,无需添加configs/前缀,直接传入模型相对路径即可,Model Zoo内部会自动处理目录结构。 - 未定义变量与数据集配置错误
train.py中cfg_save_path未赋值,需先指定保存路径;同时测试数据集应设置为LP_test而非训练集。
修复后的代码
train.py
import numpy as np from detectron2.utils.logger import setup_logger setup_logger() from detectron2.data.datasets import register_coco_instances from detectron2.engine import DefaultTrainer import os import pickle from utils import * # 移除configs/前缀 config_file_path = "COCO-Detection/faster_rcnn_R_50_FPN_3x.yaml" checkpoint_url = "COCO-Detection/faster_rcnn_R_50_FPN_3x.yaml" output_dir = "./output/object_detection" num_classes = 1 # 定义配置文件保存路径 cfg_save_path = "./config.pkl" device = "cuda" train_dataset_name = "LP_train" train_image_path = "train" train_json_annot_path = "train.json" test_dataset_name = "LP_test" test_image_path = "test" test_json_annot_path = "test.json" ############################################### register_coco_instances(name=train_dataset_name, metadata={}, json_file=train_json_annot_path, image_root=train_image_path) register_coco_instances(name=test_dataset_name, metadata={}, json_file=test_json_annot_path, image_root=test_image_path) #plot_samples(dataset_name= train_dataset_name, n = 2) ############################################### def main(): # 匹配get_train_cfg的参数列表 cfg = get_train_cfg(config_file_path, checkpoint_url, train_dataset_name, test_dataset_name, num_classes, device, output_dir) # 保存配置 with open(cfg_save_path, 'wb') as f: pickle.dump(cfg, f, protocol=pickle.HIGHEST_PROTOCOL) os.makedirs(cfg.OUTPUT_DIR, exist_ok=True) trainer = DefaultTrainer(cfg) trainer.resume_or_load(resume=False) trainer.train() if __name__ == '__main__': main()
utils.py
from detectron2.data import DatasetCatalog, MetadataCatalog from detectron2.utils.visualizer import Visualizer from detectron2.config import get_cfg from detectron2 import model_zoo from detectron2.utils.visualizer import ColorMode import random import cv2 import matplotlib.pyplot as plt def plot_samples(dataset_name, n=1): dataset_custom = DatasetCatalog.get(dataset_name) dataset_custom_metadata = MetadataCatalog.get(dataset_name) for s in random.sample(dataset_custom, n): img = cv2.imread(s["file_name"]) v = Visualizer(img[:,:,::-1], metadata=dataset_custom_metadata, scale=0.5) v = v.draw_dataset_dict(s) plt.figure(figsize=(15,20)) plt.imshow(v.get_image()) plt.show() #plt.savefig("matplotlib.png") #save config , don't show # 移除多余的self参数,添加test_dataset_name参数 def get_train_cfg(config_file_path, checkpoint_url, train_dataset_name, test_dataset_name, num_classes, device, output_dir): cfg = get_cfg() cfg.merge_from_file(model_zoo.get_config_file(config_file_path)) cfg.MODEL.WEIGHTS = model_zoo.get_checkpoint_url(checkpoint_url) cfg.DATASETS.TRAIN = (train_dataset_name,) # 修正测试数据集配置 cfg.DATASETS.TEST = (test_dataset_name,) cfg.DATALOADER.NUM_WORKERS = 5 cfg.SOLVER.IMS_PER_BATCH = 5 cfg.SOLVER.BASE_LR = 0.00025 cfg.SOLVER.MAX_ITER = 1000 cfg.SOLVER.STEPS = [] cfg.MODEL.ROI_HEADS.NUM_CLASSES = num_classes cfg.MODEL.DEVICE = device cfg.OUTPUT_DIR = output_dir return cfg
内容的提问来源于stack exchange,提问作者Ridoy Kanto Joy
相关产品推荐
相关产品推荐

