Team Ai
Apppublic

radames/Text2Human-API

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
ui_demo.py286 linesDownload Raw Back to Text2Human
1import sys2 3import cv24import numpy as np5import torch6from PIL import Image7from PyQt5.QtCore import *8from PyQt5.QtGui import *9from PyQt5.QtWidgets import *10 11from models.sample_model import SampleFromPoseModel12from ui.mouse_event import GraphicsScene13from ui.ui import Ui_Form14from utils.language_utils import (generate_shape_attributes,15                                  generate_texture_attributes)16from utils.options import dict_to_nonedict, parse17 18color_list = [(0, 0, 0), (255, 250, 250), (220, 220, 220), (250, 235, 215),19              (255, 250, 205), (211, 211, 211), (70, 130, 180),20              (127, 255, 212), (0, 100, 0), (50, 205, 50), (255, 255, 0),21              (245, 222, 179), (255, 140, 0), (255, 0, 0), (16, 78, 139),22              (144, 238, 144), (50, 205, 174), (50, 155, 250), (160, 140, 88),23              (213, 140, 88), (90, 140, 90), (185, 210, 205), (130, 165, 180),24              (225, 141, 151)]25 26 27class Ex(QWidget, Ui_Form):28 29    def __init__(self, opt):30        super(Ex, self).__init__()31        self.setupUi(self)32        self.show()33 34        self.output_img = None35 36        self.mat_img = None37 38        self.mode = 039        self.size = 640        self.mask = None41        self.mask_m = None42        self.img = None43 44        # about UI45        self.mouse_clicked = False46        self.scene = QGraphicsScene()47        self.graphicsView.setScene(self.scene)48        self.graphicsView.setAlignment(Qt.AlignTop | Qt.AlignLeft)49        self.graphicsView.setVerticalScrollBarPolicy(Qt.ScrollBarAlwaysOff)50        self.graphicsView.setHorizontalScrollBarPolicy(Qt.ScrollBarAlwaysOff)51 52        self.ref_scene = GraphicsScene(self.mode, self.size)53        self.graphicsView_2.setScene(self.ref_scene)54        self.graphicsView_2.setAlignment(Qt.AlignTop | Qt.AlignLeft)55        self.graphicsView_2.setVerticalScrollBarPolicy(Qt.ScrollBarAlwaysOff)56        self.graphicsView_2.setHorizontalScrollBarPolicy(Qt.ScrollBarAlwaysOff)57 58        self.result_scene = QGraphicsScene()59        self.graphicsView_3.setScene(self.result_scene)60        self.graphicsView_3.setAlignment(Qt.AlignTop | Qt.AlignLeft)61        self.graphicsView_3.setVerticalScrollBarPolicy(Qt.ScrollBarAlwaysOff)62        self.graphicsView_3.setHorizontalScrollBarPolicy(Qt.ScrollBarAlwaysOff)63 64        self.dlg = QColorDialog(self.graphicsView)65        self.color = None66 67        self.sample_model = SampleFromPoseModel(opt)68 69    def open_densepose(self):70        fileName, _ = QFileDialog.getOpenFileName(self, "Open File",71                                                  QDir.currentPath())72        if fileName:73            image = QPixmap(fileName)74            mat_img = Image.open(fileName)75            self.pose_img = mat_img.copy()76            if image.isNull():77                QMessageBox.information(self, "Image Viewer",78                                        "Cannot load %s." % fileName)79                return80            image = image.scaled(self.graphicsView.size(),81                                 Qt.IgnoreAspectRatio)82 83            if len(self.scene.items()) > 0:84                self.scene.removeItem(self.scene.items()[-1])85            self.scene.addPixmap(image)86 87            self.ref_scene.clear()88            self.result_scene.clear()89 90            # load pose to model91            self.pose_img = np.array(92                self.pose_img.resize(93                    size=(256, 512),94                    resample=Image.LANCZOS))[:, :, 2:].transpose(95                        2, 0, 1).astype(np.float32)96            self.pose_img = self.pose_img / 12. - 197 98            self.pose_img = torch.from_numpy(self.pose_img).unsqueeze(1)99 100            self.sample_model.feed_pose_data(self.pose_img)101 102    def generate_parsing(self):103        self.ref_scene.reset_items()104        self.ref_scene.reset()105 106        shape_texts = self.message_box_1.text()107 108        shape_attributes = generate_shape_attributes(shape_texts)109        shape_attributes = torch.LongTensor(shape_attributes).unsqueeze(0)110        self.sample_model.feed_shape_attributes(shape_attributes)111 112        self.sample_model.generate_parsing_map()113        self.sample_model.generate_quantized_segm()114 115        self.colored_segm = self.sample_model.palette_result(116            self.sample_model.segm[0].cpu())117 118        self.mask_m = cv2.cvtColor(119            cv2.cvtColor(self.colored_segm, cv2.COLOR_RGB2BGR),120            cv2.COLOR_BGR2RGB)121 122        qim = QImage(self.colored_segm.data.tobytes(),123                     self.colored_segm.shape[1], self.colored_segm.shape[0],124                     QImage.Format_RGB888)125 126        image = QPixmap.fromImage(qim)127 128        image = image.scaled(self.graphicsView.size(), Qt.IgnoreAspectRatio)129 130        if len(self.ref_scene.items()) > 0:131            self.ref_scene.removeItem(self.ref_scene.items()[-1])132        self.ref_scene.addPixmap(image)133 134        self.result_scene.clear()135 136    def generate_human(self):137        for i in range(24):138            self.mask_m = self.make_mask(self.mask_m,139                                         self.ref_scene.mask_points[i],140                                         self.ref_scene.size_points[i],141                                         color_list[i])142 143        seg_map = np.full(self.mask_m.shape[:-1], -1)144 145        # convert rgb to num146        for index, color in enumerate(color_list):147            seg_map[np.sum(self.mask_m == color, axis=2) == 3] = index148        assert (seg_map != -1).all()149 150        self.sample_model.segm = torch.from_numpy(seg_map).unsqueeze(151            0).unsqueeze(0).to(self.sample_model.device)152        self.sample_model.generate_quantized_segm()153 154        texture_texts = self.message_box_2.text()155        texture_attributes = generate_texture_attributes(texture_texts)156 157        texture_attributes = torch.LongTensor(texture_attributes)158 159        self.sample_model.feed_texture_attributes(texture_attributes)160 161        self.sample_model.generate_texture_map()162        result = self.sample_model.sample_and_refine()163        result = result.permute(0, 2, 3, 1)164        result = result.detach().cpu().numpy()165        result = result * 255166 167        result = np.asarray(result[0, :, :, :], dtype=np.uint8)168 169        self.output_img = result170 171        qim = QImage(result.data.tobytes(), result.shape[1], result.shape[0],172                     QImage.Format_RGB888)173        image = QPixmap.fromImage(qim)174 175        image = image.scaled(self.graphicsView.size(), Qt.IgnoreAspectRatio)176 177        if len(self.result_scene.items()) > 0:178            self.result_scene.removeItem(self.result_scene.items()[-1])179        self.result_scene.addPixmap(image)180 181    def top_mode(self):182        self.ref_scene.mode = 1183 184    def skin_mode(self):185        self.ref_scene.mode = 15186 187    def outer_mode(self):188        self.ref_scene.mode = 2189 190    def face_mode(self):191        self.ref_scene.mode = 14192 193    def skirt_mode(self):194        self.ref_scene.mode = 3195 196    def hair_mode(self):197        self.ref_scene.mode = 13198 199    def dress_mode(self):200        self.ref_scene.mode = 4201 202    def headwear_mode(self):203        self.ref_scene.mode = 7204 205    def pants_mode(self):206        self.ref_scene.mode = 5207 208    def eyeglass_mode(self):209        self.ref_scene.mode = 8210 211    def rompers_mode(self):212        self.ref_scene.mode = 21213 214    def footwear_mode(self):215        self.ref_scene.mode = 11216 217    def leggings_mode(self):218        self.ref_scene.mode = 6219 220    def ring_mode(self):221        self.ref_scene.mode = 16222 223    def belt_mode(self):224        self.ref_scene.mode = 10225 226    def neckwear_mode(self):227        self.ref_scene.mode = 9228 229    def wrist_mode(self):230        self.ref_scene.mode = 17231 232    def socks_mode(self):233        self.ref_scene.mode = 18234 235    def tie_mode(self):236        self.ref_scene.mode = 23237 238    def earstuds_mode(self):239        self.ref_scene.mode = 22240 241    def necklace_mode(self):242        self.ref_scene.mode = 20243 244    def bag_mode(self):245        self.ref_scene.mode = 12246 247    def glove_mode(self):248        self.ref_scene.mode = 19249 250    def background_mode(self):251        self.ref_scene.mode = 0252 253    def make_mask(self, mask, pts, sizes, color):254        if len(pts) > 0:255            for idx, pt in enumerate(pts):256                cv2.line(mask, pt['prev'], pt['curr'], color, sizes[idx])257        return mask258 259    def save_img(self):260        if type(self.output_img):261            fileName, _ = QFileDialog.getSaveFileName(self, "Save File",262                                                      QDir.currentPath())263            cv2.imwrite(fileName + '.png', self.output_img[:, :, ::-1])264 265    def undo(self):266        self.scene.undo()267 268    def clear(self):269 270        self.ref_scene.reset_items()271        self.ref_scene.reset()272 273        self.ref_scene.clear()274 275        self.result_scene.clear()276 277 278if __name__ == '__main__':279 280    app = QApplication(sys.argv)281    opt = './configs/sample_from_pose.yml'282    opt = parse(opt, is_train=False)283    opt = dict_to_nonedict(opt)284    ex = Ex(opt)285    sys.exit(app.exec_())286