Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1#!/usr/bin/env python2# Copyright (c) Facebook, Inc. and its affiliates.3"""4Point supervision Training Script.5 6This script is a simplified version of the training script in detectron2/tools.7"""8 9import os10 11import detectron2.utils.comm as comm12from detectron2.checkpoint import DetectionCheckpointer13from detectron2.config import get_cfg14from detectron2.data import MetadataCatalog, build_detection_train_loader15from detectron2.engine import DefaultTrainer, default_argument_parser, default_setup, launch16from detectron2.evaluation import COCOEvaluator, DatasetEvaluators, verify_results17from detectron2.projects.point_rend import add_pointrend_config18from detectron2.utils.logger import setup_logger19 20from point_sup import PointSupDatasetMapper, add_point_sup_config21 22 23class Trainer(DefaultTrainer):24 """25 We use the "DefaultTrainer" which contains pre-defined default logic for26 standard training workflow. They may not work for you, especially if you27 are working on a new research project. In that case you can write your28 own training loop. You can use "tools/plain_train_net.py" as an example.29 """30 31 @classmethod32 def build_evaluator(cls, cfg, dataset_name, output_folder=None):33 """34 Create evaluator(s) for a given dataset.35 This uses the special metadata "evaluator_type" associated with each builtin dataset.36 For your own dataset, you can simply create an evaluator manually in your37 script and do not have to worry about the hacky if-else logic here.38 """39 if output_folder is None:40 output_folder = os.path.join(cfg.OUTPUT_DIR, "inference")41 evaluator_list = []42 evaluator_type = MetadataCatalog.get(dataset_name).evaluator_type43 if evaluator_type == "coco":44 evaluator_list.append(COCOEvaluator(dataset_name, output_dir=output_folder))45 if len(evaluator_list) == 0:46 raise NotImplementedError(47 "no Evaluator for the dataset {} with the type {}".format(48 dataset_name, evaluator_type49 )50 )51 elif len(evaluator_list) == 1:52 return evaluator_list[0]53 return DatasetEvaluators(evaluator_list)54 55 @classmethod56 def build_train_loader(cls, cfg):57 if cfg.INPUT.POINT_SUP:58 mapper = PointSupDatasetMapper(cfg, is_train=True)59 else:60 mapper = None61 return build_detection_train_loader(cfg, mapper=mapper)62 63 64def setup(args):65 """66 Create configs and perform basic setups.67 """68 cfg = get_cfg()69 add_pointrend_config(cfg)70 add_point_sup_config(cfg)71 cfg.merge_from_file(args.config_file)72 cfg.merge_from_list(args.opts)73 cfg.freeze()74 default_setup(cfg, args)75 # Setup logger for "point_sup" module76 setup_logger(output=cfg.OUTPUT_DIR, distributed_rank=comm.get_rank(), name="point_sup")77 return cfg78 79 80def main(args):81 cfg = setup(args)82 83 if args.eval_only:84 model = Trainer.build_model(cfg)85 DetectionCheckpointer(model, save_dir=cfg.OUTPUT_DIR).resume_or_load(86 cfg.MODEL.WEIGHTS, resume=args.resume87 )88 res = Trainer.test(cfg, model)89 if cfg.TEST.AUG.ENABLED:90 res.update(Trainer.test_with_TTA(cfg, model))91 if comm.is_main_process():92 verify_results(cfg, res)93 return res94 95 """96 If you'd like to do anything fancier than the standard training logic,97 consider writing your own training loop (see plain_train_net.py) or98 subclassing the trainer.99 """100 trainer = Trainer(cfg)101 trainer.resume_or_load(resume=args.resume)102 return trainer.train()103 104 105if __name__ == "__main__":106 args = default_argument_parser().parse_args()107 print("Command Line Args:", args)108 launch(109 main,110 args.num_gpus,111 num_machines=args.num_machines,112 machine_rank=args.machine_rank,113 dist_url=args.dist_url,114 args=(args,),115 )116 