Add MediaPipe as alternative face detector

This commit is contained in:
kijai
2024-07-24 03:57:59 +03:00
parent 6261f4e474
commit 068ab2c280
9 changed files with 3468 additions and 14 deletions
+19
View File
@@ -74,6 +74,23 @@ def _transform_pts(pts, M):
return pts @ M[:2, :2].T + M[:2, 2]
def parse_pt2_from_pt478(pt478, use_lip=True):
"""
parsing the 2 points according to the 101 points, which cancels the roll
"""
# the former version use the eye center, but it is not robust, now use interpolation
pt_left_eye = pt478[468] # left eye center
pt_right_eye = pt478[473] # right eye center
if use_lip:
# use lip
pt_center_eye = (pt_left_eye + pt_right_eye) / 2
pt_center_lip = pt478[14]
pt2 = np.stack([pt_center_eye, pt_center_lip], axis=0)
else:
pt2 = np.stack([pt_left_eye, pt_right_eye], axis=0)
return pt2
def parse_pt2_from_pt101(pt101, use_lip=True):
"""
parsing the 2 points according to the 101 points, which cancels the roll
@@ -211,6 +228,8 @@ def parse_pt2_from_pt_x(pts, use_lip=True):
pt2 = parse_pt2_from_pt5(pts, use_lip=use_lip)
elif pts.shape[0] == 203:
pt2 = parse_pt2_from_pt203(pts, use_lip=use_lip)
elif pts.shape[0] == 478:
pt2 = parse_pt2_from_pt478(pts, use_lip=use_lip)
elif pts.shape[0] > 101:
# take the first 101 points
pt2 = parse_pt2_from_pt101(pts[:101], use_lip=use_lip)
+55 -5
View File
@@ -22,8 +22,7 @@ class Trajectory:
frame_rgb_lst: Union[Tuple, List, np.ndarray] = field(default_factory=list) # frame list
frame_rgb_crop_lst: Union[Tuple, List, np.ndarray] = field(default_factory=list) # frame crop list
class Cropper(object):
class CropperInsightFace(object):
def __init__(self, **kwargs) -> None:
device_id = kwargs.get('device_id', 0)
provider = kwargs.get('onnx_device', 'CPU')
@@ -54,9 +53,6 @@ class Cropper(object):
if len(src_face) == 0:
ret_dct = {}
return ret_dct
#raise Exception("No face detected in the source image!")
#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[face_index] # choose the index if multiple faces detected
pts = src_face.landmark_2d_106
@@ -82,4 +78,58 @@ class Cropper(object):
lmk = recon_ret['pts']
ret_dct['lmk_crop'] = lmk
return ret_dct, cropped_image_256
class CropperMediaPipe(object):
def __init__(self, **kwargs) -> None:
device_id = kwargs.get('device_id', 0)
provider = kwargs.get('onnx_device', 'CPU')
self.landmark_runner = LandmarkRunner(
ckpt_path=os.path.join(folder_paths.models_dir, 'liveportrait', 'landmark.onnx'),
onnx_provider=provider,
device_id=device_id
)
self.landmark_runner.warmup()
from ...media_pipe.mp_utils import LMKExtractor
self.lmk_extractor = LMKExtractor()
def crop_single_image(self, img_rgb, dsize, scale, vy_ratio, vx_ratio, face_index, face_index_order, rotate):
face_result = self.lmk_extractor(img_rgb)
if face_result is None:
ret_dct = {}
cropped_image_256 = None
return ret_dct, cropped_image_256
face_landmarks = face_result[face_index]
lmks = []
for index in range(len(face_landmarks)):
x = face_landmarks[index].x * img_rgb.shape[1]
y = face_landmarks[index].y * img_rgb.shape[0]
lmks.append([x, y])
pts = np.array(lmks)
# crop the face
ret_dct, image_crop = crop_image(
img_rgb, # ndarray
pts, # 106x2 or Nx2
dsize=dsize,
scale=scale,
vy_ratio=vy_ratio,
vx_ratio=vx_ratio,
rotate=rotate
)
# update a 256x256 version for network input or else
cropped_image_256 = cv2.resize(image_crop, (256, 256), interpolation=cv2.INTER_AREA)
ret_dct['pt_crop_256x256'] = ret_dct['pt_crop'] * 256 / dsize
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
return ret_dct, cropped_image_256
+1
View File
@@ -0,0 +1 @@
from .mp_utils import LMKExtractor
File diff suppressed because it is too large Load Diff
Binary file not shown.
+38
View File
@@ -0,0 +1,38 @@
import os
import mediapipe as mp
from mediapipe.tasks import python
from mediapipe.tasks.python import vision
from . import face_landmark
CUR_DIR = os.path.dirname(__file__)
class LMKExtractor():
def __init__(self):
# Create an FaceLandmarker object.
self.mode = mp.tasks.vision.FaceDetectorOptions.running_mode.IMAGE
base_options = python.BaseOptions(model_asset_path=os.path.join(CUR_DIR, 'mp_models','face_landmarker_v2_with_blendshapes.task'))
base_options.delegate = mp.tasks.BaseOptions.Delegate.CPU
options = vision.FaceLandmarkerOptions(base_options=base_options,
running_mode=self.mode,
output_face_blendshapes=False,
output_facial_transformation_matrixes=True,
num_faces=1,
min_face_detection_confidence=0.5,
min_face_presence_confidence=0.5,
min_tracking_confidence=0.5)
self.detector = face_landmark.FaceLandmarker.create_from_options(options)
det_base_options = python.BaseOptions(model_asset_path=os.path.join(CUR_DIR, 'mp_models','blaze_face_short_range.tflite'))
det_options = vision.FaceDetectorOptions(base_options=det_base_options)
self.det_detector = vision.FaceDetector.create_from_options(det_options)
def __call__(self, img):
image = mp.Image(image_format=mp.ImageFormat.SRGB, data=img)
try:
detection_result, _ = self.detector.detect(image)
except:
return None
return detection_result.face_landmarks
+48 -8
View File
@@ -12,7 +12,15 @@ import gc
script_directory = os.path.dirname(os.path.abspath(__file__))
from .liveportrait.live_portrait_pipeline import LivePortraitPipeline
from .liveportrait.utils.cropper import Cropper
try:
from .liveportrait.utils.cropper import CropperMediaPipe
except:
raise ModuleNotFoundError("Can't load MediaPipe, MediaPipeCropper not available")
try:
from .liveportrait.utils.cropper import CropperInsightFace
except:
raise ModuleNotFoundError("Can't load InsightFace, InsightFaceCropper not available")
from .liveportrait.modules.spade_generator import SPADEDecoder
from .liveportrait.modules.warping_network import WarpingNetwork
from .liveportrait.modules.motion_extractor import MotionExtractor
@@ -405,7 +413,7 @@ class LivePortraitComposite:
source_frame = _get_source_frame(source_image, i, liveportrait_out["mismatch_method"]).unsqueeze(0).to(device)
if not liveportrait_out["out_list"][i]:
composited_image_list.append(source_frame)
composited_image_list.append(source_frame.cpu())
out_mask_list.append(torch.zeros((1, 3, H, W), device="cpu"))
else:
cropped_image = torch.clamp(liveportrait_out["out_list"][i]["out"], 0, 1).permute(0, 2, 3, 1)
@@ -485,7 +493,37 @@ class LivePortraitLoadCropper:
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)
self.cropper = CropperInsightFace(**cropper_init_config)
return (self.cropper,)
class LivePortraitLoadMediaPipeCropper:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"landmarkrunner_onnx_device": (
['CPU', 'CUDA', 'ROCM', 'CoreML'], {
"default": 'CPU'
}),
"keep_model_loaded": ("BOOLEAN", {"default": True})
},
}
RETURN_TYPES = ("LPCROPPER",)
RETURN_NAMES = ("cropper",)
FUNCTION = "crop"
CATEGORY = "LivePortrait"
def crop(self, landmarkrunner_onnx_device, keep_model_loaded):
cropper_init_config = {
'keep_model_loaded': keep_model_loaded,
'onnx_device': landmarkrunner_onnx_device
}
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 = CropperMediaPipe(**cropper_init_config)
return (self.cropper,)
@@ -522,7 +560,7 @@ class LivePortraitCropper:
CATEGORY = "LivePortrait"
def process(self, pipeline, cropper, source_image, dsize, scale, vx_ratio, vy_ratio, face_index, face_index_order, rotate):
source_image_np = (source_image * 255).byte().numpy()
source_image_np = (source_image.contiguous() * 255).byte().numpy()
# Initialize lists
crop_info_list = []
@@ -718,15 +756,17 @@ NODE_CLASS_MAPPINGS = {
#"KeypointScaler": KeypointScaler,
"KeypointsToImage": KeypointsToImage,
"LivePortraitLoadCropper": LivePortraitLoadCropper,
"LivePortraitLoadMediaPipeCropper": LivePortraitLoadMediaPipeCropper,
"LivePortraitComposite": LivePortraitComposite,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"DownloadAndLoadLivePortraitModels": "(Down)Load LivePortraitModels",
"LivePortraitProcess": "LivePortraitProcess",
"LivePortraitCropper": "LivePortraitCropper",
"LivePortraitRetargeting": "LivePortraitRetargeting",
"LivePortraitProcess": "LivePortrait Process",
"LivePortraitCropper": "LivePortrait Cropper",
"LivePortraitRetargeting": "LivePortrait Retargeting",
#"KeypointScaler": "KeypointScaler",
"KeypointsToImage": "LivePortrait KeypointsToImage",
"LivePortraitLoadCropper": "LivePortrait LoadCropper",
"LivePortraitLoadCropper": "LivePortrait Load InsightFaceCropper",
"LivePortraitLoadMediaPipeCropper": "LivePortrait Load MediaPipeCropper",
"LivePortraitComposite": "LivePortrait Composite",
}
+2 -1
View File
@@ -2,4 +2,5 @@ pyyaml
numpy
opencv-python
onnxruntime-gpu
pykalman
pykalman
mediapipe