From 7b6ae6c5b2fa5652c10cb6c011e0e956f7c3bc5f Mon Sep 17 00:00:00 2001 From: Radionic Date: Fri, 24 Nov 2023 17:24:06 +0800 Subject: [PATCH] feat: auto segment face in server --- requirements.txt | 1 + sam/sam_multilayer.py | 273 ++++++++++++++++++++++++++++++++++++------ 2 files changed, 240 insertions(+), 34 deletions(-) diff --git a/requirements.txt b/requirements.txt index 98e4f81..7a88a92 100644 --- a/requirements.txt +++ b/requirements.txt @@ -7,4 +7,5 @@ bpy segment-anything tqdm python-dotenv +mediapipe # -e git+https://github.com/facebookresearch/segment-anything.git#egg=segment_anything \ No newline at end of file diff --git a/sam/sam_multilayer.py b/sam/sam_multilayer.py index eecbe92..b5e7533 100644 --- a/sam/sam_multilayer.py +++ b/sam/sam_multilayer.py @@ -7,9 +7,91 @@ import json from segment_anything import sam_model_registry, SamPredictor from einops import rearrange, repeat from PIL import Image +import mediapipe as mp +from math import sqrt +BaseOptions = mp.tasks.BaseOptions +FaceLandmarker = mp.tasks.vision.FaceLandmarker +FaceLandmarkerOptions = mp.tasks.vision.FaceLandmarkerOptions +VisionRunningMode = mp.tasks.vision.RunningMode global_predictor = None +face_landmarker = None + +# For auto-segmentation +layerMapping = { + "L_eye": { + "useMiddle": False, + "positiveOffsetX": 0, + "positiveOffsetY": 0, + "negativeOffsetX": 0, + "negativeOffsetY": 0, + "positiveScale": 0, + "negativeScale": 0.5, + "indices": mp.solutions.face_mesh.FACEMESH_LEFT_EYE, + }, + "R_eye": { + "useMiddle": False, + "positiveOffsetX": 0, + "positiveOffsetY": 0, + "negativeOffsetX": 0, + "negativeOffsetY": 0, + "positiveScale": 0, + "negativeScale": 0.5, + "indices": mp.solutions.face_mesh.FACEMESH_RIGHT_EYE, + }, + "L_iris": { + "useMiddle": False, + "positiveOffsetX": 0, + "positiveOffsetY": 0, + "negativeOffsetX": 0, + "negativeOffsetY": 0, + "positiveScale": -0.2, + "negativeScale": 0.5, + "indices": mp.solutions.face_mesh.FACEMESH_LEFT_IRIS, + }, + "R_iris": { + "useMiddle": False, + "positiveOffsetX": 0, + "positiveOffsetY": 0, + "negativeOffsetX": 0, + "negativeOffsetY": 0, + "positiveScale": -0.2, + "negativeScale": 0.5, + "indices": mp.solutions.face_mesh.FACEMESH_RIGHT_IRIS, + }, + "face": { + "useMiddle": False, + "positiveOffsetX": 0, + "positiveOffsetY": 40, + "negativeOffsetX": 0, + "negativeOffsetY": 60, + "positiveScale": 0.2, + "negativeScale": 0.6, + "indices": mp.solutions.face_mesh.FACEMESH_FACE_OVAL, + }, + "mouth": { + "useMiddle": True, + "positiveOffsetX": 0, + "positiveOffsetY": 0, + "negativeOffsetX": 0, + "negativeOffsetY": 0, + "positiveScale": 0, + "negativeScale": 0, + "indices": mp.solutions.face_mesh.FACEMESH_LIPS, + }, + "mouth_in": { + "useMiddle": True, + "positiveOffsetX": 0, + "positiveOffsetY": 0, + "negativeOffsetX": 0, + "negativeOffsetY": 0, + "positiveScale": 0, + "negativeScale": 0, + "indices": mp.solutions.face_mesh.FACEMESH_LIPS, + }, +} + class SAMMultiLayer: def __init__(self): @@ -18,12 +100,6 @@ class SAMMultiLayer: @classmethod def INPUT_TYPES(s): - input_dir = folder_paths.get_input_directory() - files = [ - f - for f in os.listdir(input_dir) - if os.path.isfile(os.path.join(input_dir, f)) - ] return { "required": { "image": ("IMAGE",), @@ -32,7 +108,6 @@ class SAMMultiLayer: "STRING", {"multiline": False, "default": "embedding"}, ), - # "image": (sorted(files), ), "image_prompts_json": ("STRING", {"multiline": False, "default": "[]"}), }, } @@ -42,9 +117,129 @@ class SAMMultiLayer: RETURN_TYPES = ("SAM_PROMPT",) FUNCTION = "load_image" + def load_models(self, ckpt, model_type): + global global_predictor, face_landmarker + + ckpt = folder_paths.get_full_path("sams", ckpt) + sam = sam_model_registry[model_type](checkpoint=ckpt) # .to("cuda") + global_predictor = SamPredictor(sam) + + model_path = "/home/avatech/Desktop/projects/ComfyUI/custom_nodes/avatar-graph-comfyui/sam/face_landmarker.task" + options = FaceLandmarkerOptions( + base_options=BaseOptions(model_asset_path=model_path), + running_mode=VisionRunningMode.IMAGE, + ) + face_landmarker = FaceLandmarker.create_from_options(options) + + return global_predictor, face_landmarker + + def auto_segment(self, image, face_landmarks): + H, W, C = image.shape + imagePromptsMulti = {} + boxesMulti = {} + + for key, value in layerMapping.items(): + positivePoints = [] + middlePoints = [] + negativePoints = [] + + for index in value["indices"]: + start, end = index + startPoint = face_landmarks[start] + + startX = startPoint.x * W + startY = startPoint.y * H + + if len(middlePoints) == 0: + middlePoints.append({"x": startX, "y": startY, "label": 1}) + else: + middlePoints[0]["x"] += startX + middlePoints[0]["y"] += startY + + positivePoints.append({"x": startX, "y": startY, "label": 1}) + + len_indices = len(value["indices"]) + middlePoints[0]["x"] /= len_indices + middlePoints[0]["y"] /= len_indices + + if value["useMiddle"]: + imagePromptsMulti[key] = middlePoints + else: + for i, index in enumerate(value["indices"]): + start, end = index + startPoint = face_landmarks[start] + + startX = startPoint.x * W + startY = startPoint.y * H + + middlePoint = middlePoints[0] + directionVector = { + "x": middlePoint["x"] - startX, + "y": middlePoint["y"] - startY, + } + directionVectorLength = sqrt( + directionVector["x"] * directionVector["x"] + + directionVector["y"] * directionVector["y"] + ) + + if value["negativeScale"] != 0: + negativePointDistance = ( + value["negativeScale"] * directionVectorLength + ) + negativePoint = { + "x": startX + - (negativePointDistance * directionVector["x"]) + / directionVectorLength + - value["negativeOffsetX"], + "y": startY + - (negativePointDistance * directionVector["y"]) + / directionVectorLength + - value["negativeOffsetY"], + "label": 0, + } + negativePoints.append(negativePoint) + + positivePointDistance = ( + value["positiveScale"] * directionVectorLength + ) + positivePoints[i] = { + "x": positivePoints[i]["x"] + - (positivePointDistance * directionVector["x"]) + / directionVectorLength + - value["positiveOffsetX"], + "y": positivePoints[i]["y"] + - (positivePointDistance * directionVector["y"]) + / directionVectorLength + - value["positiveOffsetY"], + "label": 1, + } + + imagePromptsMulti[key] = positivePoints + negativePoints + + points = negativePoints if len(negativePoints) > 0 else positivePoints + box = np.array([ + min(x["x"] for x in points), + min(x["y"] for x in points), + max(x["x"] for x in points), + max(x["y"] for x in points), + ]) + boxesMulti[key] = box + + return imagePromptsMulti, boxesMulti + + def detect_face(self, np_image): + global face_landmarker + mp_image = mp.Image( + image_format=mp.ImageFormat.SRGB, data=(np_image * 255).astype(np.uint8) + ) + face_landmarks = face_landmarker.detect(mp_image).face_landmarks[0] + imagePromptsMulti, boxesMulti = self.auto_segment(np_image, face_landmarks) + + return imagePromptsMulti, boxesMulti + def load_image(self, image, ckpt, embedding_id, image_prompts_json): image_prompts = json.loads(image_prompts_json) - + order_file = f"{self.output_dir}/segments_{embedding_id}/order.json" if os.path.exists(order_file): # Frontend uploads segments images to backend => backend reads all segments images and passes them to next nodes @@ -54,7 +249,9 @@ class SAMMultiLayer: result = [image_prompts] for segment in order: - image = Image.open(f"{self.output_dir}/segments_{embedding_id}/{segment}.png") + image = Image.open( + f"{self.output_dir}/segments_{embedding_id}/{segment}.png" + ) image = np.array(image).astype(np.float32) / 255.0 image = torch.from_numpy(image)[None,] result.append(image) @@ -62,16 +259,12 @@ class SAMMultiLayer: return result else: # Frontend uploads clicks coordinates to backend => backend runs SAM and passes the segments to next nodes - global global_predictor - model_type = re.findall(r'vit_[lbh]', ckpt)[0] + model_type = re.findall(r"vit_[lbh]", ckpt)[0] + global global_predictor, face_landmarker if global_predictor is None: - ckpt = folder_paths.get_full_path("sams", ckpt) - sam = sam_model_registry[model_type](checkpoint=ckpt) - predictor = SamPredictor(sam) - global_predictor = predictor - - predictor = global_predictor + global_predictor, face_landmarker = self.load_models(ckpt, model_type) + print(face_landmarker) if image.shape[3] == 4: image = image[:, :, :, :3] @@ -79,14 +272,16 @@ class SAMMultiLayer: emb_filename = f"{self.output_dir}/{embedding_id}_{model_type}.npy" if not os.path.exists(emb_filename): image_np = (image[0].numpy() * 255).astype(np.uint8) - predictor.set_image(image_np) - emb = predictor.get_image_embedding().cpu().numpy() + global_predictor.set_image(image_np) + emb = global_predictor.get_image_embedding().cpu().numpy() np.save(emb_filename, emb) - with open(f"{self.output_dir}/{embedding_id}_{model_type}.json", "w") as f: + with open( + f"{self.output_dir}/{embedding_id}_{model_type}.json", "w" + ) as f: data = { - "input_size": predictor.input_size, - "original_size": predictor.original_size, + "input_size": global_predictor.input_size, + "original_size": global_predictor.original_size, } json.dump(data, f) else: @@ -94,33 +289,43 @@ class SAMMultiLayer: with open(f"{self.output_dir}/{embedding_id}_{model_type}.json") as f: data = json.load(f) - predictor.input_size = data["input_size"] - predictor.features = torch.from_numpy(emb) - predictor.is_image_set = True - predictor.original_size = data["original_size"] + global_predictor.input_size = data["input_size"] + global_predictor.features = torch.from_numpy(emb) + global_predictor.is_image_set = True + global_predictor.original_size = data["original_size"] + + imagePromptsMulti, boxesMulti = self.detect_face( + image[0].numpy().astype(np.float32) + ) image_prompts = json.loads(image_prompts_json) - result = [image_prompts] if isinstance(image_prompts, list): pass elif all(isinstance(item, list) for item in image_prompts.values()): - for item in image_prompts.values(): - if (len(item) == 0): + for key, item in image_prompts.items(): + if len(item) == 0: h, w, c = image[0].shape result.append(torch.zeros(1, h, w, c)) continue - point_coords = np.array([[p['x'], p['y']] for p in item]) - point_labels = np.array([p['label'] for p in item]) - masks, _, _ = predictor.predict( + points = ( + item + imagePromptsMulti[key] + if key in imagePromptsMulti + else item + ) + point_coords = np.array([[p["x"], p["y"]] for p in points]) + point_labels = np.array([p["label"] for p in points]) + + masks, _, _ = global_predictor.predict( point_coords=point_coords, point_labels=point_labels, + box=boxesMulti[key] if key in boxesMulti else None, ) masks = torch.from_numpy(masks) - masks = rearrange(masks[0], 'h w -> 1 h w') - out_image = repeat(masks, '1 h w -> 1 h w c', c=3) * image + masks = rearrange(masks[0], "h w -> 1 h w") + out_image = repeat(masks, "1 h w -> 1 h w c", c=3) * image result.append(out_image) return result