[Added][Load Image Mask] To directly load masks
- Also added the used file name as output
This commit is contained in:
+98
-20
@@ -8,9 +8,9 @@ from . import main_logger
|
||||
|
||||
try:
|
||||
# We need to import the built-in LoadImage class for ImageDownload
|
||||
from nodes import LoadImage
|
||||
from nodes import LoadImage, LoadImageMask
|
||||
from folder_paths import get_input_directory, get_output_directory
|
||||
has_load_image = hasattr(LoadImage, "load_image")
|
||||
has_load_image = hasattr(LoadImage, "load_image") and hasattr(LoadImageMask, "load_image")
|
||||
from comfy.utils import common_upscale
|
||||
except Exception:
|
||||
has_load_image = False
|
||||
@@ -99,6 +99,27 @@ class CustomLoadImage(object):
|
||||
return (output_image, output_mask)
|
||||
|
||||
|
||||
class CustomLoadMask(object):
|
||||
def load_image(self, image, channel):
|
||||
image_path = image
|
||||
i = pillow(Image.open, image_path)
|
||||
i = pillow(ImageOps.exif_transpose, i)
|
||||
if i.getbands() != ("R", "G", "B", "A"):
|
||||
if i.mode == 'I':
|
||||
i = i.point(lambda i: i * (1 / 255))
|
||||
i = i.convert("RGBA")
|
||||
mask = None
|
||||
c = channel[0].upper()
|
||||
if c in i.getbands():
|
||||
mask = np.array(i.getchannel(c)).astype(np.float32) / 255.0
|
||||
mask = torch.from_numpy(mask)
|
||||
if c == 'A':
|
||||
mask = 1. - mask
|
||||
else:
|
||||
mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu")
|
||||
return (mask.unsqueeze(0),)
|
||||
|
||||
|
||||
def get_image_preview_info(file_name, where="input"):
|
||||
# This information is for the preview, as we are an output node and we return images
|
||||
# they will be displayed in our node. Quite simple.
|
||||
@@ -114,7 +135,9 @@ def get_image_preview_info(file_name, where="input"):
|
||||
|
||||
|
||||
def load_one_image(file_name, disp_name, embed_transparency):
|
||||
if os.path.isabs(file_name) and not os.path.exists(file_name):
|
||||
if not os.path.isabs(file_name):
|
||||
file_name = os.path.join(get_input_directory(), file_name)
|
||||
if not os.path.exists(file_name):
|
||||
raise ValueError(f"File '{file_name}' not found")
|
||||
|
||||
try:
|
||||
@@ -135,6 +158,7 @@ def load_one_image(file_name, disp_name, embed_transparency):
|
||||
|
||||
# Call the method and return its result directly
|
||||
result = loader_instance.load_image(file_name)
|
||||
result = (result[0], result[1], file_name)
|
||||
# Create an RGBA image if needed
|
||||
if embed_transparency:
|
||||
image, mask = result
|
||||
@@ -142,7 +166,7 @@ def load_one_image(file_name, disp_name, embed_transparency):
|
||||
mask = mask[..., None]
|
||||
# Concatenate image and mask into (b, h, w, 4)
|
||||
image_with_alpha = torch.cat([image, 1.0 - mask], dim=-1)
|
||||
result = (image_with_alpha, mask)
|
||||
result = (image_with_alpha, mask, file_name)
|
||||
return result
|
||||
|
||||
except Exception as e:
|
||||
@@ -152,6 +176,32 @@ def load_one_image(file_name, disp_name, embed_transparency):
|
||||
"It may be corrupt or in an unsupported format.") from e
|
||||
|
||||
|
||||
def load_one_mask(file_name, disp_name, channel='red'):
|
||||
if not os.path.isabs(file_name):
|
||||
file_name = os.path.join(get_input_directory(), file_name)
|
||||
if not os.path.exists(file_name):
|
||||
raise ValueError(f"File '{file_name}' not found")
|
||||
|
||||
try:
|
||||
if has_load_image:
|
||||
loader_instance = LoadImageMask()
|
||||
logger.debug(f"Calling built-in LoadImageMask.load_image() with filename: '{file_name}'")
|
||||
else:
|
||||
# Instantiate the built-in LoadImage node
|
||||
loader_instance = CustomLoadMask()
|
||||
logger.debug(f"Calling our CustomLoadMask.load_image() with filename: '{file_name}'")
|
||||
|
||||
# Call the method and return its result directly
|
||||
result = loader_instance.load_image(file_name, channel)
|
||||
return (result[0], file_name)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to load image '{disp_name}': {e}", exc_info=True)
|
||||
# Re-raise to make the error visible in ComfyUI
|
||||
raise IOError(f"Could not load the image file '{disp_name}'. "
|
||||
"It may be corrupt or in an unsupported format.") from e
|
||||
|
||||
|
||||
def load_image_wrapper(file_name, embed_transparency, disp_name=None, show_preview=True):
|
||||
disp_name = disp_name or file_name
|
||||
result = load_one_image(file_name, disp_name, embed_transparency)
|
||||
@@ -160,7 +210,8 @@ def load_image_wrapper(file_name, embed_transparency, disp_name=None, show_previ
|
||||
return {"ui": {"images": [get_image_preview_info(file_name)]}, "result": result}
|
||||
|
||||
|
||||
def load_images_wrapper(file_names, embed_transparency, disp_names=None, show_preview=True, batch_size=1):
|
||||
def load_images_wrapper(file_names, embed_transparency=False, disp_names=None, show_preview=True, batch_size=1,
|
||||
channel=None):
|
||||
# We work with lists
|
||||
if isinstance(file_names, str):
|
||||
file_names = [file_names]
|
||||
@@ -170,6 +221,7 @@ def load_images_wrapper(file_names, embed_transparency, disp_names=None, show_pr
|
||||
|
||||
imgs = []
|
||||
masks = []
|
||||
used_file_names = []
|
||||
all_preview_imgs = []
|
||||
total = len(file_names)
|
||||
for i in range(0, total, batch_size):
|
||||
@@ -186,42 +238,68 @@ def load_images_wrapper(file_names, embed_transparency, disp_names=None, show_pr
|
||||
file_name = file_names[i+j]
|
||||
disp_name = disp_names[i+j]
|
||||
|
||||
img, mask = load_one_image(file_name, disp_name, embed_transparency)
|
||||
max_w = max(max_w, img.shape[2])
|
||||
max_h = max(max_h, img.shape[1])
|
||||
imgs_batch.append(img)
|
||||
if channel is not None:
|
||||
# A mask
|
||||
mask, file_name = load_one_mask(file_name, disp_name, channel)
|
||||
else:
|
||||
img, mask, file_name = load_one_image(file_name, disp_name, embed_transparency)
|
||||
|
||||
max_w = max(max_w, img.shape[2])
|
||||
max_h = max(max_h, img.shape[1])
|
||||
imgs_batch.append(img)
|
||||
|
||||
max_mw = max(max_mw, mask.shape[2])
|
||||
max_mh = max(max_mh, mask.shape[1])
|
||||
masks_batch.append(mask)
|
||||
|
||||
used_file_names.append(file_name)
|
||||
|
||||
if show_preview:
|
||||
all_preview_imgs.append(get_image_preview_info(file_name))
|
||||
|
||||
for j in range(len(imgs_batch)):
|
||||
img = imgs_batch[j]
|
||||
H, W = img.shape[1:3]
|
||||
if H != max_h or W != max_w:
|
||||
logger.debug(f"Upscaling image to fit batch: {W}x{H} -> {max_w}x{max_h}")
|
||||
imgs_batch[j] = upscale_comfy(img, max_w, max_h, "bicubic")
|
||||
if channel is None:
|
||||
img = imgs_batch[j]
|
||||
H, W = img.shape[1:3]
|
||||
if H != max_h or W != max_w:
|
||||
logger.debug(f"Upscaling image to fit batch: {W}x{H} -> {max_w}x{max_h}")
|
||||
imgs_batch[j] = upscale_comfy(img, max_w, max_h, "bicubic")
|
||||
|
||||
mask = masks_batch[j]
|
||||
H, W = mask.shape[1:3]
|
||||
if H != max_mh or W != max_mw:
|
||||
logger.debug(f"Upscaling mask to fit batch: {W}x{H} -> {max_mw}x{max_mh}")
|
||||
masks_batch[j] = upscale_comfy(mask, max_mw, max_mh, "bicubic")
|
||||
imgs.append(torch.cat(imgs_batch))
|
||||
|
||||
if channel is None:
|
||||
imgs.append(torch.cat(imgs_batch))
|
||||
masks.append(torch.cat(masks_batch))
|
||||
else:
|
||||
# Add a single image
|
||||
file_name = file_names[i]
|
||||
img, mask = load_one_image(file_name, disp_names[i], embed_transparency)
|
||||
imgs.append(img)
|
||||
|
||||
if channel is not None:
|
||||
# A mask
|
||||
mask, file_name = load_one_mask(file_name, disp_names[i], channel)
|
||||
else:
|
||||
img, mask, file_name = load_one_image(file_name, disp_names[i], embed_transparency)
|
||||
imgs.append(img)
|
||||
|
||||
masks.append(mask)
|
||||
used_file_names.append(file_name)
|
||||
|
||||
if show_preview:
|
||||
all_preview_imgs.append(get_image_preview_info(file_name))
|
||||
logger.debug(f"Loaded {len(imgs)} batches:")
|
||||
|
||||
logger.debug(f"Loaded {len(masks)} batches:")
|
||||
for n, i in enumerate(imgs):
|
||||
logger.debug(f"{n}) {i.shape}")
|
||||
|
||||
result = (imgs, masks, used_file_names) if channel is None else (masks, used_file_names)
|
||||
if not show_preview:
|
||||
return (imgs, masks)
|
||||
return {"ui": {"images": all_preview_imgs}, "result": (imgs, masks)}
|
||||
return result
|
||||
|
||||
return {"ui": {"images": all_preview_imgs}, "result": result}
|
||||
|
||||
|
||||
def save_image(images, filenames, prompt=None, extra_pnginfo=None, compress_level=4, show_preview=True):
|
||||
|
||||
+53
-7
@@ -239,8 +239,8 @@ class ImageDownload:
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "MASK")
|
||||
RETURN_NAMES = ("image", "alpha_mask")
|
||||
RETURN_TYPES = ("IMAGE", "MASK", "STRING")
|
||||
RETURN_NAMES = ("image", "alpha_mask", "file_name")
|
||||
FUNCTION = "load_or_download_image"
|
||||
CATEGORY = BASE_CATEGORY + "/" + IO_CATEGORY
|
||||
DESCRIPTION = ("Downloads an image to ComfyUI's 'input' directory if it doesn't exist, then loads it using the "
|
||||
@@ -311,9 +311,9 @@ class ImageLoad:
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "MASK")
|
||||
RETURN_NAMES = ("image", "alpha_mask")
|
||||
OUTPUT_IS_LIST = (True, True)
|
||||
RETURN_TYPES = ("IMAGE", "MASK", "STRING")
|
||||
RETURN_NAMES = ("image", "alpha_mask", "file_name")
|
||||
OUTPUT_IS_LIST = (True, True, True)
|
||||
FUNCTION = "execute"
|
||||
CATEGORY = BASE_CATEGORY + "/" + IO_CATEGORY
|
||||
DESCRIPTION = ("Loads an image from any path")
|
||||
@@ -330,6 +330,48 @@ class ImageLoad:
|
||||
return load_images_wrapper(file_name, embed_transparency, show_preview=show_preview, batch_size=batch_size)
|
||||
|
||||
|
||||
class MaskLoad:
|
||||
_color_channels = ["red", "green", "blue", "alpha"]
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"file_name": ("STRING", {
|
||||
"tooltip": "The file name of the image to load"
|
||||
}),
|
||||
"batch_size": ("INT", {
|
||||
"default": 1,
|
||||
"min": 1,
|
||||
"max": 64,
|
||||
"tooltip": "The number of images to create in the batch"
|
||||
}),
|
||||
},
|
||||
"optional": {
|
||||
"channel": (cls._color_channels, ),
|
||||
"show_preview": SHOW_PREVIEW
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MASK", "STRING")
|
||||
RETURN_NAMES = ("mask", "file_name")
|
||||
OUTPUT_IS_LIST = (True, True)
|
||||
FUNCTION = "execute"
|
||||
CATEGORY = BASE_CATEGORY + "/" + IO_CATEGORY
|
||||
DESCRIPTION = ("Loads an image from any path using it as a mask")
|
||||
UNIQUE_NAME = "SET_MaskLoad"
|
||||
DISPLAY_NAME = "Load Mask from Path"
|
||||
INPUT_IS_LIST = True
|
||||
|
||||
def execute(self, file_name, batch_size, channel, show_preview):
|
||||
# Flatten arguments that aren't really expected to be lists
|
||||
batch_size = batch_size[0]
|
||||
channel = channel[0]
|
||||
show_preview = show_preview[0]
|
||||
|
||||
return load_images_wrapper(file_name, show_preview=show_preview, batch_size=batch_size, channel=channel)
|
||||
|
||||
|
||||
class ImageSave:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -759,8 +801,12 @@ class SaliencyEvaluationMetrics:
|
||||
pred = prediction[index_img].to(device)
|
||||
|
||||
# Ensure masks are normalized to [0, 1] range
|
||||
gt = batched_min_max_norm(gt, in_place=inputs_are_copies)
|
||||
pred = batched_min_max_norm(pred, in_place=inputs_are_copies)
|
||||
gt, skipped = batched_min_max_norm(gt, in_place=inputs_are_copies, ret_status=True)
|
||||
if skipped:
|
||||
logger.debug("GT mask already [0, 1]")
|
||||
pred, skipped = batched_min_max_norm(pred, in_place=inputs_are_copies, ret_status=True)
|
||||
if skipped:
|
||||
logger.debug("Prediction mask already [0, 1]")
|
||||
|
||||
for i in range(gt.shape[0]):
|
||||
# Get the next name
|
||||
|
||||
Reference in New Issue
Block a user