[Added] Support for batched load/save

- Also control over preview
- And control over metadata
This commit is contained in:
Salvador E. Tropea
2025-11-06 09:57:15 -03:00
parent e40b1f76a6
commit a58773147c
2 changed files with 195 additions and 71 deletions
+131 -31
View File
@@ -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": ()}
+64 -40
View File
@@ -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],)