[Added] Support for batched load/save
- Also control over preview - And control over metadata
This commit is contained in:
+131
-31
@@ -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
@@ -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],)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user