diff --git a/nodes/image.py b/nodes/image.py index 8867b79..9427dd4 100644 --- a/nodes/image.py +++ b/nodes/image.py @@ -13,7 +13,6 @@ from comfy_extras.nodes_images import ImageCrop from ..utils import load_image_from_path import torch -import torch.nn.functional as F import numpy as np from PIL import Image from PIL.PngImagePlugin import PngInfo @@ -28,7 +27,9 @@ import comfy.cli_args from ..utils.common import get_files_in_dir import datetime import logging -from ..utils.constants import QUICK_ASPECT_RATIOS +from ..utils.constants import QUICK_ASPECT_RATIOS, MAX_RESOLUTION + +from ..utils.common import calc_padding, image_manipulate, resize_needed class Sage_EmptyLatentImagePassthrough(ComfyNodeABC): def __init__(self): @@ -38,8 +39,8 @@ class Sage_EmptyLatentImagePassthrough(ComfyNodeABC): def INPUT_TYPES(cls) -> InputTypeDict: return { "required": { - "width": (IO.INT, {"default": 1024, "min": 16, "max": nodes.MAX_RESOLUTION, "step": 8, "tooltip": "The width of the latent images in pixels.", }), - "height": (IO.INT, {"default": 1024, "min": 16, "max": nodes.MAX_RESOLUTION, "step": 8, "tooltip": "The height of the latent images in pixels."}), + "width": (IO.INT, {"default": 1024, "min": 16, "max": MAX_RESOLUTION, "step": 8, "tooltip": "The width of the latent images in pixels.", }), + "height": (IO.INT, {"default": 1024, "min": 16, "max": MAX_RESOLUTION, "step": 8, "tooltip": "The height of the latent images in pixels."}), "batch_size": (IO.INT, { "default": 1, "min": 1, "max": 4096, "tooltip": "The number of latent images in the batch."}), "type": (IO.COMBO, {"default": "4_channel", "options": ["4_channel", "16_channel", "radiance"], "tooltip": "The type of latent to create. 4_channel is for standard latent diffusion models, 16_channel is for SD3 models, and radiance is for Chroma Radiance models."}) } @@ -360,8 +361,8 @@ class Sage_CubiqImageResize: return { "required": { "image": ("IMAGE",), - "width": ("INT", { "default": 1024, "min": 0, "max": nodes.MAX_RESOLUTION, "step": 1, }), - "height": ("INT", { "default": 1024, "min": 0, "max": nodes.MAX_RESOLUTION, "step": 1, }), + "width": ("INT", { "default": 1024, "min": 0, "max": MAX_RESOLUTION, "step": 1, }), + "height": ("INT", { "default": 1024, "min": 0, "max": MAX_RESOLUTION, "step": 1, }), "interpolation": (["nearest", "bilinear", "bicubic", "area", "nearest-exact", "lanczos", "bislerp"],), "method": (["stretch", "keep proportion", "fill / crop", "pad"],), "condition": (["always", "downscale if bigger", "upscale if smaller", "if bigger area", "if smaller area"],), @@ -373,25 +374,6 @@ class Sage_CubiqImageResize: RETURN_NAMES = ("IMAGE", "width", "height",) FUNCTION = "execute" CATEGORY = "Sage Utils/image" - - def padding(self, width, height, new_width, new_height): - """ - Calculate the padding values for left, right, top, and bottom. - """ - pad_left = (width - new_width) // 2 - pad_right = width - new_width - pad_left - pad_top = (height - new_height) // 2 - pad_bottom = height - new_height - pad_top - return pad_left, pad_right, pad_top, pad_bottom - - def resize_needed(self, condition, width, height, ow, oh): - if "always" in condition \ - or ("downscale if bigger" == condition and (oh > height or ow > width)) \ - or ("upscale if smaller" == condition and (oh < height or ow < width)) \ - or ("bigger area" in condition and (oh * ow > height * width)) \ - or ("smaller area" in condition and (oh * ow < height * width)): - return True - return False def execute(self, image, width, height, method="stretch", interpolation="lanczos", condition="always", multiple_of=64, keep_proportion=False): _, oh, ow, _ = image.shape @@ -407,22 +389,24 @@ class Sage_CubiqImageResize: height = height - (height % multiple_of) if method == 'keep proportion' or method == 'pad': - if width == 0 and oh < height: - width = nodes.MAX_RESOLUTION - elif width == 0 and oh >= height: - width = ow + if width == 0: + if oh < height: + width = MAX_RESOLUTION + else: + width = ow - if height == 0 and ow < width: - height = nodes.MAX_RESOLUTION - elif height == 0 and ow >= width: - height = oh + if height == 0: + if ow < width: + height = MAX_RESOLUTION + else: + height = oh ratio = min(width / ow, height / oh) new_width = round(ow*ratio) new_height = round(oh*ratio) if method == 'pad': - pad_left, pad_right, pad_top, pad_bottom = self.padding(width, height, new_width, new_height) + pad_left, pad_right, pad_top, pad_bottom = calc_padding(width, height, new_width, new_height) if pad_left > 0 or pad_right > 0 or pad_top > 0 or pad_bottom > 0: padding = True @@ -452,37 +436,13 @@ class Sage_CubiqImageResize: width = width if width > 0 else ow height = height if height > 0 else oh - if self.resize_needed(condition, width, height, ow, oh): - outputs = image.permute(0,3,1,2) + fill = method.startswith('fill') + resize = resize_needed(condition, width, height, ow, oh) - if interpolation == "lanczos": - outputs = comfy.utils.lanczos(outputs, width, height) - elif interpolation == "bislerp": - outputs = comfy.utils.bislerp(outputs, width, height) - else: - outputs = F.interpolate(outputs, size=(height, width), mode=interpolation) - - if padding: - outputs = F.pad(outputs, (pad_left, pad_right, pad_top, pad_bottom), value=0) - - outputs = outputs.permute(0,2,3,1) - - if method.startswith('fill'): - if x > 0 or y > 0 or x2 > 0 or y2 > 0: - outputs = outputs[:, y:y2, x:x2, :] - else: - outputs = image - - if multiple_of > 1 and (outputs.shape[2] % multiple_of != 0 or outputs.shape[1] % multiple_of != 0): - width = outputs.shape[2] - height = outputs.shape[1] - x = (width % multiple_of) // 2 - y = (height % multiple_of) // 2 - x2 = width - ((width % multiple_of) - x) - y2 = height - ((height % multiple_of) - y) - outputs = outputs[:, y:y2, x:x2, :] - - outputs = torch.clamp(outputs, 0, 1) + outputs = image_manipulate(image, width, height, interpolation, multiple_of, + padding, fill, resize, + pad_left, pad_right, pad_top, pad_bottom, + x, y, x2, y2) return(outputs, outputs.shape[2], outputs.shape[1],) diff --git a/nodes/image_v3.py b/nodes/image_v3.py index 7acfd2b..356adf7 100644 --- a/nodes/image_v3.py +++ b/nodes/image_v3.py @@ -52,8 +52,8 @@ class Sage_EmptyLatentImagePassthrough(io.ComfyNode): ], outputs=[ io.Latent.Output("latent"), - io.Int.Output("out_width"), - io.Int.Output("out_height"), + io.Int.Output("out_width", display_name="width"), + io.Int.Output("out_height", display_name="height"), ] ) @@ -91,8 +91,8 @@ class Sage_LoadImage(io.ComfyNode): ], outputs=[ io.Image.Output("image"), - io.Int.Output("out_width"), - io.Int.Output("out_height"), + io.Int.Output("out_width", display_name="width"), + io.Int.Output("out_height", display_name="height"), io.String.Output("metadata"), ] ) @@ -149,7 +149,7 @@ class Sage_CropImage(io.ComfyNode): io.Int.Input("bottom", default=0, tooltip="The bottom coordinate for cropping."), ], outputs=[ - io.Image.Output("out_image", tooltip="The cropped image."), + io.Image.Output("out_image", tooltip="The cropped image.", display_name="image"), ] ) @@ -264,9 +264,9 @@ class Sage_CubiqImageResize(io.ComfyNode): io.Int.Input("multiple_of", default=8, tooltip="Ensure dimensions are multiples of this value."), ], outputs=[ - io.Image.Output("out_image", tooltip="The resized image."), - io.Int.Output("out_width", tooltip="The new width."), - io.Int.Output("out_height", tooltip="The new height."), + io.Image.Output("out_image", tooltip="The resized image.", display_name="image"), + io.Int.Output("out_width", tooltip="The new width.", display_name="width"), + io.Int.Output("out_height", tooltip="The new height.", display_name="height"), ] ) @@ -302,7 +302,7 @@ class Sage_ReferenceImage(io.ComfyNode): io.Vae.Input("vae", tooltip="The VAE model for encoding the image."), ], outputs=[ - io.Conditioning.Output("out_conditioning", tooltip="The output conditioning."), + io.Conditioning.Output("out_conditioning", tooltip="The output conditioning.", display_name="conditioning"), io.Latent.Output("latent", tooltip="The encoded latent."), ] ) diff --git a/nodes/llm_v3.py b/nodes/llm_v3.py index 92b4b3c..3ab1b38 100644 --- a/nodes/llm_v3.py +++ b/nodes/llm_v3.py @@ -73,7 +73,7 @@ class Sage_ConstructLLMPrompt(io.ComfyNode): # TODO: Add dynamic boolean inputs for llm_prompts["extra"] options ], outputs=[ - io.String.Output("out_prompt") + io.String.Output("out_prompt", display_name="prompt") ] ) diff --git a/nodes/loader_v3.py b/nodes/loader_v3.py index 33d8fbf..b57bc16 100644 --- a/nodes/loader_v3.py +++ b/nodes/loader_v3.py @@ -190,9 +190,9 @@ class Sage_LoraStackLoader(io.ComfyNode): io.Custom("MODEL_SHIFTS").Input("model_shifts", optional=True) ], outputs=[ - io.Model.Output("out_model"), - io.Clip.Output("out_clip"), - io.Custom("LORA_STACK").Output("out_lora_stack"), + io.Model.Output("out_model", display_name="model"), + io.Clip.Output("out_clip", display_name="clip"), + io.Custom("LORA_STACK").Output("out_lora_stack", display_name="lora_stack"), io.String.Output("keywords") ] ) @@ -224,7 +224,7 @@ class Sage_ModelLoraStackLoader(io.ComfyNode): io.Model.Output("model"), io.Clip.Output("clip"), io.Vae.Output("vae"), - io.Custom("LORA_STACK").Output("out_lora_stack"), + io.Custom("LORA_STACK").Output("out_lora_stack", display_name="lora_stack"), io.String.Output("keywords") ] ) @@ -250,8 +250,8 @@ class Sage_UNETLoRALoader(io.ComfyNode): io.Custom("MODEL_SHIFTS").Input("model_shifts", optional=True) ], outputs=[ - io.Model.Output("out_model"), - io.Custom("LORA_STACK").Output("out_lora_stack"), + io.Model.Output("out_model", display_name="model"), + io.Custom("LORA_STACK").Output("out_lora_stack", display_name="lora_stack"), io.String.Output("keywords") ] ) diff --git a/nodes/sampler.py b/nodes/sampler.py index f8c21c1..3eb676f 100644 --- a/nodes/sampler.py +++ b/nodes/sampler.py @@ -8,6 +8,12 @@ import torch import comfy import nodes +from ..utils.common import ( + vae_decode, + vae_decode_tiled, + load_upscaler, + upscale_with_model +) class Sage_SamplerSelector(ComfyNodeABC): def __init__(self): @@ -189,7 +195,17 @@ class Sage_KSamplerTiledDecoder(ComfyNodeABC): if advanced_info["add_noise"] == False: disable_noise = True latent_result = nodes.common_ksampler(model, sampler_info["seed"], sampler_info["steps"], sampler_info["cfg"], sampler_info["sampler"], sampler_info["scheduler"], positive, negative, latent_image, denoise=denoise, disable_noise=disable_noise, start_step=advanced_info['start_at_step'], last_step=advanced_info['end_at_step'], force_full_denoise=force_full_denoise) - + + if tiling_info is not None: + images = vae_decode_tiled( + latent_result, + vae, + tiling_info["tile_size"], tiling_info["overlap"], + tiling_info["temporal_size"], tiling_info["temporal_overlap"]) + else: + images = vae_decode(latent_result, vae) + return (latent_result[0], images) +""" if tiling_info is not None: t_info_tile_size = tiling_info["tile_size"] t_info_overlap = tiling_info["overlap"] @@ -226,6 +242,8 @@ class Sage_KSamplerTiledDecoder(ComfyNodeABC): images = images.reshape(-1, images.shape[-3], images.shape[-2], images.shape[-1]) return (latent_result[0], images) +""" + class Sage_KSamplerAudioDecoder(ComfyNodeABC): @classmethod diff --git a/nodes/sampler_v3.py b/nodes/sampler_v3.py index 0687435..6bb5c86 100644 --- a/nodes/sampler_v3.py +++ b/nodes/sampler_v3.py @@ -59,7 +59,7 @@ class Sage_SchedulerSelector(io.ComfyNode): io.Combo.Input("scheduler_name", options=list(comfy.samplers.KSampler.SCHEDULERS), default="beta") ], outputs=[ - io.Int.Output("out_steps"), + io.Int.Output("out_steps", display_name="steps"), io.String.Output("scheduler") ] ) diff --git a/nodes/selector_v3.py b/nodes/selector_v3.py index bd16987..36e745b 100644 --- a/nodes/selector_v3.py +++ b/nodes/selector_v3.py @@ -372,20 +372,14 @@ class Sage_MultiSelectorQuadClip(io.ComfyNode): ret = (unet_info[0], clip_info[0], vae_info[0]) return io.NodeOutput(ret,) -# ============================================================================ -# PLACEHOLDER NODES - NOT YET FULLY IMPLEMENTED -# ============================================================================ -# These are placeholder implementations. The inputs/outputs match the original -# v1 nodes, but the execute methods need proper implementation. - class Sage_ModelShifts(io.ComfyNode): - """PLACEHOLDER: Get the model shifts and free_u2 settings to apply to the model.""" + """Get the model shifts and free_u2 settings to apply to the model.""" @classmethod def define_schema(cls): return io.Schema( node_id="Sage_ModelShifts", display_name="Model Shifts", - description="PLACEHOLDER: Get the model shifts and free_u2 settings to apply to the model. This is used by the model loader node.", + description="Get the model shifts and free_u2 settings to apply to the model. This is used by the model loader node.", category="Sage Utils/model", inputs=[ io.Combo.Input("shift_type", options=["None", "x1", "x1000"], default="None"), @@ -403,7 +397,6 @@ class Sage_ModelShifts(io.ComfyNode): @classmethod def execute(cls, **kwargs): - # TODO: Implement full logic from selector.py return io.NodeOutput({ "shift_type": kwargs.get("shift_type", "None"), "shift": kwargs.get("shift", 3.0), @@ -415,13 +408,13 @@ class Sage_ModelShifts(io.ComfyNode): }) class Sage_ModelShiftOnly(io.ComfyNode): - """PLACEHOLDER: Get the model shifts to apply to the model.""" + """Get the model shifts to apply to the model.""" @classmethod def define_schema(cls): return io.Schema( node_id="Sage_ModelShiftOnly", display_name="Model Shift Only", - description="PLACEHOLDER: Get the model shifts to apply to the model. This is used by the model loader node.", + description="Get the model shifts to apply to the model. This is used by the model loader node.", category="Sage Utils/model", inputs=[ io.Combo.Input("shift_type", options=["None", "x1", "x1000"], default="None"), @@ -434,7 +427,6 @@ class Sage_ModelShiftOnly(io.ComfyNode): @classmethod def execute(cls, **kwargs): - # TODO: Implement full logic from selector.py return io.NodeOutput({ "shift_type": kwargs.get("shift_type", "None"), "shift": kwargs.get("shift", 3.0), @@ -446,13 +438,13 @@ class Sage_ModelShiftOnly(io.ComfyNode): }) class Sage_FreeU2(io.ComfyNode): - """PLACEHOLDER: Get the free_u2 settings to apply to the model.""" + """Get the free_u2 settings to apply to the model.""" @classmethod def define_schema(cls): return io.Schema( node_id="Sage_FreeU2", display_name="FreeU v2", - description="PLACEHOLDER: Get the free_u2 settings to apply to the model.", + description="Get the free_u2 settings to apply to the model.", category="Sage Utils/model", inputs=[ io.Boolean.Input("freeu_v2", default=False), @@ -468,7 +460,6 @@ class Sage_FreeU2(io.ComfyNode): @classmethod def execute(cls, **kwargs): - # TODO: Implement full logic from selector.py return io.NodeOutput({ "shift_type": "None", "shift": 0, @@ -479,14 +470,49 @@ class Sage_FreeU2(io.ComfyNode): "s2": kwargs.get("s2", 0.2) }) +class Sage_TilingInfo(io.ComfyNode): + """Adds tiling information to the KSampler.""" + @classmethod + def define_schema(cls): + return io.Schema( + node_id="Sage_TilingInfo", + display_name="Tiling Info", + description="Adds tiling information to the KSampler.", + category="Sage Utils/sampler", + inputs=[ + io.Int.Input("tile_size", default=512, min=64, max=4096, step=32), + io.Int.Input("overlap", default=64, min=0, max=4096, step=32), + io.Int.Input("temporal_size", default=64, min=8, max=4096, step=4, tooltip="Only used for video VAEs: Amount of frames to decode at a time."), + io.Int.Input("temporal_overlap", default=8, min=4, max=4096, step=4, tooltip="Only used for video VAEs: Amount of frames to overlap.") + ], + outputs=[ + TilingInfo.Output("tiling_info") + ] + ) + + @classmethod + def execute(cls, **kwargs): + tile_size = kwargs.get("tile_size", 512) + overlap = kwargs.get("overlap", 64) + temporal_size = kwargs.get("temporal_size", 64) + temporal_overlap = kwargs.get("temporal_overlap", 8) + + t_info = { + "tile_size": tile_size, + "overlap": overlap, + "temporal_size": temporal_size, + "temporal_overlap": temporal_overlap + } + return io.NodeOutput(t_info) + class Sage_UnetClipVaeToModelInfo(io.ComfyNode): - """PLACEHOLDER: Convert UNET, CLIP, and VAE model info to a single model info output.""" + """Convert UNET, CLIP, and VAE model info to a single model info output.""" @classmethod def define_schema(cls): return io.Schema( node_id="Sage_UnetClipVaeToModelInfo", display_name="UNET CLIP VAE To Model Info", - description="PLACEHOLDER: Returns a list with the unets, clips, and vae in it to be loaded.", + description="Returns a list with the unets, clips, and vae in it to be loaded.", category="Sage Utils/model", inputs=[ UnetInfo.Input("unet_info"), @@ -500,20 +526,19 @@ class Sage_UnetClipVaeToModelInfo(io.ComfyNode): @classmethod def execute(cls, **kwargs): - # TODO: Implement full logic from selector.py unet_info = kwargs.get("unet_info", None) clip_info = kwargs.get("clip_info", None) vae_info = kwargs.get("vae_info", None) return io.NodeOutput((unet_info, clip_info, vae_info)) class Sage_LoraStack(io.ComfyNode): - """PLACEHOLDER: Choose a lora with weights, and add it to a lora_stack.""" + """Choose a lora with weights, and add it to a lora_stack.""" @classmethod def define_schema(cls): return io.Schema( node_id="Sage_LoraStack", display_name="Lora Stack", - description="PLACEHOLDER: Choose a lora with weights, and add it to a lora_stack. Compatible with other node packs that have lora_stacks.", + description="Choose a lora with weights, and add it to a lora_stack. Compatible with other node packs that have lora_stacks.", category="Sage Utils/lora", inputs=[ io.Boolean.Input("enabled", default=False), @@ -523,13 +548,12 @@ class Sage_LoraStack(io.ComfyNode): LoraStack.Input("lora_stack", optional=True) ], outputs=[ - LoraStack.Output("out_lora_stack") + LoraStack.Output("out_lora_stack", display_name="lora_stack") ] ) @classmethod def execute(cls, **kwargs): - # TODO: Implement full logic from selector.py lora_stack = kwargs.get("lora_stack", None) enabled = kwargs.get("enabled", False) @@ -542,6 +566,13 @@ class Sage_LoraStack(io.ComfyNode): return io.NodeOutput(lora_stack) +# ============================================================================ +# PLACEHOLDER NODES - NOT YET FULLY IMPLEMENTED +# ============================================================================ +# These are placeholder implementations. The inputs/outputs match the original +# v1 nodes, but the execute methods need proper implementation. + + class Sage_QuickLoraStack(io.ComfyNode): """PLACEHOLDER: Simplified lora stack node without clip_weight.""" @classmethod @@ -558,7 +589,7 @@ class Sage_QuickLoraStack(io.ComfyNode): LoraStack.Input("lora_stack", optional=True) ], outputs=[ - LoraStack.Output("out_lora_stack") + LoraStack.Output("out_lora_stack", display_name="lora_stack") ] ) @@ -604,7 +635,7 @@ class Sage_TripleLoraStack(io.ComfyNode): LoraStack.Input("lora_stack", optional=True) ], outputs=[ - LoraStack.Output("out_lora_stack") + LoraStack.Output("out_lora_stack", display_name="lora_stack") ] ) @@ -638,7 +669,7 @@ class Sage_SixLoraStack(io.ComfyNode): category="Sage Utils/lora", inputs=inputs, outputs=[ - LoraStack.Output("out_lora_stack") + LoraStack.Output("out_lora_stack", display_name="lora_stack") ] ) @@ -672,7 +703,7 @@ class Sage_TripleQuickLoraStack(io.ComfyNode): LoraStack.Input("lora_stack", optional=True) ], outputs=[ - LoraStack.Output("out_lora_stack") + LoraStack.Output("out_lora_stack", display_name="lora_stack") ] ) @@ -704,7 +735,7 @@ class Sage_QuickSixLoraStack(io.ComfyNode): category="Sage Utils/lora", inputs=inputs, outputs=[ - LoraStack.Output("out_lora_stack") + LoraStack.Output("out_lora_stack", display_name="lora_stack") ] ) @@ -736,7 +767,7 @@ class Sage_QuickNineLoraStack(io.ComfyNode): category="Sage Utils/lora", inputs=inputs, outputs=[ - LoraStack.Output("out_lora_stack") + LoraStack.Output("out_lora_stack", display_name="lora_stack") ] ) @@ -746,42 +777,6 @@ class Sage_QuickNineLoraStack(io.ComfyNode): lora_stack = kwargs.get("lora_stack", None) return io.NodeOutput(lora_stack) -class Sage_TilingInfo(io.ComfyNode): - """PLACEHOLDER: Adds tiling information to the KSampler.""" - @classmethod - def define_schema(cls): - return io.Schema( - node_id="Sage_TilingInfo", - display_name="Tiling Info", - description="PLACEHOLDER: Adds tiling information to the KSampler.", - category="Sage Utils/sampler", - inputs=[ - io.Int.Input("tile_size", default=512, min=64, max=4096, step=32), - io.Int.Input("overlap", default=64, min=0, max=4096, step=32), - io.Int.Input("temporal_size", default=64, min=8, max=4096, step=4, tooltip="Only used for video VAEs: Amount of frames to decode at a time."), - io.Int.Input("temporal_overlap", default=8, min=4, max=4096, step=4, tooltip="Only used for video VAEs: Amount of frames to overlap.") - ], - outputs=[ - TilingInfo.Output("tiling_info") - ] - ) - - @classmethod - def execute(cls, **kwargs): - # TODO: Implement full logic from selector.py - tile_size = kwargs.get("tile_size", 512) - overlap = kwargs.get("overlap", 64) - temporal_size = kwargs.get("temporal_size", 64) - temporal_overlap = kwargs.get("temporal_overlap", 8) - - t_info = { - "tile_size": tile_size, - "overlap": overlap, - "temporal_size": temporal_size, - "temporal_overlap": temporal_overlap - } - return io.NodeOutput(t_info) - # ============================================================================ SELECTOR_NODES = [ diff --git a/nodes/util_v3.py b/nodes/util_v3.py index e91b687..e153603 100644 --- a/nodes/util_v3.py +++ b/nodes/util_v3.py @@ -36,7 +36,7 @@ class Sage_FreeMemory(io.ComfyNode): io.AnyType.Input("value") ], outputs=[ - io.AnyType.Output("out_value") + io.AnyType.Output("out_value", display_name="value") ] ) @@ -60,7 +60,7 @@ class Sage_Halt(io.ComfyNode): io.AnyType.Input("value") ], outputs=[ - io.AnyType.Output("out_value") + io.AnyType.Output("out_value", display_name="value") ] ) diff --git a/utils/common.py b/utils/common.py index 2757f9d..ffca1ed 100644 --- a/utils/common.py +++ b/utils/common.py @@ -43,7 +43,14 @@ from .helpers_image import ( tensor_to_base64, tensor_to_temp_image, load_image_from_path, - load_image_from_url + load_image_from_url, + vae_decode, + vae_decode_tiled, + load_upscaler, + upscale_with_model, + calc_padding, + image_manipulate, + resize_needed ) # CivitAI utilities @@ -132,7 +139,8 @@ __all__ = [ # Image utilities 'blank_image', 'url_to_torch_image', 'tensor_to_base64', 'tensor_to_temp_image', - 'load_image_from_path', 'load_image_from_url', + 'load_image_from_path', 'load_image_from_url', 'vae_decode', 'vae_decode_tiled', + 'load_upscaler', 'upscale_with_model', 'calc_padding', 'image_manipulate', 'resize_needed', # CivitAI utilities 'get_civitai_model_version_json_by_hash', 'get_civitai_model_version_json_by_id', diff --git a/utils/constants.py b/utils/constants.py index 471f2f2..8829ebb 100644 --- a/utils/constants.py +++ b/utils/constants.py @@ -4,6 +4,8 @@ This module contains common constants used throughout the SageUtils project. These constants help maintain consistency and avoid hardcoded values in multiple files. """ +import nodes + # Supported model file extensions (based on ComfyUI's supported_pt_extensions) # These are the file types that ComfyUI can load as models # Additional extensions (.gguf, .nf4) supported via custom extensions @@ -123,4 +125,6 @@ QUICK_ASPECT_RATIOS = { "7:9": (896, 1152), "8:10": (1024, 1280), "13:19": (832, 1216) - } \ No newline at end of file + } + +MAX_RESOLUTION = nodes.MAX_RESOLUTION diff --git a/utils/helpers_image.py b/utils/helpers_image.py index 1d309c7..96089bd 100644 --- a/utils/helpers_image.py +++ b/utils/helpers_image.py @@ -7,11 +7,15 @@ import pathlib from PIL import Image, ImageOps, ImageSequence import numpy as np import torch +import torch.nn.functional as F import requests import node_helpers import folder_paths +import comfy.utils +import comfy.model_management as mm +from spandrel import ModelLoader, ImageModelDescriptor def blank_image(): """Create a blank 1024x1024 RGB image as a torch tensor.""" @@ -20,7 +24,6 @@ def blank_image(): img = np.array(img.convert("RGB")).astype(np.float32) / 255.0 return torch.from_numpy(img)[None, :] - def url_to_torch_image(url): """Load an image from a URL and return as a torch tensor.""" response = requests.get(url, stream=True) @@ -29,7 +32,6 @@ def url_to_torch_image(url): img = np.array(img.convert("RGB")).astype(np.float32) / 255.0 return torch.from_numpy(img)[None, :] - def tensor_to_base64(tensor): """Convert a torch tensor image batch to a list of base64-encoded PNGs.""" if tensor is None or not isinstance(tensor, torch.Tensor): @@ -45,7 +47,6 @@ def tensor_to_base64(tensor): base64_images.append(base64.b64encode(buffered.getvalue()).decode('utf-8')) return base64_images - def _load_image(img) -> tuple: """Internal helper to process a PIL image (from path or URL) into torch tensors and mask.""" output_images, output_masks = [], [] @@ -70,20 +71,17 @@ def _load_image(img) -> tuple: output_mask = torch.cat(output_masks, dim=0) if len(output_masks) > 1 and getattr(img, 'format', None) != "MPO" else output_masks[0] return output_image, output_mask, w, h, f"{getattr(img, 'info', {})}" - def load_image_from_path(image_path) -> tuple: """Load an image (and mask if present) from a file path as torch tensors.""" img = node_helpers.pillow(Image.open, image_path) return _load_image(img) - def load_image_from_url(url) -> tuple: """Load an image (and mask if present) from a URL as torch tensors.""" response = requests.get(url, stream=True) img = node_helpers.pillow(Image.open, io.BytesIO(response.content)) return _load_image(img) - def tensor_to_temp_image(tensor, filename=None): """Save a torch tensor image batch to temporary PNG files. Returns list of file paths.""" if tensor is None or not isinstance(tensor, torch.Tensor): @@ -106,3 +104,153 @@ def tensor_to_temp_image(tensor, filename=None): print(f"Saved {len(filenames)} images to {output_dir}") print(filenames) return filenames + +def calc_padding(width, height, new_width, new_height): + """ + Calculate the padding values for left, right, top, and bottom. + """ + pad_left = (width - new_width) // 2 + pad_right = width - new_width - pad_left + pad_top = (height - new_height) // 2 + pad_bottom = height - new_height - pad_top + return pad_left, pad_right, pad_top, pad_bottom + +def image_padding(image, pad_left=0, pad_right=0, pad_top=0, pad_bottom=0): + return F.pad(image, (pad_left, pad_right, pad_top, pad_bottom), value=0) + +def image_fill(image, x=0, y=0, x2=0, y2=0): + if x > 0 or y > 0 or x2 > 0 or y2 > 0: + return image[:, y:y2, x:x2, :] + return image + +def image_mult_of(outputs, multiple_of=1): + if multiple_of > 1 and (outputs.shape[2] % multiple_of != 0 or outputs.shape[1] % multiple_of != 0): + width = outputs.shape[2] + height = outputs.shape[1] + x = (width % multiple_of) // 2 + y = (height % multiple_of) // 2 + x2 = width - ((width % multiple_of) - x) + y2 = height - ((height % multiple_of) - y) + outputs = outputs[:, y:y2, x:x2, :] + return outputs + +def resize_needed(condition, width, height, ow, oh): + if "always" in condition \ + or ("downscale if bigger" == condition and (oh > height or ow > width)) \ + or ("upscale if smaller" == condition and (oh < height or ow < width)) \ + or ("bigger area" in condition and (oh * ow > height * width)) \ + or ("smaller area" in condition and (oh * ow < height * width)): + return True + return False + +def image_resize(outputs, width, height, interpolation): + if interpolation == "lanczos": + outputs = comfy.utils.lanczos(outputs, width, height) + elif interpolation == "bislerp": + outputs = comfy.utils.bislerp(outputs, width, height) + else: + outputs = F.interpolate(outputs, size=(height, width), mode=interpolation) + return outputs + +def image_manipulate(image, width, height, interpolation, multiple_of = 1, + padding = False, fill = False, resize = False, + pad_left=0, pad_right=0, pad_top=0, pad_bottom=0, + x=0, y=0, x2=0, y2=0): + outputs = image + + if resize: + outputs = outputs.permute(0,3,1,2) + outputs = image_resize(outputs, width, height, interpolation) + + if padding: + outputs = image_padding(outputs, pad_left, pad_right, pad_top, pad_bottom) + + outputs = outputs.permute(0,2,3,1) + + if fill: + outputs = image_fill(outputs, x, y, x2, y2) + + outputs = image_mult_of(outputs, multiple_of) + + outputs = torch.clamp(outputs, 0, 1) + return outputs + +def vae_decode(latent_result, vae): + latent = latent_result[0]["samples"] + images = vae.decode(latent) + + if len(images.shape) == 5: #Combine batches + images = images.reshape(-1, images.shape[-3], images.shape[-2], images.shape[-1]) + + return images + +def vae_decode_tiled(latent_result, vae, tile_size, overlap, temporal_size, temporal_overlap): + latent = latent_result[0]["samples"] + + if tile_size < overlap * 4: + overlap = tile_size // 4 + if temporal_size < temporal_overlap * 2: + temporal_overlap = temporal_overlap // 2 + temporal_compression = vae.temporal_compression_decode() + + if temporal_compression is not None: + temporal_size = max(2, temporal_size // temporal_overlap) + temporal_overlap = max(1, min(temporal_size // 2, temporal_overlap // temporal_compression)) + else: + temporal_size = None + temporal_overlap = None + + compression = vae.spacial_compression_decode() + + images = vae.decode_tiled( + latent, + tile_x=tile_size // compression, tile_y=tile_size // compression, + overlap=overlap // compression, + tile_t=temporal_size, + overlap_t= temporal_overlap) + + if len(images.shape) == 5: #Combine batches + images = images.reshape(-1, images.shape[-3], images.shape[-2], images.shape[-1]) + + return images + +def load_upscaler(model_path): + sd = comfy.utils.load_torch_file(model_path, safe_load=True) + if "module.layers.0.residual_group.blocks.0.norm1.weight" in sd: + sd = comfy.utils.state_dict_prefix_replace(sd, {"module.":""}) + upload_model = ModelLoader().load_from_state_dict(sd).eval() + + if not isinstance(upload_model, ImageModelDescriptor): + raise Exception("Upscale model must be a single-image model.") + + return upload_model + +def upscale_with_model(upscale_model, image, tile = 512, overlap = 32): + device = mm.get_torch_device() + + memory_required = mm.module_size(upscale_model.model) + memory_required += ((512 * 512 * 3 * 384.0) * max(upscale_model.scale, 1.0) + image.nelement()) * image.element_size() + mm.free_memory(memory_required, device) + + upscale_model.to(device) + in_img = image.movedim(-1,-3).to(device) + + oom = True + scaled = None + while oom: + try: + steps = in_img.shape[0] * comfy.utils.get_tiled_scale_steps(in_img.shape[3], in_img.shape[2], tile_x=tile, tile_y=tile, overlap=overlap) + pbar = comfy.utils.ProgressBar(steps) + scaled = comfy.utils.tiled_scale(in_img, lambda a: upscale_model(a), tile_x=tile, tile_y=tile, overlap=overlap, upscale_amount=upscale_model.scale, pbar=pbar) + oom = False + except mm.OOM_EXCEPTION as e: + tile //= 2 + if tile <= overlap: + raise e + + upscale_model.to("cpu") + image = None + if scaled: + image = torch.clamp(scaled.movedim(-3,-1), min=0, max=1.0) + return image +