Initial commit

Maybe it will not work at all...
This commit is contained in:
Patrice Ferlet
2024-09-11 22:09:11 +02:00
parent 34b7f003e3
commit 2a3fe2fb10
10 changed files with 288 additions and 0 deletions
+46
View File
@@ -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)
+12
View File
@@ -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",
}
View File
+86
View File
@@ -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)
Binary file not shown.

After

Width:  |  Height:  |  Size: 324 KiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 46 KiB

+35
View File
@@ -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)
+92
View File
@@ -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,)
+3
View File
@@ -0,0 +1,3 @@
onnxruntime
numpy
Pillow
+14
View File
@@ -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")