Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model
0
1#!/usr/bin/env python2# Copyright (c) Facebook, Inc. and its affiliates.3import argparse4import os5from itertools import chain6import cv27import tqdm8 9from detectron2.config import get_cfg10from detectron2.data import DatasetCatalog, MetadataCatalog, build_detection_train_loader11from detectron2.data import detection_utils as utils12from detectron2.data.build import filter_images_with_few_keypoints13from detectron2.utils.logger import setup_logger14from detectron2.utils.visualizer import Visualizer15 16 17def setup(args):18 cfg = get_cfg()19 if args.config_file:20 cfg.merge_from_file(args.config_file)21 cfg.merge_from_list(args.opts)22 cfg.DATALOADER.NUM_WORKERS = 023 cfg.freeze()24 return cfg25 26 27def parse_args(in_args=None):28 parser = argparse.ArgumentParser(description="Visualize ground-truth data")29 parser.add_argument(30 "--source",31 choices=["annotation", "dataloader"],32 required=True,33 help="visualize the annotations or the data loader (with pre-processing)",34 )35 parser.add_argument("--config-file", metavar="FILE", help="path to config file")36 parser.add_argument("--output-dir", default="./", help="path to output directory")37 parser.add_argument("--show", action="store_true", help="show output in a window")38 parser.add_argument(39 "opts",40 help="Modify config options using the command-line",41 default=None,42 nargs=argparse.REMAINDER,43 )44 return parser.parse_args(in_args)45 46 47if __name__ == "__main__":48 args = parse_args()49 logger = setup_logger()50 logger.info("Arguments: " + str(args))51 cfg = setup(args)52 53 dirname = args.output_dir54 os.makedirs(dirname, exist_ok=True)55 metadata = MetadataCatalog.get(cfg.DATASETS.TRAIN[0])56 57 def output(vis, fname):58 if args.show:59 print(fname)60 cv2.imshow("window", vis.get_image()[:, :, ::-1])61 cv2.waitKey()62 else:63 filepath = os.path.join(dirname, fname)64 print("Saving to {} ...".format(filepath))65 vis.save(filepath)66 67 scale = 1.068 if args.source == "dataloader":69 train_data_loader = build_detection_train_loader(cfg)70 for batch in train_data_loader:71 for per_image in batch:72 # Pytorch tensor is in (C, H, W) format73 img = per_image["image"].permute(1, 2, 0).cpu().detach().numpy()74 img = utils.convert_image_to_rgb(img, cfg.INPUT.FORMAT)75 76 visualizer = Visualizer(img, metadata=metadata, scale=scale)77 target_fields = per_image["instances"].get_fields()78 labels = [metadata.thing_classes[i] for i in target_fields["gt_classes"]]79 vis = visualizer.overlay_instances(80 labels=labels,81 boxes=target_fields.get("gt_boxes", None),82 masks=target_fields.get("gt_masks", None),83 keypoints=target_fields.get("gt_keypoints", None),84 )85 output(vis, str(per_image["image_id"]) + ".jpg")86 else:87 dicts = list(chain.from_iterable([DatasetCatalog.get(k) for k in cfg.DATASETS.TRAIN]))88 if cfg.MODEL.KEYPOINT_ON:89 dicts = filter_images_with_few_keypoints(dicts, 1)90 for dic in tqdm.tqdm(dicts):91 img = utils.read_image(dic["file_name"], "RGB")92 visualizer = Visualizer(img, metadata=metadata, scale=scale)93 vis = visualizer.draw_dataset_dict(dic)94 output(vis, os.path.basename(dic["file_name"]))95 