Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
apply_net.py354 linesDownload Raw Back to DensePose
1#!/usr/bin/env python32# Copyright (c) Facebook, Inc. and its affiliates.3 4import argparse5import glob6import logging7import os8import sys9from typing import Any, ClassVar, Dict, List10import torch11 12from detectron2.config import CfgNode, get_cfg13from detectron2.data.detection_utils import read_image14from detectron2.engine.defaults import DefaultPredictor15from detectron2.structures.instances import Instances16from detectron2.utils.logger import setup_logger17 18from densepose import add_densepose_config19from densepose.structures import DensePoseChartPredictorOutput, DensePoseEmbeddingPredictorOutput20from densepose.utils.logger import verbosity_to_level21from densepose.vis.base import CompoundVisualizer22from densepose.vis.bounding_box import ScoredBoundingBoxVisualizer23from densepose.vis.densepose_outputs_vertex import (24    DensePoseOutputsTextureVisualizer,25    DensePoseOutputsVertexVisualizer,26    get_texture_atlases,27)28from densepose.vis.densepose_results import (29    DensePoseResultsContourVisualizer,30    DensePoseResultsFineSegmentationVisualizer,31    DensePoseResultsUVisualizer,32    DensePoseResultsVVisualizer,33)34from densepose.vis.densepose_results_textures import (35    DensePoseResultsVisualizerWithTexture,36    get_texture_atlas,37)38from densepose.vis.extractor import (39    CompoundExtractor,40    DensePoseOutputsExtractor,41    DensePoseResultExtractor,42    create_extractor,43)44 45DOC = """Apply Net - a tool to print / visualize DensePose results46"""47 48LOGGER_NAME = "apply_net"49logger = logging.getLogger(LOGGER_NAME)50 51_ACTION_REGISTRY: Dict[str, "Action"] = {}52 53 54class Action(object):55    @classmethod56    def add_arguments(cls: type, parser: argparse.ArgumentParser):57        parser.add_argument(58            "-v",59            "--verbosity",60            action="count",61            help="Verbose mode. Multiple -v options increase the verbosity.",62        )63 64 65def register_action(cls: type):66    """67    Decorator for action classes to automate action registration68    """69    global _ACTION_REGISTRY70    _ACTION_REGISTRY[cls.COMMAND] = cls71    return cls72 73 74class InferenceAction(Action):75    @classmethod76    def add_arguments(cls: type, parser: argparse.ArgumentParser):77        super(InferenceAction, cls).add_arguments(parser)78        parser.add_argument("cfg", metavar="<config>", help="Config file")79        parser.add_argument("model", metavar="<model>", help="Model file")80        parser.add_argument("input", metavar="<input>", help="Input data")81        parser.add_argument(82            "--opts",83            help="Modify config options using the command-line 'KEY VALUE' pairs",84            default=[],85            nargs=argparse.REMAINDER,86        )87 88    @classmethod89    def execute(cls: type, args: argparse.Namespace):90        logger.info(f"Loading config from {args.cfg}")91        opts = []92        cfg = cls.setup_config(args.cfg, args.model, args, opts)93        logger.info(f"Loading model from {args.model}")94        predictor = DefaultPredictor(cfg)95        logger.info(f"Loading data from {args.input}")96        file_list = cls._get_input_file_list(args.input)97        if len(file_list) == 0:98            logger.warning(f"No input images for {args.input}")99            return100        context = cls.create_context(args, cfg)101        for file_name in file_list:102            img = read_image(file_name, format="BGR")  # predictor expects BGR image.103            with torch.no_grad():104                outputs = predictor(img)["instances"]105                cls.execute_on_outputs(context, {"file_name": file_name, "image": img}, outputs)106        cls.postexecute(context)107 108    @classmethod109    def setup_config(110        cls: type, config_fpath: str, model_fpath: str, args: argparse.Namespace, opts: List[str]111    ):112        cfg = get_cfg()113        add_densepose_config(cfg)114        cfg.merge_from_file(config_fpath)115        cfg.merge_from_list(args.opts)116        if opts:117            cfg.merge_from_list(opts)118        cfg.MODEL.WEIGHTS = model_fpath119        cfg.freeze()120        return cfg121 122    @classmethod123    def _get_input_file_list(cls: type, input_spec: str):124        if os.path.isdir(input_spec):125            file_list = [126                os.path.join(input_spec, fname)127                for fname in os.listdir(input_spec)128                if os.path.isfile(os.path.join(input_spec, fname))129            ]130        elif os.path.isfile(input_spec):131            file_list = [input_spec]132        else:133            file_list = glob.glob(input_spec)134        return file_list135 136 137@register_action138class DumpAction(InferenceAction):139    """140    Dump action that outputs results to a pickle file141    """142 143    COMMAND: ClassVar[str] = "dump"144 145    @classmethod146    def add_parser(cls: type, subparsers: argparse._SubParsersAction):147        parser = subparsers.add_parser(cls.COMMAND, help="Dump model outputs to a file.")148        cls.add_arguments(parser)149        parser.set_defaults(func=cls.execute)150 151    @classmethod152    def add_arguments(cls: type, parser: argparse.ArgumentParser):153        super(DumpAction, cls).add_arguments(parser)154        parser.add_argument(155            "--output",156            metavar="<dump_file>",157            default="results.pkl",158            help="File name to save dump to",159        )160 161    @classmethod162    def execute_on_outputs(163        cls: type, context: Dict[str, Any], entry: Dict[str, Any], outputs: Instances164    ):165        image_fpath = entry["file_name"]166        logger.info(f"Processing {image_fpath}")167        result = {"file_name": image_fpath}168        if outputs.has("scores"):169            result["scores"] = outputs.get("scores").cpu()170        if outputs.has("pred_boxes"):171            result["pred_boxes_XYXY"] = outputs.get("pred_boxes").tensor.cpu()172            if outputs.has("pred_densepose"):173                if isinstance(outputs.pred_densepose, DensePoseChartPredictorOutput):174                    extractor = DensePoseResultExtractor()175                elif isinstance(outputs.pred_densepose, DensePoseEmbeddingPredictorOutput):176                    extractor = DensePoseOutputsExtractor()177                result["pred_densepose"] = extractor(outputs)[0]178        context["results"].append(result)179 180    @classmethod181    def create_context(cls: type, args: argparse.Namespace, cfg: CfgNode):182        context = {"results": [], "out_fname": args.output}183        return context184 185    @classmethod186    def postexecute(cls: type, context: Dict[str, Any]):187        out_fname = context["out_fname"]188        out_dir = os.path.dirname(out_fname)189        if len(out_dir) > 0 and not os.path.exists(out_dir):190            os.makedirs(out_dir)191        with open(out_fname, "wb") as hFile:192            torch.save(context["results"], hFile)193            logger.info(f"Output saved to {out_fname}")194 195 196@register_action197class ShowAction(InferenceAction):198    """199    Show action that visualizes selected entries on an image200    """201 202    COMMAND: ClassVar[str] = "show"203    VISUALIZERS: ClassVar[Dict[str, object]] = {204        "dp_contour": DensePoseResultsContourVisualizer,205        "dp_segm": DensePoseResultsFineSegmentationVisualizer,206        "dp_u": DensePoseResultsUVisualizer,207        "dp_v": DensePoseResultsVVisualizer,208        "dp_iuv_texture": DensePoseResultsVisualizerWithTexture,209        "dp_cse_texture": DensePoseOutputsTextureVisualizer,210        "dp_vertex": DensePoseOutputsVertexVisualizer,211        "bbox": ScoredBoundingBoxVisualizer,212    }213 214    @classmethod215    def add_parser(cls: type, subparsers: argparse._SubParsersAction):216        parser = subparsers.add_parser(cls.COMMAND, help="Visualize selected entries")217        cls.add_arguments(parser)218        parser.set_defaults(func=cls.execute)219 220    @classmethod221    def add_arguments(cls: type, parser: argparse.ArgumentParser):222        super(ShowAction, cls).add_arguments(parser)223        parser.add_argument(224            "visualizations",225            metavar="<visualizations>",226            help="Comma separated list of visualizations, possible values: "227            "[{}]".format(",".join(sorted(cls.VISUALIZERS.keys()))),228        )229        parser.add_argument(230            "--min_score",231            metavar="<score>",232            default=0.8,233            type=float,234            help="Minimum detection score to visualize",235        )236        parser.add_argument(237            "--nms_thresh", metavar="<threshold>", default=None, type=float, help="NMS threshold"238        )239        parser.add_argument(240            "--texture_atlas",241            metavar="<texture_atlas>",242            default=None,243            help="Texture atlas file (for IUV texture transfer)",244        )245        parser.add_argument(246            "--texture_atlases_map",247            metavar="<texture_atlases_map>",248            default=None,249            help="JSON string of a dict containing texture atlas files for each mesh",250        )251        parser.add_argument(252            "--output",253            metavar="<image_file>",254            default="outputres.png",255            help="File name to save output to",256        )257 258    @classmethod259    def setup_config(260        cls: type, config_fpath: str, model_fpath: str, args: argparse.Namespace, opts: List[str]261    ):262        opts.append("MODEL.ROI_HEADS.SCORE_THRESH_TEST")263        opts.append(str(args.min_score))264        if args.nms_thresh is not None:265            opts.append("MODEL.ROI_HEADS.NMS_THRESH_TEST")266            opts.append(str(args.nms_thresh))267        cfg = super(ShowAction, cls).setup_config(config_fpath, model_fpath, args, opts)268        return cfg269 270    @classmethod271    def execute_on_outputs(272        cls: type, context: Dict[str, Any], entry: Dict[str, Any], outputs: Instances273    ):274        import cv2275        import numpy as np276 277        visualizer = context["visualizer"]278        extractor = context["extractor"]279        image_fpath = entry["file_name"]280        logger.info(f"Processing {image_fpath}")281        image = cv2.cvtColor(entry["image"], cv2.COLOR_BGR2GRAY)282        image = np.tile(image[:, :, np.newaxis], [1, 1, 3])283        data = extractor(outputs)284        image_vis = visualizer.visualize(image, data)285        entry_idx = context["entry_idx"] + 1286        out_fname = cls._get_out_fname(entry_idx, context["out_fname"])287        out_dir = os.path.dirname(out_fname)288        if len(out_dir) > 0 and not os.path.exists(out_dir):289            os.makedirs(out_dir)290        cv2.imwrite(out_fname, image_vis)291        logger.info(f"Output saved to {out_fname}")292        context["entry_idx"] += 1293 294    @classmethod295    def postexecute(cls: type, context: Dict[str, Any]):296        pass297 298    @classmethod299    def _get_out_fname(cls: type, entry_idx: int, fname_base: str):300        base, ext = os.path.splitext(fname_base)301        return base + ".{0:04d}".format(entry_idx) + ext302 303    @classmethod304    def create_context(cls: type, args: argparse.Namespace, cfg: CfgNode) -> Dict[str, Any]:305        vis_specs = args.visualizations.split(",")306        visualizers = []307        extractors = []308        for vis_spec in vis_specs:309            texture_atlas = get_texture_atlas(args.texture_atlas)310            texture_atlases_dict = get_texture_atlases(args.texture_atlases_map)311            vis = cls.VISUALIZERS[vis_spec](312                cfg=cfg,313                texture_atlas=texture_atlas,314                texture_atlases_dict=texture_atlases_dict,315            )316            visualizers.append(vis)317            extractor = create_extractor(vis)318            extractors.append(extractor)319        visualizer = CompoundVisualizer(visualizers)320        extractor = CompoundExtractor(extractors)321        context = {322            "extractor": extractor,323            "visualizer": visualizer,324            "out_fname": args.output,325            "entry_idx": 0,326        }327        return context328 329 330def create_argument_parser() -> argparse.ArgumentParser:331    parser = argparse.ArgumentParser(332        description=DOC,333        formatter_class=lambda prog: argparse.HelpFormatter(prog, max_help_position=120),334    )335    parser.set_defaults(func=lambda _: parser.print_help(sys.stdout))336    subparsers = parser.add_subparsers(title="Actions")337    for _, action in _ACTION_REGISTRY.items():338        action.add_parser(subparsers)339    return parser340 341 342def main():343    parser = create_argument_parser()344    args = parser.parse_args()345    verbosity = getattr(args, "verbosity", None)346    global logger347    logger = setup_logger(name=LOGGER_NAME)348    logger.setLevel(verbosity_to_level(verbosity))349    args.func(args)350 351 352if __name__ == "__main__":353    main()354