From 68d0ddf72a95d2d45eb50357b32643f02510ae7a Mon Sep 17 00:00:00 2001 From: Mel Massadian Date: Mon, 8 Jul 2024 21:22:50 +0200 Subject: [PATCH] remove reference frame attempt also use batches for driving when either retargetting is enabled --- liveportrait/live_portrait_pipeline.py | 122 ++++++++++++++++++++----- 1 file changed, 99 insertions(+), 23 deletions(-) diff --git a/liveportrait/live_portrait_pipeline.py b/liveportrait/live_portrait_pipeline.py index 4375364..89eaf45 100644 --- a/liveportrait/live_portrait_pipeline.py +++ b/liveportrait/live_portrait_pipeline.py @@ -73,24 +73,17 @@ class LivePortraitPipeline(object): self, source_np, driving_images_np, mismatch_method="repeat", reference_frame=0 ): inference_cfg = self.live_portrait_wrapper.cfg - is_video = source_np.shape[0] > 1 I_p_lst = [] I_p_paste_lst = [] + driving_lmk_lst = [] total_frames = driving_images_np.shape[0] pbar = comfy.utils.ProgressBar(total_frames) - ref_frame = self._get_source_frame( - source_np, reference_frame, total_frames, mismatch_method - ) - rcrop_info = self.cropper.crop_single_image(ref_frame) - rsource_lmk = rcrop_info["lmk_crop"] - rimg_crop, ref_crop_256x256 = ( - rcrop_info["img_crop"], - rcrop_info["img_crop_256x256"], - ) + if inference_cfg.flag_eye_retargeting or inference_cfg.flag_lip_retargeting: + driving_lmk_lst = self.cropper.get_retargeting_lmk_info(driving_images_np) for i in range(total_frames): source_frame_rgb = self._get_source_frame( @@ -100,7 +93,7 @@ class LivePortraitPipeline(object): crop_info = self.cropper.crop_single_image(source_frame_rgb) source_lmk = crop_info["lmk_crop"] - img_crop, img_crop_256x256 = ( + _, img_crop_256x256 = ( crop_info["img_crop"], crop_info["img_crop_256x256"], ) @@ -110,9 +103,6 @@ class LivePortraitPipeline(object): else: I_s = self.live_portrait_wrapper.prepare_source(source_frame_rgb) - rel_s_info = self.live_portrait_wrapper.get_kp_info( - self.live_portrait_wrapper.prepare_source(ref_crop_256x256) - ) x_s_info = self.live_portrait_wrapper.get_kp_info(I_s) x_c_s = x_s_info["kp"] R_s = get_rotation_matrix( @@ -148,7 +138,7 @@ class LivePortraitPipeline(object): )[0] if inference_cfg.flag_eye_retargeting or inference_cfg.flag_lip_retargeting: - driving_lmk_lst = self.cropper.get_retargeting_lmk_info([driving_frame]) + # 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 @@ -162,11 +152,9 @@ class LivePortraitPipeline(object): if inference_cfg.flag_relative: R_new = R_d @ R_s - delta_new = rel_s_info["exp"] + (x_d_info["exp"] - rel_s_info["exp"]) - scale_new = rel_s_info["scale"] * ( - x_d_info["scale"] / rel_s_info["scale"] - ) - t_new = rel_s_info["t"] + (x_d_info["t"] - rel_s_info["t"]) + 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 delta_new = x_d_info["exp"] @@ -174,12 +162,100 @@ class LivePortraitPipeline(object): t_new = x_d_info["t"] t_new[..., 2].fill_(0) # zero tz - x_d_new = scale_new * (x_c_s @ R_new + delta_new) + t_new + x_d_i_new = scale_new * (x_c_s @ R_new + delta_new) + t_new + 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_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) if inference_cfg.flag_stitching: - x_d_new = self.live_portrait_wrapper.stitching(x_s, x_d_new) + 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_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] I_p_lst.append(I_p_i)