Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1#!/usr/bin/env python32# Copyright (c) Facebook, Inc. and its affiliates.3 4"""5TensorMask Training Script.6 7This script is a simplified version of the training script in detectron2/tools.8"""9 10import os11 12import detectron2.utils.comm as comm13from detectron2.checkpoint import DetectionCheckpointer14from detectron2.config import get_cfg15from detectron2.engine import DefaultTrainer, default_argument_parser, default_setup, launch16from detectron2.evaluation import COCOEvaluator, verify_results17 18from tensormask import add_tensormask_config19 20 21class Trainer(DefaultTrainer):22 @classmethod23 def build_evaluator(cls, cfg, dataset_name, output_folder=None):24 if output_folder is None:25 output_folder = os.path.join(cfg.OUTPUT_DIR, "inference")26 return COCOEvaluator(dataset_name, output_dir=output_folder)27 28 29def setup(args):30 """31 Create configs and perform basic setups.32 """33 cfg = get_cfg()34 add_tensormask_config(cfg)35 cfg.merge_from_file(args.config_file)36 cfg.merge_from_list(args.opts)37 cfg.freeze()38 default_setup(cfg, args)39 return cfg40 41 42def main(args):43 cfg = setup(args)44 45 if args.eval_only:46 model = Trainer.build_model(cfg)47 DetectionCheckpointer(model, save_dir=cfg.OUTPUT_DIR).resume_or_load(48 cfg.MODEL.WEIGHTS, resume=args.resume49 )50 res = Trainer.test(cfg, model)51 if comm.is_main_process():52 verify_results(cfg, res)53 return res54 55 trainer = Trainer(cfg)56 trainer.resume_or_load(resume=args.resume)57 return trainer.train()58 59 60if __name__ == "__main__":61 args = default_argument_parser().parse_args()62 print("Command Line Args:", args)63 launch(64 main,65 args.num_gpus,66 num_machines=args.num_machines,67 machine_rank=args.machine_rank,68 dist_url=args.dist_url,69 args=(args,),70 )71 