change detector's output type from IMAGE to MASK

This commit is contained in:
Dr.Lt.Data
2023-03-31 19:38:34 +09:00
parent 34a3dce952
commit d0ccce9264
4 changed files with 826 additions and 216 deletions
+2 -2
View File
@@ -7,8 +7,8 @@
* SegmDetector - Detect segmentation from input image.
* BboxDetector - Detect bbox from input image.
### ![example](https://user-images.githubusercontent.com/128333288/228896955-d84e164e-af87-4c4d-a377-662a6363381f.png)
Detected mask example.
### ![example](https://user-images.githubusercontent.com/128333288/229093882-239c4a2a-0c94-4b77-b0fa-8ac07eda17b9.png)
Head detection example.
# Installation
+92 -21
View File
@@ -14,18 +14,23 @@ def packages_mim():
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'])
print("Loading: ComfyUI-Impact-Pack")
print("### ComfyUI-Impact-Pack: Check dependencies")
installed_pip = packages_pip()
if "mmcv-full" not in packages_mim():
if "openmim" not in installed_pip:
subprocess.check_call([sys.executable, '-m', 'pip', 'install', '-U', 'openmim'])
installed_mim = packages_mim()
if "mmcv-full" not in installed_mim:
subprocess.check_call([sys.executable, '-m', 'mim', 'install', 'mmcv-full==1.7.0'])
if "mmdet" not in packages_mim():
if "mmdet" not in installed_mim:
subprocess.check_call([sys.executable, '-m', 'mim', 'install', 'mmdet==2.28.2'])
# Download model
print("### Check basic models")
print("### ComfyUI-Impact-Pack: Check basic models")
if os.path.realpath("..").endswith("custom_nodes"):
# For user
@@ -66,7 +71,6 @@ 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
@@ -78,9 +82,7 @@ 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)
segmasks.append(segms[i].astype(np.float32))
return segmasks
def combine_masks(masks):
@@ -90,8 +92,35 @@ def combine_masks(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
# combined_mask = Image.fromarray(combined_cv2_mask)
# return combined_mask
mask = torch.from_numpy(combined_cv2_mask)
return mask
def bitwise_and_masks(mask1, mask2):
cv2_mask1 = np.array(mask1)
cv2_mask2 = np.array(mask2)
cv2_mask = cv2.bitwise_and(cv2_mask1, cv2_mask2)
mask = torch.from_numpy(cv2_mask)
return mask
def dilate_masks(masks, dilation_factor, iter=1):
if dilation_factor == 0:
return masks
dilated_masks = []
kernel = np.ones((dilation_factor,dilation_factor), np.uint8)
for i in range(len(masks)):
cv2_mask = masks[i]
dilated_mask = cv2.dilate(cv2_mask, kernel, iter)
dilated_masks.append(dilated_mask)
return dilated_masks
def subtract_masks(mask1, mask2):
cv2_mask1 = np.array(mask1) * 255
cv2_mask2 = np.array(mask2) * 255
cv2_mask = cv2.subtract(cv2_mask1, cv2_mask2)
mask = torch.from_numpy(cv2_mask) / 255.0
return mask
def inference_segm(model, image, conf_threshold):
image = image.numpy()[0] * 255
@@ -192,21 +221,22 @@ class SegmDetector:
"segm_model": ("SEGM_MODEL", ),
"image": ("IMAGE", ),
"threshold": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
"dilation": ("INT", {"default": 0, "min": 0, "max": 255, "step": 1}),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_TYPES = ("MASK",)
FUNCTION = "doit"
CATEGORY = "ImpactPack"
def doit(self, segm_model, image, threshold):
def doit(self, segm_model, image, threshold, dilation):
mmdet_results = inference_segm(segm_model, image, threshold)
segmasks = create_segmasks(mmdet_results)
if dilation > 0:
segmasks = dilate_masks(segmasks, dilation)
mask = combine_masks(segmasks)
image = pil2tensor(mask)
return (image,)
return (mask,)
class BboxDetector(SegmDetector):
@classmethod
@@ -216,21 +246,62 @@ class BboxDetector(SegmDetector):
"bbox_model": ("BBOX_MODEL", ),
"image": ("IMAGE", ),
"threshold": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
"dilation": ("INT", {"default": 4, "min": 0, "max": 255, "step": 1}),
}
}
def doit(self, bbox_model, image, threshold):
def doit(self, bbox_model, image, threshold, dilation):
mmdet_results = inference_bbox(bbox_model, image, threshold)
segmasks = create_segmasks(mmdet_results)
if dilation > 0:
segmasks = dilate_masks(segmasks, dilation)
mask = combine_masks(segmasks)
return (mask,)
image = pil2tensor(mask)
return (image,)
class BitwiseAndMask:
@classmethod
def INPUT_TYPES(s):
return {"required":
{
"mask1": ("MASK", ),
"mask2": ("MASK", ),
}
}
RETURN_TYPES = ("MASK",)
FUNCTION = "doit"
CATEGORY = "ImpactPack"
def doit(self, mask1, mask2):
mask = bitwise_and_masks(mask1, mask2)
return (mask,)
class SubtractMask:
@classmethod
def INPUT_TYPES(s):
return {"required":
{
"mask1": ("MASK", ),
"mask2": ("MASK", ),
}
}
RETURN_TYPES = ("MASK",)
FUNCTION = "doit"
CATEGORY = "ImpactPack"
def doit(self, mask1, mask2):
mask = subtract_masks(mask1, mask2)
return (mask,)
NODE_CLASS_MAPPINGS = {
"MMDetLoader": MMDetLoader,
"BboxDetector": BboxDetector,
"SegmDetector": SegmDetector,
}
"BitwiseAndMask": BitwiseAndMask,
"SubtractMask": SubtractMask,
}
+732 -193
View File
File diff suppressed because it is too large Load Diff
BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 204 KiB

After

Width:  |  Height:  |  Size: 553 KiB