diff --git a/README.md b/README.md new file mode 100644 index 0000000..6704719 --- /dev/null +++ b/README.md @@ -0,0 +1,46 @@ +# Yet another custom node to detect human parts + +Detect human parts using the DeepLabV3+ ResNet50 model from Keras-io. You can extract hair, arms, legs, and other parts +with ease and with small memory usage. + +This node aims to detect human parts using the model created by +[Keras-io](https://huggingface.co/keras-io/deeplabv3p-resnet50). Their "[Space](https://huggingface.co/spaces/keras-io/Human-Part-Segmentation)" was impressive, and I wanted to use the +model. + +Unfortunately, the model uses an old Keras version, and there were no PyTorch implementation. + +So I decided to convert the model to [ONNX](https://onnx.ai/) format and to create my [HugginFace +repository](https://huggingface.co/Metal3d/deeplabv3p-resnet50-human) to share the model with the community. + +> Fortunately, Keras provides the model with a CC1.0 license, thank you guys to allow us to use it without any +> restriction. + +## Example + +You can drag and drop the following image to try: + +![Example workflow](./images/Human Parts.png) + +## DeepLabV3+ ResNet50 for Human + +Actually, all the model I found was not trained to detect human parts, but to detect some objects or urban elements. The +Keras model is the only one I found that works! + +## Installation + +I strongly recommend to use ComfyUI-Manager to install the node. It will install the dependencies and the model. + +If you're using the command line, you can install the node with: + +```bash +cd /path/to/your/ComfyUI/custom_nodes +git clone ... +cd HumanParts +pip install -r requirements.txt +# or +python -m pip install -r requirements.txt +``` + +Then, restart ComfyUI and you may find the "HumanParts" node. + +![The node](./images/node.png) diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..ee69680 --- /dev/null +++ b/__init__.py @@ -0,0 +1,12 @@ +__all__ = ["HumanParts"] + +from .nodes import HumanParts + +NODE_CLASS_MAPPINGS = { + "HumanParts": HumanParts, +} + +# A dictionary that contains the friendly/humanly readable titles for the nodes +NODE_DISPLAY_NAME_MAPPINGS = { + "HumanParts": "🧍 Human Parts mask generator", +} diff --git a/detector/__init__.py b/detector/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/detector/human_parts.py b/detector/human_parts.py new file mode 100644 index 0000000..ff1a228 --- /dev/null +++ b/detector/human_parts.py @@ -0,0 +1,86 @@ +from typing import Tuple + +import numpy as np +import torch +from onnxruntime import InferenceSession +from PIL import Image + +# classes used in the model +classes = { + "background": 0, + "hair": 2, + "glasses": 4, + "top-clothes": 5, + "bottom-clothes": 9, + "torso-skin": 10, + "face": 13, + "left-arm": 14, + "right-arm": 15, + "left-leg": 16, + "right-leg": 17, + "left-foot": 18, + "right-foot": 19, +} + + +def get_mask( + image: torch.Tensor, model: InferenceSession, rotation: float, **kwargs +) -> Tuple[torch.Tensor, int]: + """ + Return a Tensor with the mask of the human parts in the image. + + The rotation parameter is not used for now. The idea is to propose rotation to help + the model to detect the human parts in the image if the character is not in a casual position. + Several tests have been done, but the model seems to fail to detect the human parts in these cases, + and the rotation does not help. + """ + + image = image.squeeze(0) + image_np = image.numpy() * 255 + + pil_image = Image.fromarray(image_np.astype(np.uint8)) + original_size = pil_image.size # to resize the mask later + # resize to 512x512 as the model expects + pil_image = pil_image.resize((512, 512)) + center = (256, 256) + + if rotation != 0: + pil_image = pil_image.rotate(rotation, center=center) + + # normalize the image + image_np = np.array(pil_image).astype(np.float32) / 127.5 - 1 + image_np = np.expand_dims(image_np, axis=0) + + # use the onnx model to get the mask + input_name = model.get_inputs()[0].name + output_name = model.get_outputs()[0].name + result = model.run([output_name], {input_name: image_np}) + result = np.array(result[0]).argmax(axis=3).squeeze(0) + + score: int = 0 + + mask = np.zeros_like(result) + for class_name, enabled in kwargs.items(): + if enabled and class_name in classes: + class_index = classes[class_name] + detected = result == class_index + mask[detected] = 255 + score += mask.sum() + + # back to the original size + mask_image = Image.fromarray(mask.astype(np.uint8), mode="L") + if rotation != 0: + mask_image = mask_image.rotate(-rotation, center=center) + + mask_image = mask_image.resize(original_size) + + # and back to numpy... + mask = np.array(mask_image).astype(np.float32) / 255 + + # add 2 dimensions to match the expected output + mask = np.expand_dims(mask, axis=0) + mask = np.expand_dims(mask, axis=0) + # ensure to return a "binary mask_image" + + del image_np, result # free up memory, maybe not necessary + return (torch.from_numpy(mask.astype(np.uint8)), score) diff --git a/images/Human Parts.png b/images/Human Parts.png new file mode 100644 index 0000000..4ddc08b Binary files /dev/null and b/images/Human Parts.png differ diff --git a/images/node.png b/images/node.png new file mode 100644 index 0000000..8761312 Binary files /dev/null and b/images/node.png differ diff --git a/install.py b/install.py new file mode 100644 index 0000000..e5162d0 --- /dev/null +++ b/install.py @@ -0,0 +1,35 @@ +import os +import urllib.request + +from tqdm import tqdm + +try: + from .utils import model_name, model_path, model_url, models_dir_path +except ImportError: + from utils import model_name, model_path, model_url, models_dir_path + + +def download(url, path, name): + request = urllib.request.urlopen(url) + total = int(request.headers.get("Content-Length", 0)) + with tqdm( + total=total, + desc=f"[HumanParts] Downloading {name} to {path}", + unit="B", + unit_scale=True, + unit_divisor=1024, + ) as progress: + urllib.request.urlretrieve( + url, + path, + reporthook=lambda count, block_size, total_size: progress.update( + block_size + ), + ) + + +if not os.path.exists(models_dir_path): + os.makedirs(models_dir_path) + +if not os.path.exists(model_path): + download(model_url, model_path, model_name) diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..fb42dc5 --- /dev/null +++ b/nodes.py @@ -0,0 +1,92 @@ +from typing import Tuple + +import onnxruntime as ort +import torch + +from .utils import model_path +from .detector.human_parts import get_mask + + +class HumanParts: + """ + This node is used to get a mask of the human parts in the image. + + The model used is DeepLabV3+ with a ResNet50 backbone trained + by Keras-io, converted to ONNX format. + + """ + + RETURN_TYPES = ("MASK",) + RETURN_NAMES = ("mask",) + FUNCTION = "get_mask" + CATEGORY = "Metal3d" + + @classmethod + def INPUT_TYPES(cls): + def _bool_widget(is_enabled=False, tooltip: str | None = None): + """Helper function to create a boolean widget""" + return ( + "BOOLEAN", + { + "default": is_enabled, + "label_on": "Enabled", + "label_off": "Disabled", + "tooltip": tooltip, + }, + ) + + return { + "required": { + "image": ("IMAGE",), + "background": _bool_widget( + tooltip="Background, excluding human parts, invert this mask to get the human parts", + ), + "face": _bool_widget( + is_enabled=True, + tooltip="Face, including eyes, mouth, etc.", + ), + "hair": _bool_widget( + tooltip="Hair, including beard, mustache, etc.", + ), + "glasses": _bool_widget( + tooltip="Glasses, sunglasses, etc. Eyes can be included" + ), + "top-clothes": _bool_widget( + tooltip="Shirt, T-shirt, etc.", + ), + "bottom-clothes": _bool_widget( + tooltip="Pants, shorts, etc.", + ), + "torso-skin": _bool_widget( + tooltip="Skin of the torso, excluding clothes. Neck can be included" + ), + "left-arm": _bool_widget( + tooltip="Left arm, excluding clothes, hand can be included" + ), + "right-arm": _bool_widget( + tooltip="Right arm, excluding clothes, hand can be included" + ), + "left-leg": _bool_widget( + tooltip="Left leg, excluding clothes, foot can be included" + ), + "right-leg": _bool_widget( + tooltip="Right leg, excluding clothes, foot can be included" + ), + "left-foot": _bool_widget( + tooltip="Left foot, excluding shoes", + ), + "right-foot": _bool_widget( + tooltip="Right foot, excluding shoes", + ), + } + } + + def get_mask(self, image: torch.Tensor, **kwargs) -> Tuple[torch.Tensor]: + """ + Return a Tensor with the mask of the human parts in the image. + """ + + model = ort.InferenceSession(model_path) + ret_tensor, _ = get_mask(image, model=model, rotation=0, **kwargs) + + return (ret_tensor,) diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..7cabeb7 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,3 @@ +onnxruntime +numpy +Pillow diff --git a/utils.py b/utils.py new file mode 100644 index 0000000..48de26f --- /dev/null +++ b/utils.py @@ -0,0 +1,14 @@ +import os + +# get the model paths +try: + from folder_paths import models_dir # pyright: ignore +except ImportError: + from pathlib import Path + + models_dir = os.path.join(Path(__file__).parents[2], "models") + +models_dir_path = os.path.join(models_dir, "onnx", "human-parts") +model_url = "https://huggingface.co/Metal3d/deeplabv3p-resnet50-human/resolve/main/deeplabv3p-resnet50-human.onnx" +model_name = os.path.basename(model_url) +model_path = os.path.join(models_dir_path, "deeplabv3p-resnet50-human.onnx")