radames/Text2Human-API
1
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 