Add MediaPipe as alternative face detector
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
@@ -0,0 +1 @@
|
||||
from .mp_utils import LMKExtractor
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
Binary file not shown.
@@ -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
|
||||
|
||||
@@ -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
@@ -2,4 +2,5 @@ pyyaml
|
||||
numpy
|
||||
opencv-python
|
||||
onnxruntime-gpu
|
||||
pykalman
|
||||
pykalman
|
||||
mediapipe
|
||||
Reference in New Issue
Block a user