expression friendly method for driving single images

This commit is contained in:
kijai
2024-08-02 19:50:24 +03:00
parent f3916f522a
commit c2bb34d4f8
3 changed files with 44 additions and 3 deletions
+11 -1
View File
@@ -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
+24 -1
View File
@@ -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
+9 -1
View File
@@ -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"])