eye/lip retargeting fixes

This commit is contained in:
kijai
2024-07-09 21:49:54 +03:00
parent 2e40fe3820
commit c0959056ae
2 changed files with 29 additions and 18 deletions
+6 -7
View File
@@ -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
+23 -11
View File
@@ -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