Work on v3. Work on refactoring image resize and vae tiling. Start work on upscaling nodes.
This commit is contained in:
+24
-64
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user