fix: return 3d mask as possible

- MaskComposite requires 3d mask
This commit is contained in:
Dr.Lt.Data
2023-12-26 23:47:42 +09:00
parent da9d04d9a2
commit 7f70bf0eee
6 changed files with 26 additions and 18 deletions
+2
View File
@@ -310,6 +310,8 @@ NODE_CLASS_MAPPINGS = {
NODE_DISPLAY_NAME_MAPPINGS = {
"SAMLoader": "SAMLoader (Impact)",
"BboxDetectorSEGS": "BBOX Detector (SEGS)",
"SegmDetectorSEGS": "SEGM Detector (SEGS)",
"ONNXDetectorSEGS": "ONNX Detector (SEGS/legacy) - use BBOXDetector",
+1 -1
View File
@@ -2,7 +2,7 @@ import configparser
import os
version = "V4.53"
version = "V4.53.1"
dependency_version = 19
+1
View File
@@ -610,6 +610,7 @@ def make_sam_mask(sam_model, segs, image, detection_hint, dilation,
else:
mask = torch.zeros((8, 8), dtype=torch.float32, device="cpu") # empty mask
mask = utils.make_3d_mask(mask)
return mask
+2 -2
View File
@@ -203,8 +203,8 @@ class SimpleDetectorForEach:
@staticmethod
def detect(bbox_detector, image, bbox_threshold, bbox_dilation, crop_factor, drop_size,
sub_threshold, sub_dilation, sub_bbox_expansion,
sam_mask_hint_threshold, post_dilation=0, sam_model_opt=None, segm_detector_opt=None):
sub_threshold, sub_dilation, sub_bbox_expansion,
sam_mask_hint_threshold, post_dilation=0, sam_model_opt=None, segm_detector_opt=None):
if len(image) > 1:
raise Exception('[Impact Pack] ERROR: SimpleDetectorForEach does not allow image batches.\nPlease refer to https://github.com/ltdrdata/ComfyUI-extension-tutorials/blob/Main/ComfyUI-Impact-Pack/tutorial/batching-detailer.md for more information.')
+10 -10
View File
@@ -81,9 +81,10 @@ class CLIPSegDetectorProvider:
class SAMLoader:
@classmethod
def INPUT_TYPES(cls):
models = [x for x in folder_paths.get_filename_list("sams") if 'hq' not in x]
return {
"required": {
"model_name": (folder_paths.get_filename_list("sams"), ),
"model_name": (models, ),
"device_mode": (["AUTO", "Prefer GPU", "CPU"],),
}
}
@@ -1559,7 +1560,7 @@ class MasksToMaskList:
def doit(self, masks):
if masks is None:
empty_mask = torch.zeros((64,64), dtype=torch.float32, device="cpu")
empty_mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu")
return ([empty_mask], )
res = []
@@ -1569,6 +1570,8 @@ class MasksToMaskList:
print(f"mask len: {len(res)}")
res = [make_3d_mask(x) for x in res]
return (res, )
@@ -1589,23 +1592,20 @@ class MaskListToMaskBatch:
def doit(self, mask):
if len(mask) == 1:
if len(mask[0].shape) == 2:
mask = mask[0].unsqueeze(0)
mask = make_3d_mask(mask[0])
return (mask,)
elif len(mask) > 1:
mask1 = mask[0]
if len(mask1.shape) == 2:
mask1 = mask1.unsqueeze(0)
mask1 = make_3d_mask(mask[0])
for mask2 in mask[1:]:
if len(mask2.shape) == 2:
mask2 = mask2.unsqueeze(0)
mask2 = make_3d_mask(mask2)
if mask1.shape[1:] != mask2.shape[1:]:
mask2 = comfy.utils.common_upscale(mask2.movedim(-1, 1), mask1.shape[2], mask1.shape[1], "lanczos", "center").movedim(1, -1)
mask1 = torch.cat((mask1, mask2), dim=0)
return (mask1,)
else:
empty_mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu").unsqueeze(0)
empty_mask = torch.zeros((1, 64, 64), dtype=torch.float32, device="cpu").unsqueeze(0)
return (empty_mask,)
+10 -5
View File
@@ -643,6 +643,7 @@ class SEGSToMaskList:
if len(masks) == 0:
empty_mask = torch.zeros(segs[0], dtype=torch.float32, device="cpu")
masks = [empty_mask]
masks = [utils.make_3d_mask(mask) for mask in masks]
return (masks,)
@@ -655,14 +656,14 @@ class SEGSToMaskBatch:
}
RETURN_TYPES = ("MASK",)
OUTPUT_IS_LIST = (True,)
FUNCTION = "doit"
CATEGORY = "ImpactPack/Util"
def doit(self, segs):
masks = core.segs_to_masklist(segs)
mask_batch = torch.stack(masks, dim=0)
masks = [utils.make_3d_mask(mask) for mask in masks]
mask_batch = torch.concat(masks)
return (mask_batch,)
@@ -826,7 +827,9 @@ class DilateMask:
def doit(self, mask, dilation):
mask = core.dilate_mask(mask.numpy(), dilation)
return (torch.from_numpy(mask), )
mask = torch.from_numpy(mask)
mask = utils.make_3d_mask(mask)
return (mask, )
class GaussianBlurMask:
@@ -987,7 +990,7 @@ class EmptySEGS:
class SegsToCombinedMask:
@classmethod
def INPUT_TYPES(s):
return {"required": { "segs": ("SEGS",), } }
return {"required": {"segs": ("SEGS",), }}
RETURN_TYPES = ("MASK",)
FUNCTION = "doit"
@@ -995,7 +998,9 @@ class SegsToCombinedMask:
CATEGORY = "ImpactPack/Operation"
def doit(self, segs):
return (core.segs_to_combined_mask(segs),)
mask = core.segs_to_combined_mask(segs)
mask = utils.make_3d_mask(mask)
return (mask,)
class MediaPipeFaceMeshToSEGS: