diff --git a/nodes/image_saver.py b/nodes/image_saver.py index a51c79d..a8c7286 100644 --- a/nodes/image_saver.py +++ b/nodes/image_saver.py @@ -13,6 +13,8 @@ class OCS_ImageSaver: + INPUT_IS_LIST = True + def __init__(self): self.output_dir = folder_paths.get_output_directory() self.type = "output" @@ -85,6 +87,23 @@ def save_images( EXIF_UserComment: str = "", extra_pnginfo=None, ): + + flat_images = self._flatten_images(images) + if not flat_images: + raise ValueError("No images provided to OCS_ImageSaver") + + # unpack widget scalars when Comfy wraps them in single-element lists + seed = self._unwrap_scalar(seed) + filename = self._unwrap_scalar(filename) + path = self._unwrap_scalar(path) + image_format = self._unwrap_scalar(image_format) + lossless_webp = self._unwrap_scalar(lossless_webp) + jpg_webp_quality = self._unwrap_scalar(jpg_webp_quality) + date_format = self._unwrap_scalar(date_format) + time_format = self._unwrap_scalar(time_format) + embed_workflow = self._unwrap_scalar(embed_workflow) + EXIF_UserComment = self._unwrap_scalar(EXIF_UserComment) + ( full_output_folder, filename_alt, @@ -94,8 +113,8 @@ def save_images( ) = folder_paths.get_save_image_path( self.prefix_append, self.output_dir, - images[0].shape[1], - images[0].shape[0], + flat_images[0].shape[1], + flat_images[0].shape[0], ) output_folder = Path(full_output_folder) @@ -110,7 +129,7 @@ def save_images( saved_filenames, saved_paths, ui_images = [], [], [] - for (batch_number, image) in enumerate(images): + for (batch_number, image) in enumerate(flat_images): var_map = base_vars.copy() var_map["%counter"] = f"{counter_base + batch_number:05}" @@ -159,12 +178,45 @@ def save_images( def _single_or_list(lst): return lst[0] if len(lst) == 1 else lst + @staticmethod + def _unwrap_scalar(value): + if isinstance(value, (list, tuple)): + return value[0] if value else None + return value + @staticmethod def _replace_tokens(template: str, mapping: dict) -> str: for k, v in mapping.items(): template = template.replace(k, str(v)) return template.strip("/") + @staticmethod + def _flatten_images(images): + """Normalise IMAGE, batch, or (possibly nested) lists into a flat list.""" + + if isinstance(images, torch.Tensor): + # single image or explicit batch tensor + return list(images) if images.ndim == 4 else [images] + + flat = [] + + def _walk(item): + if isinstance(item, torch.Tensor): + flat.extend(list(item) if item.ndim == 4 else [item]) + elif isinstance(item, (list, tuple)): + for sub in item: + _walk(sub) + elif item is None: + # ignore empty placeholders often produced by optional inputs + return + else: + raise TypeError( + "OCS_ImageSaver expects IMAGE tensors, batches, or lists thereof" + ) + + _walk(images) + return flat + @staticmethod def process_image( img: Image.Image,