Work on v3. Work on refactoring image resize and vae tiling. Start work on upscaling nodes.

This commit is contained in:
Shanoah Alkire
2025-11-18 20:58:09 -08:00
parent 3385a35320
commit 20eee2ce91
11 changed files with 290 additions and 157 deletions
+24 -64
View File
@@ -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],)
+9 -9
View File
@@ -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."),
]
)
+1 -1
View File
@@ -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")
]
)
+6 -6
View File
@@ -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")
]
)
+19 -1
View File
@@ -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
+1 -1
View File
@@ -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")
]
)
+59 -64
View File
@@ -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 = [
+2 -2
View File
@@ -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")
]
)
+10 -2
View File
@@ -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',
+5 -1
View File
@@ -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)
}
}
MAX_RESOLUTION = nodes.MAX_RESOLUTION
+154 -6
View File
@@ -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