Download baselines/CenterNet2/CenterNet2.py from stereoid/Orienter: direct link, hf CLI and curl.
- Browser
- Download file 2.29 kB
-
https://huggingface.co/stereoid/Orienter/resolve/main/baselines/CenterNet2/CenterNet2.py
- Command line
-
hf download hf://stereoid/Orienter/baselines/CenterNet2/CenterNet2.py
-
curl -L -o CenterNet2.py https://huggingface.co/stereoid/Orienter/resolve/main/baselines/CenterNet2/CenterNet2.py
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() | |