[Image/Mask loaders] Allow image recycle

So when we have multiple instances of the same image we load just once
This commit is contained in:
Salvador E. Tropea
2025-12-29 14:05:52 -03:00
parent 9ccef62bc6
commit 7e5f7c3c0a
2 changed files with 25 additions and 8 deletions
+20 -8
View File
@@ -147,13 +147,16 @@ def get_image_preview_info(file_name, where="input"):
return {"filename": fname, "subfolder": dname, "type": where}
def load_one_image(file_name, disp_name, embed_transparency):
def load_one_image(file_name, disp_name, embed_transparency, known=None):
if file_name is None:
return (empty_image(b=1, c=3), empty_image(b=1), 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")
if known is not None and file_name in known:
logger.debug('Recycling '+file_name)
return (known[file_name][0], known[file_name][1], file_name)
try:
if has_load_image:
@@ -182,6 +185,8 @@ def load_one_image(file_name, disp_name, embed_transparency):
# 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, file_name)
if known is not None:
known[file_name] = (result[0], result[1])
return result
except Exception as e:
@@ -191,13 +196,16 @@ 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'):
def load_one_mask(file_name, disp_name, channel='red', known=None):
if file_name is None:
return (empty_image(b=1), 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")
if known is not None and file_name in known:
logger.debug('Recycling '+file_name)
return (known[file_name], file_name)
try:
if has_load_image:
@@ -210,6 +218,8 @@ def load_one_mask(file_name, disp_name, channel='red'):
# Call the method and return its result directly
result = loader_instance.load_image(file_name, channel)
if known is not None:
known[file_name] = result[0]
return (result[0], file_name)
except Exception as e:
@@ -241,6 +251,8 @@ def load_images_wrapper(file_names, embed_transparency=False, disp_names=None, s
imgs = []
masks = []
known_imgs = {}
known_masks = {}
is_mask = channel is not None
used_file_names = []
all_preview_imgs = []
@@ -261,9 +273,9 @@ def load_images_wrapper(file_names, embed_transparency=False, disp_names=None, s
if is_mask:
# A mask
mask, file_name = load_one_mask(file_name, disp_name, channel)
mask, file_name = load_one_mask(file_name, disp_name, channel, known=known_masks)
else:
img, mask, file_name = load_one_image(file_name, disp_name, embed_transparency)
img, mask, file_name = load_one_image(file_name, disp_name, embed_transparency, known=known_imgs)
max_w = max(max_w, img.shape[2])
max_h = max(max_h, img.shape[1])
@@ -279,7 +291,7 @@ def load_images_wrapper(file_names, embed_transparency=False, disp_names=None, s
all_preview_imgs.append(get_image_preview_info(file_name))
for j in range(len(imgs_batch)):
if channel is None:
if not is_mask:
img = imgs_batch[j]
H, W = img.shape[1:3]
if H != max_h or W != max_w:
@@ -292,7 +304,7 @@ def load_images_wrapper(file_names, embed_transparency=False, disp_names=None, s
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")
if channel is None:
if not is_mask:
imgs.append(torch.cat(imgs_batch))
masks.append(torch.cat(masks_batch))
else:
@@ -301,9 +313,9 @@ def load_images_wrapper(file_names, embed_transparency=False, disp_names=None, s
if is_mask:
# A mask
mask, file_name = load_one_mask(file_name, disp_names[i], channel)
mask, file_name = load_one_mask(file_name, disp_names[i], channel, known=known_masks)
else:
img, mask, file_name = load_one_image(file_name, disp_names[i], embed_transparency)
img, mask, file_name = load_one_image(file_name, disp_names[i], embed_transparency, known=known_imgs)
imgs.append(img)
masks.append(mask)
+5
View File
@@ -3238,6 +3238,7 @@ class ImageMemoryEstimator:
BYTES_PER_PIXEL = 4 * channels
logger.debug(f"BYTES_PER_PIXEL {BYTES_PER_PIXEL}")
known = set()
for fname in file_name:
if fname is None:
continue
@@ -3248,6 +3249,9 @@ class ImageMemoryEstimator:
logger.warning(f"Could not find file: {fname}")
continue
if file_path in known:
continue
try:
# PIL.Image.open is lazy; it reads metadata without loading pixel data
with Image.open(file_path) as img:
@@ -3264,6 +3268,7 @@ class ImageMemoryEstimator:
max_dim = (w, h)
valid_count += 1
known.add(file_path)
except Exception as e:
logger.error(f"Failed to read metadata for {fname}: {e}")