CoolFace
Apppublic

MinhQuangIntercom/tryon

sourceHugging Facecc-by-nc-sa-4.0updated 2y agoView on Hugging Face
0likes
apply_net.py360 linesDownload Raw Back to root
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: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(81            "--opts",82            help="Modify config options using the command-line 'KEY VALUE' pairs",83            default=[],84            nargs=argparse.REMAINDER,85        )86 87    @classmethod88    def execute(cls: type, args: argparse.Namespace, human_img):89        logger.info(f"Loading config from {args.cfg}")90        opts = []91        cfg = cls.setup_config(args.cfg, args.model, args, opts)92        logger.info(f"Loading model from {args.model}")93        predictor = DefaultPredictor(cfg)94        # logger.info(f"Loading data from {args.input}")95        # file_list = cls._get_input_file_list(args.input)96        # if len(file_list) == 0:97        #     logger.warning(f"No input images for {args.input}")98        #     return99        context = cls.create_context(args, cfg)100        # for file_name in file_list:101        #     img = read_image(file_name, format="BGR")  # predictor expects BGR image.102        with torch.no_grad():103            outputs = predictor(human_img)["instances"]104            out_pose = cls.execute_on_outputs(context, {"image": human_img}, outputs)105        cls.postexecute(context)106        return out_pose107 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        visualizer = context["visualizer"]277        extractor = context["extractor"]278        # image_fpath = entry["file_name"]279        # logger.info(f"Processing {image_fpath}")280        image = cv2.cvtColor(entry["image"], cv2.COLOR_BGR2GRAY)281        image = np.tile(image[:, :, np.newaxis], [1, 1, 3])282        data = extractor(outputs)283        image_vis = visualizer.visualize(image, data)284 285        return image_vis286        entry_idx = context["entry_idx"] + 1287        out_fname = './image-densepose/' + image_fpath.split('/')[-1]288        out_dir = './image-densepose'289        out_dir = os.path.dirname(out_fname)290        if len(out_dir) > 0 and not os.path.exists(out_dir):291            os.makedirs(out_dir)292        cv2.imwrite(out_fname, image_vis)293        logger.info(f"Output saved to {out_fname}")294        context["entry_idx"] += 1295 296    @classmethod297    def postexecute(cls: type, context: Dict[str, Any]):298        pass299# python ./apply_net.py show ./configs/densepose_rcnn_R_50_FPN_s1x.yaml https://dl.fbaipublicfiles.com/densepose/densepose_rcnn_R_50_FPN_s1x/165712039/model_final_162be9.pkl /home/alin0222/DressCode/upper_body/images dp_segm -v --opts MODEL.DEVICE cpu300 301    @classmethod302    def _get_out_fname(cls: type, entry_idx: int, fname_base: str):303        base, ext = os.path.splitext(fname_base)304        return base + ".{0:04d}".format(entry_idx) + ext305 306    @classmethod307    def create_context(cls: type, args: argparse.Namespace, cfg: CfgNode) -> Dict[str, Any]:308        vis_specs = args.visualizations.split(",")309        visualizers = []310        extractors = []311        for vis_spec in vis_specs:312            texture_atlas = get_texture_atlas(args.texture_atlas)313            texture_atlases_dict = get_texture_atlases(args.texture_atlases_map)314            vis = cls.VISUALIZERS[vis_spec](315                cfg=cfg,316                texture_atlas=texture_atlas,317                texture_atlases_dict=texture_atlases_dict,318            )319            visualizers.append(vis)320            extractor = create_extractor(vis)321            extractors.append(extractor)322        visualizer = CompoundVisualizer(visualizers)323        extractor = CompoundExtractor(extractors)324        context = {325            "extractor": extractor,326            "visualizer": visualizer,327            "out_fname": args.output,328            "entry_idx": 0,329        }330        return context331 332 333def create_argument_parser() -> argparse.ArgumentParser:334    parser = argparse.ArgumentParser(335        description=DOC,336        formatter_class=lambda prog: argparse.HelpFormatter(prog, max_help_position=120),337    )338    parser.set_defaults(func=lambda _: parser.print_help(sys.stdout))339    subparsers = parser.add_subparsers(title="Actions")340    for _, action in _ACTION_REGISTRY.items():341        action.add_parser(subparsers)342    return parser343 344 345def main():346    parser = create_argument_parser()347    args = parser.parse_args()348    verbosity = getattr(args, "verbosity", None)349    global logger350    logger = setup_logger(name=LOGGER_NAME)351    logger.setLevel(verbosity_to_level(verbosity))352    args.func(args)353 354 355if __name__ == "__main__":356    main()357 358 359# python ./apply_net.py show ./configs/densepose_rcnn_R_50_FPN_s1x.yaml https://dl.fbaipublicfiles.com/densepose/densepose_rcnn_R_50_FPN_s1x/165712039/model_final_162be9.pkl /home/alin0222/Dresscode/dresses/humanonly dp_segm -v --opts MODEL.DEVICE cuda360