From c0959056ae415c4b56ea7f2e0168acdd4be05e3c Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 9 Jul 2024 21:49:54 +0300 Subject: [PATCH] eye/lip retargeting fixes --- liveportrait/live_portrait_pipeline.py | 13 +++++----- nodes.py | 34 +++++++++++++++++--------- 2 files changed, 29 insertions(+), 18 deletions(-) diff --git a/liveportrait/live_portrait_pipeline.py b/liveportrait/live_portrait_pipeline.py index 246c1fb..bfd6fc4 100644 --- a/liveportrait/live_portrait_pipeline.py +++ b/liveportrait/live_portrait_pipeline.py @@ -72,7 +72,7 @@ class LivePortraitPipeline(object): pbar = comfy.utils.ProgressBar(total_frames) if inference_cfg.flag_eye_retargeting or inference_cfg.flag_lip_retargeting: - driving_landmark_list = self.cropper.get_retargeting_lmk_info(driving_images_np) + driving_landmark_list = crop_info["driving_landmark_list"] for i in tqdm(range(total_frames), desc='Animating...', total=total_frames): source_frame_rgb = self._get_source_frame( @@ -80,9 +80,9 @@ class LivePortraitPipeline(object): ) driving_frame = driving_images_np[i] - safe_index = min(i, len(crop_info) - 1) - source_lmk = crop_info[safe_index]["lmk_crop"] - img_crop_256x256 = crop_info[safe_index]["img_crop_256x256"] + safe_index = min(i, len(crop_info["crop_info_list"]) - 1) + source_lmk = crop_info["crop_info_list"][safe_index]["lmk_crop"] + img_crop_256x256 = crop_info["crop_info_list"][safe_index]["img_crop_256x256"] if inference_cfg.flag_do_crop: I_s = self.live_portrait_wrapper.prepare_source(img_crop_256x256) @@ -123,7 +123,6 @@ class LivePortraitPipeline(object): )[0] if inference_cfg.flag_eye_retargeting or inference_cfg.flag_lip_retargeting: - # driving_landmark_list = 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_landmark_list @@ -253,7 +252,7 @@ class LivePortraitPipeline(object): # Transform and blend I_p_i_to_ori = _transform_img( I_p_i, - crop_info[safe_index]["M_c2o"], + crop_info["crop_info_list"][safe_index]["M_c2o"], dsize=(source_frame_rgb.shape[1], source_frame_rgb.shape[0]), ) @@ -262,7 +261,7 @@ class LivePortraitPipeline(object): inference_cfg.mask_crop = cv2.imread(os.path.join(script_directory, "utils", "resources", "mask_template.png"), cv2.IMREAD_COLOR) mask_ori = _transform_img( inference_cfg.mask_crop, - crop_info[safe_index]["M_c2o"], + crop_info["crop_info_list"][safe_index]["M_c2o"], dsize=(source_frame_rgb.shape[1], source_frame_rgb.shape[0]), ) mask_ori = mask_ori.astype(np.float32) / 255.0 diff --git a/nodes.py b/nodes.py index f59612a..a8cb318 100644 --- a/nodes.py +++ b/nodes.py @@ -339,6 +339,9 @@ class LivePortraitCropper: }), "keep_model_loaded": ("BOOLEAN", {"default": True}) }, + "optional": { + "opt_driving_images": ("IMAGE",), + } } @@ -347,31 +350,33 @@ class LivePortraitCropper: FUNCTION = "process" CATEGORY = "LivePortrait" - def process(self, source_image, dsize, scale, vx_ratio, vy_ratio, face_index, keep_model_loaded, onnx_device='CUDA'): + def process(self, source_image, dsize, scale, vx_ratio, vy_ratio, face_index, keep_model_loaded, onnx_device='CUDA', opt_driving_images=None): source_image_np = (source_image * 255).byte().numpy() cropper_init_config = { 'keep_model_loaded': keep_model_loaded, 'onnx_device': onnx_device } - crop_config = { - "dsize" : dsize, - "scale" : scale, - "vx_ratio" : vx_ratio, - "vy_ratio" : vy_ratio, - "face_index" : face_index, - } if not hasattr(self, 'cropper') or self.cropper is None or self.current_config != cropper_init_config: self.current_config = cropper_init_config self.cropper = Cropper(**cropper_init_config) crop_info_list = [] + + if opt_driving_images is not None: + driving_images_np = (opt_driving_images * 255).byte().numpy() + driving_landmark_list = [] pbar = comfy.utils.ProgressBar(len(source_image_np)) for i in tqdm(range(len(source_image_np)), desc='Detecting and cropping..', total=len(source_image_np)): crop_info = self.cropper.crop_single_image(source_image_np[i], dsize, scale, vy_ratio, vx_ratio, face_index) crop_info_list.append(crop_info) + + if opt_driving_images is not None: + driving_crop_dict = self.cropper.crop_single_image(driving_images_np[i], dsize, scale, vy_ratio, vx_ratio, face_index) + driving_landmark_list.append(driving_crop_dict['lmk_crop']) + pbar.update(1) if not keep_model_loaded: @@ -382,7 +387,14 @@ class LivePortraitCropper: cropped_tensors = torch.from_numpy(cropped_image) / 255 cropped_tensors = cropped_tensors.unsqueeze(0).cpu().float() - return (cropped_tensors, crop_info_list) + crop_info_dict = { + 'crop_info_list': crop_info_list + } + + if opt_driving_images is not None: + crop_info_dict['driving_landmark_list'] = driving_landmark_list + + return (cropped_tensors, crop_info_dict) class KeypointsToImage: @classmethod @@ -398,10 +410,10 @@ class KeypointsToImage: CATEGORY = "LivePortrait" def drawkeypoints(self, crop_info): - height, width = crop_info[0]['input_image_size'] + height, width = crop_info["crop_info_list"][0]['input_image_size'] keypoints_img_list = [] pbar = comfy.utils.ProgressBar(len(crop_info)) - for crop in crop_info: + for crop in crop_info["crop_info_list"]: keypoints = crop['lmk_crop'].copy() # Draw each landmark as a circle blank_image = np.zeros((height, width, 3), dtype=np.uint8) * 255