Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
train_net.py172 linesDownload Raw Back to Panoptic-DeepLab
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