diff --git a/src/nodes/helpers.py b/src/nodes/helpers.py index 8843821..36751ef 100644 --- a/src/nodes/helpers.py +++ b/src/nodes/helpers.py @@ -9,14 +9,32 @@ from . import main_logger try: # We need to import the built-in LoadImage class for ImageDownload from nodes import LoadImage - from folder_paths import get_input_directory + from folder_paths import get_input_directory, get_output_directory has_load_image = hasattr(LoadImage, "load_image") + from comfy.utils import common_upscale except Exception: has_load_image = False logger = main_logger +def upscale(image, width, height, upscale_method): + # return F.interpolate(image, size=(height, width), mode=upscale_method) + return common_upscale(image, width, height, upscale_method, crop="disabled") + + +def upscale_comfy(image, width, height, upscale_method, crop="disabled"): + if image.dim == 3: + # A mask + if upscale_method == "lanczos": + # Lanczos needs an RGB image + return upscale(image.unsqueeze(1).repeat(1, 3, 1, 1), width, height, upscale_method, + crop=crop).movedim(1, -1)[:, :, :, 0] + else: + return upscale(image.unsqueeze(1), width, height, upscale_method, crop=crop).squeeze(1) + return common_upscale(image.movedim(-1, 1), width, height, upscale_method, crop=crop).movedim(1, -1) + + def pillow(fn, arg): prev_value = None try: @@ -81,12 +99,27 @@ class CustomLoadImage(object): return (output_image, output_mask) -def load_image_wrapper(file_name, embed_transparency, disp_name=None, show_preview=True): - disp_name = disp_name or file_name +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. + if os.path.isabs(file_name): + base = get_input_directory() if where == "input" else get_output_directory() + ff_name = os.path.relpath(file_name, base) + fname = os.path.basename(ff_name) + dname = os.path.dirname(ff_name) + else: + fname = file_name + dname = "" + return {"filename": fname, "subfolder": dname, "type": where} + + +def load_one_image(file_name, disp_name, embed_transparency): + if os.path.isabs(file_name) and not os.path.exists(file_name): + raise ValueError(f"File '{file_name}' not found") - # --- REUSE ComfyUI's LoadImage LOGIC --- try: if has_load_image: + # --- REUSE ComfyUI's LoadImage LOGIC --- # Instantiate the built-in LoadImage node loader_instance = LoadImage() @@ -110,43 +143,103 @@ def load_image_wrapper(file_name, embed_transparency, disp_name=None, show_previ # 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) - # No preview - if not show_preview: - return result - # 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. - if os.path.isabs(file_name): - ff_name = os.path.relpath(file_name, get_input_directory()) - fname = os.path.basename(ff_name) - dname = os.path.dirname(ff_name) - else: - fname = file_name - dname = "" - downloaded_file = { - "images": [{ - "filename": fname, - "subfolder": dname, - "type": "input" # We stored the file in the "input" folder - }] - } - return {"ui": downloaded_file, "result": result} + return result except Exception as e: - logger.error(f"Failed to load image '{disp_name}' using built-in LoadImage node: {e}", exc_info=True) + 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}' using the standard loader. " + raise IOError(f"Could not load the image file '{disp_name}'. " "It may be corrupt or in an unsupported format.") from e -def save_image(images, filenames, prompt=None, extra_pnginfo=None, compress_level=4): +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) + if not show_preview: + return result + 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): + # We work with lists + if isinstance(file_names, str): + file_names = [file_names] + disp_names = disp_names or file_names + if isinstance(disp_names, str): + disp_names = [disp_names] + + imgs = [] + masks = [] + all_preview_imgs = [] + total = len(file_names) + for i in range(0, total, batch_size): + if batch_size > 1: + # Add upto batch_size images in a batch + imgs_batch = [] + masks_batch = [] + max_w = max_h = 0 + max_mw = max_mh = 0 + for j in range(batch_size): + index = i+j + if index >= total: + continue + 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) + max_mw = max(max_mw, mask.shape[2]) + max_mh = max(max_mh, mask.shape[1]) + masks_batch.append(mask) + 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") + 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)) + 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) + masks.append(mask) + if show_preview: + all_preview_imgs.append(get_image_preview_info(file_name)) + logger.debug(f"Loaded {len(imgs)} batches:") + for n, i in enumerate(imgs): + logger.debug(f"{n}) {i.shape}") + if not show_preview: + return (imgs, masks) + return {"ui": {"images": all_preview_imgs}, "result": (imgs, masks)} + + +def save_image(images, filenames, prompt=None, extra_pnginfo=None, compress_level=4, show_preview=True): if isinstance(filenames, str): if not filenames: raise ValueError("You must provide a file name") filenames = [filenames] - B = images.shape[0] - if len(filenames) != B: - raise ValueError(f"{B} images provided but only {len(filenames)} file names") - for batch_number, (image, filename) in enumerate(zip(images, filenames)): + + # Make a list with all the images, unroll the batches + imgs = [] + for img in images: + imgs.extend([i for i in img]) + + if len(filenames) != len(imgs): + raise ValueError(f"{len(imgs)} images provided but only {len(filenames)} file names") + + all_preview_imgs = [] + for index, (image, filename) in enumerate(zip(imgs, filenames)): i = 255. * image.cpu().numpy() img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8)) metadata = PngInfo() @@ -158,3 +251,10 @@ def save_image(images, filenames, prompt=None, extra_pnginfo=None, compress_leve logger.debug(f"Saving {image.shape[1]}x{image.shape[0]} image to {filename}") img.save(filename, pnginfo=metadata, compress_level=compress_level) + + if show_preview: + all_preview_imgs.append(get_image_preview_info(filename, where="output")) + + if not show_preview: + return () + return {"ui": {"images": all_preview_imgs}, "result": ()} diff --git a/src/nodes/nodes_img.py b/src/nodes/nodes_img.py index b8f93ed..5d8eb4e 100644 --- a/src/nodes/nodes_img.py +++ b/src/nodes/nodes_img.py @@ -26,11 +26,10 @@ from typing import Optional # We are the main source, so we use the main_logger from . import main_logger -from .helpers import load_image_wrapper, save_image +from .helpers import load_image_wrapper, load_images_wrapper, save_image, upscale, upscale_comfy try: from folder_paths import get_input_directory, get_output_directory from comfy import model_management - from comfy.utils import common_upscale except ModuleNotFoundError: # No ComfyUI, this is a test environment def get_input_directory(): @@ -97,6 +96,18 @@ NORM_PARAM = ("FLOAT", { "step": 0.1, "display": "number"}) MAX_FILES = 0xffffffffffffffff +EMBED_TRANSPARENCY = ("BOOLEAN", { + "default": False, + "tooltip": "Create RGBA images when they have transparency."}) +SAVE_PROMPT = ("BOOLEAN", { + "default": False, + "tooltip": "Save prompt submitted to ComfyUI"}) +SAVE_WORKFLOW = ("BOOLEAN", { + "default": False, + "tooltip": "Save the ComfyUI workflow"}) +SHOW_PREVIEW = ("BOOLEAN", { + "default": True, + "tooltip": "Show a preview of the images"}) # A dictionary to cache loaded fonts font_cache = {} @@ -117,11 +128,6 @@ def pil_to_tensor(pil_image: Image.Image) -> torch.Tensor: return torch.from_numpy(np_image) -def upscale(image, width, height, upscale_method): - # return F.interpolate(image, size=(height, width), mode=upscale_method) - return common_upscale(image, width, height, upscale_method, crop="disabled") - - def parse_size(size_str, reference_dim): """Parses a size string which can be pixels or a percentage.""" size_str = size_str.strip() @@ -211,10 +217,7 @@ class ImageDownload: "default": "", "tooltip": "The name used locally. Leave empty to use `filename`" }), - "embed_transparency": ("BOOLEAN", { - "default": False, - "tooltip": "Create RGBA images when they have transparency." - }), + "embed_transparency": EMBED_TRANSPARENCY, } } @@ -277,28 +280,36 @@ class ImageLoad: "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": { - "embed_transparency": ("BOOLEAN", { - "default": False, - "tooltip": "Create RGBA images when they have transparency." - }), + "embed_transparency": EMBED_TRANSPARENCY, + "show_preview": SHOW_PREVIEW } } RETURN_TYPES = ("IMAGE", "MASK") RETURN_NAMES = ("image", "alpha_mask") + OUTPUT_IS_LIST = (True, True) FUNCTION = "execute" CATEGORY = BASE_CATEGORY + "/" + IO_CATEGORY DESCRIPTION = ("Loads an image from any path") UNIQUE_NAME = "SET_ImageLoad" DISPLAY_NAME = "Load Image from Path" + INPUT_IS_LIST = True - def execute(self, file_name: str, embed_transparency: bool = False): - if not os.path.exists(file_name): - raise ValueError(f"File '{file_name}' not found") + def execute(self, file_name, batch_size, embed_transparency, show_preview): + # Flatten arguments that aren't really expected to be lists + batch_size = batch_size[0] + embed_transparency = embed_transparency[0] + show_preview = show_preview[0] - return load_image_wrapper(file_name, embed_transparency, show_preview=False) + return load_images_wrapper(file_name, embed_transparency, show_preview=show_preview, batch_size=batch_size) class ImageSave: @@ -309,8 +320,14 @@ class ImageSave: "image": ("IMAGE", {"tooltip": "The images to save."}), "filename": ("STRING", {"default": "", "tooltip": "The file name for the image"}) }, + "optional": { + "show_preview": SHOW_PREVIEW, + "save_prompt": SAVE_PROMPT, + "save_workflow": SAVE_WORKFLOW + }, "hidden": { - "prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO" + "prompt": "PROMPT", + "extra_pnginfo": "EXTRA_PNGINFO" }, } @@ -318,15 +335,27 @@ class ImageSave: FUNCTION = "execute" OUTPUT_NODE = True + INPUT_IS_LIST = True CATEGORY = BASE_CATEGORY + "/" + IO_CATEGORY DESCRIPTION = ("Saves an image to an arbitrary path") UNIQUE_NAME = "SET_ImageSave" DISPLAY_NAME = "Save Image to Path" - def execute(self, image, filename, prompt=None, extra_pnginfo=None): - save_image(image, filename, prompt, extra_pnginfo) - return () + def execute(self, image, filename, show_preview, save_prompt, save_workflow, prompt=None, + extra_pnginfo=None): + # Flatten arguments that aren't really expected to be lists + if not save_prompt[0]: + prompt = None + elif prompt is not None: + prompt = prompt[0] + if not save_workflow[0]: + extra_pnginfo = None + elif extra_pnginfo is not None: + extra_pnginfo = extra_pnginfo[0] + show_preview = show_preview[0] + + return save_image(image, filename, prompt, extra_pnginfo, show_preview=show_preview) class MaskSave: @@ -337,21 +366,24 @@ class MaskSave: "mask": ("MASK", {"tooltip": "The mask to save."}), "filename": ("STRING", {"default": "", "tooltip": "The file name for the image"}) }, + "optional": { + "show_preview": SHOW_PREVIEW, + }, } RETURN_TYPES = () FUNCTION = "execute" OUTPUT_NODE = True + INPUT_IS_LIST = True CATEGORY = BASE_CATEGORY + "/" + IO_CATEGORY DESCRIPTION = ("Saves a mask to an arbitrary path") UNIQUE_NAME = "SET_MaskSave" DISPLAY_NAME = "Save Mask to Path" - def execute(self, mask, filename): - save_image(mask, filename) - return () + def execute(self, mask, filename, show_preview): + return save_image(mask, filename, show_preview=show_preview[0]) class ImageDataset: @@ -442,9 +474,9 @@ class ImageDataset: images, _, _ = cls.generate_lists(source, pattern, destination, dest_ext, reference, sort_method, MAX_FILES, skip_first_images, select_every_nth, random_seed) except ValueError: - logger.debug(f"IS_CHANGED -> 0 ValueError") + logger.debug("No more images, we got ValueError (IS_CHANGED -> 0 )") return 0 - logger.debug(f"IS_CHANGED -> {len(images)}") + logger.debug(f"We have more images to process (IS_CHANGED -> {len(images)})") return len(images) def execute(self, source, pattern, destination, dest_ext, reference=None, sort_method="None", @@ -1153,7 +1185,7 @@ class ImagePad: target_height = H if extra_padding > 0: - image = upscale(image.movedim(-1, 1), W - extra_padding, H - extra_padding, BEST_UPSCALE).movedim(1, -1) + image = upscale_comfy(image, W - extra_padding, H - extra_padding, BEST_UPSCALE) B, H, W, C = image.shape padded_width = target_width @@ -1469,14 +1501,10 @@ class ImageResize: out_mask = out_mask.narrow(-1, x, crop_w).narrow(-2, y, crop_h) # Resize the image - out_image = upscale(out_image.movedim(-1, 1), width, height, upscale_method).movedim(1, -1) + out_image = upscale_comfy(out_image.movedim, width, height, upscale_method) if out_mask is not None: - if upscale_method == "lanczos": - out_mask = upscale(out_mask.unsqueeze(1).repeat(1, 3, 1, 1), width, height, - upscale_method).movedim(1, -1)[:, :, :, 0] - else: - out_mask = upscale(out_mask.unsqueeze(1), width, height, upscale_method).squeeze(1) + out_mask = upscale_comfy(out_mask, width, height, upscale_method) # Pad logic if (keep_proportion.startswith("pad") or pillarbox_blur) and (pad_left > 0 or pad_right > 0 or pad_top > 0 @@ -1603,11 +1631,7 @@ class ResizeMask: width = round(ow*ratio) height = round(oh*ratio) - if upscale_method == "lanczos": - out_mask = common_upscale(mask.unsqueeze(1).repeat(1, 3, 1, 1), width, height, upscale_method, - crop=crop).movedim(1, -1)[:, :, :, 0] - else: - out_mask = common_upscale(mask.unsqueeze(1), width, height, upscale_method, crop=crop).squeeze(1) + out_mask = upscale_comfy(mask, width, height, upscale_method, crop=crop) return (out_mask, out_mask.shape[2], out_mask.shape[1],)