Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
query_db.py251 linesDownload Raw Back to DensePose
1#!/usr/bin/env python32# Copyright (c) Facebook, Inc. and its affiliates.3 4import argparse5import logging6import os7import sys8from timeit import default_timer as timer9from typing import Any, ClassVar, Dict, List10import torch11 12from detectron2.data.catalog import DatasetCatalog13from detectron2.utils.file_io import PathManager14from detectron2.utils.logger import setup_logger15 16from densepose.structures import DensePoseDataRelative17from densepose.utils.dbhelper import EntrySelector18from densepose.utils.logger import verbosity_to_level19from densepose.vis.base import CompoundVisualizer20from densepose.vis.bounding_box import BoundingBoxVisualizer21from densepose.vis.densepose_data_points import (22    DensePoseDataCoarseSegmentationVisualizer,23    DensePoseDataPointsIVisualizer,24    DensePoseDataPointsUVisualizer,25    DensePoseDataPointsVisualizer,26    DensePoseDataPointsVVisualizer,27)28 29DOC = """Query DB - a tool to print / visualize data from a database30"""31 32LOGGER_NAME = "query_db"33 34logger = logging.getLogger(LOGGER_NAME)35 36_ACTION_REGISTRY: Dict[str, "Action"] = {}37 38 39class Action(object):40    @classmethod41    def add_arguments(cls: type, parser: argparse.ArgumentParser):42        parser.add_argument(43            "-v",44            "--verbosity",45            action="count",46            help="Verbose mode. Multiple -v options increase the verbosity.",47        )48 49 50def register_action(cls: type):51    """52    Decorator for action classes to automate action registration53    """54    global _ACTION_REGISTRY55    _ACTION_REGISTRY[cls.COMMAND] = cls56    return cls57 58 59class EntrywiseAction(Action):60    @classmethod61    def add_arguments(cls: type, parser: argparse.ArgumentParser):62        super(EntrywiseAction, cls).add_arguments(parser)63        parser.add_argument(64            "dataset", metavar="<dataset>", help="Dataset name (e.g. densepose_coco_2014_train)"65        )66        parser.add_argument(67            "selector",68            metavar="<selector>",69            help="Dataset entry selector in the form field1[:type]=value1[,"70            "field2[:type]=value_min-value_max...] which selects all "71            "entries from the dataset that satisfy the constraints",72        )73        parser.add_argument(74            "--max-entries", metavar="N", help="Maximum number of entries to process", type=int75        )76 77    @classmethod78    def execute(cls: type, args: argparse.Namespace):79        dataset = setup_dataset(args.dataset)80        entry_selector = EntrySelector.from_string(args.selector)81        context = cls.create_context(args)82        if args.max_entries is not None:83            for _, entry in zip(range(args.max_entries), dataset):84                if entry_selector(entry):85                    cls.execute_on_entry(entry, context)86        else:87            for entry in dataset:88                if entry_selector(entry):89                    cls.execute_on_entry(entry, context)90 91    @classmethod92    def create_context(cls: type, args: argparse.Namespace) -> Dict[str, Any]:93        context = {}94        return context95 96 97@register_action98class PrintAction(EntrywiseAction):99    """100    Print action that outputs selected entries to stdout101    """102 103    COMMAND: ClassVar[str] = "print"104 105    @classmethod106    def add_parser(cls: type, subparsers: argparse._SubParsersAction):107        parser = subparsers.add_parser(cls.COMMAND, help="Output selected entries to stdout. ")108        cls.add_arguments(parser)109        parser.set_defaults(func=cls.execute)110 111    @classmethod112    def add_arguments(cls: type, parser: argparse.ArgumentParser):113        super(PrintAction, cls).add_arguments(parser)114 115    @classmethod116    def execute_on_entry(cls: type, entry: Dict[str, Any], context: Dict[str, Any]):117        import pprint118 119        printer = pprint.PrettyPrinter(indent=2, width=200, compact=True)120        printer.pprint(entry)121 122 123@register_action124class ShowAction(EntrywiseAction):125    """126    Show action that visualizes selected entries on an image127    """128 129    COMMAND: ClassVar[str] = "show"130    VISUALIZERS: ClassVar[Dict[str, object]] = {131        "dp_segm": DensePoseDataCoarseSegmentationVisualizer(),132        "dp_i": DensePoseDataPointsIVisualizer(),133        "dp_u": DensePoseDataPointsUVisualizer(),134        "dp_v": DensePoseDataPointsVVisualizer(),135        "dp_pts": DensePoseDataPointsVisualizer(),136        "bbox": BoundingBoxVisualizer(),137    }138 139    @classmethod140    def add_parser(cls: type, subparsers: argparse._SubParsersAction):141        parser = subparsers.add_parser(cls.COMMAND, help="Visualize selected entries")142        cls.add_arguments(parser)143        parser.set_defaults(func=cls.execute)144 145    @classmethod146    def add_arguments(cls: type, parser: argparse.ArgumentParser):147        super(ShowAction, cls).add_arguments(parser)148        parser.add_argument(149            "visualizations",150            metavar="<visualizations>",151            help="Comma separated list of visualizations, possible values: "152            "[{}]".format(",".join(sorted(cls.VISUALIZERS.keys()))),153        )154        parser.add_argument(155            "--output",156            metavar="<image_file>",157            default="output.png",158            help="File name to save output to",159        )160 161    @classmethod162    def execute_on_entry(cls: type, entry: Dict[str, Any], context: Dict[str, Any]):163        import cv2164        import numpy as np165 166        image_fpath = PathManager.get_local_path(entry["file_name"])167        image = cv2.imread(image_fpath, cv2.IMREAD_GRAYSCALE)168        image = np.tile(image[:, :, np.newaxis], [1, 1, 3])169        datas = cls._extract_data_for_visualizers_from_entry(context["vis_specs"], entry)170        visualizer = context["visualizer"]171        image_vis = visualizer.visualize(image, datas)172        entry_idx = context["entry_idx"] + 1173        out_fname = cls._get_out_fname(entry_idx, context["out_fname"])174        cv2.imwrite(out_fname, image_vis)175        logger.info(f"Output saved to {out_fname}")176        context["entry_idx"] += 1177 178    @classmethod179    def _get_out_fname(cls: type, entry_idx: int, fname_base: str):180        base, ext = os.path.splitext(fname_base)181        return base + ".{0:04d}".format(entry_idx) + ext182 183    @classmethod184    def create_context(cls: type, args: argparse.Namespace) -> Dict[str, Any]:185        vis_specs = args.visualizations.split(",")186        visualizers = []187        for vis_spec in vis_specs:188            vis = cls.VISUALIZERS[vis_spec]189            visualizers.append(vis)190        context = {191            "vis_specs": vis_specs,192            "visualizer": CompoundVisualizer(visualizers),193            "out_fname": args.output,194            "entry_idx": 0,195        }196        return context197 198    @classmethod199    def _extract_data_for_visualizers_from_entry(200        cls: type, vis_specs: List[str], entry: Dict[str, Any]201    ):202        dp_list = []203        bbox_list = []204        for annotation in entry["annotations"]:205            is_valid, _ = DensePoseDataRelative.validate_annotation(annotation)206            if not is_valid:207                continue208            bbox = torch.as_tensor(annotation["bbox"])209            bbox_list.append(bbox)210            dp_data = DensePoseDataRelative(annotation)211            dp_list.append(dp_data)212        datas = []213        for vis_spec in vis_specs:214            datas.append(bbox_list if "bbox" == vis_spec else (bbox_list, dp_list))215        return datas216 217 218def setup_dataset(dataset_name):219    logger.info("Loading dataset {}".format(dataset_name))220    start = timer()221    dataset = DatasetCatalog.get(dataset_name)222    stop = timer()223    logger.info("Loaded dataset {} in {:.3f}s".format(dataset_name, stop - start))224    return dataset225 226 227def create_argument_parser() -> argparse.ArgumentParser:228    parser = argparse.ArgumentParser(229        description=DOC,230        formatter_class=lambda prog: argparse.HelpFormatter(prog, max_help_position=120),231    )232    parser.set_defaults(func=lambda _: parser.print_help(sys.stdout))233    subparsers = parser.add_subparsers(title="Actions")234    for _, action in _ACTION_REGISTRY.items():235        action.add_parser(subparsers)236    return parser237 238 239def main():240    parser = create_argument_parser()241    args = parser.parse_args()242    verbosity = getattr(args, "verbosity", None)243    global logger244    logger = setup_logger(name=LOGGER_NAME)245    logger.setLevel(verbosity_to_level(verbosity))246    args.func(args)247 248 249if __name__ == "__main__":250    main()251