Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1#!/usr/bin/env python32# Copyright (c) Facebook, Inc. and its affiliates.3 4"""5Panoptic-DeepLab Training Script.6This script is a simplified version of the training script in detectron2/tools.7"""8 9import os10import torch11 12import detectron2.data.transforms as T13from detectron2.checkpoint import DetectionCheckpointer14from detectron2.config import get_cfg15from detectron2.data import MetadataCatalog, build_detection_train_loader16from detectron2.engine import DefaultTrainer, default_argument_parser, default_setup, launch17from detectron2.evaluation import (18 CityscapesInstanceEvaluator,19 CityscapesSemSegEvaluator,20 COCOEvaluator,21 COCOPanopticEvaluator,22 DatasetEvaluators,23)24from detectron2.projects.deeplab import build_lr_scheduler25from detectron2.projects.panoptic_deeplab import (26 PanopticDeeplabDatasetMapper,27 add_panoptic_deeplab_config,28)29from detectron2.solver import get_default_optimizer_params30from detectron2.solver.build import maybe_add_gradient_clipping31 32 33def build_sem_seg_train_aug(cfg):34 augs = [35 T.ResizeShortestEdge(36 cfg.INPUT.MIN_SIZE_TRAIN, cfg.INPUT.MAX_SIZE_TRAIN, cfg.INPUT.MIN_SIZE_TRAIN_SAMPLING37 )38 ]39 if cfg.INPUT.CROP.ENABLED:40 augs.append(T.RandomCrop(cfg.INPUT.CROP.TYPE, cfg.INPUT.CROP.SIZE))41 augs.append(T.RandomFlip())42 return augs43 44 45class Trainer(DefaultTrainer):46 """47 We use the "DefaultTrainer" which contains a number pre-defined logic for48 standard training workflow. They may not work for you, especially if you49 are working on a new research project. In that case you can use the cleaner50 "SimpleTrainer", or write your own training loop.51 """52 53 @classmethod54 def build_evaluator(cls, cfg, dataset_name, output_folder=None):55 """56 Create evaluator(s) for a given dataset.57 This uses the special metadata "evaluator_type" associated with each builtin dataset.58 For your own dataset, you can simply create an evaluator manually in your59 script and do not have to worry about the hacky if-else logic here.60 """61 if cfg.MODEL.PANOPTIC_DEEPLAB.BENCHMARK_NETWORK_SPEED:62 return None63 if output_folder is None:64 output_folder = os.path.join(cfg.OUTPUT_DIR, "inference")65 evaluator_list = []66 evaluator_type = MetadataCatalog.get(dataset_name).evaluator_type67 if evaluator_type in ["cityscapes_panoptic_seg", "coco_panoptic_seg"]:68 evaluator_list.append(COCOPanopticEvaluator(dataset_name, output_folder))69 if evaluator_type == "cityscapes_panoptic_seg":70 evaluator_list.append(CityscapesSemSegEvaluator(dataset_name))71 evaluator_list.append(CityscapesInstanceEvaluator(dataset_name))72 if evaluator_type == "coco_panoptic_seg":73 # `thing_classes` in COCO panoptic metadata includes both thing and74 # stuff classes for visualization. COCOEvaluator requires metadata75 # which only contains thing classes, thus we map the name of76 # panoptic datasets to their corresponding instance datasets.77 dataset_name_mapper = {78 "coco_2017_val_panoptic": "coco_2017_val",79 "coco_2017_val_100_panoptic": "coco_2017_val_100",80 }81 evaluator_list.append(82 COCOEvaluator(dataset_name_mapper[dataset_name], output_dir=output_folder)83 )84 if len(evaluator_list) == 0:85 raise NotImplementedError(86 "no Evaluator for the dataset {} with the type {}".format(87 dataset_name, evaluator_type88 )89 )90 elif len(evaluator_list) == 1:91 return evaluator_list[0]92 return DatasetEvaluators(evaluator_list)93 94 @classmethod95 def build_train_loader(cls, cfg):96 mapper = PanopticDeeplabDatasetMapper(cfg, augmentations=build_sem_seg_train_aug(cfg))97 return build_detection_train_loader(cfg, mapper=mapper)98 99 @classmethod100 def build_lr_scheduler(cls, cfg, optimizer):101 """102 It now calls :func:`detectron2.solver.build_lr_scheduler`.103 Overwrite it if you'd like a different scheduler.104 """105 return build_lr_scheduler(cfg, optimizer)106 107 @classmethod108 def build_optimizer(cls, cfg, model):109 """110 Build an optimizer from config.111 """112 params = get_default_optimizer_params(113 model,114 weight_decay=cfg.SOLVER.WEIGHT_DECAY,115 weight_decay_norm=cfg.SOLVER.WEIGHT_DECAY_NORM,116 )117 118 optimizer_type = cfg.SOLVER.OPTIMIZER119 if optimizer_type == "SGD":120 return maybe_add_gradient_clipping(cfg, torch.optim.SGD)(121 params,122 cfg.SOLVER.BASE_LR,123 momentum=cfg.SOLVER.MOMENTUM,124 nesterov=cfg.SOLVER.NESTEROV,125 )126 elif optimizer_type == "ADAM":127 return maybe_add_gradient_clipping(cfg, torch.optim.Adam)(params, cfg.SOLVER.BASE_LR)128 else:129 raise NotImplementedError(f"no optimizer type {optimizer_type}")130 131 132def setup(args):133 """134 Create configs and perform basic setups.135 """136 cfg = get_cfg()137 add_panoptic_deeplab_config(cfg)138 cfg.merge_from_file(args.config_file)139 cfg.merge_from_list(args.opts)140 cfg.freeze()141 default_setup(cfg, args)142 return cfg143 144 145def main(args):146 cfg = setup(args)147 148 if args.eval_only:149 model = Trainer.build_model(cfg)150 DetectionCheckpointer(model, save_dir=cfg.OUTPUT_DIR).resume_or_load(151 cfg.MODEL.WEIGHTS, resume=args.resume152 )153 res = Trainer.test(cfg, model)154 return res155 156 trainer = Trainer(cfg)157 trainer.resume_or_load(resume=args.resume)158 return trainer.train()159 160 161if __name__ == "__main__":162 args = default_argument_parser().parse_args()163 print("Command Line Args:", args)164 launch(165 main,166 args.num_gpus,167 num_machines=args.num_machines,168 machine_rank=args.machine_rank,169 dist_url=args.dist_url,170 args=(args,),171 )172 