Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1#!/usr/bin/env python32# Copyright (c) Facebook, Inc. and its affiliates.3 4"""5TridentNet Training Script.6 7This script is a simplified version of the training script in detectron2/tools.8"""9 10import os11 12from detectron2.checkpoint import DetectionCheckpointer13from detectron2.config import get_cfg14from detectron2.engine import DefaultTrainer, default_argument_parser, default_setup, launch15from detectron2.evaluation import COCOEvaluator16 17from tridentnet import add_tridentnet_config18 19 20class Trainer(DefaultTrainer):21 @classmethod22 def build_evaluator(cls, cfg, dataset_name, output_folder=None):23 if output_folder is None:24 output_folder = os.path.join(cfg.OUTPUT_DIR, "inference")25 return COCOEvaluator(dataset_name, output_dir=output_folder)26 27 28def setup(args):29 """30 Create configs and perform basic setups.31 """32 cfg = get_cfg()33 add_tridentnet_config(cfg)34 cfg.merge_from_file(args.config_file)35 cfg.merge_from_list(args.opts)36 cfg.freeze()37 default_setup(cfg, args)38 return cfg39 40 41def main(args):42 cfg = setup(args)43 44 if args.eval_only:45 model = Trainer.build_model(cfg)46 DetectionCheckpointer(model, save_dir=cfg.OUTPUT_DIR).resume_or_load(47 cfg.MODEL.WEIGHTS, resume=args.resume48 )49 res = Trainer.test(cfg, model)50 return res51 52 trainer = Trainer(cfg)53 trainer.resume_or_load(resume=args.resume)54 return trainer.train()55 56 57if __name__ == "__main__":58 args = default_argument_parser().parse_args()59 print("Command Line Args:", args)60 launch(61 main,62 args.num_gpus,63 num_machines=args.num_machines,64 machine_rank=args.machine_rank,65 dist_url=args.dist_url,66 args=(args,),67 )68 