Team Ai
Apppublic

Shellbrady/LivePortrait5

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
cropper.py146 linesDownload Raw Back to utils
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