iti/HandMesh
0
1import os2import torch3from utils.vis import cnt_area4import numpy as np5import cv26from utils.vis import registration, map2uv, inv_base_tranmsform, base_transform, tensor2array7from utils.draw3d import save_a_image_with_mesh_joints, draw_2d_skeleton, draw_3d_skeleton8from utils.read import save_mesh9import json10from utils import utils, writer11from datasets.FreiHAND.kinematics import mano_to_mpii12from utils.progress.bar import Bar13from termcolor import colored, cprint14import pickle15import time16from PIL import Image17from utils.transforms import rigid_align18import gradio as gr19from options.base_options import BaseOptions20import os.path as osp21from mobrecon.mobrecon_densestack import MobRecon22from utils.read import spiral_tramsform23 24 25class Runner(object):26 def __init__(self, args, model, faces, device):27 super(Runner, self).__init__()28 self.args = args29 self.model = model30 self.faces = faces31 self.device = device32 self.face = torch.from_numpy(self.faces[0].astype(np.int64)).to(self.device)33 34 def set_demo(self, args):35 with open(os.path.join(args.work_dir, 'template', 'MANO_RIGHT.pkl'), 'rb') as f:36 mano = pickle.load(f, encoding='latin1')37 self.j_regressor = np.zeros([21, 778])38 self.j_regressor[:16] = mano['J_regressor'].toarray()39 for k, v in {16: 333, 17: 444, 18: 672, 19: 555, 20: 744}.items():40 self.j_regressor[k, v] = 141 self.std = torch.tensor(0.20)42 43 def poseEstimator(self, image):44 args = self.args45 self.model.eval()46 image_fp = os.path.join(args.work_dir, 'images')47 image_files = [os.path.join(image_fp, i) for i in os.listdir(image_fp) if '_img.jpg' in i]48 49 with torch.no_grad():50 image = Image.fromarray(np.uint8(image)) #.convert('RGB')51 image.resize(size=(args.size, args.size))52 image = np.array(image)53 input = torch.from_numpy(base_transform(image, size=args.size)).unsqueeze(0).to(self.device) # A tensor with shape (1, 224, 224, 3)54 K = np.array([[500, 0, 128], [0, 500, 128], [0, 0, 1]])55 56 K[0, 0] = K[0, 0] / 224 * args.size57 K[1, 1] = K[1, 1] / 224 * args.size58 K[0, 2] = args.size // 259 K[1, 2] = args.size // 260 61 out = self.model(input)62 63 # silhouette64 mask_pred = out.get('mask_pred')65 if mask_pred is not None:66 mask_pred = (mask_pred[0] > 0.3).cpu().numpy().astype(np.uint8)67 mask_pred = cv2.resize(mask_pred, (input.size(3), input.size(2)))68 try:69 contours, _ = cv2.findContours(mask_pred, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)70 contours.sort(key=cnt_area, reverse=True)71 poly = contours[0].transpose(1, 0, 2).astype(np.int32)72 except:73 poly = None74 else:75 mask_pred = np.zeros([input.size(3), input.size(2)])76 poly = None77 78 79 # vertex80 pred = out['mesh_pred'][0] if isinstance(out['mesh_pred'], list) else out['mesh_pred']81 82 vertex = (pred[0].cpu() * self.std.cpu()).numpy() # Shape (778, 3)83 84 uv_pred = out['uv_pred']85 if uv_pred.ndim == 4:86 uv_point_pred, uv_pred_conf = map2uv(uv_pred.cpu().numpy(), (input.size(2), input.size(3)))87 else:88 uv_point_pred, uv_pred_conf = (uv_pred * args.size).cpu().numpy(), [None,]89 vertex, align_state = registration(vertex, uv_point_pred[0], self.j_regressor, K, args.size, uv_conf=uv_pred_conf[0], poly=poly)90 91 vertex2xyz = mano_to_mpii(np.matmul(self.j_regressor, vertex))92 skeleton_overlay = draw_2d_skeleton(image[..., ::-1], uv_point_pred[0])93 frame = skeleton_overlay[..., ::-1]94 return frame95 96# get config97args = BaseOptions().parse()98 99# dir prepare100args.work_dir = osp.dirname(osp.realpath(__file__))101data_fp = osp.join(args.work_dir, 'data', args.dataset)102args.out_dir = osp.join(args.work_dir, 'out', args.dataset, args.exp_name)103args.checkpoints_dir = osp.join(args.out_dir, 'checkpoints')104utils.makedirs(osp.join(args.out_dir, args.phase))105utils.makedirs(args.out_dir)106utils.makedirs(args.checkpoints_dir)107 108template_fp = osp.join(args.work_dir, 'template', 'template.ply')109transform_fp = osp.join(args.work_dir, 'template', 'transform.pkl')110spiral_indices_list, down_transform_list, up_transform_list, tmp = spiral_tramsform(transform_fp, template_fp, args.ds_factors, args.seq_length, args.dilation)111 112for i in range(len(up_transform_list)):113 up_transform_list[i] = (*up_transform_list[i]._indices(), up_transform_list[i]._values())114 model = MobRecon(args, spiral_indices_list, up_transform_list)115 116device = torch.device('cpu')117torch.set_num_threads(args.n_threads)118runner = Runner(args, model, tmp['face'], device)119runner.set_demo(args)120 121iface = gr.Interface(runner.poseEstimator, gr.inputs.Image(shape=(224, 224)), "image")122iface.launch()