Author SHA1 Message Date
Acly 00ccb7dcf0 Add LoadImageCached and SaveImageCached (renamed from SendImageHTTP)
* short-lived in-memory cache for image transfers
* upload images via HTTP to cache and load/reference them in workflows
* save/store images in workflows to cache and download them via HTTP
2025-10-20 13:43:25 +02:00
Acly dfe014ae88 Fix region attention mask not being applied 2025-10-19 12:45:21 +02:00
Acly 6a7ae5ab70 Add SendImageHTTP node
* alternative to SendImageWebSocket
* requires an extra step, but transfers are much faster for large images
* doesn't involve saving files to disk
2025-10-19 12:06:19 +02:00
Acly d4cac6ac95 Change node definitions to "V3" schema, remove CropImage node 2025-10-18 23:59:05 +02:00
Aoi 929fdfcc13 Comment out translation package download print statement
Comment out print statement to prevent encoding errors.
2025-10-18 10:55:56 +02:00
10 changed files with 686 additions and 507 deletions
+42 -1
View File
@@ -33,7 +33,7 @@ Loads a mask (single channel) from a PNG embedded into the prompt as base64 stri
### Send Image (WebSocket)
Sends an output image over the client WebSocket connection as PNG binary data.
* Inputs: the image (RGB or RGBA)
* Inputs: the image (RGB or RGBA), supports batches
This will first send one binary message for each image in the batch via WebSocket:
```
@@ -44,6 +44,47 @@ That is two 32-bit integers (big endian) with values 1 and 2 followed by the PNG
{'type': 'executed', 'data': {'node': '<node ID>', 'output': {'images': [{'source': 'websocket', 'content-type': 'image/png', 'type': 'output'}, ...]}, 'prompt_id': '<prompt ID>}}
```
### Load Image from Cache
Loads an image or mask that has been uploaded previously into the workflow.
Uploaded images are temporarily stored in RAM rather than written to disk. This
method has less overhead compared to embedding images as base64 into the prompt,
but is more complex to implement.
* Inputs: id of an image that was uploaded previously
* Outputs: image (RGB) and mask (A of RGBA input, or first channel if no alpha present).
To upload an image, upload the _bytes_ of a PNG via a HTTP PUT request to
`/api/etn/image/{id}`. JPEG or other formats also work. Choose any `id` which
does not clash with other images you upload, and reference it in the node. The
request returns `201` if the image was uploaded and `200` if it was already
cached.
### Save Image to Cache
Stores an output image in RAM temporarily and allows retrieval over HTTP.
This is typically faster than WebSocket, especially for large images.
* Inputs: the image (RGB or RGBA). Batches are supported.
This node will send a JSON message over WebSocket when an image is ready:
```json
{
'type': 'executed',
'data': {
'node': '<node ID>',
'output': {
'images': [
{'source': 'http', 'id': '<image ID>', 'content-type': 'image/png', 'type': 'output'}
]
},
'prompt_id': 'prompt ID'
}
}
```
To download the images, send a HTTP GET request to `/api/etn/image/{id}` with
the image IDs from the message. Images will be cached for a few minutes.
## <a id="regions" href="#toc">Regions</a>
These nodes implement attention masking for arbitrary number of image regions. Text prompts only apply to the masked area.
+39 -56
View File
@@ -1,59 +1,42 @@
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.LoadImageCache,
nodes.SaveImageCache,
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"
+41
View File
@@ -15,6 +15,7 @@ import server
from .translation import available_languages, translate
from .krita import WorkflowExchange
from .nodes import image_cache
input_block_name = "model.diffusion_model.input_blocks.0.0.weight"
@@ -243,6 +244,13 @@ def has_invalid_filename(filename: str):
return None
async def image_sender(data: bytes):
mem = memoryview(data)
csize = 2**14
for i in range(0, len(mem), csize):
yield mem[i : i + csize]
_server: server.PromptServer | None = getattr(server.PromptServer, "instance", None)
if _server is not None:
_workflow_exchange = WorkflowExchange(_server)
@@ -277,6 +285,39 @@ if _server is not None:
except Exception as e:
return web.json_response(dict(error=str(e)), status=500)
@_server.routes.get("/api/etn/image/{id}")
async def get_image(request: web.Request):
try:
id = request.match_info.get("id", "")
data, content_type = image_cache.get(id)
if data is None or content_type is None:
return web.json_response(dict(error="Image not found"), status=404)
response = web.Response(
body=image_sender(data),
content_type=content_type,
headers={"Content-Length": str(len(data))},
)
return response
except Exception as e:
return web.json_response(dict(error=str(e)), status=500)
@_server.routes.put("/api/etn/image/{id}")
async def put_image(request: web.Request):
try:
id = request.match_info.get("id", "")
if id in image_cache:
return web.json_response(dict(status="cached"), status=200)
content_type = request.headers.get("Content-Type", "application/octet-stream")
data = bytearray()
async for chunk, _ in request.content.iter_chunks():
data.extend(chunk)
image_cache.insert(id, bytes(data), content_type)
return web.json_response(dict(status="success"), status=201)
except Exception as e:
return web.json_response(dict(error=str(e)), status=500)
@_server.routes.put("/api/etn/upload/{folder_name}/{filename}")
async def upload(request: web.Request):
folder_name = request.match_info.get("folder_name", "")
+139 -137
View File
@@ -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!")
+206 -110
View File
@@ -1,6 +1,9 @@
from __future__ import annotations
from copy import copy
from dataclasses import dataclass
import time
from typing import NamedTuple
from uuid import uuid4
from PIL import Image
import numpy as np
import base64
@@ -11,18 +14,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 +47,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 +70,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 +103,129 @@ class SendImageWebSocket:
"type": "output",
})
return {"ui": {"images": results}}
return io.NodeOutput(ui={"images": results})
class CropImage:
"""Deprecated, ComfyUI has an ImageCrop node now which does the same."""
class ImageCache:
@dataclass
class Entry:
data: bytes
content_type: str
timestamp: float
retrieved: int
def __init__(self):
self.images: dict[str, ImageCache.Entry] = {}
def add(self, image: Image.Image, format: str):
key = uuid4().hex
with BytesIO() as output:
image.save(output, format=format, quality=95, compress_level=1)
image_data = output.getvalue()
self.insert(key, image_data, f"image/{format.lower()}")
return key
def insert(self, key: str, data: bytes, content_type: str):
self.images[key] = ImageCache.Entry(
data=data,
content_type=content_type,
timestamp=time.time(),
retrieved=0,
)
def get(self, key: str, extend: bool = False):
entry = self.images.get(key)
if entry is None:
return None, None
entry.retrieved += 1
if extend:
entry.timestamp = time.time()
self.prune()
return entry.data, entry.content_type
def prune(self):
now = time.time()
keys_to_delete = []
for key, entry in self.images.items():
d = now - entry.timestamp
if (d > 60 and entry.retrieved > 1) or d > 600:
keys_to_delete.append(key)
for key in keys_to_delete:
del self.images[key]
def __contains__(self, key: str):
return key in self.images
image_cache = ImageCache()
class LoadImageCache(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_LoadImageCache",
display_name="Load Image from Cache",
category="external_tooling",
inputs=[io.String.Input("id", multiline=False)],
outputs=[io.Image.Output(display_name="image"), io.Mask.Output(display_name="mask")],
)
@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},
),
}
}
def execute(cls, id: str):
image_data, content_type = image_cache.get(id, extend=True)
if image_data is None:
raise ValueError(f"Image with ID {id} not found in cache.")
CATEGORY = "external_tooling"
RETURN_TYPES = ("IMAGE",)
FUNCTION = "crop"
img = Image.open(BytesIO(image_data))
w, h = img.size
c = len(img.getbands())
normalized = np.array(img).astype(np.float32) / 255.0
tensor = torch.from_numpy(normalized).reshape(1, h, w, c)
match c:
case 1:
image = tensor.expand(1, h, w, 3)
mask = tensor.reshape(1, h, w)
case 3:
image = tensor
mask = tensor[..., 0]
case 4:
image = tensor[..., :3]
mask = tensor[..., 3]
def crop(self, image, x, y, width, height):
out = image[:, y : y + height, x : x + width, :]
return (out,)
return io.NodeOutput(image, mask)
class SaveImageCache(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_SaveImageCache",
display_name="Save Image to Cache",
category="external_tooling",
inputs=[
io.Image.Input("images"),
io.Combo.Input("format", options=["PNG", "JPEG"], default="PNG"),
],
is_output_node=True,
)
@classmethod
def execute(cls, images: torch.Tensor, format: str):
results = []
for tensor in images:
array = 255.0 * tensor.cpu().numpy()
image = Image.fromarray(np.clip(array, 0, 255).astype(np.uint8))
key = image_cache.add(image, format)
results.append({
"source": "http",
"id": key,
"content-type": f"image/{format.lower()}",
"type": "output",
})
return io.NodeOutput(ui={"images": results})
def to_bchw(image: torch.Tensor):
@@ -148,21 +244,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 +286,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 +317,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,
+23 -27
View File
@@ -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))
+1 -1
View File
@@ -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]
+85 -73
View File
@@ -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,104 +66,112 @@ 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):
new_model = model.clone()
region_list = regions.preprocess()
num_conds = len(region_list)
@classmethod
def execute(cls, model: ModelPatcher, regions: Region):
return io.NodeOutput(AttentionMaskPatch.apply(model, regions))
class AttentionMaskPatch:
def __init__(self, region_list: list[Region]):
mask = torch.stack([r.mask for r in region_list], dim=0)
mask_sum = mask.sum(dim=0, keepdim=True)
assert mask_sum.sum() > 0, "There are areas that are zero in all masks."
self.mask = mask / mask_sum
self.conds = [r.conditioning[0][0] for r in region_list]
num_tokens = [cond.shape[1] for cond in self.conds]
self.num_tokens = [cond.shape[1] for cond in self.conds]
self.num_conds = len(region_list)
self.batch_size = 0
@staticmethod
def apply(model: ModelPatcher, regions: Region):
patch = AttentionMaskPatch(regions.preprocess())
def attn2_patch(q: Tensor, k: Tensor, v: Tensor, extra_options: dict):
assert k.mean() == v.mean(), "k and v must be the same."
device, dtype = q.device, q.dtype
if self.conds[0].device != device or self.conds[0].dtype != dtype:
self.conds = [cond.to(device, dtype=dtype) for cond in self.conds]
if self.mask.device != device or self.mask.dtype != dtype:
self.mask = self.mask.to(device, dtype=dtype)
if patch.conds[0].device != device or patch.conds[0].dtype != dtype:
patch.conds = [cond.to(device, dtype=dtype) for cond in patch.conds]
if patch.mask.device != device or patch.mask.dtype != dtype:
patch.mask = patch.mask.to(device, dtype=dtype)
cond_or_unconds = extra_options["cond_or_uncond"]
num_chunks = len(cond_or_unconds)
self.batch_size = q.shape[0] // num_chunks
patch.batch_size = q.shape[0] // num_chunks
q_chunks = q.chunk(num_chunks, dim=0)
k_chunks = k.chunk(num_chunks, dim=0)
lcm_tokens = lcm_for_list(num_tokens + [k.shape[1]])
lcm_tokens = lcm_for_list(patch.num_tokens + [k.shape[1]])
conds_tensor = [
cond.repeat(self.batch_size, lcm_tokens // num_tokens[i], 1)
for i, cond in enumerate(self.conds)
cond.repeat(patch.batch_size, lcm_tokens // patch.num_tokens[i], 1)
for i, cond in enumerate(patch.conds)
]
conds_tensor = torch.cat(conds_tensor, dim=0)
@@ -173,9 +182,9 @@ class AttentionMask:
qs.insert(0, q_chunks[i])
ks.insert(0, k_target)
else:
qs.insert(0, q_chunks[i].repeat(num_conds, 1, 1))
qs.insert(0, q_chunks[i].repeat(patch.num_conds, 1, 1))
ks.insert(0, conds_tensor)
for _ in range(num_conds - 1):
for _ in range(patch.num_conds - 1):
cond_or_unconds.insert(i, 0)
qs = torch.cat(qs, dim=0)
@@ -183,29 +192,32 @@ class AttentionMask:
return qs, ks, ks
def attn2_output_patch(out: Tensor, extra_options: dict):
num_conds = patch.num_conds
cond_or_unconds = extra_options["cond_or_uncond"]
mask_downsample = downsample_mask(
self.mask, self.batch_size, out.shape[1], extra_options["original_shape"]
patch.mask, patch.batch_size, out.shape[1], extra_options["original_shape"]
)
outputs: list[Tensor] = []
pos = 0
i = 0
while i < len(cond_or_unconds):
if cond_or_unconds[i] == 1: # uncond
outputs.append(out[pos : pos + self.batch_size])
pos += self.batch_size
outputs.append(out[pos : pos + patch.batch_size])
pos += patch.batch_size
else:
masked = out[pos : pos + num_conds * self.batch_size] * mask_downsample
masked = masked.view(num_conds, self.batch_size, out.shape[1], out.shape[2])
masked = out[pos : pos + num_conds * patch.batch_size] * mask_downsample
masked = masked.view(num_conds, patch.batch_size, out.shape[1], out.shape[2])
masked = masked.sum(dim=0)
outputs.append(masked)
pos += num_conds * self.batch_size
pos += num_conds * patch.batch_size
for _ in range(num_conds - 1):
cond_or_unconds.pop(i)
i += 1
return torch.cat(outputs, dim=0)
new_model = model.clone()
new_model.set_model_attn2_patch(attn2_patch)
new_model.set_model_attn2_output_patch(attn2_output_patch)
return (new_model,)
new_model.set_attachments("etn_attention_mask", patch)
return new_model
+95 -91
View File
@@ -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)
+15 -11
View File
@@ -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
@@ -43,7 +44,7 @@ def translate_chunk(text: str, language: str):
(p for p in available if p.from_code == language and p.to_code == target), None
)
assert pkg, f"Couldn't find package for translation from {language}"
print("Downloading and installing translation package", pkg)
# print("Downloading and installing translation package", pkg) # this will cause encoding errors
pkg.install()
text, embeddings = _extract_embeddings(text)
@@ -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)")