Don't draw keypoints for every frame by default

This commit is contained in:
kijai
2024-07-09 14:25:39 +03:00
parent a284bb52b2
commit c21705edb5
3 changed files with 17 additions and 13 deletions
+1 -1
View File
@@ -83,7 +83,7 @@ class LivePortraitPipeline(object):
)
driving_frame = driving_images_np[i]
crop_info, _ = self.cropper.crop_single_image(source_frame_rgb)
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"],
+7 -3
View File
@@ -59,7 +59,7 @@ class Cropper(object):
if hasattr(self.crop_cfg, k):
setattr(self.crop_cfg, k, v)
def crop_single_image(self, obj, **kwargs):
def crop_single_image(self, obj, draw_keypoints, **kwargs):
direction = kwargs.get('direction', 'large-small')
# crop and align a single image
@@ -77,8 +77,8 @@ class Cropper(object):
if len(src_face) == 0:
log('No face detected in the source image.')
raise Exception("No face detected in the source image!")
elif len(src_face) > 1:
log(f'More than one face detected in the image, only pick one face by rule {direction}.')
#elif len(src_face) > 1:
# log(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]
pts = src_face.landmark_2d_106
@@ -101,6 +101,8 @@ class Cropper(object):
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:
@@ -111,6 +113,8 @@ class Cropper(object):
keypoints_image = cv2.cvtColor(blank_image, cv2.COLOR_BGR2RGB)
return ret_dct, keypoints_image
else:
return ret_dct
def get_retargeting_lmk_info(self, driving_rgb_lst):
# TODO: implement a tracking-based version
+1 -1
View File
@@ -405,7 +405,7 @@ class LivePortraitCropper:
)
cropper = Cropper(crop_cfg=crop_cfg, provider=onnx_device)
crop_info, keypoints_img = cropper.crop_single_image(source_image_np[0])
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()