Shellbrady/LivePortrait5
0
1# coding: utf-82 3import gradio as gr4import numpy as np5import os.path as osp6from typing import List, Union, Tuple7from dataclasses import dataclass, field8import cv2; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False)9 10from .landmark_runner import LandmarkRunner11from .face_analysis_diy import FaceAnalysisDIY12from .helper import prefix13from .crop import crop_image, crop_image_by_bbox, parse_bbox_from_landmark, average_bbox_lst14from .timer import Timer15from .rprint import rlog as log16from .io import load_image_rgb17from .video import VideoWriter, get_fps, change_video_fps18 19 20def make_abs_path(fn):21 return osp.join(osp.dirname(osp.realpath(__file__)), fn)22 23 24@dataclass25class Trajectory:26 start: int = -1 # 起始帧 闭区间27 end: int = -1 # 结束帧 闭区间28 lmk_lst: Union[Tuple, List, np.ndarray] = field(default_factory=list) # lmk list29 bbox_lst: Union[Tuple, List, np.ndarray] = field(default_factory=list) # bbox list30 frame_rgb_lst: Union[Tuple, List, np.ndarray] = field(default_factory=list) # frame list31 frame_rgb_crop_lst: Union[Tuple, List, np.ndarray] = field(default_factory=list) # frame crop list32 33 34class Cropper(object):35 def __init__(self, **kwargs) -> None:36 device_id = kwargs.get('device_id', 0)37 self.landmark_runner = LandmarkRunner(38 ckpt_path=make_abs_path('../../pretrained_weights/liveportrait/landmark.onnx'),39 onnx_provider='cpu',40 device_id=device_id41 )42 self.landmark_runner.warmup()43 44 self.face_analysis_wrapper = FaceAnalysisDIY(45 name='buffalo_l',46 root=make_abs_path('../../pretrained_weights/insightface'),47 providers=["CPUExecutionProvider"]48 )49 self.face_analysis_wrapper.prepare(ctx_id=device_id, det_size=(512, 512))50 self.face_analysis_wrapper.warmup()51 52 self.crop_cfg = kwargs.get('crop_cfg', None)53 54 def update_config(self, user_args):55 for k, v in user_args.items():56 if hasattr(self.crop_cfg, k):57 setattr(self.crop_cfg, k, v)58 59 def crop_single_image(self, obj, **kwargs):60 direction = kwargs.get('direction', 'large-small')61 62 # crop and align a single image63 if isinstance(obj, str):64 img_rgb = load_image_rgb(obj)65 elif isinstance(obj, np.ndarray):66 img_rgb = obj67 68 src_face = self.face_analysis_wrapper.get(69 img_rgb,70 flag_do_landmark_2d_106=True,71 direction=direction72 )73 74 if len(src_face) == 0:75 log('No face detected in the source image.')76 raise gr.Error("No face detected in the source image 💥!", duration=5)77 raise Exception("No face detected in the source image!")78 elif len(src_face) > 1:79 log(f'More than one face detected in the image, only pick one face by rule {direction}.')80 81 src_face = src_face[0]82 pts = src_face.landmark_2d_10683 84 # crop the face85 ret_dct = crop_image(86 img_rgb, # ndarray87 pts, # 106x2 or Nx288 dsize=kwargs.get('dsize', 512),89 scale=kwargs.get('scale', 2.3),90 vy_ratio=kwargs.get('vy_ratio', -0.15),91 )92 # update a 256x256 version for network input or else93 ret_dct['img_crop_256x256'] = cv2.resize(ret_dct['img_crop'], (256, 256), interpolation=cv2.INTER_AREA)94 ret_dct['pt_crop_256x256'] = ret_dct['pt_crop'] * 256 / kwargs.get('dsize', 512)95 96 recon_ret = self.landmark_runner.run(img_rgb, pts)97 lmk = recon_ret['pts']98 ret_dct['lmk_crop'] = lmk99 100 return ret_dct101 102 def get_retargeting_lmk_info(self, driving_rgb_lst):103 # TODO: implement a tracking-based version104 driving_lmk_lst = []105 for driving_image in driving_rgb_lst:106 ret_dct = self.crop_single_image(driving_image)107 driving_lmk_lst.append(ret_dct['lmk_crop'])108 return driving_lmk_lst109 110 def make_video_clip(self, driving_rgb_lst, output_path, output_fps=30, **kwargs):111 trajectory = Trajectory()112 direction = kwargs.get('direction', 'large-small')113 for idx, driving_image in enumerate(driving_rgb_lst):114 if idx == 0 or trajectory.start == -1:115 src_face = self.face_analysis_wrapper.get(116 driving_image,117 flag_do_landmark_2d_106=True,118 direction=direction119 )120 if len(src_face) == 0:121 # No face detected in the driving_image122 continue123 elif len(src_face) > 1:124 log(f'More than one face detected in the driving frame_{idx}, only pick one face by rule {direction}.')125 src_face = src_face[0]126 pts = src_face.landmark_2d_106127 lmk_203 = self.landmark_runner(driving_image, pts)['pts']128 trajectory.start, trajectory.end = idx, idx129 else:130 lmk_203 = self.face_recon_wrapper(driving_image, trajectory.lmk_lst[-1])['pts']131 trajectory.end = idx132 133 trajectory.lmk_lst.append(lmk_203)134 ret_bbox = parse_bbox_from_landmark(lmk_203, scale=self.crop_cfg.globalscale, vy_ratio=elf.crop_cfg.vy_ratio)['bbox']135 bbox = [ret_bbox[0, 0], ret_bbox[0, 1], ret_bbox[2, 0], ret_bbox[2, 1]] # 4,136 trajectory.bbox_lst.append(bbox) # bbox137 trajectory.frame_rgb_lst.append(driving_image)138 139 global_bbox = average_bbox_lst(trajectory.bbox_lst)140 for idx, (frame_rgb, lmk) in enumerate(zip(trajectory.frame_rgb_lst, trajectory.lmk_lst)):141 ret_dct = crop_image_by_bbox(142 frame_rgb, global_bbox, lmk=lmk,143 dsize=self.video_crop_cfg.dsize, flag_rot=self.video_crop_cfg.flag_rot, borderValue=self.video_crop_cfg.borderValue144 )145 frame_rgb_crop = ret_dct['img_crop']146 