Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model
0
1#!/usr/bin/env python2# Copyright (c) Facebook, Inc. and its affiliates.3"""4A main training script.5 6This scripts reads a given config file and runs the training or evaluation.7It is an entry point that is made to train standard models in detectron2.8 9In order to let one script support training of many models,10this script contains logic that are specific to these built-in models and therefore11may not be suitable for your own project.12For example, your research project perhaps only needs a single "evaluator".13 14Therefore, we recommend you to use detectron2 as an library and take15this file as an example of how to use the library.16You may want to write your own script with your datasets and other customizations.17"""18 19import logging20import os21from collections import OrderedDict22 23import detectron2.utils.comm as comm24from detectron2.checkpoint import DetectionCheckpointer25from detectron2.config import get_cfg26from detectron2.data import MetadataCatalog27from detectron2.engine import DefaultTrainer, default_argument_parser, default_setup, hooks, launch28from detectron2.evaluation import (29 CityscapesInstanceEvaluator,30 CityscapesSemSegEvaluator,31 COCOEvaluator,32 COCOPanopticEvaluator,33 DatasetEvaluators,34 LVISEvaluator,35 PascalVOCDetectionEvaluator,36 SemSegEvaluator,37 verify_results,38)39from detectron2.modeling import GeneralizedRCNNWithTTA40 41 42def build_evaluator(cfg, dataset_name, output_folder=None):43 """44 Create evaluator(s) for a given dataset.45 This uses the special metadata "evaluator_type" associated with each builtin dataset.46 For your own dataset, you can simply create an evaluator manually in your47 script and do not have to worry about the hacky if-else logic here.48 """49 if output_folder is None:50 output_folder = os.path.join(cfg.OUTPUT_DIR, "inference")51 evaluator_list = []52 evaluator_type = MetadataCatalog.get(dataset_name).evaluator_type53 if evaluator_type in ["sem_seg", "coco_panoptic_seg"]:54 evaluator_list.append(55 SemSegEvaluator(56 dataset_name,57 distributed=True,58 output_dir=output_folder,59 )60 )61 if evaluator_type in ["coco", "coco_panoptic_seg"]:62 evaluator_list.append(COCOEvaluator(dataset_name, output_dir=output_folder))63 if evaluator_type == "coco_panoptic_seg":64 evaluator_list.append(COCOPanopticEvaluator(dataset_name, output_folder))65 if evaluator_type == "cityscapes_instance":66 return CityscapesInstanceEvaluator(dataset_name)67 if evaluator_type == "cityscapes_sem_seg":68 return CityscapesSemSegEvaluator(dataset_name)69 elif evaluator_type == "pascal_voc":70 return PascalVOCDetectionEvaluator(dataset_name)71 elif evaluator_type == "lvis":72 return LVISEvaluator(dataset_name, output_dir=output_folder)73 if len(evaluator_list) == 0:74 raise NotImplementedError(75 "no Evaluator for the dataset {} with the type {}".format(dataset_name, evaluator_type)76 )77 elif len(evaluator_list) == 1:78 return evaluator_list[0]79 return DatasetEvaluators(evaluator_list)80 81 82class Trainer(DefaultTrainer):83 """84 We use the "DefaultTrainer" which contains pre-defined default logic for85 standard training workflow. They may not work for you, especially if you86 are working on a new research project. In that case you can write your87 own training loop. You can use "tools/plain_train_net.py" as an example.88 """89 90 @classmethod91 def build_evaluator(cls, cfg, dataset_name, output_folder=None):92 return build_evaluator(cfg, dataset_name, output_folder)93 94 @classmethod95 def test_with_TTA(cls, cfg, model):96 logger = logging.getLogger("detectron2.trainer")97 # In the end of training, run an evaluation with TTA98 # Only support some R-CNN models.99 logger.info("Running inference with test-time augmentation ...")100 model = GeneralizedRCNNWithTTA(cfg, model)101 evaluators = [102 cls.build_evaluator(103 cfg, name, output_folder=os.path.join(cfg.OUTPUT_DIR, "inference_TTA")104 )105 for name in cfg.DATASETS.TEST106 ]107 res = cls.test(cfg, model, evaluators)108 res = OrderedDict({k + "_TTA": v for k, v in res.items()})109 return res110 111 112def setup(args):113 """114 Create configs and perform basic setups.115 """116 cfg = get_cfg()117 cfg.merge_from_file(args.config_file)118 cfg.merge_from_list(args.opts)119 cfg.freeze()120 default_setup(cfg, args)121 return cfg122 123 124def main(args):125 cfg = setup(args)126 127 if args.eval_only:128 model = Trainer.build_model(cfg)129 DetectionCheckpointer(model, save_dir=cfg.OUTPUT_DIR).resume_or_load(130 cfg.MODEL.WEIGHTS, resume=args.resume131 )132 res = Trainer.test(cfg, model)133 if cfg.TEST.AUG.ENABLED:134 res.update(Trainer.test_with_TTA(cfg, model))135 if comm.is_main_process():136 verify_results(cfg, res)137 return res138 139 """140 If you'd like to do anything fancier than the standard training logic,141 consider writing your own training loop (see plain_train_net.py) or142 subclassing the trainer.143 """144 trainer = Trainer(cfg)145 trainer.resume_or_load(resume=args.resume)146 if cfg.TEST.AUG.ENABLED:147 trainer.register_hooks(148 [hooks.EvalHook(0, lambda: trainer.test_with_TTA(cfg, trainer.model))]149 )150 return trainer.train()151 152 153if __name__ == "__main__":154 args = default_argument_parser().parse_args()155 print("Command Line Args:", args)156 launch(157 main,158 args.num_gpus,159 num_machines=args.num_machines,160 machine_rank=args.machine_rank,161 dist_url=args.dist_url,162 args=(args,),163 )164 