diff --git a/liveportrait/live_portrait_pipeline.py b/liveportrait/live_portrait_pipeline.py index 3540d91..07d6ed6 100644 --- a/liveportrait/live_portrait_pipeline.py +++ b/liveportrait/live_portrait_pipeline.py @@ -61,9 +61,9 @@ class LivePortraitPipeline(object): ): inference_cfg = self.live_portrait_wrapper.cfg - I_p_lst = [] - I_p_paste_lst = [] - driving_lmk_lst = [] + cropped_image_list = [] + composited_image_list = [] + driving_landmark_list = [] out_mask_list = [] R_d_0, x_d_0_info = None, None @@ -72,20 +72,16 @@ class LivePortraitPipeline(object): pbar = comfy.utils.ProgressBar(total_frames) if inference_cfg.flag_eye_retargeting or inference_cfg.flag_lip_retargeting: - driving_lmk_lst = self.cropper.get_retargeting_lmk_info(driving_images_np) + driving_landmark_list = self.cropper.get_retargeting_lmk_info(driving_images_np) for i in tqdm(range(total_frames), desc='Animating...', total=total_frames): source_frame_rgb = self._get_source_frame( source_np, i, total_frames, mismatch_method ) driving_frame = driving_images_np[i] - - crop_info = self.cropper.crop_single_image(source_frame_rgb, draw_keypoints=False) - source_lmk = crop_info["lmk_crop"] - _, img_crop_256x256 = ( - crop_info["img_crop"], - crop_info["img_crop_256x256"], - ) + + source_lmk = crop_info[i]["lmk_crop"] + img_crop_256x256 = crop_info[i]["img_crop_256x256"] if inference_cfg.flag_do_crop: I_s = self.live_portrait_wrapper.prepare_source(img_crop_256x256) @@ -126,10 +122,10 @@ 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_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_lmk_lst + source_lmk, driving_landmark_list ) ) @@ -251,12 +247,12 @@ class LivePortraitPipeline(object): 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) + cropped_image_list.append(I_p_i) # Transform and blend I_p_i_to_ori = _transform_img( I_p_i, - crop_info["M_c2o"], + crop_info[i]["M_c2o"], dsize=(source_frame_rgb.shape[1], source_frame_rgb.shape[0]), ) @@ -265,18 +261,16 @@ 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["M_c2o"], + crop_info[i]["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) + composited_image_list.append(I_p_i_to_ori_blend) out_mask_list.append(mask_ori) pbar.update(1) - return I_p_lst, I_p_paste_lst, out_mask_list + return cropped_image_list, composited_image_list, out_mask_list diff --git a/liveportrait/utils/crop.py b/liveportrait/utils/crop.py index ac2e725..b859f26 100644 --- a/liveportrait/utils/crop.py +++ b/liveportrait/utils/crop.py @@ -4,7 +4,7 @@ cropping function and the related preprocess functions for cropping """ -import cv2; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False) # NOTE: enforce single thread +import cv2#; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False) # NOTE: enforce single thread import numpy as np from math import sin, cos, acos, degrees diff --git a/liveportrait/utils/cropper.py b/liveportrait/utils/cropper.py index 945c52b..412dab8 100644 --- a/liveportrait/utils/cropper.py +++ b/liveportrait/utils/cropper.py @@ -3,11 +3,11 @@ import numpy as np from typing import List, Union, Tuple from dataclasses import dataclass, field -import cv2; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False) +import cv2#; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False) from .landmark_runner import LandmarkRunner from .face_analysis_diy import FaceAnalysisDIY -from .crop import crop_image, crop_image_by_bbox, parse_bbox_from_landmark, average_bbox_lst +from .crop import crop_image import folder_paths import os @@ -15,8 +15,8 @@ script_directory = os.path.dirname(os.path.abspath(__file__)) @dataclass class Trajectory: - start: int = -1 # 起始帧 闭区间 - end: int = -1 # 结束帧 闭区间 + start: int = -1 + end: int = -1 lmk_lst: Union[Tuple, List, np.ndarray] = field(default_factory=list) # lmk list bbox_lst: Union[Tuple, List, np.ndarray] = field(default_factory=list) # bbox list frame_rgb_lst: Union[Tuple, List, np.ndarray] = field(default_factory=list) # frame list @@ -48,9 +48,11 @@ class Cropper(object): if hasattr(self.crop_cfg, k): setattr(self.crop_cfg, k, v) - def crop_single_image(self, img_rgb, draw_keypoints, **kwargs): + def crop_single_image(self, img_rgb, **kwargs): direction = kwargs.get('direction', 'large-small') + + src_face = self.face_analysis_wrapper.get( img_rgb, flag_do_landmark_2d_106=True, @@ -62,10 +64,9 @@ class Cropper(object): #elif len(src_face) > 1: # print(f'More than one face detected in the image, only pick one face by rule {direction}.') - src_face = src_face[self.crop_cfg.face_index] + src_face = src_face[self.crop_cfg.face_index] # choose the index if multiple faces detected pts = src_face.landmark_2d_106 - # crop the face ret_dct = crop_image( img_rgb, # ndarray @@ -78,67 +79,19 @@ class Cropper(object): ret_dct['img_crop_256x256'] = cv2.resize(ret_dct['img_crop'], (256, 256), interpolation=cv2.INTER_AREA) ret_dct['pt_crop_256x256'] = ret_dct['pt_crop'] * 256 / kwargs.get('dsize', 512) + input_image_size = img_rgb.shape[:2] + ret_dct['input_image_size'] = input_image_size + recon_ret = self.landmark_runner.run(img_rgb, pts) lmk = recon_ret['pts'] ret_dct['lmk_crop'] = lmk - # Draw each landmark as a circle - if draw_keypoints: - print("Drawing keypoints...") - height, width = img_rgb.shape[:2] - blank_image = np.zeros((height, width, 3), dtype=np.uint8) * 255 - for (x, y) in lmk: - # Ensure the coordinates are within the dimensions of the blank image - if 0 <= x < width and 0 <= y < height: - cv2.circle(blank_image, (int(x), int(y)), radius=2, color=(0, 0, 255)) - - keypoints_image = cv2.cvtColor(blank_image, cv2.COLOR_BGR2RGB) - - return ret_dct, keypoints_image - else: - return ret_dct + return ret_dct def get_retargeting_lmk_info(self, driving_rgb_lst): # TODO: implement a tracking-based version driving_lmk_lst = [] for driving_image in driving_rgb_lst: - ret_dct = self.crop_single_image(driving_image, draw_keypoints=False) + ret_dct = self.crop_single_image(driving_image) driving_lmk_lst.append(ret_dct['lmk_crop']) return driving_lmk_lst - - def make_video_clip(self, driving_rgb_lst, output_path, output_fps=30, **kwargs): - trajectory = Trajectory() - direction = kwargs.get('direction', 'large-small') - for idx, driving_image in enumerate(driving_rgb_lst): - if idx == 0 or trajectory.start == -1: - src_face = self.face_analysis_wrapper.get( - driving_image, - flag_do_landmark_2d_106=True, - direction=direction - ) - if len(src_face) == 0: - # No face detected in the driving_image - continue - elif len(src_face) > 1: - print(f'More than one face detected in the driving frame_{idx}, only pick one face by rule {direction}.') - src_face = src_face[0] - pts = src_face.landmark_2d_106 - lmk_203 = self.landmark_runner(driving_image, pts)['pts'] - trajectory.start, trajectory.end = idx, idx - else: - lmk_203 = self.face_recon_wrapper(driving_image, trajectory.lmk_lst[-1])['pts'] - trajectory.end = idx - - trajectory.lmk_lst.append(lmk_203) - ret_bbox = parse_bbox_from_landmark(lmk_203, scale=self.crop_cfg.globalscale, vy_ratio=elf.crop_cfg.vy_ratio)['bbox'] - bbox = [ret_bbox[0, 0], ret_bbox[0, 1], ret_bbox[2, 0], ret_bbox[2, 1]] # 4, - trajectory.bbox_lst.append(bbox) # bbox - trajectory.frame_rgb_lst.append(driving_image) - - global_bbox = average_bbox_lst(trajectory.bbox_lst) - for idx, (frame_rgb, lmk) in enumerate(zip(trajectory.frame_rgb_lst, trajectory.lmk_lst)): - ret_dct = crop_image_by_bbox( - frame_rgb, global_bbox, lmk=lmk, - dsize=self.video_crop_cfg.dsize, flag_rot=self.video_crop_cfg.flag_rot, borderValue=self.video_crop_cfg.borderValue - ) - frame_rgb_crop = ret_dct['img_crop'] diff --git a/liveportrait/utils/face_analysis_diy.py b/liveportrait/utils/face_analysis_diy.py index 883992b..ac29ace 100644 --- a/liveportrait/utils/face_analysis_diy.py +++ b/liveportrait/utils/face_analysis_diy.py @@ -1,15 +1,30 @@ # coding: utf-8 """ -face detectoin and alignment using InsightFace +face detection and alignment using InsightFace """ +from insightface.utils import transform + +#patch Insightface function to get rid of the annoying warnings +def patched_estimate_affine_matrix_3d23d(X, Y): + ''' Using least-squares solution + Args: + X: [n, 3]. 3d points(fixed) + Y: [n, 3]. corresponding 3d points(moving). Y = PX + Returns: + P_Affine: (3, 4). Affine camera matrix (the third row is [0, 0, 0, 1]). + ''' + X_homo = np.hstack((X, np.ones([X.shape[0],1]))) # n x 4 + P = np.linalg.lstsq(X_homo, Y, rcond=None)[0].T # Affine matrix. 3 x 4 + return P + +transform.estimate_affine_matrix_3d23d = patched_estimate_affine_matrix_3d23d import numpy as np from insightface.app import FaceAnalysis from insightface.app.common import Face from .timer import Timer - def sort_by_direction(faces, direction: str = 'large-small', face_center=None): if len(faces) <= 0: return faces diff --git a/liveportrait/utils/landmark_runner.py b/liveportrait/utils/landmark_runner.py index 57bb143..fb3f11e 100644 --- a/liveportrait/utils/landmark_runner.py +++ b/liveportrait/utils/landmark_runner.py @@ -1,6 +1,6 @@ # coding: utf-8 -import cv2; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False) +import cv2#; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False) import torch import numpy as np import onnxruntime diff --git a/nodes.py b/nodes.py index e39c8dd..884cb41 100644 --- a/nodes.py +++ b/nodes.py @@ -6,6 +6,7 @@ import comfy.model_management as mm import comfy.utils import numpy as np import cv2 +from tqdm import tqdm script_directory = os.path.dirname(os.path.abspath(__file__)) @@ -224,7 +225,7 @@ class LivePortraitProcess: return {"required": { "pipeline": ("LIVEPORTRAITPIPE",), - "crop_info": ("CROPINFO", {"default": {}}), + "crop_info": ("CROPINFO", {"default": []}), "source_image": ("IMAGE",), "driving_images": ("IMAGE",), "lip_zero": ("BOOLEAN", {"default": True}), @@ -267,7 +268,7 @@ class LivePortraitProcess: self, source_image: torch.Tensor, driving_images: torch.Tensor, - crop_info: dict, + crop_info: list, pipeline: LivePortraitPipeline, lip_zero: bool, lip_zero_threshold: float, @@ -305,13 +306,11 @@ class LivePortraitProcess: crop_mask = np.repeat(np.atleast_3d(crop_mask), 3, axis=2) pipeline.live_portrait_wrapper.cfg.mask_crop = crop_mask - pipeline.cropper = crop_info['cropper'] - cropped_out_list = [] full_out_list = [] cropped_out_list, full_out_list, out_mask_list = pipeline.execute( - source_np, driving_images_np, crop_info['crop_info'], mismatch_method + source_np, driving_images_np, crop_info, mismatch_method ) cropped_tensors_out = ( torch.stack([torch.from_numpy(np_array) for np_array in cropped_out_list]) @@ -343,20 +342,15 @@ class LivePortraitCropper: "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}), "face_index": ("INT", {"default": 0, "min": 0, "max": 100}), - }, - "optional": { - "onnx_device": ( - [ - 'CPU', - 'CUDA', - ], { + "onnx_device": ( + ['CPU', 'CUDA', 'ROCM'], { "default": 'CPU' }), - } + }, } - RETURN_TYPES = ("IMAGE", "CROPINFO", "IMAGE",) - RETURN_NAMES = ("cropped_image", "crop_info", "keypoints_image",) + RETURN_TYPES = ("IMAGE", "CROPINFO",) + RETURN_NAMES = ("cropped_image", "crop_info",) FUNCTION = "process" CATEGORY = "LivePortrait" @@ -372,21 +366,56 @@ class LivePortraitCropper: ) cropper = Cropper(crop_cfg=crop_cfg, provider=onnx_device) - crop_info, keypoints_img = cropper.crop_single_image(source_image_np[0], draw_keypoints=True) - keypoints_image_tensor = torch.from_numpy(keypoints_img) / 255 - keypoints_image_tensor = keypoints_image_tensor.unsqueeze(0).cpu().float() + crop_info_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 = cropper.crop_single_image(source_image_np[i]) + crop_info_list.append(crop_info) + pbar.update(1) cropped_image = crop_info['img_crop_256x256'] cropped_tensors = torch.from_numpy(cropped_image) / 255 cropped_tensors = cropped_tensors.unsqueeze(0).cpu().float() - cropper_dict = { - "cropper": cropper, - "crop_info": crop_info, + return (cropped_tensors, crop_info_list) + +class KeypointsToImage: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "crop_info": ("CROPINFO", {"default": []}), + }, } - return (cropped_tensors, cropper_dict, keypoints_image_tensor) + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("keypoints_image",) + FUNCTION = "drawkeypoints" + CATEGORY = "LivePortrait" + + def drawkeypoints(self, crop_info): + height, width = crop_info[0]['input_image_size'] + keypoints_img_list = [] + pbar = comfy.utils.ProgressBar(len(crop_info)) + for crop in crop_info: + keypoints = crop['lmk_crop'].copy() + # Draw each landmark as a circle + blank_image = np.zeros((height, width, 3), dtype=np.uint8) * 255 + for (x, y) in keypoints: + # Ensure the coordinates are within the dimensions of the blank image + if 0 <= x < width and 0 <= y < height: + cv2.circle(blank_image, (int(x), int(y)), radius=2, color=(0, 0, 255)) + + keypoints_image = cv2.cvtColor(blank_image, cv2.COLOR_BGR2RGB) + keypoints_img_list.append(keypoints_image) + pbar.update(1) + + keypoints_img_tensor = ( + torch.stack([torch.from_numpy(np_array) for np_array in keypoints_img_list]) / 255).float() + + + return (keypoints_img_tensor,) class KeypointScaler: @classmethod @@ -442,11 +471,13 @@ NODE_CLASS_MAPPINGS = { "DownloadAndLoadLivePortraitModels": DownloadAndLoadLivePortraitModels, "LivePortraitProcess": LivePortraitProcess, "LivePortraitCropper": LivePortraitCropper, - "KeypointScaler": KeypointScaler + "KeypointScaler": KeypointScaler, + "KeypointsToImage": KeypointsToImage } NODE_DISPLAY_NAME_MAPPINGS = { "DownloadAndLoadLivePortraitModels": "(Down)Load LivePortraitModels", "LivePortraitProcess": "LivePortraitProcess", "LivePortraitCropper": "LivePortraitCropper", - "KeypointScaler": "KeypointScaler" + "KeypointScaler": "KeypointScaler", + "KeypointsToImage": "LivePortrait KeypointsToImage" } \ No newline at end of file