From c2bb34d4f8f9652081c90ac61fc70574e99fe60f Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 2 Aug 2024 19:50:24 +0300 Subject: [PATCH] expression friendly method for driving single images --- liveportrait/live_portrait_pipeline.py | 12 +++++++++++- liveportrait/utils/helper.py | 25 ++++++++++++++++++++++++- nodes.py | 10 +++++++++- 3 files changed, 44 insertions(+), 3 deletions(-) diff --git a/liveportrait/live_portrait_pipeline.py b/liveportrait/live_portrait_pipeline.py index 68dccc6..202481a 100644 --- a/liveportrait/live_portrait_pipeline.py +++ b/liveportrait/live_portrait_pipeline.py @@ -14,6 +14,7 @@ from .utils.camera import get_rotation_matrix from .live_portrait_wrapper import LivePortraitWrapper from .utils.retargeting_utils import calc_eye_close_ratio, calc_lip_close_ratio from .utils.filter import smooth +from .utils.helper import calc_motion_multiplier import os script_directory = os.path.dirname(os.path.abspath(__file__)) @@ -38,7 +39,7 @@ class LivePortraitPipeline(object): ) def execute( - self, driving_images, crop_info, driving_landmarks, delta_multiplier, relative_motion_mode, driving_smooth_observation_variance, mismatch_method="constant", + self, driving_images, crop_info, driving_landmarks, delta_multiplier, relative_motion_mode, driving_smooth_observation_variance, mismatch_method="constant", expression_friendly=False, driving_multiplier=1.0, ): inference_cfg = self.live_portrait_wrapper.cfg device = inference_cfg.device_id @@ -174,6 +175,15 @@ class LivePortraitPipeline(object): delta_new = delta_new * delta_multiplier x_d_i_new = scale_new * (x_c_s @ R_new + delta_new) + t_new + + if expression_friendly: + if i == 0: + x_d_0_new = x_d_i_new + motion_multiplier = calc_motion_multiplier(x_s, x_d_0_new) + motion_multiplier *= driving_multiplier + x_d_diff = (x_d_i_new - x_d_0_new) * motion_multiplier + x_d_i_new = x_d_diff + x_s + if ( not inference_cfg.flag_stitching and not inference_cfg.flag_eye_retargeting diff --git a/liveportrait/utils/helper.py b/liveportrait/utils/helper.py index 98b2cc5..fd46387 100644 --- a/liveportrait/utils/helper.py +++ b/liveportrait/utils/helper.py @@ -4,10 +4,12 @@ utility functions and classes to handle feature extraction and model loading """ -import os.path as osp import cv2 import torch +import numpy as np +from typing import Union from collections import OrderedDict +from scipy.spatial import ConvexHull # pylint: disable=E0401,E0611 def squeeze_tensor_to_numpy(tensor): out = tensor.data.squeeze(0).cpu().numpy() @@ -71,3 +73,24 @@ def resize_to_limit(img, max_dim=1280, n=2): if new_h != img.shape[0] or new_w != img.shape[1]: img = img[:new_h, :new_w] return img + +def tensor_to_numpy(data: Union[np.ndarray, torch.Tensor]) -> np.ndarray: + """transform torch.Tensor into numpy.ndarray""" + if isinstance(data, torch.Tensor): + return data.data.cpu().numpy() + return data + +def calc_motion_multiplier( + kp_source: Union[np.ndarray, torch.Tensor], + kp_driving_initial: Union[np.ndarray, torch.Tensor] +) -> float: + """calculate motion_multiplier based on the source image and the first driving frame""" + kp_source_np = tensor_to_numpy(kp_source) + kp_driving_initial_np = tensor_to_numpy(kp_driving_initial) + + source_area = ConvexHull(kp_source_np.squeeze(0)).volume + driving_area = ConvexHull(kp_driving_initial_np.squeeze(0)).volume + motion_multiplier = np.sqrt(source_area) / np.sqrt(driving_area) + # motion_multiplier = np.cbrt(source_area) / np.cbrt(driving_area) + + return motion_multiplier \ No newline at end of file diff --git a/nodes.py b/nodes.py index 645998f..cd494f2 100644 --- a/nodes.py +++ b/nodes.py @@ -280,6 +280,8 @@ class LivePortraitProcess: "optional": { "opt_retargeting_info": ("RETARGETINGINFO", {"default": None}), + "expression_friendly": ("BOOLEAN", {"default": False}), + "expression_friendly_multiplier": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 100.0, "step": 0.001}), } } @@ -308,9 +310,13 @@ class LivePortraitProcess: delta_multiplier: float = 1.0, mismatch_method: str = "constant", opt_retargeting_info: dict = None, + expression_friendly: bool = False, + expression_friendly_multiplier: float = 1.0, ): if driving_images.shape[0] < source_image.shape[0]: raise ValueError("The number of driving images should be larger than the number of source images.") + if expression_friendly and source_image.shape[0] > 1: + raise ValueError("expression_friendly works only with single source image") if opt_retargeting_info is not None: pipeline.live_portrait_wrapper.cfg.flag_eye_retargeting = opt_retargeting_info["eye_retargeting"] @@ -347,7 +353,9 @@ class LivePortraitProcess: delta_multiplier, relative_motion_mode, driving_smooth_observation_variance, - mismatch_method + mismatch_method, + expression_friendly=expression_friendly, + driving_multiplier=expression_friendly_multiplier, ) total_frames = len(out["out_list"])