[Added][Load Image Mask] To directly load masks

- Also added the used file name as output
This commit is contained in:
Salvador E. Tropea
2025-11-10 13:59:41 -03:00
parent 64eafc6402
commit 83264cccab
2 changed files with 151 additions and 27 deletions
+98 -20
View File
@@ -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
View File
@@ -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