import numpy as np import os, json, cv2, random import detectron2 from detectron2.utils.logger import setup_logger from detectron2.engine import DefaultTrainer, DefaultPredictor from detectron2.config import get_cfg from centernet.config import add_centernet_config from detectron2.checkpoint import DetectionCheckpointer, PeriodicCheckpointer from detectron2.data.datasets import register_coco_instances MODEL_CONFIG_PATH = './configs/CenterNet2_R50_1x.yaml' MODEL_WEIGHTS_PATH = './models/CenterNet2_R50_1x.pth' TRAIN_ANN_PATH = './datasets/coco/annotations/instances_train2017.json' TRAIN_IMG_DIR = './datasets/coco/train2017/' VAL_ANN_PATH = './datasets/coco/annotations/instances_val2017.json' VAL_IMG_DIR = './datasets/coco/val2017/' LR = 0.00025 MAX_ITER = 300 BATCH_SIZE = 2 # NUM_CLASSES = 39 NUM_CLASSES = 80 DATALOADER_NUM_WORKERS = 2 def do_validate(cfg): DetectionCheckpointer(model, save_dir=cfg.OUTPUT_DIR).resume_or_load( cfg.MODEL.WEIGHTS, resume=False ) def do_predict(): pass def do_train(cfg): setup_logger() os.makedirs(cfg.OUTPUT_DIR, exist_ok=True) trainer = DefaultTrainer(cfg) trainer.resume_or_load(resume=False) trainer.train() def main(): register_coco_instances("train", {}, TRAIN_ANN_PATH, TRAIN_IMG_DIR) register_coco_instances("val", {}, VAL_ANN_PATH, VAL_IMG_DIR) cfg = get_cfg() add_centernet_config(cfg) cfg.merge_from_file(MODEL_CONFIG_PATH) cfg.MODEL.WEIGHTS = MODEL_WEIGHTS_PATH cfg.DATASETS.TRAIN = "train" cfg.DATASETS.TEST = "val" cfg.DATALOADER.NUM_WORKERS = DATALOADER_NUM_WORKERS cfg.SOLVER.IMS_PER_BATCH = BATCH_SIZE # This is the real "batch size" commonly known to deep learning people cfg.SOLVER.BASE_LR = LR # pick a good LR cfg.SOLVER.MAX_ITER = MAX_ITER cfg.SOLVER.STEPS = [] # do not decay learning rate cfg.MODEL.ROI_HEADS.BATCH_SIZE_PER_IMAGE = 128 # The "RoIHead batch size". 128 is faster, and good enough for this toy dataset (default: 512) cfg.MODEL.ROI_HEADS.NUM_CLASSES = NUM_CLASSES # only has one class (ballon). (see https://detectron2.readthedocs.io/tutorials/datasets.html#update-the-config-for-new-datasets) # do_train(cfg) do_validate(cfg) if __name__ == '__main__': main()