Orienter / baselines /CenterNet2 /CenterNet2.py
stereoid's picture
Add files using upload-large-folder tool
1da285f verified
Raw
History Blame Contribute Delete
2.29 kB
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()