Team Ai
Apppublic

Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
visualize_data.py95 linesDownload Raw Back to tools
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