From d4cac6ac9584861cc3df5cbc977f3c99e80060dc Mon Sep 17 00:00:00 2001 From: Acly Date: Sat, 18 Oct 2025 23:58:51 +0200 Subject: [PATCH] Change node definitions to "V3" schema, remove CropImage node --- __init__.py | 93 +++++++---------- krita.py | 276 +++++++++++++++++++++++++------------------------ nodes.py | 203 ++++++++++++++++-------------------- nsfw.py | 50 +++++---- pyproject.toml | 2 +- region.py | 112 ++++++++++---------- tile.py | 186 +++++++++++++++++---------------- translation.py | 24 +++-- 8 files changed, 456 insertions(+), 490 deletions(-) diff --git a/__init__.py b/__init__.py index bf61fc1..259f990 100644 --- a/__init__.py +++ b/__init__.py @@ -1,59 +1,40 @@ +from comfy_api.latest import ComfyExtension, io from . import api as api, nodes, tile, region, nsfw, translation, krita -NODE_CLASS_MAPPINGS = { - "ETN_LoadImageBase64": nodes.LoadImageBase64, - "ETN_LoadMaskBase64": nodes.LoadMaskBase64, - "ETN_SendImageWebSocket": nodes.SendImageWebSocket, - "ETN_CropImage": nodes.CropImage, - "ETN_ApplyMaskToImage": nodes.ApplyMaskToImage, - "ETN_ReferenceImage": nodes.ReferenceImage, - "ETN_ApplyReferenceImages": nodes.ApplyReferenceImages, - "ETN_TileLayout": tile.TileLayout, - "ETN_ExtractImageTile": tile.ExtractImageTile, - "ETN_ExtractMaskTile": tile.ExtractMaskTile, - "ETN_GenerateTileMask": tile.GenerateTileMask, - "ETN_MergeImageTile": tile.MergeImageTile, - "ETN_BackgroundRegion": region.BackgroundRegion, - "ETN_DefineRegion": region.DefineRegion, - "ETN_ListRegionMasks": region.ListRegionMasks, - "ETN_AttentionMask": region.AttentionMask, - "ETN_NSFWFilter": nsfw.NSFWFilter, - "ETN_Translate": translation.Translate, - "ETN_KritaOutput": krita.KritaOutput, - "ETN_KritaSendText": krita.KritaSendText, - "ETN_KritaCanvas": krita.KritaCanvas, - "ETN_KritaSelection": krita.KritaSelection, - "ETN_KritaImageLayer": krita.KritaImageLayer, - "ETN_KritaMaskLayer": krita.KritaMaskLayer, - "ETN_Parameter": krita.Parameter, - "ETN_KritaStyle": krita.KritaStyle, -} -NODE_DISPLAY_NAME_MAPPINGS = { - "ETN_LoadImageBase64": "Load Image (Base64)", - "ETN_LoadMaskBase64": "Load Mask (Base64)", - "ETN_SendImageWebSocket": "Send Image (WebSocket)", - "ETN_CropImage": "Crop Image", - "ETN_ApplyMaskToImage": "Apply Mask to Image", - "ETN_ReferenceImage": "Reference Image", - "ETN_ApplyReferenceImages": "Apply Reference Images", - "ETN_TileLayout": "Create Tile Layout", - "ETN_ExtractImageTile": "Extract Image Tile", - "ETN_ExtractMaskTile": "Extract Mask Tile", - "ETN_MergeImageTile": "Merge Image Tile", - "ETN_GenerateTileMask": "Generate Tile Mask", - "ETN_BackgroundRegion": "Background Region", - "ETN_DefineRegion": "Define Region", - "ETN_ListRegionMasks": "List Region Masks", - "ETN_AttentionMask": "Regions Attention Mask", - "ETN_NSFWFilter": "NSFW Filter", - "ETN_Translate": "Translate Text", - "ETN_KritaOutput": "Krita Output", - "ETN_KritaSendText": "Send Text", - "ETN_KritaCanvas": "Krita Canvas", - "ETN_KritaSelection": "Krita Selection", - "ETN_KritaImageLayer": "Krita Image Layer", - "ETN_KritaMaskLayer": "Krita Mask Layer", - "ETN_Parameter": "Parameter", - "ETN_KritaStyle": "Krita Style", -} + +class ExternalToolingNodes(ComfyExtension): + async def get_node_list(self) -> list[type[io.ComfyNode]]: + return [ + nodes.LoadImageBase64, + nodes.LoadMaskBase64, + nodes.SendImageWebSocket, + nodes.ApplyMaskToImage, + nodes.ReferenceImage, + nodes.ApplyReferenceImages, + tile.CreateTileLayout, + tile.ExtractImageTile, + tile.ExtractMaskTile, + tile.GenerateTileMask, + tile.MergeImageTile, + region.BackgroundRegion, + region.DefineRegion, + region.ListRegionMasks, + region.AttentionMask, + nsfw.NSFWFilter, + translation.Translate, + krita.KritaOutput, + krita.KritaSendText, + krita.KritaCanvas, + krita.KritaSelection, + krita.KritaImageLayer, + krita.KritaMaskLayer, + krita.Parameter, + krita.KritaStyle, + ] + + +async def comfy_entrypoint(): + return ExternalToolingNodes() + + WEB_DIRECTORY = "./js" diff --git a/krita.py b/krita.py index 77ccd42..70b7836 100644 --- a/krita.py +++ b/krita.py @@ -8,6 +8,7 @@ from PIL import Image import server import comfy.samplers from comfy.comfy_types.node_typing import IO +from comfy_api.latest import io from .nodes import SendImageWebSocket @@ -72,37 +73,39 @@ class _BasicTypes(str): BasicTypes = _BasicTypes("BASIC") -class KritaOutput: +class KritaOutput(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return {"required": {"images": ("IMAGE",)}} + def define_schema(cls): + return io.Schema( + node_id="ETN_KritaOutput", + display_name="Krita Output", + category="krita", + inputs=[io.Image.Input("images")], + is_output_node=True, + ) - RETURN_TYPES = () - FUNCTION = "send_images" - OUTPUT_NODE = True - CATEGORY = "krita" - - def send_images(self, images): - return SendImageWebSocket().send_images(images, "PNG") - - -class KritaSendText: @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "value": (IO.ANY, {}), - "name": ("STRING", {"default": "Output"}), - "type": (["text", "markdown", "html"], {"default": "text"}), - } - } + def execute(cls, images: torch.Tensor): + return SendImageWebSocket.execute(images, "PNG") - RETURN_TYPES = () - FUNCTION = "send" - OUTPUT_NODE = True - CATEGORY = "krita" - def send(self, value: Any, name: str, type: str): +class KritaSendText(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="ETN_KritaSendText", + display_name="Send Text", + category="krita", + inputs=[ + io.AnyType.Input("value"), + io.String.Input("name", default="Output"), + io.Combo.Input("type", options=["text", "markdown", "html"], default="text"), + ], + is_output_node=True, + ) + + @classmethod + def execute(cls, value: Any, name: str, type: str): mime = { "text": "text/plain", "markdown": "text/markdown", @@ -115,72 +118,79 @@ class KritaSendText: except Exception as e: text = f"Could not convert to text: {e}" - print(f"Sending text: {name} = {text}") - return {"ui": {"text": [{"name": name, "text": text, "content-type": mime}]}} + return io.NodeOutput(ui={"text": [{"name": name, "text": text, "content-type": mime}]}) -class KritaCanvas: +class KritaCanvas(io.ComfyNode): @classmethod - def INPUT_TYPES(cls): - return {} + def define_schema(cls): + return io.Schema( + node_id="ETN_KritaCanvas", + display_name="Krita Canvas", + category="krita", + outputs=[ + io.Image.Output(display_name="image"), + io.Int.Output(display_name="width"), + io.Int.Output(display_name="height"), + io.Int.Output(display_name="seed"), + ], + ) - RETURN_TYPES = ("IMAGE", "INT", "INT", "INT") - RETURN_NAMES = ("image", "width", "height", "seed") - FUNCTION = "placeholder" - CATEGORY = "krita" - - def placeholder(self): - return (_placeholder_image(), 512, 512, 0) - - -class KritaSelection: @classmethod - def INPUT_TYPES(cls): - return {} - - RETURN_TYPES = (IO.MASK, IO.BOOLEAN) - RETURN_NAMES = ("mask", "active") - FUNCTION = "placeholder" - CATEGORY = "krita" - - def placeholder(self): - return (torch.ones(1, 512, 512), False) + def execute(cls): + return io.NodeOutput(_placeholder_image(), 512, 512, 0) -class KritaImageLayer: +class KritaSelection(io.ComfyNode): @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "name": ("STRING", {"default": "Image"}), - } - } + def define_schema(cls): + return io.Schema( + node_id="ETN_KritaSelection", + display_name="Krita Selection", + category="krita", + outputs=[io.Mask.Output(display_name="mask"), io.Boolean.Output(display_name="active")], + ) - RETURN_TYPES = ("IMAGE", "MASK") - RETURN_NAMES = ("image", "mask") - FUNCTION = "placeholder" - CATEGORY = "krita" - - def placeholder(self, name: str): - return (_placeholder_image(), torch.ones(1, 512, 512)) - - -class KritaMaskLayer: @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "name": ("STRING", {"default": "Mask"}), - } - } + def execute(cls): + return io.NodeOutput(torch.ones(1, 512, 512), False) - RETURN_TYPES = ("MASK",) - RETURN_NAMES = ("mask",) - FUNCTION = "placeholder" - CATEGORY = "krita" - def placeholder(self, name: str): - return (torch.ones(1, 512, 512),) +class KritaImageLayer(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="ETN_KritaImageLayer", + display_name="Krita Image Layer", + category="krita", + inputs=[io.String.Input("name", default="Image")], + outputs=[ + io.Image.Output(display_name="image"), + io.Mask.Output(display_name="mask"), + ], + ) + + @classmethod + def execute(cls, name: str): + return io.NodeOutput(_placeholder_image(), torch.ones(1, 512, 512)) + + +class KritaMaskLayer(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="ETN_KritaMaskLayer", + display_name="Krita Mask Layer", + category="krita", + inputs=[io.String.Input("name", default="Mask")], + outputs=[ + io.Mask.Output(display_name="mask"), + ], + ) + + @classmethod + def execute(cls, name: str): + return io.NodeOutput(torch.ones(1, 512, 512)) _param_types = [ @@ -193,71 +203,63 @@ _param_types = [ "prompt (positive)", "prompt (negative)", ] -_any_float = {"default": 0.0, "min": -sys.float_info.max, "max": sys.float_info.max} +_fmax = sys.float_info.max -class Parameter: +class Parameter(io.ComfyNode): @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "name": ("STRING", {"default": "Parameter"}), - "type": (_param_types, {"default": "auto"}), - "default": ("STRING", {"default": ""}), - }, - "optional": { - "min": ("FLOAT", _any_float), - "max": ("FLOAT", _any_float), - }, - } + def define_schema(cls): + return io.Schema( + node_id="ETN_Parameter", + display_name="Parameter", + category="krita", + inputs=[ + io.String.Input("name", default="Parameter"), + io.Combo.Input("type", options=_param_types, default="auto"), + io.String.Input("default", default=""), + io.Float.Input("min", default=0.0, min=-_fmax, max=_fmax, optional=True), + io.Float.Input("max", default=1.0, min=-_fmax, max=_fmax, optional=True), + ], + outputs=[io.AnyType.Output(display_name="value")], + ) - RETURN_TYPES = (BasicTypes,) - RETURN_NAMES = ("value",) - FUNCTION = "placeholder" - CATEGORY = "krita" - - def placeholder(self, name: str, type: str, default, min=0.0, max=1.0): + @classmethod + def execute(cls, name: str, type: str, default, min=0.0, max=1.0): if type == "number": - return (float(default),) + return io.NodeOutput(float(default)) elif type == "number (integer)": - return (int(default),) - return (default,) + return io.NodeOutput(int(default)) + return io.NodeOutput(default) -class KritaStyle: +class KritaStyle(io.ComfyNode): @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "name": ("STRING", {"default": "Style"}), - "sampler_preset": (["auto", "regular", "live"],), - } - } + def define_schema(cls): + return io.Schema( + node_id="ETN_KritaStyle", + display_name="Krita Style", + category="krita", + inputs=[ + io.String.Input("name", default="Style"), + io.Combo.Input("sampler_preset", options=["auto", "regular", "live"]), + ], + outputs=[ + io.Model.Output(display_name="model"), + io.Clip.Output(display_name="clip"), + io.Vae.Output(display_name="vae"), + io.String.Output(display_name="positive prompt"), + io.String.Output(display_name="negative prompt"), + io.Combo.Output( + display_name="sampler name", options=comfy.samplers.KSampler.SAMPLERS + ), + io.Combo.Output( + display_name="scheduler", options=comfy.samplers.KSampler.SCHEDULERS + ), + io.Int.Output(display_name="steps"), + io.Float.Output(display_name="guidance"), + ], + ) - RETURN_TYPES = ( - "MODEL", - "CLIP", - "VAE", - "STRING", - "STRING", - comfy.samplers.KSampler.SAMPLERS, - comfy.samplers.KSampler.SCHEDULERS, - "INT", - "FLOAT", - ) - RETURN_NAMES = ( - "model", - "clip", - "vae", - "positive prompt", - "negative prompt", - "sampler name", - "scheduler", - "steps", - "guidance", - ) - FUNCTION = "placeholder" - CATEGORY = "krita" - - def placeholder(self, name: str, sampler_preset: str): + @classmethod + def execute(cls, name: str, sampler_preset: str): raise NotImplementedError("This workflow must be started from Krita!") diff --git a/nodes.py b/nodes.py index bfbbdc8..5ca3883 100644 --- a/nodes.py +++ b/nodes.py @@ -11,18 +11,22 @@ from server import PromptServer, BinaryEventTypes from comfy.clip_vision import ClipVisionModel from comfy.sd import StyleModel +from comfy_api.latest import io -class LoadImageBase64: +class LoadImageBase64(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return {"required": {"image": ("STRING", {"multiline": False})}} + def define_schema(cls): + return io.Schema( + node_id="ETN_LoadImageBase64", + display_name="Load Image (Base64)", + category="external_tooling", + inputs=[io.String.Input("image", multiline=False)], + outputs=[io.Image.Output(display_name="image"), io.Mask.Output(display_name="mask")], + ) - RETURN_TYPES = ("IMAGE", "MASK") - CATEGORY = "external_tooling" - FUNCTION = "load_image" - - def load_image(self, image: str): + @classmethod + def execute(cls, image: str): _strip_prefix(image, "data:image/png;base64,") imgdata = base64.b64decode(image) img = Image.open(BytesIO(imgdata)) @@ -40,16 +44,19 @@ class LoadImageBase64: return (img, mask) -class LoadMaskBase64: +class LoadMaskBase64(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return {"required": {"mask": ("STRING", {"multiline": False})}} + def define_schema(cls): + return io.Schema( + node_id="ETN_LoadMaskBase64", + display_name="Load Mask (Base64)", + category="external_tooling", + inputs=[io.String.Input("mask", multiline=False)], + outputs=[io.Mask.Output(display_name="mask")], + ) - RETURN_TYPES = ("MASK",) - CATEGORY = "external_tooling" - FUNCTION = "load_mask" - - def load_mask(self, mask: str): + @classmethod + def execute(cls, mask: str): _strip_prefix(mask, "data:image/png;base64,") imgdata = base64.b64decode(mask) img = Image.open(BytesIO(imgdata)) @@ -60,22 +67,22 @@ class LoadMaskBase64: return (img.unsqueeze(0),) -class SendImageWebSocket: +class SendImageWebSocket(io.ComfyNode): @classmethod - def INPUT_TYPES(s): - return { - "required": { - "images": ("IMAGE",), - "format": (["PNG", "JPEG"], {"default": "PNG"}), - } - } + def define_schema(cls): + return io.Schema( + node_id="ETN_SendImageWebSocket", + display_name="Send Image (WebSocket)", + category="external_tooling", + inputs=[ + io.Image.Input("images"), + io.Combo.Input("format", options=["PNG", "JPEG"], default="PNG"), + ], + is_output_node=True, + ) - RETURN_TYPES = () - FUNCTION = "send_images" - OUTPUT_NODE = True - CATEGORY = "external_tooling" - - def send_images(self, images, format): + @classmethod + def execute(cls, images: torch.Tensor, format: str): results = [] for tensor in images: array = 255.0 * tensor.cpu().numpy() @@ -93,43 +100,7 @@ class SendImageWebSocket: "type": "output", }) - return {"ui": {"images": results}} - - -class CropImage: - """Deprecated, ComfyUI has an ImageCrop node now which does the same.""" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "image": ("IMAGE",), - "x": ( - "INT", - {"default": 0, "min": 0, "max": 8192, "step": 1}, - ), - "y": ( - "INT", - {"default": 0, "min": 0, "max": 8192, "step": 1}, - ), - "width": ( - "INT", - {"default": 512, "min": 1, "max": 8192, "step": 1}, - ), - "height": ( - "INT", - {"default": 512, "min": 1, "max": 8192, "step": 1}, - ), - } - } - - CATEGORY = "external_tooling" - RETURN_TYPES = ("IMAGE",) - FUNCTION = "crop" - - def crop(self, image, x, y, width, height): - out = image[:, y : y + height, x : x + width, :] - return (out,) + return io.NodeOutput(ui={"images": results}) def to_bchw(image: torch.Tensor): @@ -148,21 +119,22 @@ def mask_batch(mask: torch.Tensor): return mask -class ApplyMaskToImage: +class ApplyMaskToImage(io.ComfyNode): @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "image": ("IMAGE",), - "mask": ("MASK",), - } - } + def define_schema(cls): + return io.Schema( + node_id="ETN_ApplyMaskToImage", + display_name="Apply Mask to Image", + category="external_tooling", + inputs=[ + io.Image.Input("image"), + io.Mask.Input("mask"), + ], + outputs=[io.Image.Output(display_name="masked")], + ) - CATEGORY = "external_tooling" - RETURN_TYPES = ("IMAGE",) - FUNCTION = "apply_mask" - - def apply_mask(self, image: torch.Tensor, mask: torch.Tensor): + @classmethod + def execute(cls, image: torch.Tensor, mask: torch.Tensor): out = to_bchw(image) if out.shape[1] == 3: # Assuming RGB images out = torch.cat([out, torch.ones_like(out[:, :1, :, :])], dim=1) @@ -189,28 +161,26 @@ class _ReferenceImageData(NamedTuple): range: tuple[float, float] -class ReferenceImage: +class ReferenceImage(io.ComfyNode): @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "image": ("IMAGE",), - "weight": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0}), - "range_start": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0}), - "range_end": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0}), - }, - "optional": { - "reference_images": ("REFERENCE_IMAGE",), - }, - } + def define_schema(cls): + return io.Schema( + node_id="ETN_ReferenceImage", + display_name="Reference Image", + category="external_tooling", + inputs=[ + io.Image.Input("image"), + io.Float.Input("weight", default=1.0, min=0.0, max=10.0), + io.Float.Input("range_start", default=0.0, min=0.0, max=1.0), + io.Float.Input("range_end", default=1.0, min=0.0, max=1.0), + io.Custom("ReferenceImage").Input("reference_images", optional=True), + ], + outputs=[io.Custom("ReferenceImage").Output(display_name="reference_images")], + ) - CATEGORY = "external_tooling" - RETURN_TYPES = ("REFERENCE_IMAGE",) - RETURN_NAMES = ("reference_images",) - FUNCTION = "append" - - def append( - self, + @classmethod + def execute( + cls, image: torch.Tensor, weight: float, range_start: float, @@ -222,24 +192,25 @@ class ReferenceImage: return (imgs,) -class ApplyReferenceImages: +class ApplyReferenceImages(io.ComfyNode): @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "conditioning": ("CONDITIONING",), - "clip_vision": ("CLIP_VISION",), - "style_model": ("STYLE_MODEL",), - "references": ("REFERENCE_IMAGE",), - } - } + def define_schema(cls): + return io.Schema( + node_id="ETN_ApplyReferenceImages", + display_name="Apply Reference Images", + category="external_tooling", + inputs=[ + io.Conditioning.Input("conditioning"), + io.ClipVision.Input("clip_vision"), + io.StyleModel.Input("style_model"), + io.Custom("ReferenceImage").Input("references"), + ], + outputs=[io.Conditioning.Output(display_name="conditioning")], + ) - CATEGORY = "external_tooling" - RETURN_TYPES = ("CONDITIONING",) - FUNCTION = "apply" - - def apply( - self, + @classmethod + def execute( + cls, conditioning: list[list], clip_vision: ClipVisionModel, style_model: StyleModel, diff --git a/nsfw.py b/nsfw.py index 68b5f42..3bbd659 100644 --- a/nsfw.py +++ b/nsfw.py @@ -8,6 +8,7 @@ import torch.nn.functional as F from torch import Tensor from transformers import CLIPImageProcessor, CLIPConfig, CLIPVisionModel, PreTrainedModel from kornia.filters import box_blur +from comfy_api.latest import io from .nodes import to_bchw, to_bhwc @@ -76,7 +77,7 @@ class CLIPSafetyChecker(PreTrainedModel): class CachedModels: - _instance: WeakRef | None = None + _instance: CachedModels | None = None def __init__(self): model_dir = Path(__file__).parent / "safetychecker" @@ -91,11 +92,9 @@ class CachedModels: @classmethod def load(cls): - models = cls._instance and cls._instance() - if models is None: - models = cls() - cls._instance = WeakRef(models) - return models + if cls._instance is None: + cls._instance = CachedModels() + return cls._instance def download(self, url: str, target: Path): import requests @@ -118,29 +117,26 @@ class CachedModels: ) from e -class NSFWFilter: - models: CachedModels +class NSFWFilter(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="ETN_NSFWFilter", + display_name="NSFW Filter", + category="external_tooling", + inputs=[ + io.Image.Input("image"), + io.Float.Input("sensitivity", default=0.5, min=0.0, max=1.0, step=0.1), + ], + outputs=[io.Image.Output(display_name="image")], + ) @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "image": ("IMAGE",), - "sensitivity": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.10}), - }, - } - - RETURN_TYPES = ("IMAGE",) - FUNCTION = "check" - CATEGORY = "external_tooling" - - def __init__(self): - self.models = CachedModels.load() - - def check(self, image, sensitivity): + def execute(cls, image: Tensor, sensitivity: float): + models = CachedModels.load() image = to_bchw(image) - input = self.models.feature_extractor(image, do_rescale=False, return_tensors="pt") - filtered = self.models.safety_checker( + input = models.feature_extractor(image, do_rescale=False, return_tensors="pt") + filtered = models.safety_checker( images=image, clip_input=input.pixel_values, sensitivity=sensitivity ) - return (to_bhwc(filtered),) + return io.NodeOutput(to_bhwc(filtered)) diff --git a/pyproject.toml b/pyproject.toml index 6106704..40b7903 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "comfyui-tooling-nodes" description = "Provides nodes and server API extensions geared towards using ComfyUI as a backend for external tools." -version = "2.0.6" +version = "3.0.0" license = { file = "LICENSE" } [project.urls] diff --git a/region.py b/region.py index 2602d5f..bb33991 100644 --- a/region.py +++ b/region.py @@ -8,6 +8,7 @@ import torch.nn.functional as F import math from torch import Tensor, Size from comfy.model_patcher import ModelPatcher +from comfy_api.latest import io def downsample_mask(mask: Tensor, batch: int, target_size: int, original_shape: Size) -> Tensor: @@ -65,74 +66,82 @@ class Region(NamedTuple): return result -class BackgroundRegion: +Regions = io.Custom("Regions") + + +class BackgroundRegion(io.ComfyNode): @classmethod - def INPUT_TYPES(cls): - return {"required": {"conditioning": ("CONDITIONING",)}} + def define_schema(cls): + return io.Schema( + node_id="ETN_BackgroundRegion", + display_name="Background Region", + category="external_tooling/regions", + inputs=[io.Conditioning.Input("conditioning")], + outputs=[Regions.Output(display_name="regions")], + ) - CATEGORY = "external_tooling/regions" - RETURN_TYPES = ("REGIONS",) - FUNCTION = "define" - - def define(self, conditioning: list): + @classmethod + def execute(cls, conditioning: list): return (Region(None, None, conditioning),) -class DefineRegion: +class DefineRegion(io.ComfyNode): @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "mask": ("MASK",), - "conditioning": ("CONDITIONING",), - }, - "optional": { - "regions": ("REGIONS",), - }, - } + def define_schema(cls): + return io.Schema( + node_id="ETN_DefineRegion", + display_name="Define Region", + category="external_tooling/regions", + inputs=[ + io.Mask.Input("mask"), + io.Conditioning.Input("conditioning"), + Regions.Input("regions", optional=True), + ], + outputs=[Regions.Output(display_name="regions")], + ) - CATEGORY = "external_tooling/regions" - RETURN_TYPES = ("REGIONS",) - FUNCTION = "define" - - def define(self, mask: Tensor, conditioning: list, regions: Region | None = None): + @classmethod + def execute(cls, mask: Tensor, conditioning: list, regions: Region | None = None): if mask.dim() < 3: mask = mask.unsqueeze(0) - return (Region(regions, mask, conditioning),) + return io.NodeOutput(Region(regions, mask, conditioning)) -class ListRegionMasks: +class ListRegionMasks(io.ComfyNode): @classmethod - def INPUT_TYPES(cls): - return {"required": {"regions": ("REGIONS",)}} + def define_schema(cls): + return io.Schema( + node_id="ETN_ListRegionMasks", + display_name="List Region Masks", + category="external_tooling/regions", + inputs=[Regions.Input("regions")], + outputs=[io.Mask.Output(display_name="masks")], + ) - CATEGORY = "external_tooling/regions" - RETURN_TYPES = ("MASK",) - FUNCTION = "get_masks" - - def get_masks(self, regions: Region): - return (torch.stack([r.mask for r in regions.preprocess()], dim=0),) - - -class AttentionMask: @classmethod - def INPUT_TYPES(s): - return { - "required": { - "model": ("MODEL",), - "regions": ("REGIONS",), - } - } + def execute(cls, regions: Region): + return io.NodeOutput(torch.stack([r.mask for r in regions.preprocess()], dim=0)) - RETURN_TYPES = ("MODEL",) - FUNCTION = "attention_mask" - CATEGORY = "external_tooling/regions" - mask: Tensor - conds: list[Tensor] - batch_size: int +class AttentionMask(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="ETN_AttentionMask", + display_name="Regions Attention Mask", + category="external_tooling/regions", + inputs=[io.Model.Input("model"), Regions.Input("regions")], + outputs=[io.Model.Output(display_name="model")], + ) - def attention_mask(self, model: ModelPatcher, regions: Region): + @classmethod + def execute(cls, model: ModelPatcher, regions: Region): + AttentionMaskPatch(model, regions) + return io.NodeOutput(model) + + +class AttentionMaskPatch: + def __init__(self, model: ModelPatcher, regions: Region): new_model = model.clone() region_list = regions.preprocess() num_conds = len(region_list) @@ -208,4 +217,3 @@ class AttentionMask: new_model.set_model_attn2_patch(attn2_patch) new_model.set_model_attn2_output_patch(attn2_output_patch) - return (new_model,) diff --git a/tile.py b/tile.py index 0844e25..9567200 100644 --- a/tile.py +++ b/tile.py @@ -3,49 +3,25 @@ import numpy as np import numpy.typing as npt import torch from torch import Tensor +from comfy_api.latest import io IntArray = npt.NDArray[np.int_] class TileLayout: - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "image": ("IMAGE",), - "min_tile_size": ("INT", {"default": 512, "min": 64, "max": 8192, "step": 8}), - "padding": ("INT", {"default": 32, "min": 0, "max": 8192, "step": 8}), - "blending": ("INT", {"default": 8, "min": 0, "max": 256, "step": 8}), - } - } - - CATEGORY = "external_tooling/tiles" - RETURN_TYPES = ("TILE_LAYOUT",) - FUNCTION = "node" - - image_size: IntArray - tile_size: IntArray - padding: int - blending: int - tile_count: IntArray - - def node(self, image: Tensor, min_tile_size: int, padding: int, blending: int): - self.init(image, min_tile_size, padding, blending) - return (self,) - - def init(self, image: Tensor, min_tile_size: int, padding: int, blending: int): + def __init__(self, image: Tensor, min_tile_size: int, padding: int, blending: int): assert all([x % 8 == 0 for x in image.shape[-3:-1]]), "Image size must be divisible by 8" assert min_tile_size % 8 == 0, "Tile size must be divisible by 8" assert blending <= padding, "Blending must be smaller than padding" - self.image_size = np.array(image.shape[-3:-1]) - self.padding = padding - self.blending = blending - self.tile_count = np.maximum(1, self.image_size // (min_tile_size - 2 * padding)) + self.image_size: IntArray = np.array(image.shape[-3:-1]) + self.padding: int = padding + self.blending: int = blending + self.tile_count: IntArray = np.maximum(1, self.image_size // (min_tile_size - 2 * padding)) image_size_with_overlap = self.image_size + (self.tile_count - 1) * 2 * padding tile_size = np.ceil(image_size_with_overlap / self.tile_count) - self.tile_size = (np.ceil(tile_size / 8) * 8).astype(int) + self.tile_size: IntArray = (np.ceil(tile_size / 8) * 8).astype(int) def size(self, coord: IntArray): return self.end(coord) - self.start(coord) @@ -96,80 +72,108 @@ class TileLayout: image[rect] = (1 - mask) * image[rect] + mask * tile -class ExtractImageTile: +class CreateTileLayout(io.ComfyNode): @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "image": ("IMAGE",), - "layout": ("TILE_LAYOUT",), - "index": ("INT", {"min": 0}), - } - } + def define_schema(cls): + return io.Schema( + node_id="ETN_TileLayout", + display_name="Create Tile Layout", + category="external_tooling/tiles", + inputs=[ + io.Image.Input("image"), + io.Int.Input("min_tile_size", default=512, min=64, max=8192, step=8), + io.Int.Input("padding", default=32, min=0, max=8192, step=8), + io.Int.Input("blending", default=8, min=0, max=256, step=8), + ], + outputs=[io.Custom("TileLayout").Output(display_name="layout")], + ) - CATEGORY = "external_tooling/tiles" - RETURN_TYPES = ("IMAGE",) - FUNCTION = "slice" - - def slice(self, image: Tensor, layout: TileLayout, index: int): - return (layout.tile(image, index),) - - -class ExtractMaskTile: @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "mask": ("MASK",), - "layout": ("TILE_LAYOUT",), - "index": ("INT", {"min": 0}), - } - } + def execute(cls, image: Tensor, min_tile_size: int, padding: int, blending: int): + return io.NodeOutput(TileLayout(image, min_tile_size, padding, blending)) - CATEGORY = "external_tooling/tiles" - RETURN_TYPES = ("MASK",) - FUNCTION = "slice" - def slice(self, mask: Tensor, layout: TileLayout, index: int): +class ExtractImageTile(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="ETN_ExtractImageTile", + display_name="Extract Image Tile", + category="external_tooling/tiles", + inputs=[ + io.Image.Input("image"), + io.Custom("TileLayout").Input("layout"), + io.Int.Input("index", default=0, min=0), + ], + outputs=[io.Image.Output(display_name="tile")], + ) + + @classmethod + def execute(cls, image: Tensor, layout: TileLayout, index: int): + return io.NodeOutput(layout.tile(image, index)) + + +class ExtractMaskTile(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="ETN_ExtractMaskTile", + display_name="Extract Mask Tile", + category="external_tooling/tiles", + inputs=[ + io.Mask.Input("mask"), + io.Custom("TileLayout").Input("layout"), + io.Int.Input("index", default=0, min=0), + ], + outputs=[io.Mask.Output(display_name="tile")], + ) + + @classmethod + def execute(cls, mask: Tensor, layout: TileLayout, index: int): tile = layout.tile(mask.unsqueeze(3), index) - return (tile.squeeze(3),) + return io.NodeOutput(tile.squeeze(3)) -class GenerateTileMask: +class GenerateTileMask(io.ComfyNode): @classmethod - def INPUT_TYPES(cls): - return { - "required": {"layout": ("TILE_LAYOUT",), "index": ("INT", {"min": 0})}, - "optional": {"blend": ("BOOLEAN",)}, - } + def define_schema(cls): + return io.Schema( + node_id="ETN_GenerateTileMask", + display_name="Generate Tile Mask", + category="external_tooling/tiles", + inputs=[ + io.Custom("TileLayout").Input("layout"), + io.Int.Input("index", default=0, min=0), + io.Boolean.Input("blend", default=False, optional=True), + ], + outputs=[io.Mask.Output(display_name="mask")], + ) - CATEGORY = "external_tooling/tiles" - RETURN_TYPES = ("MASK",) - FUNCTION = "generate" - - def generate(self, layout: TileLayout, index: int, blend: bool = False): - return (layout.mask(layout.coord(index), blend=blend),) - - -class MergeImageTile: @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "image": ("IMAGE",), - "layout": ("TILE_LAYOUT",), - "index": ("INT", {"min": 0}), - "tile": ("IMAGE",), - } - } + def execute(cls, layout: TileLayout, index: int, blend: bool = False): + return io.NodeOutput(layout.mask(layout.coord(index), blend=blend)) - CATEGORY = "external_tooling/tiles" - RETURN_TYPES = ("IMAGE",) - FUNCTION = "merge" - def merge(self, image: Tensor, layout: TileLayout, index: int, tile: Tensor): +class MergeImageTile(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="ETN_MergeImageTile", + display_name="Merge Image Tile", + category="external_tooling/tiles", + inputs=[ + io.Image.Input("image"), + io.Custom("TileLayout").Input("layout"), + io.Int.Input("index", default=0, min=0), + io.Image.Input("tile"), + ], + outputs=[io.Image.Output(display_name="image")], + ) + + @classmethod + def execute(cls, image: Tensor, layout: TileLayout, index: int, tile: Tensor): assert index < layout.total_count, f"Index {index} out of range" if index == 0: image = image.clone() layout.merge(image, index, tile) - return (image,) + return io.NodeOutput(image) diff --git a/translation.py b/translation.py index d3e3520..8302ae4 100644 --- a/translation.py +++ b/translation.py @@ -10,6 +10,7 @@ from __future__ import annotations import re from functools import cache from typing import NamedTuple +from comfy_api.latest import io @cache @@ -61,17 +62,20 @@ def translate(text: str): return " ".join(translate_chunk(c.text, c.lang) for c in chunks) -class Translate: - @staticmethod - def INPUT_TYPES(): - return {"required": {"text": ("STRING", {"multiline": True})}} +class Translate(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="ETN_Translate", + display_name="Translate Text", + category="external_tooling", + inputs=[io.String.Input("text", multiline=True)], + outputs=[io.String.Output(display_name="translation")], + ) - CATEGORY = "external_tooling" - RETURN_TYPES = ("STRING",) - FUNCTION = "translate" - - def translate(self, text: str): - return (translate(text),) + @classmethod + def execute(cls, text: str): + return io.NodeOutput(translate(text)) _lang_regex = re.compile(r"(lang:\w\w)")