File size: 2,294 Bytes
1da285f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
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()