diff --git a/liveportrait/utils/cropper.py b/liveportrait/utils/cropper.py index 185c43b..599c988 100644 --- a/liveportrait/utils/cropper.py +++ b/liveportrait/utils/cropper.py @@ -101,7 +101,7 @@ class Cropper(object): ret_dct['lmk_crop'] = lmk # Draw each landmark as a circle - width, height = 512, 512 + 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 diff --git a/nodes.py b/nodes.py index c3b299f..2d4a5be 100644 --- a/nodes.py +++ b/nodes.py @@ -272,6 +272,9 @@ class LivePortraitProcess: "stitching": ("BOOLEAN", {"default": True}), "relative": ("BOOLEAN", {"default": True}), }, + "optional": { + "mask": ("MASK", {"default": None}), + } } RETURN_TYPES = ( @@ -299,6 +302,7 @@ class LivePortraitProcess: eyes_retargeting_multiplier: float, lip_retargeting_multiplier: float, mismatch_method: str = "repeat", + mask: torch.Tensor = None, ): source_np = (source_image * 255).byte().numpy() driving_images_np = (driving_images * 255).byte().numpy() @@ -315,6 +319,12 @@ class LivePortraitProcess: pipeline.live_portrait_wrapper.cfg.flag_relative = relative pipeline.live_portrait_wrapper.cfg.flag_lip_zero = lip_zero + if mask is not None: + crop_mask = mask[0].cpu().numpy() + crop_mask = (crop_mask * 255).astype(np.uint8) + 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 = []