feat: auto segment face in server

This commit is contained in:
Radionic
2023-11-24 17:24:06 +08:00
parent 5d15f76f46
commit 7b6ae6c5b2
2 changed files with 240 additions and 34 deletions
+1
View File
@@ -7,4 +7,5 @@ bpy
segment-anything
tqdm
python-dotenv
mediapipe
# -e git+https://github.com/facebookresearch/segment-anything.git#egg=segment_anything
+239 -34
View File
@@ -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