Team Ai
Apppublic

Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
train_net.py68 linesDownload Raw Back to TridentNet
1#!/usr/bin/env python32# Copyright (c) Facebook, Inc. and its affiliates.3 4"""5TridentNet Training Script.6 7This script is a simplified version of the training script in detectron2/tools.8"""9 10import os11 12from detectron2.checkpoint import DetectionCheckpointer13from detectron2.config import get_cfg14from detectron2.engine import DefaultTrainer, default_argument_parser, default_setup, launch15from detectron2.evaluation import COCOEvaluator16 17from tridentnet import add_tridentnet_config18 19 20class Trainer(DefaultTrainer):21    @classmethod22    def build_evaluator(cls, cfg, dataset_name, output_folder=None):23        if output_folder is None:24            output_folder = os.path.join(cfg.OUTPUT_DIR, "inference")25        return COCOEvaluator(dataset_name, output_dir=output_folder)26 27 28def setup(args):29    """30    Create configs and perform basic setups.31    """32    cfg = get_cfg()33    add_tridentnet_config(cfg)34    cfg.merge_from_file(args.config_file)35    cfg.merge_from_list(args.opts)36    cfg.freeze()37    default_setup(cfg, args)38    return cfg39 40 41def main(args):42    cfg = setup(args)43 44    if args.eval_only:45        model = Trainer.build_model(cfg)46        DetectionCheckpointer(model, save_dir=cfg.OUTPUT_DIR).resume_or_load(47            cfg.MODEL.WEIGHTS, resume=args.resume48        )49        res = Trainer.test(cfg, model)50        return res51 52    trainer = Trainer(cfg)53    trainer.resume_or_load(resume=args.resume)54    return trainer.train()55 56 57if __name__ == "__main__":58    args = default_argument_parser().parse_args()59    print("Command Line Args:", args)60    launch(61        main,62        args.num_gpus,63        num_machines=args.num_machines,64        machine_rank=args.machine_rank,65        dist_url=args.dist_url,66        args=(args,),67    )68