too much diff due to formatting
This commit is contained in:
Mel Massadian
2024-07-08 16:48:15 +02:00
parent c87ecc6e04
commit 72bb6910e9
2 changed files with 366 additions and 267 deletions
+162 -149
View File
@@ -5,183 +5,196 @@ Pipeline of LivePortrait
"""
import cv2
import numpy as np
import os.path as osp
from rich.progress import track
import comfy.utils
import os.path as osp
import numpy as np
from .config.inference_config import InferenceConfig
#from .utils.cropper import Cropper
# from .utils.cropper import Cropper
from .utils.camera import get_rotation_matrix
#from .utils.video import images2video, concat_frames
from .utils.crop import _transform_img
#from .utils.retargeting_utils import calc_lip_close_ratio
#from .utils.io import load_image_rgb, load_driving_info
#from .utils.helper import mkdir, basename, dct2cuda, is_video, is_template, resize_to_limit
# from .utils.video import images2video, concat_frames
# from .utils.retargeting_utils import calc_lip_close_ratio
# from .utils.io import load_image_rgb, load_driving_info
# from .utils.helper import mkdir, basename, dct2cuda, is_video, is_template, resize_to_limit
from .utils.helper import resize_to_limit
#from .utils.rprint import rlog as log
# from .utils.rprint import rlog as log
from .live_portrait_wrapper import LivePortraitWrapper
import comfy.utils
def make_abs_path(fn):
return osp.join(osp.dirname(osp.realpath(__file__)), fn)
class LivePortraitPipeline(object):
def __init__(self, appearance_feature_extractor, motion_extractor, warping_module,
spade_generator, stitching_retargeting_module, inference_cfg: InferenceConfig):
def __init__(
self,
appearance_feature_extractor,
motion_extractor,
warping_module,
spade_generator,
stitching_retargeting_module,
inference_cfg: InferenceConfig,
):
self.live_portrait_wrapper: LivePortraitWrapper = LivePortraitWrapper(
appearance_feature_extractor, motion_extractor, warping_module,
spade_generator, stitching_retargeting_module, cfg=inference_cfg)
appearance_feature_extractor,
motion_extractor,
warping_module,
spade_generator,
stitching_retargeting_module,
cfg=inference_cfg,
)
def execute(self, img_rgb, driving_images_np):
inference_cfg = self.live_portrait_wrapper.cfg # for convenience
######## process reference portrait ########
#img_rgb = load_image_rgb(args.source_image)
img_rgb = resize_to_limit(img_rgb, inference_cfg.ref_max_shape, inference_cfg.ref_shape_n)
#log(f"Load source image from {args.source_image}")
crop_info = self.cropper.crop_single_image(img_rgb)
source_lmk = crop_info['lmk_crop']
img_crop, img_crop_256x256 = crop_info['img_crop'], crop_info['img_crop_256x256']
if inference_cfg.flag_do_crop:
I_s = self.live_portrait_wrapper.prepare_source(img_crop_256x256)
else:
I_s = self.live_portrait_wrapper.prepare_source(img_rgb)
x_s_info = self.live_portrait_wrapper.get_kp_info(I_s)
x_c_s = x_s_info['kp']
R_s = get_rotation_matrix(x_s_info['pitch'], x_s_info['yaw'], x_s_info['roll'])
f_s = self.live_portrait_wrapper.extract_feature_3d(I_s)
x_s = self.live_portrait_wrapper.transform_keypoint(x_s_info)
def _get_source_frame(self, source_np, idx, total_frames, method):
if source_np.shape[0] == 1:
return source_np[0]
if inference_cfg.flag_lip_zero:
# let lip-open scalar to be 0 at first
c_d_lip_before_animation = [0.]
combined_lip_ratio_tensor_before_animation = self.live_portrait_wrapper.calc_combined_lip_ratio(c_d_lip_before_animation, source_lmk)
if combined_lip_ratio_tensor_before_animation[0][0] < inference_cfg.lip_zero_threshold:
inference_cfg.flag_lip_zero = False
else:
lip_delta_before_animation = self.live_portrait_wrapper.retarget_lip(x_s, combined_lip_ratio_tensor_before_animation)
############################################
if method == "repeat":
return source_np[min(idx, source_np.shape[0] - 1)]
elif method == "cycle":
return source_np[idx % source_np.shape[0]]
elif method == "mirror":
cycle_length = 2 * source_np.shape[0] - 2
mirror_idx = idx % cycle_length
if mirror_idx >= source_np.shape[0]:
mirror_idx = cycle_length - mirror_idx
return source_np[mirror_idx]
elif method == "nearest":
ratio = idx / (total_frames - 1)
return source_np[
min(int(ratio * (source_np.shape[0] - 1)), source_np.shape[0] - 1)
]
######## process driving info ########
#if is_video(args.driving_info):
#log(f"Load from video file (mp4 mov avi etc...): {args.driving_info}")
# TODO: 这里track一下驱动视频 -> 构建模板
#driving_rgb_lst = load_driving_info(args.driving_info)
driving_rgb_lst = driving_images_np
driving_rgb_lst_256 = [cv2.resize(_, (256, 256)) for _ in driving_rgb_lst]
I_d_lst = self.live_portrait_wrapper.prepare_driving_videos(driving_rgb_lst_256)
n_frames = I_d_lst.shape[0]
if inference_cfg.flag_eye_retargeting or inference_cfg.flag_lip_retargeting:
driving_lmk_lst = self.cropper.get_retargeting_lmk_info(driving_rgb_lst)
input_eye_ratio_lst, input_lip_ratio_lst = self.live_portrait_wrapper.calc_retargeting_ratio(source_lmk, driving_lmk_lst)
# elif is_template(args.driving_info):
# log(f"Load from video templates {args.driving_info}")
# with open(args.driving_info, 'rb') as f:
# template_lst, driving_lmk_lst = pickle.load(f)
# n_frames = template_lst[0]['n_frames']
# input_eye_ratio_lst, input_lip_ratio_lst = self.live_portrait_wrapper.calc_retargeting_ratio(source_lmk, driving_lmk_lst)
# else:
# raise Exception("Unsupported driving types!")
#########################################
######## prepare for pasteback ########
if inference_cfg.flag_pasteback:
if inference_cfg.mask_crop is None:
inference_cfg.mask_crop = cv2.imread(make_abs_path('./utils/resources/mask_template.png'), cv2.IMREAD_COLOR)
mask_ori = _transform_img(inference_cfg.mask_crop, crop_info['M_c2o'], dsize=(img_rgb.shape[1], img_rgb.shape[0]))
mask_ori = mask_ori.astype(np.float32) / 255.
I_p_paste_lst = []
#########################################
def execute(self, source_np, driving_images_np, mismatch_method="repeat"):
inference_cfg = self.live_portrait_wrapper.cfg
is_video = source_np.shape[0] > 1
I_p_lst = []
R_d_0, x_d_0_info = None, None
pbar = comfy.utils.ProgressBar(n_frames)
for i in track(range(n_frames), description='Animating...', total=n_frames):
#if is_video(args.driving_info):
# extract kp info by M
I_d_i = I_d_lst[i]
x_d_i_info = self.live_portrait_wrapper.get_kp_info(I_d_i)
R_d_i = get_rotation_matrix(x_d_i_info['pitch'], x_d_i_info['yaw'], x_d_i_info['roll'])
# else:
# # from template
# x_d_i_info = template_lst[i]
# x_d_i_info = dct2cuda(x_d_i_info, inference_cfg.device_id)
# R_d_i = x_d_i_info['R_d']
I_p_paste_lst = []
if i == 0:
R_d_0 = R_d_i
x_d_0_info = x_d_i_info
total_frames = driving_images_np.shape[0]
pbar = comfy.utils.ProgressBar(total_frames)
for i in range(total_frames):
source_frame = self._get_source_frame(
source_np, i, total_frames, mismatch_method
)
driving_frame = driving_images_np[i]
source_frame_rgb = resize_to_limit(
source_frame, inference_cfg.ref_max_shape, inference_cfg.ref_shape_n
)
crop_info = self.cropper.crop_single_image(source_frame_rgb)
source_lmk = crop_info["lmk_crop"]
img_crop, img_crop_256x256 = (
crop_info["img_crop"],
crop_info["img_crop_256x256"],
)
if inference_cfg.flag_do_crop:
I_s = self.live_portrait_wrapper.prepare_source(img_crop_256x256)
else:
I_s = self.live_portrait_wrapper.prepare_source(source_frame_rgb)
x_s_info = self.live_portrait_wrapper.get_kp_info(I_s)
x_c_s = x_s_info["kp"]
R_s = get_rotation_matrix(
x_s_info["pitch"], x_s_info["yaw"], x_s_info["roll"]
)
f_s = self.live_portrait_wrapper.extract_feature_3d(I_s)
x_s = self.live_portrait_wrapper.transform_keypoint(x_s_info)
if inference_cfg.flag_lip_zero:
c_d_lip_before_animation = [0.0]
combined_lip_ratio_tensor_before_animation = (
self.live_portrait_wrapper.calc_combined_lip_ratio(
c_d_lip_before_animation, source_lmk
)
)
# TODO: expose lip_zero_threshold
if (
combined_lip_ratio_tensor_before_animation[0][0]
< inference_cfg.lip_zero_threshold
):
inference_cfg.flag_lip_zero = False
else:
lip_delta_before_animation = (
self.live_portrait_wrapper.retarget_lip(
x_s, combined_lip_ratio_tensor_before_animation
)
)
# driving_frame_rgb = cv2.cvtColor(driving_frame, cv2.COLOR_BGR2RGB)
driving_frame_256 = cv2.resize(driving_frame, (256, 256))
I_d = self.live_portrait_wrapper.prepare_driving_videos(
[driving_frame_256]
)[0]
if inference_cfg.flag_eye_retargeting or inference_cfg.flag_lip_retargeting:
driving_lmk_lst = self.cropper.get_retargeting_lmk_info([driving_frame])
input_eye_ratio_lst, input_lip_ratio_lst = (
self.live_portrait_wrapper.calc_retargeting_ratio(
source_lmk, driving_lmk_lst
)
)
x_d_info = self.live_portrait_wrapper.get_kp_info(I_d)
R_d = get_rotation_matrix(
x_d_info["pitch"], x_d_info["yaw"], x_d_info["roll"]
)
if inference_cfg.flag_relative:
R_new = (R_d_i @ R_d_0.permute(0, 2, 1)) @ R_s
delta_new = x_s_info['exp'] + (x_d_i_info['exp'] - x_d_0_info['exp'])
scale_new = x_s_info['scale'] * (x_d_i_info['scale'] / x_d_0_info['scale'])
t_new = x_s_info['t'] + (x_d_i_info['t'] - x_d_0_info['t'])
R_new = R_d @ R_s
delta_new = x_s_info["exp"] + (x_d_info["exp"] - x_s_info["exp"])
scale_new = x_s_info["scale"] * (x_d_info["scale"] / x_s_info["scale"])
t_new = x_s_info["t"] + (x_d_info["t"] - x_s_info["t"])
else:
R_new = R_d_i
delta_new = x_d_i_info['exp']
scale_new = x_s_info['scale']
t_new = x_d_i_info['t']
R_new = R_d
delta_new = x_d_info["exp"]
scale_new = x_s_info["scale"]
t_new = x_d_info["t"]
t_new[..., 2].fill_(0) # zero tz
x_d_i_new = scale_new * (x_c_s @ R_new + delta_new) + t_new
t_new[..., 2].fill_(0) # zero tz
x_d_new = scale_new * (x_c_s @ R_new + delta_new) + t_new
# Algorithm 1:
if not inference_cfg.flag_stitching and not inference_cfg.flag_eye_retargeting and not inference_cfg.flag_lip_retargeting:
# without stitching or retargeting
if inference_cfg.flag_lip_zero:
x_d_i_new += lip_delta_before_animation.reshape(-1, x_s.shape[1], 3)
else:
pass
elif inference_cfg.flag_stitching and not inference_cfg.flag_eye_retargeting and not inference_cfg.flag_lip_retargeting:
# with stitching and without retargeting
if inference_cfg.flag_lip_zero:
x_d_i_new = self.live_portrait_wrapper.stitching(x_s, x_d_i_new) + lip_delta_before_animation.reshape(-1, x_s.shape[1], 3)
else:
x_d_i_new = self.live_portrait_wrapper.stitching(x_s, x_d_i_new)
else:
eyes_delta, lip_delta = None, None
if inference_cfg.flag_eye_retargeting:
c_d_eyes_i = input_eye_ratio_lst[i]
combined_eye_ratio_tensor = self.live_portrait_wrapper.calc_combined_eye_ratio(c_d_eyes_i, source_lmk)
combined_eye_ratio_tensor = combined_eye_ratio_tensor * inference_cfg.eyes_retargeting_multiplier
# ∆_eyes,i = R_eyes(x_s; c_s,eyes, c_d,eyes,i)
eyes_delta = self.live_portrait_wrapper.retarget_eye(x_s, combined_eye_ratio_tensor)
if inference_cfg.flag_lip_retargeting:
c_d_lip_i = input_lip_ratio_lst[i]
combined_lip_ratio_tensor = self.live_portrait_wrapper.calc_combined_lip_ratio(c_d_lip_i, source_lmk)
combined_lip_ratio_tensor = combined_lip_ratio_tensor * inference_cfg.lip_retargeting_multiplier
# ∆_lip,i = R_lip(x_s; c_s,lip, c_d,lip,i)
lip_delta = self.live_portrait_wrapper.retarget_lip(x_s, combined_lip_ratio_tensor)
if inference_cfg.flag_stitching:
x_d_new = self.live_portrait_wrapper.stitching(x_s, x_d_new)
if inference_cfg.flag_relative: # use x_s
x_d_i_new = x_s + \
(eyes_delta.reshape(-1, x_s.shape[1], 3) if eyes_delta is not None else 0) + \
(lip_delta.reshape(-1, x_s.shape[1], 3) if lip_delta is not None else 0)
else: # use x_d,i
x_d_i_new = x_d_i_new + \
(eyes_delta.reshape(-1, x_s.shape[1], 3) if eyes_delta is not None else 0) + \
(lip_delta.reshape(-1, x_s.shape[1], 3) if lip_delta is not None else 0)
if inference_cfg.flag_stitching:
x_d_i_new = self.live_portrait_wrapper.stitching(x_s, x_d_i_new)
out = self.live_portrait_wrapper.warp_decode(f_s, x_s, x_d_i_new)
I_p_i = self.live_portrait_wrapper.parse_output(out['out'])[0]
out = self.live_portrait_wrapper.warp_decode(f_s, x_s, x_d_new)
I_p_i = self.live_portrait_wrapper.parse_output(out["out"])[0]
I_p_lst.append(I_p_i)
# Transform and blend
I_p_i_to_ori = _transform_img(
I_p_i,
crop_info["M_c2o"],
dsize=(source_frame_rgb.shape[1], source_frame_rgb.shape[0]),
)
if inference_cfg.flag_pasteback:
if inference_cfg.mask_crop is None:
inference_cfg.mask_crop = cv2.imread(
make_abs_path("./utils/resources/mask_template.png"),
cv2.IMREAD_COLOR,
)
mask_ori = _transform_img(
inference_cfg.mask_crop,
crop_info["M_c2o"],
dsize=(source_frame_rgb.shape[1], source_frame_rgb.shape[0]),
)
mask_ori = mask_ori.astype(np.float32) / 255.0
I_p_i_to_ori_blend = np.clip(
mask_ori * I_p_i_to_ori + (1 - mask_ori) * source_frame_rgb, 0, 255
).astype(np.uint8)
else:
I_p_i_to_ori_blend = I_p_i_to_ori
I_p_paste_lst.append(I_p_i_to_ori_blend)
pbar.update(1)
#if inference_cfg.flag_pasteback:
I_p_i_to_ori = _transform_img(I_p_i, crop_info['M_c2o'], dsize=(img_rgb.shape[1], img_rgb.shape[0]))
I_p_i_to_ori_blend = np.clip(mask_ori * I_p_i_to_ori + (1 - mask_ori) * img_rgb, 0, 255).astype(np.uint8)
out = np.hstack([I_p_i_to_ori, I_p_i_to_ori_blend])
I_p_paste_lst.append(I_p_i_to_ori_blend)
return I_p_lst, I_p_paste_lst
+204 -118
View File
@@ -13,28 +13,35 @@ from .liveportrait.utils.cropper import Cropper
from .liveportrait.modules.spade_generator import SPADEDecoder
from .liveportrait.modules.warping_network import WarpingNetwork
from .liveportrait.modules.motion_extractor import MotionExtractor
from .liveportrait.modules.appearance_feature_extractor import AppearanceFeatureExtractor
from .liveportrait.modules.stitching_retargeting_network import StitchingRetargetingNetwork
from .liveportrait.modules.appearance_feature_extractor import (
AppearanceFeatureExtractor,
)
from .liveportrait.modules.stitching_retargeting_network import (
StitchingRetargetingNetwork,
)
class InferenceConfig:
def __init__(self,
mask_crop = None,
flag_use_half_precision=True,
flag_lip_zero=True,
lip_zero_threshold=0.03,
flag_eye_retargeting=False,
flag_lip_retargeting=False,
flag_stitching=True,
flag_relative=True,
anchor_frame=0,
input_shape=(256, 256),
flag_write_result=True,
flag_pasteback=True,
ref_max_shape=1280,
ref_shape_n=2,
device_id=0,
flag_do_crop=True,
flag_do_rot=True):
def __init__(
self,
mask_crop=None,
flag_use_half_precision=True,
flag_lip_zero=True,
lip_zero_threshold=0.03,
flag_eye_retargeting=False,
flag_lip_retargeting=False,
flag_stitching=True,
flag_relative=True,
anchor_frame=0,
input_shape=(256, 256),
flag_write_result=True,
flag_pasteback=True,
ref_max_shape=1280,
ref_shape_n=2,
device_id=0,
flag_do_crop=True,
flag_do_rot=True,
):
self.flag_use_half_precision = flag_use_half_precision
self.flag_lip_zero = flag_lip_zero
self.lip_zero_threshold = lip_zero_threshold
@@ -51,7 +58,8 @@ class InferenceConfig:
self.device_id = device_id
self.flag_do_crop = flag_do_crop
self.flag_do_rot = flag_do_rot
self.mask_crop=mask_crop
self.mask_crop = mask_crop
class CropConfig:
def __init__(self, dsize=512, scale=2.3, vx_ratio=0, vy_ratio=-0.125):
@@ -60,22 +68,24 @@ class CropConfig:
self.vx_ratio = vx_ratio
self.vy_ratio = vy_ratio
class ArgumentConfig:
def __init__(self,
device_id=0,
flag_lip_zero=True,
flag_eye_retargeting=False,
flag_lip_retargeting=False,
flag_stitching=True,
flag_relative=True,
flag_pasteback=True,
flag_do_crop=True,
flag_do_rot=True,
dsize=512,
scale=2.3,
vx_ratio=0,
vy_ratio=-0.125,
):
def __init__(
self,
device_id=0,
flag_lip_zero=True,
flag_eye_retargeting=False,
flag_lip_retargeting=False,
flag_stitching=True,
flag_relative=True,
flag_pasteback=True,
flag_do_crop=True,
flag_do_rot=True,
dsize=512,
scale=2.3,
vx_ratio=0,
vy_ratio=-0.125,
):
self.device_id = device_id
self.flag_lip_zero = flag_lip_zero
self.flag_eye_retargeting = flag_eye_retargeting
@@ -90,11 +100,12 @@ class ArgumentConfig:
self.vx_ratio = vx_ratio
self.vy_ratio = vy_ratio
class DownloadAndLoadLivePortraitModels:
@classmethod
def INPUT_TYPES(s):
return {"required": {
},
return {
"required": {},
}
RETURN_TYPES = ("LIVEPORTRAITPIPE",)
@@ -114,84 +125,109 @@ class DownloadAndLoadLivePortraitModels:
if not os.path.exists(model_path):
print(f"Downloading model to: {model_path}")
from huggingface_hub import snapshot_download
snapshot_download(repo_id="Kijai/LivePortrait_safetensors",
local_dir=download_path,
local_dir_use_symlinks=False)
model_config_path = os.path.join(script_directory, 'liveportrait', 'config', 'models.yaml')
with open(model_config_path, 'r') as file:
snapshot_download(
repo_id="Kijai/LivePortrait_safetensors",
local_dir=download_path,
local_dir_use_symlinks=False,
)
model_config_path = os.path.join(
script_directory, "liveportrait", "config", "models.yaml"
)
with open(model_config_path, "r") as file:
model_config = yaml.safe_load(file)
feature_extractor_path = os.path.join(model_path, 'appearance_feature_extractor.safetensors')
motion_extractor_path = os.path.join(model_path, 'motion_extractor.safetensors')
warping_module_path = os.path.join(model_path, 'warping_module.safetensors')
spade_generator_path = os.path.join(model_path, 'spade_generator.safetensors')
stitching_retargeting_path = os.path.join(model_path, 'stitching_retargeting_module.safetensors')
feature_extractor_path = os.path.join(
model_path, "appearance_feature_extractor.safetensors"
)
motion_extractor_path = os.path.join(model_path, "motion_extractor.safetensors")
warping_module_path = os.path.join(model_path, "warping_module.safetensors")
spade_generator_path = os.path.join(model_path, "spade_generator.safetensors")
stitching_retargeting_path = os.path.join(
model_path, "stitching_retargeting_module.safetensors"
)
# init F
model_params = model_config['model_params']['appearance_feature_extractor_params']
self.appearance_feature_extractor = AppearanceFeatureExtractor(**model_params).to(device)
self.appearance_feature_extractor.load_state_dict(comfy.utils.load_torch_file(feature_extractor_path))
model_params = model_config["model_params"][
"appearance_feature_extractor_params"
]
self.appearance_feature_extractor = AppearanceFeatureExtractor(
**model_params
).to(device)
self.appearance_feature_extractor.load_state_dict(
comfy.utils.load_torch_file(feature_extractor_path)
)
self.appearance_feature_extractor.eval()
print('Load appearance_feature_extractor done.')
print("Load appearance_feature_extractor done.")
pbar.update(1)
# init M
model_params = model_config['model_params']['motion_extractor_params']
model_params = model_config["model_params"]["motion_extractor_params"]
self.motion_extractor = MotionExtractor(**model_params).to(device)
self.motion_extractor.load_state_dict(comfy.utils.load_torch_file(motion_extractor_path))
self.motion_extractor.load_state_dict(
comfy.utils.load_torch_file(motion_extractor_path)
)
self.motion_extractor.eval()
print('Load motion_extractor done.')
print("Load motion_extractor done.")
pbar.update(1)
# init W
model_params = model_config['model_params']['warping_module_params']
model_params = model_config["model_params"]["warping_module_params"]
self.warping_module = WarpingNetwork(**model_params).to(device)
self.warping_module.load_state_dict(comfy.utils.load_torch_file(warping_module_path))
self.warping_module.load_state_dict(
comfy.utils.load_torch_file(warping_module_path)
)
self.warping_module.eval()
print('Load warping_module done.')
print("Load warping_module done.")
pbar.update(1)
# init G
model_params = model_config['model_params']['spade_generator_params']
model_params = model_config["model_params"]["spade_generator_params"]
self.spade_generator = SPADEDecoder(**model_params).to(device)
self.spade_generator.load_state_dict(comfy.utils.load_torch_file(spade_generator_path))
self.spade_generator.load_state_dict(
comfy.utils.load_torch_file(spade_generator_path)
)
self.spade_generator.eval()
print('Load spade_generator done.')
print("Load spade_generator done.")
pbar.update(1)
def filter_checkpoint_for_model(checkpoint, prefix):
"""Filter and adjust the checkpoint dictionary for a specific model based on the prefix."""
# Create a new dictionary where keys are adjusted by removing the prefix and the model name
filtered_checkpoint = {key.replace(prefix + "_module.", ""): value for key, value in checkpoint.items() if key.startswith(prefix)}
filtered_checkpoint = {
key.replace(prefix + "_module.", ""): value
for key, value in checkpoint.items()
if key.startswith(prefix)
}
return filtered_checkpoint
config = model_config['model_params']['stitching_retargeting_module_params']
config = model_config["model_params"]["stitching_retargeting_module_params"]
checkpoint = comfy.utils.load_torch_file(stitching_retargeting_path)
stitcher_prefix = 'retarget_shoulder'
stitcher_prefix = "retarget_shoulder"
stitcher_checkpoint = filter_checkpoint_for_model(checkpoint, stitcher_prefix)
stitcher = StitchingRetargetingNetwork(**config.get('stitching'))
stitcher = StitchingRetargetingNetwork(**config.get("stitching"))
stitcher.load_state_dict(stitcher_checkpoint)
stitcher = stitcher.to(device)
stitcher.eval()
lip_prefix = 'retarget_mouth'
lip_prefix = "retarget_mouth"
lip_checkpoint = filter_checkpoint_for_model(checkpoint, lip_prefix)
retargetor_lip = StitchingRetargetingNetwork(**config.get('lip'))
retargetor_lip = StitchingRetargetingNetwork(**config.get("lip"))
retargetor_lip.load_state_dict(lip_checkpoint)
retargetor_lip = retargetor_lip.to(device)
retargetor_lip.eval()
eye_prefix = 'retarget_eye'
eye_prefix = "retarget_eye"
eye_checkpoint = filter_checkpoint_for_model(checkpoint, eye_prefix)
retargetor_eye = StitchingRetargetingNetwork(**config.get('eye'))
retargetor_eye = StitchingRetargetingNetwork(**config.get("eye"))
retargetor_eye.load_state_dict(eye_checkpoint)
retargetor_eye = retargetor_eye.to(device)
retargetor_eye.eval()
print('Load stitching_retargeting_module done.')
print("Load stitching_retargeting_module done.")
self.stich_retargeting_module = {
'stitching': stitcher,
'lip': retargetor_lip,
'eye': retargetor_eye
"stitching": stitcher,
"lip": retargetor_lip,
"eye": retargetor_eye,
}
pipeline = LivePortraitPipeline(
@@ -200,79 +236,128 @@ class DownloadAndLoadLivePortraitModels:
self.warping_module,
self.spade_generator,
self.stich_retargeting_module,
InferenceConfig()
InferenceConfig(),
)
return (pipeline,)
# OUR CURRENT NODE
class LivePortraitProcess:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"pipeline": ("LIVEPORTRAITPIPE",),
"source_image": ("IMAGE",),
"driving_images": ("IMAGE",),
"dsize": ("INT", {"default": 512, "min": 64, "max": 2048}),
"scale": ("FLOAT", {"default": 2.3, "min": 1.0, "max": 4.0, "step": 0.01}),
"vx_ratio": ("FLOAT", {"default": 0.0, "min": -1.0, "max": 1.0, "step": 0.01}),
"vy_ratio": ("FLOAT", {"default": -0.125, "min": -1.0, "max": 1.0, "step": 0.01}),
"lip_zero": ("BOOLEAN", {"default": True}),
"eye_retargeting": ("BOOLEAN", {"default": False}),
"eyes_retargeting_multiplier": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.001}),
"lip_retargeting": ("BOOLEAN", {"default": False}),
"lip_retargeting_multiplier": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.001}),
"stitching": ("BOOLEAN", {"default": True}),
"relative": ("BOOLEAN", {"default": True}),
return {
"required": {
"pipeline": ("LIVEPORTRAITPIPE",),
"source_image": ("IMAGE",),
"driving_images": ("IMAGE",),
"dsize": ("INT", {"default": 512, "min": 64, "max": 2048}),
"scale": (
"FLOAT",
{"default": 2.3, "min": 1.0, "max": 4.0, "step": 0.01},
),
"vx_ratio": (
"FLOAT",
{"default": 0.0, "min": -1.0, "max": 1.0, "step": 0.01},
),
"vy_ratio": (
"FLOAT",
{"default": -0.125, "min": -1.0, "max": 1.0, "step": 0.01},
),
"lip_zero": ("BOOLEAN", {"default": True}),
"eye_retargeting": ("BOOLEAN", {"default": False}),
"eyes_retargeting_multiplier": (
"FLOAT",
{"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.001},
),
"lip_retargeting": ("BOOLEAN", {"default": False}),
"lip_retargeting_multiplier": (
"FLOAT",
{"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.001},
),
"stitching": ("BOOLEAN", {"default": True}),
"relative": ("BOOLEAN", {"default": False}),
},
"optional": {
"mismatch_method": (
["repeat", "cycle", "mirror", "nearest"],
{"default": "repeat"},
)
},
}
RETURN_TYPES = ("IMAGE", "IMAGE",)
RETURN_NAMES = ("cropped_images", "full_images",)
RETURN_TYPES = (
"IMAGE",
"IMAGE",
)
RETURN_NAMES = (
"cropped_images",
"full_images",
)
FUNCTION = "process"
CATEGORY = "LivePortrait"
def process(self, source_image, driving_images, dsize, scale, vx_ratio, vy_ratio, pipeline,
lip_zero, eye_retargeting, lip_retargeting, stitching, relative, eyes_retargeting_multiplier, lip_retargeting_multiplier):
def process(
self,
source_image: torch.Tensor,
driving_images: torch.Tensor,
dsize: int,
scale: float,
vx_ratio: float,
vy_ratio: float,
pipeline: LivePortraitPipeline,
lip_zero: bool,
eye_retargeting: bool,
lip_retargeting: bool,
stitching: bool,
relative: bool,
eyes_retargeting_multiplier: float,
lip_retargeting_multiplier: float,
mismatch_method: str = "repeat",
):
source_image_np = (source_image * 255).byte().numpy()
driving_images_np = (driving_images * 255).byte().numpy()
crop_cfg = CropConfig(
dsize = dsize,
scale = scale,
vx_ratio = vx_ratio,
vy_ratio = vy_ratio,
)
dsize=dsize,
scale=scale,
vx_ratio=vx_ratio,
vy_ratio=vy_ratio,
)
cropper = Cropper(crop_cfg=crop_cfg)
pipeline.cropper = cropper
pipeline.live_portrait_wrapper.cfg.flag_eye_retargeting = eye_retargeting
pipeline.live_portrait_wrapper.cfg.eyes_retargeting_multiplier = eyes_retargeting_multiplier
pipeline.live_portrait_wrapper.cfg.eyes_retargeting_multiplier = (
eyes_retargeting_multiplier
)
pipeline.live_portrait_wrapper.cfg.flag_lip_retargeting = lip_retargeting
pipeline.live_portrait_wrapper.cfg.lip_retargeting_multiplier = lip_retargeting_multiplier
pipeline.live_portrait_wrapper.cfg.lip_retargeting_multiplier = (
lip_retargeting_multiplier
)
pipeline.live_portrait_wrapper.cfg.flag_stitching = stitching
pipeline.live_portrait_wrapper.cfg.flag_relative = relative
pipeline.live_portrait_wrapper.cfg.flag_lip_zero = lip_zero
cropped_out_list = []
full_out_list = []
for img in source_image_np:
cropped_frames, full_frame = pipeline.execute(img, driving_images_np)
cropped_tensors = [torch.from_numpy(np_array) for np_array in cropped_frames]
cropped_tensors_out = torch.stack(cropped_tensors) / 255
cropped_tensors_out = cropped_tensors_out.cpu().float()
full_tensors = [torch.from_numpy(np_array) for np_array in full_frame]
full_tensors_out = torch.stack(full_tensors) / 255
full_tensors_out = full_tensors_out.cpu().float()
source_np = (source_image * 255).byte().numpy()
cropped_out_list.append(cropped_tensors_out)
full_out_list.append(full_tensors_out)
cropped_out_list, full_out_list = pipeline.execute(
source_np, driving_images_np, mismatch_method
)
cropped_tensors_out = (
torch.stack([torch.from_numpy(np_array) for np_array in cropped_out_list])
/ 255
)
full_tensors_out = (
torch.stack([torch.from_numpy(np_array) for np_array in full_out_list])
/ 255
)
cropped_tensors_out = torch.cat(cropped_out_list, dim=0)
full_tensors_out = torch.cat(full_out_list, dim=0)
return (cropped_tensors_out.cpu().float(), full_tensors_out.cpu().float())
return (cropped_tensors_out, full_tensors_out)
NODE_CLASS_MAPPINGS = {
"DownloadAndLoadLivePortraitModels": DownloadAndLoadLivePortraitModels,
@@ -281,4 +366,5 @@ NODE_CLASS_MAPPINGS = {
NODE_DISPLAY_NAME_MAPPINGS = {
"DownloadAndLoadLivePortraitModels": "(Down)Load LivePortraitModels",
"LivePortraitProcess": "LivePortraitProcess",
}
}