Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
train_net.py116 linesDownload Raw Back to PointSup
1#!/usr/bin/env python2# Copyright (c) Facebook, Inc. and its affiliates.3"""4Point supervision Training Script.5 6This script is a simplified version of the training script in detectron2/tools.7"""8 9import os10 11import detectron2.utils.comm as comm12from detectron2.checkpoint import DetectionCheckpointer13from detectron2.config import get_cfg14from detectron2.data import MetadataCatalog, build_detection_train_loader15from detectron2.engine import DefaultTrainer, default_argument_parser, default_setup, launch16from detectron2.evaluation import COCOEvaluator, DatasetEvaluators, verify_results17from detectron2.projects.point_rend import add_pointrend_config18from detectron2.utils.logger import setup_logger19 20from point_sup import PointSupDatasetMapper, add_point_sup_config21 22 23class Trainer(DefaultTrainer):24    """25    We use the "DefaultTrainer" which contains pre-defined default logic for26    standard training workflow. They may not work for you, especially if you27    are working on a new research project. In that case you can write your28    own training loop. You can use "tools/plain_train_net.py" as an example.29    """30 31    @classmethod32    def build_evaluator(cls, cfg, dataset_name, output_folder=None):33        """34        Create evaluator(s) for a given dataset.35        This uses the special metadata "evaluator_type" associated with each builtin dataset.36        For your own dataset, you can simply create an evaluator manually in your37        script and do not have to worry about the hacky if-else logic here.38        """39        if output_folder is None:40            output_folder = os.path.join(cfg.OUTPUT_DIR, "inference")41        evaluator_list = []42        evaluator_type = MetadataCatalog.get(dataset_name).evaluator_type43        if evaluator_type == "coco":44            evaluator_list.append(COCOEvaluator(dataset_name, output_dir=output_folder))45        if len(evaluator_list) == 0:46            raise NotImplementedError(47                "no Evaluator for the dataset {} with the type {}".format(48                    dataset_name, evaluator_type49                )50            )51        elif len(evaluator_list) == 1:52            return evaluator_list[0]53        return DatasetEvaluators(evaluator_list)54 55    @classmethod56    def build_train_loader(cls, cfg):57        if cfg.INPUT.POINT_SUP:58            mapper = PointSupDatasetMapper(cfg, is_train=True)59        else:60            mapper = None61        return build_detection_train_loader(cfg, mapper=mapper)62 63 64def setup(args):65    """66    Create configs and perform basic setups.67    """68    cfg = get_cfg()69    add_pointrend_config(cfg)70    add_point_sup_config(cfg)71    cfg.merge_from_file(args.config_file)72    cfg.merge_from_list(args.opts)73    cfg.freeze()74    default_setup(cfg, args)75    # Setup logger for "point_sup" module76    setup_logger(output=cfg.OUTPUT_DIR, distributed_rank=comm.get_rank(), name="point_sup")77    return cfg78 79 80def main(args):81    cfg = setup(args)82 83    if args.eval_only:84        model = Trainer.build_model(cfg)85        DetectionCheckpointer(model, save_dir=cfg.OUTPUT_DIR).resume_or_load(86            cfg.MODEL.WEIGHTS, resume=args.resume87        )88        res = Trainer.test(cfg, model)89        if cfg.TEST.AUG.ENABLED:90            res.update(Trainer.test_with_TTA(cfg, model))91        if comm.is_main_process():92            verify_results(cfg, res)93        return res94 95    """96    If you'd like to do anything fancier than the standard training logic,97    consider writing your own training loop (see plain_train_net.py) or98    subclassing the trainer.99    """100    trainer = Trainer(cfg)101    trainer.resume_or_load(resume=args.resume)102    return trainer.train()103 104 105if __name__ == "__main__":106    args = default_argument_parser().parse_args()107    print("Command Line Args:", args)108    launch(109        main,110        args.num_gpus,111        num_machines=args.num_machines,112        machine_rank=args.machine_rank,113        dist_url=args.dist_url,114        args=(args,),115    )116