From 72bb6910e9bf8ae63bf98666bf4caebbb18725b7 Mon Sep 17 00:00:00 2001 From: Mel Massadian Date: Mon, 8 Jul 2024 16:48:15 +0200 Subject: [PATCH] initial too much diff due to formatting --- liveportrait/live_portrait_pipeline.py | 311 ++++++++++++------------ nodes.py | 322 ++++++++++++++++--------- 2 files changed, 366 insertions(+), 267 deletions(-) diff --git a/liveportrait/live_portrait_pipeline.py b/liveportrait/live_portrait_pipeline.py index 7293cab..71885c8 100644 --- a/liveportrait/live_portrait_pipeline.py +++ b/liveportrait/live_portrait_pipeline.py @@ -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 diff --git a/nodes.py b/nodes.py index 0d831c6..aa56394 100644 --- a/nodes.py +++ b/nodes.py @@ -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", - } \ No newline at end of file +} +