Team Ai
Apppublic

radames/Text2Human-API

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
model.py146 linesDownload Raw Back to root
1from __future__ import annotations2 3import os4import pathlib5import sys6import zipfile7 8import huggingface_hub9import numpy as np10import PIL.Image11import torch12 13sys.path.insert(0, 'Text2Human')14 15from models.sample_model import SampleFromPoseModel16from utils.language_utils import (generate_shape_attributes,17                                  generate_texture_attributes)18from utils.options import dict_to_nonedict, parse19from utils.util import set_random_seed20 21COLOR_LIST = [22    (0, 0, 0),23    (255, 250, 250),24    (220, 220, 220),25    (250, 235, 215),26    (255, 250, 205),27    (211, 211, 211),28    (70, 130, 180),29    (127, 255, 212),30    (0, 100, 0),31    (50, 205, 50),32    (255, 255, 0),33    (245, 222, 179),34    (255, 140, 0),35    (255, 0, 0),36    (16, 78, 139),37    (144, 238, 144),38    (50, 205, 174),39    (50, 155, 250),40    (160, 140, 88),41    (213, 140, 88),42    (90, 140, 90),43    (185, 210, 205),44    (130, 165, 180),45    (225, 141, 151),46]47 48 49class Model:50    def __init__(self, device: str):51        self.config = self._load_config()52        self.config['device'] = device53        self._download_models()54        self.model = SampleFromPoseModel(self.config)55        self.model.batch_size = 156 57    def _load_config(self) -> dict:58        path = 'Text2Human/configs/sample_from_pose.yml'59        config = parse(path, is_train=False)60        config = dict_to_nonedict(config)61        return config62 63    def _download_models(self) -> None:64        model_dir = pathlib.Path('pretrained_models')65        if model_dir.exists():66            return67        token = os.getenv('HF_TOKEN')68        path = huggingface_hub.hf_hub_download('yumingj/Text2Human_SSHQ',69                                               'pretrained_models.zip',70                                               use_auth_token=token)71        model_dir.mkdir()72        with zipfile.ZipFile(path) as f:73            f.extractall(model_dir)74 75    @staticmethod76    def preprocess_pose_image(image: PIL.Image.Image) -> torch.Tensor:77        image = np.array(78            image.resize(79                size=(256, 512),80                resample=PIL.Image.Resampling.LANCZOS))[:, :, 2:].transpose(81                    2, 0, 1).astype(np.float32)82        image = image / 12. - 183        data = torch.from_numpy(image).unsqueeze(1)84        return data85 86    @staticmethod87    def process_mask(mask: np.ndarray) -> np.ndarray:88        if mask.shape != (512, 256, 3):89            return None90        seg_map = np.full(mask.shape[:-1], -1)91        for index, color in enumerate(COLOR_LIST):92            seg_map[np.sum(mask == color, axis=2) == 3] = index93        return seg_map94 95    @staticmethod96    def postprocess(result: torch.Tensor) -> np.ndarray:97        result = result.permute(0, 2, 3, 1)98        result = result.detach().cpu().numpy()99        result = result * 255100        result = np.asarray(result[0, :, :, :], dtype=np.uint8)101        return result102 103    def process_pose_image(self, pose_image: PIL.Image.Image) -> torch.Tensor:104        if pose_image is None:105            return106        data = self.preprocess_pose_image(pose_image)107        self.model.feed_pose_data(data)108        return data109 110    def generate_label_image(self, pose_data: torch.Tensor,111                             shape_text: str) -> np.ndarray:112        if pose_data is None:113            return114        self.model.feed_pose_data(pose_data)115        shape_attributes = generate_shape_attributes(shape_text)116        shape_attributes = torch.LongTensor(shape_attributes).unsqueeze(0)117        self.model.feed_shape_attributes(shape_attributes)118        self.model.generate_parsing_map()119        self.model.generate_quantized_segm()120        colored_segm = self.model.palette_result(self.model.segm[0].cpu())121        return colored_segm122 123    def generate_human(self, label_image: np.ndarray, texture_text: str,124                       sample_steps: int, seed: int) -> np.ndarray:125        if label_image is None:126            return127        mask = label_image.copy()128        seg_map = self.process_mask(mask)129        if seg_map is None:130            return131        self.model.segm = torch.from_numpy(seg_map).unsqueeze(0).unsqueeze(132            0).to(self.model.device)133        self.model.generate_quantized_segm()134 135        set_random_seed(seed)136 137        texture_attributes = generate_texture_attributes(texture_text)138        texture_attributes = torch.LongTensor(texture_attributes)139        self.model.feed_texture_attributes(texture_attributes)140        self.model.generate_texture_map()141 142        self.model.sample_steps = sample_steps143        out = self.model.sample_and_refine()144        res = self.postprocess(out)145        return res146