Team Ai
Apppublic

Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
train_net.py164 linesDownload Raw Back to tools
1#!/usr/bin/env python2# Copyright (c) Facebook, Inc. and its affiliates.3"""4A main training script.5 6This scripts reads a given config file and runs the training or evaluation.7It is an entry point that is made to train standard models in detectron2.8 9In order to let one script support training of many models,10this script contains logic that are specific to these built-in models and therefore11may not be suitable for your own project.12For example, your research project perhaps only needs a single "evaluator".13 14Therefore, we recommend you to use detectron2 as an library and take15this file as an example of how to use the library.16You may want to write your own script with your datasets and other customizations.17"""18 19import logging20import os21from collections import OrderedDict22 23import detectron2.utils.comm as comm24from detectron2.checkpoint import DetectionCheckpointer25from detectron2.config import get_cfg26from detectron2.data import MetadataCatalog27from detectron2.engine import DefaultTrainer, default_argument_parser, default_setup, hooks, launch28from detectron2.evaluation import (29    CityscapesInstanceEvaluator,30    CityscapesSemSegEvaluator,31    COCOEvaluator,32    COCOPanopticEvaluator,33    DatasetEvaluators,34    LVISEvaluator,35    PascalVOCDetectionEvaluator,36    SemSegEvaluator,37    verify_results,38)39from detectron2.modeling import GeneralizedRCNNWithTTA40 41 42def build_evaluator(cfg, dataset_name, output_folder=None):43    """44    Create evaluator(s) for a given dataset.45    This uses the special metadata "evaluator_type" associated with each builtin dataset.46    For your own dataset, you can simply create an evaluator manually in your47    script and do not have to worry about the hacky if-else logic here.48    """49    if output_folder is None:50        output_folder = os.path.join(cfg.OUTPUT_DIR, "inference")51    evaluator_list = []52    evaluator_type = MetadataCatalog.get(dataset_name).evaluator_type53    if evaluator_type in ["sem_seg", "coco_panoptic_seg"]:54        evaluator_list.append(55            SemSegEvaluator(56                dataset_name,57                distributed=True,58                output_dir=output_folder,59            )60        )61    if evaluator_type in ["coco", "coco_panoptic_seg"]:62        evaluator_list.append(COCOEvaluator(dataset_name, output_dir=output_folder))63    if evaluator_type == "coco_panoptic_seg":64        evaluator_list.append(COCOPanopticEvaluator(dataset_name, output_folder))65    if evaluator_type == "cityscapes_instance":66        return CityscapesInstanceEvaluator(dataset_name)67    if evaluator_type == "cityscapes_sem_seg":68        return CityscapesSemSegEvaluator(dataset_name)69    elif evaluator_type == "pascal_voc":70        return PascalVOCDetectionEvaluator(dataset_name)71    elif evaluator_type == "lvis":72        return LVISEvaluator(dataset_name, output_dir=output_folder)73    if len(evaluator_list) == 0:74        raise NotImplementedError(75            "no Evaluator for the dataset {} with the type {}".format(dataset_name, evaluator_type)76        )77    elif len(evaluator_list) == 1:78        return evaluator_list[0]79    return DatasetEvaluators(evaluator_list)80 81 82class Trainer(DefaultTrainer):83    """84    We use the "DefaultTrainer" which contains pre-defined default logic for85    standard training workflow. They may not work for you, especially if you86    are working on a new research project. In that case you can write your87    own training loop. You can use "tools/plain_train_net.py" as an example.88    """89 90    @classmethod91    def build_evaluator(cls, cfg, dataset_name, output_folder=None):92        return build_evaluator(cfg, dataset_name, output_folder)93 94    @classmethod95    def test_with_TTA(cls, cfg, model):96        logger = logging.getLogger("detectron2.trainer")97        # In the end of training, run an evaluation with TTA98        # Only support some R-CNN models.99        logger.info("Running inference with test-time augmentation ...")100        model = GeneralizedRCNNWithTTA(cfg, model)101        evaluators = [102            cls.build_evaluator(103                cfg, name, output_folder=os.path.join(cfg.OUTPUT_DIR, "inference_TTA")104            )105            for name in cfg.DATASETS.TEST106        ]107        res = cls.test(cfg, model, evaluators)108        res = OrderedDict({k + "_TTA": v for k, v in res.items()})109        return res110 111 112def setup(args):113    """114    Create configs and perform basic setups.115    """116    cfg = get_cfg()117    cfg.merge_from_file(args.config_file)118    cfg.merge_from_list(args.opts)119    cfg.freeze()120    default_setup(cfg, args)121    return cfg122 123 124def main(args):125    cfg = setup(args)126 127    if args.eval_only:128        model = Trainer.build_model(cfg)129        DetectionCheckpointer(model, save_dir=cfg.OUTPUT_DIR).resume_or_load(130            cfg.MODEL.WEIGHTS, resume=args.resume131        )132        res = Trainer.test(cfg, model)133        if cfg.TEST.AUG.ENABLED:134            res.update(Trainer.test_with_TTA(cfg, model))135        if comm.is_main_process():136            verify_results(cfg, res)137        return res138 139    """140    If you'd like to do anything fancier than the standard training logic,141    consider writing your own training loop (see plain_train_net.py) or142    subclassing the trainer.143    """144    trainer = Trainer(cfg)145    trainer.resume_or_load(resume=args.resume)146    if cfg.TEST.AUG.ENABLED:147        trainer.register_hooks(148            [hooks.EvalHook(0, lambda: trainer.test_with_TTA(cfg, trainer.model))]149        )150    return trainer.train()151 152 153if __name__ == "__main__":154    args = default_argument_parser().parse_args()155    print("Command Line Args:", args)156    launch(157        main,158        args.num_gpus,159        num_machines=args.num_machines,160        machine_rank=args.machine_rank,161        dist_url=args.dist_url,162        args=(args,),163    )164