Files
ltdrdata-ComfyUI-Impact-Pack/comfyui-impact-pack.py
T
Dr.Lt.Data 873bbebc93 split output model to bbox_model, segm_model
show error message for invalid model
2023-03-31 01:03:56 +09:00

236 lines
7.8 KiB
Python

import os, sys, subprocess
from torchvision.datasets.utils import download_url
# ----- SETUP --------------------------------------------------------------
sys.path.insert(0, os.path.join(os.path.dirname(os.path.realpath(__file__)), "comfy"))
sys.path.append('../ComfyUI')
def packages_pip():
import sys, subprocess
return [r.decode().split('==')[0] for r in subprocess.check_output([sys.executable, '-m', 'pip', 'freeze']).split()]
def packages_mim():
import sys, subprocess
return [r.decode().split('==')[0] for r in subprocess.check_output([sys.executable, '-m', 'mim', 'list']).split()]
# INSTALL
print("### Check dependencies")
if "openmim" not in packages_pip():
subprocess.check_call([sys.executable, '-m', 'pip', '-U', 'install', 'openmim'])
if "mmcv-full" not in packages_mim():
subprocess.check_call([sys.executable, '-m', 'mim', 'install', 'mmcv-full==1.7.0'])
if "mmdet" not in packages_mim():
subprocess.check_call([sys.executable, '-m', 'mim', 'install', 'mmdet==2.28.2'])
# Download model
print("### Check basic models")
if os.path.realpath("..").endswith("custom_nodes"):
# For user
comfy_path = os.path.realpath("../..")
else:
# For development
comfy_path = os.path.realpath("../ComfyUI")
model_path = os.path.join(comfy_path, "models")
bbox_path = os.path.join(model_path, "mmdets", "bbox")
segm_path = os.path.join(model_path, "mmdets", "segm")
if not os.path.exists(os.path.join(bbox_path, "mmdet_anime-face_yolov3.pth")):
download_url("https://huggingface.co/dustysys/ddetailer/resolve/main/mmdet/bbox/mmdet_anime-face_yolov3.pth", bbox_path)
if not os.path.exists(os.path.join(bbox_path, "mmdet_anime-face_yolov3.py")):
download_url("https://huggingface.co/dustysys/ddetailer/raw/main/mmdet/bbox/mmdet_anime-face_yolov3.py", bbox_path)
if not os.path.exists(os.path.join(segm_path, "mmdet_dd-person_mask2former.pth")):
download_url("https://huggingface.co/dustysys/ddetailer/resolve/main/mmdet/segm/mmdet_dd-person_mask2former.pth", segm_path)
if not os.path.exists(os.path.join(segm_path, "mmdet_dd-person_mask2former.py")):
download_url("https://huggingface.co/dustysys/ddetailer/raw/main/mmdet/segm/mmdet_dd-person_mask2former.py", segm_path)
# ----- MAIN CODE --------------------------------------------------------------
# Core
import torch
import cv2
import mmcv
import numpy as np
from mmdet.core import get_classes
from mmdet.apis import (inference_detector,
init_detector)
from PIL import Image
import model_management
def load_mmdet(model_path):
print(model_management.vram_state)
model_config = os.path.splitext(model_path)[0] + ".py"
model = init_detector(model_config, model_path, device="cpu")
return model
def pil2tensor(image):
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
def create_segmasks(results):
segms = results[2]
segmasks = []
for i in range(len(segms)):
cv2_mask = segms[i].astype(np.uint8) * 255
mask = Image.fromarray(cv2_mask)
segmasks.append(mask)
return segmasks
def combine_masks(masks):
initial_cv2_mask = np.array(masks[0])
combined_cv2_mask = initial_cv2_mask
for i in range(1, len(masks)):
cv2_mask = np.array(masks[i])
combined_cv2_mask = cv2.bitwise_or(combined_cv2_mask, cv2_mask)
combined_mask = Image.fromarray(combined_cv2_mask)
return combined_mask
def inference_segm(model, image, conf_threshold):
image = image.numpy()[0] * 255
mmdet_results = inference_detector(model, image)
bbox_results, segm_results = mmdet_results
label = "A"
classes = get_classes("coco")
labels = [
np.full(bbox.shape[0], i, dtype=np.int32)
for i, bbox in enumerate(bbox_results)
]
n,m = bbox_results[0].shape
if (n == 0):
return [[],[],[]]
labels = np.concatenate(labels)
bboxes = np.vstack(bbox_results)
segms = mmcv.concat_list(segm_results)
filter_inds = np.where(bboxes[:,-1] > conf_threshold)[0]
results = [[],[],[]]
for i in filter_inds:
results[0].append(label + "-" + classes[labels[i]])
results[1].append(bboxes[i])
results[2].append(segms[i])
return results
return mmdet_results
def inference_bbox(model, image, conf_threshold):
image = image.numpy()[0] * 255
label = "A"
results = inference_detector(model, image)
cv2_image = np.array(image)
cv2_image = cv2_image[:, :, ::-1].copy()
cv2_gray = cv2.cvtColor(cv2_image, cv2.COLOR_BGR2GRAY)
segms = []
for (x0, y0, x1, y1, conf) in results[0]:
cv2_mask = np.zeros((cv2_gray.shape), np.uint8)
cv2.rectangle(cv2_mask, (int(x0), int(y0)), (int(x1), int(y1)), 255, -1)
cv2_mask_bool = cv2_mask.astype(bool)
segms.append(cv2_mask_bool)
n, m = results[0].shape
if (n == 0):
return [[], [], []]
bboxes = np.vstack(results[0])
filter_inds = np.where(bboxes[:, -1] > conf_threshold)[0]
results = [[], [], []]
for i in filter_inds:
results[0].append(label)
results[1].append(bboxes[i])
results[2].append(segms[i])
return results
# Nodes
import folder_paths
# folder_paths.supported_pt_extensions
folder_paths.folder_names_and_paths["mmdets_bbox"] = ([os.path.join(model_path, "mmdets", "bbox")], folder_paths.supported_pt_extensions)
folder_paths.folder_names_and_paths["mmdets_segm"] = ([os.path.join(model_path, "mmdets", "segm")], folder_paths.supported_pt_extensions)
folder_paths.folder_names_and_paths["mmdets"] = ([os.path.join(model_path, "mmdets")], folder_paths.supported_pt_extensions)
class NO_BBOX_MODEL:
ERROR = ""
class NO_SEGM_MODEL:
ERROR = ""
class MMDetLoader:
@classmethod
def INPUT_TYPES(s):
bboxs = [ "bbox/"+x for x in folder_paths.get_filename_list("mmdets_bbox") ]
segms = [ "segm/"+x for x in folder_paths.get_filename_list("mmdets_segm") ]
return {"required": { "model_name": (bboxs + segms, )}}
RETURN_TYPES = ("BBOX_MODEL", "SEGM_MODEL")
FUNCTION = "load_mmdet"
CATEGORY = "ImpactPack"
def load_mmdet(self, model_name):
mmdet_path = folder_paths.get_full_path("mmdets", model_name)
model = load_mmdet(mmdet_path)
if model_name.startswith("bbox"):
return (model, NO_SEGM_MODEL())
else:
return (NO_BBOX_MODEL(), model)
class SegmDetector:
input_dir = os.path.join(os.path.dirname(os.path.realpath(__file__)), "input")
@classmethod
def INPUT_TYPES(s):
return {"required":
{
"segm_model": ("SEGM_MODEL", ),
"image": ("IMAGE", ),
"threshold": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
}
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "doit"
CATEGORY = "ImpactPack"
def doit(self, segm_model, image, threshold):
mmdet_results = inference_segm(segm_model, image, threshold)
segmasks = create_segmasks(mmdet_results)
mask = combine_masks(segmasks)
image = pil2tensor(mask)
return (image,)
class BboxDetector(SegmDetector):
@classmethod
def INPUT_TYPES(s):
return {"required":
{
"bbox_model": ("BBOX_MODEL", ),
"image": ("IMAGE", ),
"threshold": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
}
}
def doit(self, bbox_model, image, threshold):
mmdet_results = inference_bbox(bbox_model, image, threshold)
segmasks = create_segmasks(mmdet_results)
mask = combine_masks(segmasks)
image = pil2tensor(mask)
return (image,)
NODE_CLASS_MAPPINGS = {
"MMDetLoader": MMDetLoader,
"BboxDetector": BboxDetector,
"SegmDetector": SegmDetector,
}