43 Commits
Author SHA1 Message Date
Acly 0abe742480 Version 3.4.0 2026-10-04 17:59:27 +09:00
Acly b3ae4aa2d9 Fix missing maskb batch dim; Make ApplyMaskToImage respect existing alpha 2026-09-26 18:44:00 +09:00
fukc-gihtub 1b8d81ce5a Add QwenImage21 model type 2026-09-23 11:33:15 +09:00
Sen-sou ca01116495 fix: add compatibility with ComfyUI INT8 models 2026-08-19 10:27:17 +02:00
Acly 5d3194f4d4 Version 3.3.0 2026-06-28 11:22:20 +09:00
fukc-gihtub d5812f900b Add Krea2 model type 2026-06-27 13:22:23 +09:00
Mutive 7064288fbe Add Anima region attention support (#67)
* Add Anima region attention support

* Added adaptation acknowledgement for Anima Attention couple implementation

* Address Anima attention cleanup feedback
2026-06-27 13:21:43 +09:00
Acly d82675092e Version 3.2.0 2026-05-31 11:15:36 +02:00
Acly a1e51904de Fix some type errors 2026-05-30 13:50:52 +02:00
Acly ffa130239b Add Anima Control LLLite nodes
* from kohya-ss/ComfyUI-Anima-LLLite
* split nodes into load/apply to cache model loading
2026-05-30 13:50:34 +02:00
Acly 2fd51d0d47 Workaround for import failure in transformers with certain pytorch installs #66
* seems to affects pytorch compiled with DISTRIBUTED=0, eg. Windows ROCm
* should probably be fixed in transformers somehow?
2026-05-30 11:49:19 +02:00
Acly d3c75155b4 Fix already cached image PUT requests (#65)
* Support already cached image PUT requests without aborting the connection
* conditionally send 100 Continue instead
2026-05-11 09:14:53 +02:00
Acly cbaef8d9c5 Version 3.1.4, fix type checks 2026-05-03 18:15:02 +02:00
FeepingCreature b2783d82a6 Add Anima model type 2026-05-01 17:53:32 +02:00
VERIGEN 09759222de Add ERNIE Image model detection 2026-05-01 17:53:02 +02:00
Acly 7fc3df1174 Version 3.1.3 2026-02-21 20:33:28 +01:00
Acly ed99942f86 Add mask output to KritaCanvas 2026-02-20 13:08:25 +01:00
Acly 2d395424ea Fix nsfw filter with transformers>5 2026-02-05 16:28:32 +01:00
Acly 9b9ea62dd8 Version 3.1.2 2026-01-31 18:01:29 +01:00
Acly ad36f89af3 Tiles: make multiple for tile layout configurable
* can now ensure eg. multiple of 16 tiles to be compatible with flux2 latent downsample factor
* default is 8, which matches previous hardcoded value
2026-01-26 17:28:46 +01:00
Alex 7130dcb2df Add KritaStyleAndPrompt node for synced prompts across workspaces
New node ETN_KritaStyleAndPrompt that works like KritaStyle but:
- Prompts and style sync between Generate/Live/Animation/Graph workspaces
- Outputs fully prepared prompts (wildcards evaluated, style merged)
- Model output includes extracted LoRAs from prompts
2026-01-24 13:15:12 +01:00
Acly 77186eda87 Model inspection: detect Flux 2 klein GGUF variants 2026-01-20 17:11:14 +01:00
Acly 24a7bd1a77 Version 3.1.1 2026-01-18 21:02:30 +01:00
Acly 2d14a03ad8 Model inspection: detect variants of Flux 2 (Klein-4B, Klein-9B) 2026-01-16 19:28:58 +01:00
Acly 9d2e03e8d5 Version 3.1.0 2026-01-05 10:18:20 +01:00
Acly c5606f8e8f API: print filename with stack traces when there is an error during inspection 2026-01-05 10:16:54 +01:00
Acly 79e9b6426f Support transmitting partial tiles/crops of the canvas in Krita workflows 2025-12-30 23:47:29 +01:00
Acly ad320a218c Support additional output info for Krita workflows: name, animation, layers 2025-12-29 21:20:57 +01:00
Jax a310f4593b Transmit request to resize the canvas with Krita Output node (#52)
* Added Krita Resize node for plugin
* Registered canvas resize node to the __init__.py
* Removed resizenode and integrated it into the "Krita output"
* Quick update to simplify the node to return the re-sized image to krita instead of a json
2025-12-29 19:35:06 +01:00
Acly 7d957dcfa7 Model inspection: support Z-Image SVDQ (Nunchaku) files 2025-12-22 11:25:09 +01:00
Acly 22cfd71f95 Fix detection of integer widget for Parameter node 2025-12-17 11:05:00 +01:00
Acly 21a2f44d4c Fix Parameter node min/max being reset to default when it's set to 0 #53
* Use a different default than 0 as workaround
* Don't want to change type of min/max as that would break workflows
2025-12-17 10:40:58 +01:00
Acly 0220252912 Version 3.0.1 2025-12-01 09:36:18 +01:00
Acly f447ef70fa Model inspection: support Z-Image GGUFs 2025-11-29 20:26:00 +01:00
Acly fb27a5bda8 Model inspection: support Lumina2, Z-Image, Flux2 2025-11-28 20:00:00 +01:00
Acly aa83259e66 Change image cache to take size into account 2025-11-09 15:01:52 +01:00
Acly 75c632df4b Version 3.0.0 2025-11-03 10:57:01 +01:00
Acly a088a2dde2 API: support pagination for /api/etn/model_info 2025-10-23 14:41:24 +02:00
Acly fbf99f2a08 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 14:39:42 +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
12 changed files with 2029 additions and 545 deletions
+69 -2
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.
@@ -164,6 +205,8 @@ There are various types of models that can be loaded as checkpoint, LoRA, Contro
#### Paramters
* `folder_name`: sub-directory in ComfyUI's models folder.
Supported model types: `checkpoints`, `diffusion_models`, `unet`, `unet_gguf`
* `limit=n`: (query parameter, optional) inspect at `n` models
* `offset=i`: (query parameter, optional) start with the `i`th model
#### Output
Lists available models with additional classification info:
@@ -177,7 +220,7 @@ Lists available models with additional classification info:
...
}
```
Possible values for base model: `sd15, sd20, sd21, sd3, sdxl, sdxl-refiner, ssd1b, svd, cascade-b, cascade-c, aura-flow, hunyuan-dit, flux, flux-schnell, lumina2, chroma, qwen-image`
Possible values for base model: `sd15, sd20, sd21, sd3, sdxl, sdxl-refiner, ssd1b, svd, cascade-b, cascade-c, aura-flow, hunyuan-dit, flux, flux-schnell, flux2, lumina2, z-image, chroma, qwen-image`
If base model is `sdxl`, the `type` attribute is set with possible values: `eps, edm, v-prediction, v-prediction-edm`
@@ -187,6 +230,24 @@ Detection supports quantized models:
Returns an entry `{"base_model": "unknown"}` for models with unknown format or which do not match any of the known base models.
#### Pagination
The query parameters limit and offset allow inspecting a subset of models per request.
Usually inspection is quite fast (it only looks at model headers), but it can be slow
in some cases due to anti-virus or slow harddrives.
```
GET /api/etn/model_info/checkpoints?limit=10&offset=20
```
This will return at most 10 models, starting with the 20th model in the list.
It also returns a special `_meta` entry in the output JSON:
```json
{
"checkpoint_20.safetensors": { ... },
"_meta": { "offset": 20, "count": 1, "total": 21 }
}
```
### GET /api/etn/languages
Returns a list of available languages for translation.
@@ -232,3 +293,9 @@ git clone https://github.com/Acly/comfyui-tooling-nodes.git
```
Restart ComfyUI and the nodes are functional.
## Acknowledgements
* Region nodes adapted from [laksjdjf/cgem156-ComfyUI](https://github.com/laksjdjf/cgem156-ComfyUI/blob/main/scripts/attention_couple/node.py)
* Control nodes adapted from [kohya-ss/ComfyUI-Anima-LLLite](https://github.com/kohya-ss/ComfyUI-Anima-LLLite)
+55 -57
View File
@@ -1,59 +1,57 @@
from . import api as api, nodes, tile, region, nsfw, translation, krita
from comfy_api.latest import ComfyExtension, io
from . import api as api
from . import control, krita, nodes, region, tile, translation
class ExternalToolingNodes(ComfyExtension):
async def get_node_list(self) -> list[type[io.ComfyNode]]:
node_list = [
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,
translation.Translate,
krita.KritaOutput,
krita.KritaSendText,
krita.KritaCanvas,
krita.KritaSelection,
krita.KritaImageLayer,
krita.KritaMaskLayer,
krita.Parameter,
krita.KritaStyle,
krita.KritaStyleAndPrompt,
control.ControlApply,
control.ControlLoad,
]
try: # see #66
from . import nsfw
node_list.append(nsfw.NSFWFilter)
except (ImportError, ModuleNotFoundError):
import traceback
print("[comfyui-tooling-nodes] WARNING: Could not import all nodes.")
traceback.print_exc()
return node_list
async def comfy_entrypoint():
return ExternalToolingNodes()
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",
}
WEB_DIRECTORY = "./js"
+115 -14
View File
@@ -1,6 +1,6 @@
from __future__ import annotations
from aiohttp import web
from typing import NamedTuple
from typing import Any, NamedTuple
from pathlib import Path
import json
import traceback
@@ -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"
@@ -43,6 +44,8 @@ model_names = {
"CosmosI2V": "cosmos",
"CosmosT2IPredict2": "cosmos-predict2",
"CosmosI2VPredict2": "cosmos-predict2",
"ZImage": "z-image",
"Lumina2": "lumina2",
"WAN21_T2V": "wan21",
"WAN21_I2V": "wan21",
"WAN21_FunControl2V": "wan21-fun",
@@ -53,6 +56,11 @@ model_names = {
"ACEStep": "ace-step",
"Omnigen2": "omnigen2",
"QwenImage": "qwen-image",
"QwenImage21": "qwen-image21",
"ErnieImage": "ernie-image",
"Flux2": "flux2",
"Anima": "anima",
"Krea2": "krea2",
}
gguf_architectures = {
@@ -117,12 +125,15 @@ def inspect_safetensors(filename: str, model_type: str, is_checkpoint: bool):
raw_name = base_model.__class__.__name__
if raw_name == "SDXL":
model_type = base_model.model_type(cfg).name.lower().replace("_", "-")
if raw_name == "Flux2":
hidden_size = unet_config.get("hidden_size", 0)
model_type = {3072: "klein-4b", 4096: "klein-9b"}.get(hidden_size, "dev")
if not raw_name:
return {"base_model": "unknown"}
base_model_name = model_names.get(raw_name, "unknown")
result = {"base_model": base_model_name}
result: dict[str, Any] = {"base_model": base_model_name}
result["is_inpaint"] = (
base_model_name in ["sd15", "sdxl"] and input_count > 4
) or raw_name == "FluxInpaint"
@@ -141,6 +152,7 @@ def inspect_safetensors(filename: str, model_type: str, is_checkpoint: bool):
return result
return {"base_model": "unknown"}
except Exception as e:
print("[comfyui-tooling-nodes] Error inspecting file", filename)
traceback.print_exc()
return {"base_model": "unknown", "error": f"Failed to detect base model: {e}"}
@@ -150,12 +162,16 @@ def detect_svdq(cfg: dict) -> str | None:
if comfy_config := md.get("comfy_config"):
if isinstance(comfy_config, str):
comfy_config = json.loads(comfy_config)
return comfy_config.get("model_class")
model_class = md.get("model_class")
if model_class == "NunchakuFluxTransformer2dModel":
return "Flux"
if model_class == "NunchakuQwenImageTransformer2DModel":
return "QwenImage"
if model_class := comfy_config.get("model_class"):
return model_class
match md.get("model_class"):
case "NunchakuFluxTransformer2dModel":
return "Flux"
case "NunchakuQwenImageTransformer2DModel":
return "QwenImage"
case "NunchakuZImageTransformer2DModel":
return "ZImage"
return None
@@ -167,6 +183,9 @@ def inspect_gguf(filename: str, model_type: str):
try:
path = folder_paths.get_full_path(model_type, filename)
if path is None:
raise Exception(f"Could not find full path for {model_type}/{filename}")
reader = gguf.GGUFReader(path)
arch_field = reader.get_field("general.architecture")
if arch_field is not None:
@@ -177,19 +196,45 @@ def inspect_gguf(filename: str, model_type: str):
arch_str = str(arch_field.parts[arch_field.data[-1]], encoding="utf-8")
else: # stable-diffusion.cpp, requires conversion. not handled for now
return {"base_model": "flux", "is_inpaint": False}
if arch_str == "flux" and any(
t.name.startswith("distilled_guidance_layer")
for t in itertools.islice(reader.tensors, 5)
):
arch_str = "chroma"
# Detect Z-Image (modified Lumina2)
if arch_str == "lumina2":
for t in reader.tensors:
if t.name == "cap_embedder.1.bias" and t.shape[0] == 3840:
arch_str = "z-image"
break
# Detect Flux variants
result_type = None
if arch_str == "flux":
for t in reader.tensors:
if t.name.startswith("distilled_guidance_layer"):
arch_str = "chroma"
break
elif t.name == "double_stream_modulation_img.lin.weight":
arch_str = "flux2"
if t.shape[0] == 3072:
result_type = "klein-4b"
elif t.shape[0] == 4096:
result_type = "klein-9b"
break
result = {
"base_model": gguf_architectures.get(arch_str, arch_str),
"is_inpaint": False,
}
if result_type is not None:
result["type"] = result_type
try:
result["quant"] = reader.get_field("general.file_type").lower()
except Exception as e:
if file_type := reader.get_field("general.file_type"):
result["quant"] = file_type.contents().lower()
except Exception:
result["quant"] = "gguf"
return result
@@ -204,17 +249,22 @@ def inspect_diffusion_model(filename: str, model_type: str, is_checkpoint: bool)
return inspect_safetensors(filename, model_type, is_checkpoint)
def inspect_models(model_type: str):
def inspect_models(model_type: str, params: dict[str, str]):
try:
try:
files = folder_paths.get_filename_list(model_type)
except KeyError:
return web.json_response({"error": f"Model folder not found: {model_type}"})
limit = int(params.get("limit", "1000"))
offset = int(params.get("offset", "0"))
files_range = files[offset : offset + limit]
is_checkpoint = model_type == "checkpoints"
info = {
filename: inspect_diffusion_model(filename, model_type, is_checkpoint)
for filename in files
for filename in files_range
}
if "limit" in params:
info["_meta"] = dict(offset=offset, count=len(files_range), total=len(files))
return web.json_response(info)
except Exception as e:
traceback.print_exc()
@@ -243,6 +293,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)
@@ -253,11 +310,11 @@ if _server is not None:
error = has_invalid_folder_name(folder_name)
if error is not None:
return error
return inspect_models(folder_name)
return inspect_models(folder_name, request.rel_url.query)
@_server.routes.get("/api/etn/model_info")
async def api_model_info(request):
return inspect_models("checkpoints")
return inspect_models("checkpoints", request.rel_url.query)
@_server.routes.get("/api/etn/languages")
async def languages(request):
@@ -277,6 +334,50 @@ 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)
async def put_image(request: web.Request):
try:
id = request.match_info.get("id", "")
if id in image_cache:
await request.release() # Consume and discard the data to avoid connection abort
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)
async def _put_image_expect_handler(request: web.Request):
if request.match_info.get("id", "") in image_cache:
# Skip "100 Continue" since we don't need the data, return 200 immediately.
return web.json_response(dict(status="cached"), status=200)
# otherwise run default aiohttp handler
return None
_server.app.router.add_route(
"PUT", "/api/etn/image/{id}", put_image, expect_handler=_put_image_expect_handler
)
@_server.routes.put("/api/etn/upload/{folder_name}/{filename}")
async def upload(request: web.Request):
folder_name = request.match_info.get("folder_name", "")
+859
View File
@@ -0,0 +1,859 @@
"""ControlNet-LLLite for Anima (DiT) — ComfyUI port (v2 architecture).
Adapted from kohya-ss/ComfyUI-Anima-LLLite
https://github.com/kohya-ss/ComfyUI-Anima-LLLite
Apache-2.0 license
Adapted from kohya-ss/sd-scripts. The on-disk weight format is the v2
named-key format (per-module key prefix = lllite_name, shared encoder under
``lllite_conditioning1.*``, depth embedding split per-module as
``{name}.depth_embed``); legacy ``lllite_modules.*`` files are rejected.
Differences vs. the sd-scripts reference (``networks/control_net_lllite_anima.py``):
* No dependency on ``library.utils`` — uses stdlib logging.
* Module discovery filters the LLM-Adapter sub-tree by class identity in
addition to the path-based check (ComfyUI ships two distinct ``Attention``
classes that share the bare class name).
* ``LLLiteModuleDiT`` keeps a ``restore()`` method (and an idempotent
``apply_to()``); ComfyUI patches/unpatches the original Linear around
every sampler call via ``set_model_unet_function_wrapper``.
* Forward pass casts ``x`` and ``cond_emb`` to the LLLite parameter dtype
so autocast / mixed-precision flows that hand us a different dtype than
the LLLite weights still work.
* CFG batch-size and sequence-length mismatches fall back to identity
instead of asserting, so a slightly-off cond image cannot abort sampling.
* The training-side ``AnimaControlNetLLLiteWrapper`` is omitted; ComfyUI
integrates via ``model_function_wrapper`` in nodes.py instead.
"""
from __future__ import annotations
from copy import copy
import logging
import os
from dataclasses import dataclass
from typing import Any
import folder_paths
import safetensors
import safetensors.torch
import torch
import torch.nn.functional as F
from comfy.model_patcher import ModelPatcher
from comfy_api.latest import io
from torch import nn
logger = logging.getLogger("comfyui-tooling-nodes")
# Class names of the modules that LLLite injects into. The LLM-Adapter uses
# a different ``Attention`` class with the same bare name; we filter it by
# path (``llm_adapter`` in the qualified name) and by the ``is_selfattn``
# attribute presence.
TARGET_ATTENTION_CLASS = "Attention"
TARGET_MLP_CLASS = "GPT2FeedForward"
LLM_ADAPTER_NAME = "llm_adapter"
LLLITE_ARCH_VERSION = "2"
# ----------------------------------------------------------------------------
# target_layers: atomic specifiers and presets
# ----------------------------------------------------------------------------
ATOMIC_SPECIFIERS: tuple[str, ...] = (
"self_attn_q_pre",
"self_attn_kv_pre",
"cross_attn_q_pre",
"mlp_fc1_pre",
)
PRESETS: dict = {
"self_attn_q": ("self_attn_q_pre",),
"self_attn_qkv": ("self_attn_q_pre", "self_attn_kv_pre"),
"self_attn_qkv_cross_q": ("self_attn_q_pre", "self_attn_kv_pre", "cross_attn_q_pre"),
}
def parse_target_layers(spec: str) -> tuple[str, ...]:
"""Resolve a ``target_layers`` spec to a canonical atomic tuple.
Accepts a preset name (``"self_attn_qkv"``) or a comma-separated list of
atomic specifiers (``"self_attn_q_pre,mlp_fc1_pre"``). Returns the atomics
in ``ATOMIC_SPECIFIERS`` order with duplicates removed.
"""
if not isinstance(spec, str):
raise TypeError(f"target_layers must be str, got {type(spec).__name__}")
spec = spec.strip()
if not spec:
raise ValueError("target_layers spec is empty")
if spec in PRESETS:
parts = list(PRESETS[spec])
else:
parts = [p.strip() for p in spec.split(",") if p.strip()]
bad = [p for p in parts if p not in ATOMIC_SPECIFIERS]
if bad:
raise ValueError(
f"unknown target_layers atomic specifier(s): {bad}. "
f"valid atomic={list(ATOMIC_SPECIFIERS)}, presets={list(PRESETS)}"
)
return tuple(a for a in ATOMIC_SPECIFIERS if a in parts)
# ----------------------------------------------------------------------------
# Conditioning1 trunk (v2)
# ----------------------------------------------------------------------------
def _gn(channels: int) -> nn.GroupNorm:
g = 8
while g > 1 and channels % g != 0:
g //= 2
return nn.GroupNorm(g, channels)
class _ResBlock(nn.Module):
def __init__(self, ch: int):
super().__init__()
self.norm1 = _gn(ch)
self.conv1 = nn.Conv2d(ch, ch, kernel_size=3, padding=1)
self.norm2 = _gn(ch)
self.conv2 = nn.Conv2d(ch, ch, kernel_size=3, padding=1)
def forward(self, x: torch.Tensor) -> torch.Tensor:
h = self.conv1(F.silu(self.norm1(x)))
h = self.conv2(F.silu(self.norm2(h)))
return x + h
ASPP_DEFAULT_DILATIONS: tuple[int, ...] = (1, 2, 4, 8)
class _ASPP(nn.Module):
def __init__(self, ch: int, dilations: tuple[int, ...] = ASPP_DEFAULT_DILATIONS):
super().__init__()
assert len(dilations) >= 1, "ASPP needs at least one dilation"
branches = []
for d in dilations:
if d == 1:
conv = nn.Conv2d(ch, ch, kernel_size=1)
else:
conv = nn.Conv2d(ch, ch, kernel_size=3, padding=d, dilation=d)
branches.append(nn.Sequential(conv, _gn(ch), nn.SiLU()))
self.branches = nn.ModuleList(branches)
self.global_pool = nn.AdaptiveAvgPool2d(1)
self.global_conv = nn.Sequential(nn.Conv2d(ch, ch, kernel_size=1), _gn(ch), nn.SiLU())
n_branches = len(dilations) + 1
self.proj = nn.Sequential(nn.Conv2d(ch * n_branches, ch, kernel_size=1), _gn(ch), nn.SiLU())
def forward(self, x: torch.Tensor) -> torch.Tensor:
h, w = x.shape[-2:]
outs = [b(x) for b in self.branches]
g = self.global_conv(self.global_pool(x))
g = F.interpolate(g, size=(h, w), mode="bilinear", align_corners=False)
outs.append(g)
return self.proj(torch.cat(outs, dim=1))
class _Conditioning1(nn.Module):
def __init__(
self,
cond_dim: int,
cond_emb_dim: int,
n_resblocks: int,
use_aspp: bool = False,
aspp_dilations: tuple[int, ...] = ASPP_DEFAULT_DILATIONS,
cond_in_channels: int = 3,
):
super().__init__()
assert cond_dim % 2 == 0, f"cond_dim must be even, got {cond_dim}"
assert cond_in_channels >= 1, f"cond_in_channels must be >= 1, got {cond_in_channels}"
ch_half = cond_dim // 2
self.cond_in_channels = cond_in_channels
self.conv1 = nn.Conv2d(cond_in_channels, ch_half, kernel_size=4, stride=4, padding=0)
self.norm1 = _gn(ch_half)
self.conv2 = nn.Conv2d(ch_half, ch_half, kernel_size=3, stride=1, padding=1)
self.norm2 = _gn(ch_half)
self.conv3 = nn.Conv2d(ch_half, cond_dim, kernel_size=4, stride=4, padding=0)
self.norm3 = _gn(cond_dim)
self.resblocks = nn.ModuleList([_ResBlock(cond_dim) for _ in range(n_resblocks)])
self.aspp = _ASPP(cond_dim, aspp_dilations) if use_aspp else None
self.proj = nn.Conv2d(cond_dim, cond_emb_dim, kernel_size=1)
self.out_norm = nn.LayerNorm(cond_emb_dim)
def forward(self, x: torch.Tensor) -> torch.Tensor:
h = F.silu(self.norm1(self.conv1(x)))
h = F.silu(self.norm2(self.conv2(h)))
h = F.silu(self.norm3(self.conv3(h)))
for rb in self.resblocks:
h = rb(h)
if self.aspp is not None:
h = self.aspp(h)
h = self.proj(h)
b, c, hh, ww = h.shape
h = h.view(b, c, hh * ww).permute(0, 2, 1).contiguous()
h = self.out_norm(h)
return h
# ----------------------------------------------------------------------------
# LLLite module (v2: FiLM + SiLU + 5D path + depth embedding)
# ----------------------------------------------------------------------------
class LLLiteModuleDiT(nn.Module):
def __init__(
self,
name: str,
org_module: nn.Linear,
cond_emb_dim: int,
mlp_dim: int,
dropout: float | None = None,
multiplier: float = 1.0,
):
super().__init__()
self.lllite_name = name
# Wrap in a list so the original Linear is not registered as a submodule
# and its weights stay out of state_dict.
self.org_module = [org_module]
self.cond_emb_dim = cond_emb_dim
self.mlp_dim = mlp_dim
self.dropout = dropout
self.multiplier = multiplier
in_dim = org_module.in_features
self.down = nn.Linear(in_dim, mlp_dim)
self.mid = nn.Linear(mlp_dim + cond_emb_dim, mlp_dim)
# FiLM: cond_local -> (gamma, beta), zero-init for identity at start.
self.cond_to_film = nn.Linear(cond_emb_dim, 2 * mlp_dim)
nn.init.zeros_(self.cond_to_film.weight)
nn.init.zeros_(self.cond_to_film.bias)
self.up = nn.Linear(mlp_dim, in_dim)
nn.init.zeros_(self.up.weight)
nn.init.zeros_(self.up.bias)
self.cond_emb: torch.Tensor | None = None
self.org_forward = None
# Set by the parent ControlNetLLLiteDiT after construction.
self.layer_idx: int = -1
self._depth_embeds_ref: list[nn.Parameter] = []
def apply_to(self):
if self.org_forward is None:
self.org_forward = self.org_module[0].forward
self.org_module[0].forward = self.forward
def restore(self):
if self.org_forward is not None:
self.org_module[0].forward = self.org_forward
self.org_forward = None
def forward(self, x: torch.Tensor) -> torch.Tensor:
# Input layouts:
# self/cross attention q/k/v: (B, S, D) — already flattened in the Anima block
# mlp.layer1: (B, T, H, W, D) — passed un-flattened
# Flatten the 5D case to 3D for the LLLite path and reshape on exit.
if self.multiplier == 0.0 or self.cond_emb is None:
return self.org_forward(x)
orig_shape = x.shape
is_5d = x.dim() == 5
if is_5d:
B, T, H, W, D = orig_shape
x = x.reshape(B, T * H * W, D)
cx = self.cond_emb # (B_c, S, cond_emb_dim)
# Broadcast cond_emb to the runtime batch (CFG cond+uncond, multi-cond).
if x.shape[0] != cx.shape[0]:
if x.shape[0] % cx.shape[0] != 0:
return self.org_forward(x.reshape(orig_shape) if is_5d else x)
cx = cx.repeat(x.shape[0] // cx.shape[0], 1, 1)
if x.shape[1] != cx.shape[1]:
return self.org_forward(x.reshape(orig_shape) if is_5d else x)
# Run the LLLite mini-MLP in its own parameter dtype, then cast the
# correction back to ``x``'s dtype before adding. Robust to autocast
# flows where x and LLLite weights have different dtypes.
param_dtype = self.down.weight.dtype
x_proc = x if x.dtype == param_dtype else x.to(param_dtype)
if cx.dtype != param_dtype or cx.device != x.device:
cx = cx.to(device=x.device, dtype=param_dtype)
# Per-module depth embedding (zero-init so it's a no-op at train start).
if self._depth_embeds_ref:
depth_e = self._depth_embeds_ref[0][self.layer_idx]
if depth_e.dtype != param_dtype or depth_e.device != x.device:
depth_e = depth_e.to(device=x.device, dtype=param_dtype)
cond_local = cx + depth_e
else:
cond_local = cx
h = F.silu(self.down(x_proc))
gb = self.cond_to_film(cond_local)
gamma, beta = gb.chunk(2, dim=-1)
m = self.mid(torch.cat([cond_local, h], dim=-1))
m = m * (1 + gamma) + beta
m = F.silu(m)
if self.dropout is not None and self.training:
m = F.dropout(m, p=self.dropout)
out = self.up(m) * self.multiplier
if out.dtype != x.dtype:
out = out.to(x.dtype)
y = self.org_forward(x + out)
if is_5d:
# org Linear out_features may differ from in_features — recover with -1.
y = y.reshape(orig_shape[0], orig_shape[1], orig_shape[2], orig_shape[3], -1)
return y
# ----------------------------------------------------------------------------
# ControlNetLLLiteDiT
# ----------------------------------------------------------------------------
class ControlNetLLLiteDiT(nn.Module):
def __init__(
self,
dit: nn.Module,
cond_emb_dim: int = 32,
mlp_dim: int = 64,
target_layers: str = "self_attn_q",
dropout: float | None = None,
multiplier: float = 1.0,
cond_dim: int = 64,
cond_resblocks: int = 1,
use_aspp: bool = False,
aspp_dilations: tuple[int, ...] = ASPP_DEFAULT_DILATIONS,
cond_in_channels: int = 3,
inpaint_masked_input: bool = False,
):
super().__init__()
atomics = parse_target_layers(target_layers)
self.cond_emb_dim = cond_emb_dim
self.mlp_dim = mlp_dim
self.target_layers = target_layers
self.target_atomics = atomics
self.dropout = dropout
self.multiplier = multiplier
self.cond_dim = cond_dim
self.cond_resblocks = cond_resblocks
self.use_aspp = use_aspp
self.aspp_dilations = tuple(aspp_dilations) if use_aspp else ()
# 4ch (RGB+mask) inpainting metadata. `inpaint_masked_input` records the training-time
# RGB-masking policy for cond_image preparation; it does not alter the forward pass here.
self.cond_in_channels = cond_in_channels
self.inpaint_masked_input = inpaint_masked_input
self.conditioning1 = _Conditioning1(
cond_dim,
cond_emb_dim,
cond_resblocks,
use_aspp=use_aspp,
aspp_dilations=aspp_dilations,
cond_in_channels=cond_in_channels,
)
modules = self._create_modules(dit, cond_emb_dim, mlp_dim, atomics, dropout, multiplier)
self.lllite_modules = nn.ModuleList(modules)
n = len(self.lllite_modules)
self.depth_embeds = nn.Parameter(torch.zeros(n, cond_emb_dim))
for i, m in enumerate(self.lllite_modules):
m.layer_idx = i
m._depth_embeds_ref = [self.depth_embeds]
aspp_info = f"aspp={'on' + str(list(self.aspp_dilations)) if use_aspp else 'off'}"
inpaint_info = (
f", inpaint=on(masked_input={inpaint_masked_input})" if cond_in_channels != 3 else ""
)
logger.info(
"ControlNet-LLLite (Anima v%s): created %d modules for target=%r "
"(atomics=%s), cond_in_channels=%d, cond_dim=%d, cond_resblocks=%d, %s, "
"cond_emb_dim=%d, mlp_dim=%d%s",
LLLITE_ARCH_VERSION,
n,
target_layers,
list(atomics),
cond_in_channels,
cond_dim,
cond_resblocks,
aspp_info,
cond_emb_dim,
mlp_dim,
inpaint_info,
)
@staticmethod
def _attn_atomic_match(is_self_attn: bool, child_name: str, atomics: tuple[str, ...]) -> bool:
if "output_proj" in child_name:
return False
if is_self_attn:
if child_name == "q_proj":
return "self_attn_q_pre" in atomics
if child_name in ("k_proj", "v_proj"):
return "self_attn_kv_pre" in atomics
return False
else:
if child_name == "q_proj":
return "cross_attn_q_pre" in atomics
return False # cross_attn K,V live in text-embedding space
def _create_modules(
self,
dit: nn.Module,
cond_emb_dim: int,
mlp_dim: int,
atomics: tuple[str, ...],
dropout: float | None,
multiplier: float,
) -> list[LLLiteModuleDiT]:
modules: list[LLLiteModuleDiT] = []
want_mlp_fc1 = "mlp_fc1_pre" in atomics
any_attn = any(
a in atomics for a in ("self_attn_q_pre", "self_attn_kv_pre", "cross_attn_q_pre")
)
for name, module in dit.named_modules():
if LLM_ADAPTER_NAME in name:
continue
cls = module.__class__.__name__
def _is_linear_like(module):
return (
hasattr(module, "in_features")
and hasattr(module, "out_features")
and callable(getattr(module, "forward", None))
)
if any_attn and cls == TARGET_ATTENTION_CLASS:
# The Anima-block Attention exposes is_selfattn; the LLM-Adapter
# Attention does not — skip the latter even if path filter misses.
if not hasattr(module, "is_selfattn"):
continue
is_self_attn = bool(module.is_selfattn)
for child_name, child in module.named_children():
if not _is_linear_like(child):
continue
if not self._attn_atomic_match(is_self_attn, child_name, atomics):
continue
full_name = f"lllite_dit.{name}.{child_name}".replace(".", "_")
modules.append(
LLLiteModuleDiT(
full_name, child, cond_emb_dim, mlp_dim, dropout, multiplier
)
)
elif want_mlp_fc1 and cls == TARGET_MLP_CLASS:
child = getattr(module, "layer1", None)
if not _is_linear_like(child):
continue
full_name = f"lllite_dit.{name}.layer1".replace(".", "_")
modules.append(
LLLiteModuleDiT(full_name, child, cond_emb_dim, mlp_dim, dropout, multiplier)
)
return modules
def set_cond_image(self, cond_image: torch.Tensor | None):
"""cond_image: (B, 3, H*16, W*16) in [-1, 1]; ``None`` clears."""
if cond_image is None:
for m in self.lllite_modules:
m.cond_emb = None
return
cx = self.conditioning1(cond_image) # (B, S, cond_emb_dim)
for m in self.lllite_modules:
m.cond_emb = cx
def clear_cond_image(self):
self.set_cond_image(None)
def set_multiplier(self, multiplier: float):
self.multiplier = multiplier
for m in self.lllite_modules:
m.multiplier = multiplier
def apply_to(self):
for m in self.lllite_modules:
m.apply_to()
def restore(self):
for m in self.lllite_modules:
m.restore()
# ----------------------------------------------------------------------------
# Save / load (named-key format; legacy lllite_modules.* is rejected)
# ----------------------------------------------------------------------------
_INTERNAL_MODULES_PREFIX = "lllite_modules."
_INTERNAL_COND_PREFIX = "conditioning1."
_INTERNAL_DEPTH_KEY = "depth_embeds"
_SAVED_COND_PREFIX = "lllite_conditioning1."
_SAVED_DEPTH_SUFFIX = ".depth_embed"
def _from_saved_state_dict(lllite: ControlNetLLLiteDiT, weights_sd: dict) -> dict:
"""Rewrite a v2 named-key state dict back to the internal layout."""
name_to_idx = {m.lllite_name: i for i, m in enumerate(lllite.lllite_modules)}
n_modules = len(name_to_idx)
out: dict = {}
depth_slices: dict = {}
for k, v in weights_sd.items():
if k.startswith(_SAVED_COND_PREFIX):
out[_INTERNAL_COND_PREFIX + k[len(_SAVED_COND_PREFIX) :]] = v
continue
if k.endswith(_SAVED_DEPTH_SUFFIX):
name = k[: -len(_SAVED_DEPTH_SUFFIX)]
if name in name_to_idx:
depth_slices[name_to_idx[name]] = v
continue
head, dot, tail = k.partition(".")
if dot and head in name_to_idx:
out[f"{_INTERNAL_MODULES_PREFIX}{name_to_idx[head]}.{tail}"] = v
continue
out[k] = v
if depth_slices:
missing = [i for i in range(n_modules) if i not in depth_slices]
if missing:
raise RuntimeError(f"depth_embed slices missing for module idx(es) {missing}")
out[_INTERNAL_DEPTH_KEY] = torch.stack([depth_slices[i] for i in range(n_modules)], dim=0)
return out
def load_lllite_weights(lllite: ControlNetLLLiteDiT, file: str, strict: bool = False):
weights_sd = safetensors.torch.load_file(file)
if any(k.startswith(_INTERNAL_MODULES_PREFIX) for k in weights_sd):
raise RuntimeError(
f"weights at {file} appear to be in a legacy ControlNet-LLLite weight format "
f"(keys starting with '{_INTERNAL_MODULES_PREFIX}'). The current code uses a "
f"named-key format (per-module key prefix = lllite_name, e.g. "
f"'lllite_dit_blocks_0_self_attn_q_proj.down.weight'). Re-train with the current codebase."
)
converted = _from_saved_state_dict(lllite, weights_sd)
info = lllite.load_state_dict(converted, strict=strict)
logger.info("loaded LLLite weights from %s: %s", file, info)
return info
def read_lllite_metadata(file: str) -> dict:
if os.path.splitext(file)[1] != ".safetensors":
raise RuntimeError(f"Must use .safetensors files, got {file}")
with safetensors.safe_open(file, framework="pt") as f:
return f.metadata() or {}
# ----------------------------------------------------------------------------
# ComfyUI nodes for Anima ControlNet-LLLite
# ----------------------------------------------------------------------------
def _get_inner_dit(model) -> torch.nn.Module:
"""Reach the underlying Anima DiT (nn.Module) from a ComfyUI ModelPatcher."""
inner = getattr(model, "model", None)
if inner is None:
raise RuntimeError("Input MODEL has no .model attribute (not a ModelPatcher?)")
dit = getattr(inner, "diffusion_model", None)
if dit is None:
raise RuntimeError("MODEL.model has no .diffusion_model — not a UNet/DiT model?")
return dit
def _target_cond_hw(latent_h: int, latent_w: int, patch_spatial: int = 2) -> tuple[int, int]:
"""Return the (H, W) the cond image / mask must be resized to.
The LLLite ``conditioning1`` Conv has stride 16, so the cond image must be
sized to ``latent_HW * 8`` in input pixel space (= ``token_HW * 16`` after
DiT patchify with patch_spatial=2). The DiT internally pads the latent up
to a multiple of ``patch_spatial`` (see ``MiniTrainDIT.forward`` →
``pad_to_patch_size``), so we mirror that rounding here — otherwise odd
latent dims (e.g. 1032 px → 129 latent) yield a token-count mismatch that
silently bypasses every LLLite module.
"""
padded_h = ((latent_h + patch_spatial - 1) // patch_spatial) * patch_spatial
padded_w = ((latent_w + patch_spatial - 1) // patch_spatial) * patch_spatial
return padded_h * 8, padded_w * 8
def _prepare_cond_image(
image: torch.Tensor,
latent_h: int,
latent_w: int,
device: torch.device,
dtype: torch.dtype,
patch_spatial: int = 2,
) -> torch.Tensor:
"""ComfyUI IMAGE (B,H,W,3) in [0,1] → (1,3,H*8,W*8) in [-1,1]."""
if image.ndim == 4 and image.shape[-1] == 3:
# (B, H, W, 3) -> (B, 3, H, W)
img = image.permute(0, 3, 1, 2).contiguous()
else:
raise ValueError(f"Unexpected cond image shape: {tuple(image.shape)} (expected B,H,W,3)")
img = img[:1] # use first frame only
target_h, target_w = _target_cond_hw(latent_h, latent_w, patch_spatial)
if img.shape[-2] != target_h or img.shape[-1] != target_w:
img = F.interpolate(img, size=(target_h, target_w), mode="bicubic", align_corners=False)
img = img.clamp(0.0, 1.0)
img = img * 2.0 - 1.0
return img.to(device=device, dtype=dtype)
def _prepare_mask(
mask: torch.Tensor,
latent_h: int,
latent_w: int,
device: torch.device,
dtype: torch.dtype,
patch_spatial: int = 2,
) -> torch.Tensor:
"""ComfyUI MASK (B,H,W) in [0,1] → (1,1,H*8,W*8) binarized at 0.5.
Returns the mask in ``{0.0, 1.0}`` (1 = inpaint area, 0 = keep). The caller
is responsible for the ``*2-1`` rescale before concat with RGB.
"""
if mask.ndim == 3:
m = mask.unsqueeze(1) # (B, 1, H, W)
elif mask.ndim == 4 and mask.shape[1] == 1:
m = mask
else:
raise ValueError(f"Unexpected mask shape: {tuple(mask.shape)} (expected B,H,W or B,1,H,W)")
m = m[:1]
target_h, target_w = _target_cond_hw(latent_h, latent_w, patch_spatial)
if m.shape[-2] != target_h or m.shape[-1] != target_w:
m = F.interpolate(m.float(), size=(target_h, target_w), mode="nearest")
m = (m >= 0.5).to(dtype=dtype)
return m.to(device=device)
def _build_inpaint_cond_image(
rgb_pm1: torch.Tensor, mask01: torch.Tensor, masked_input: bool
) -> torch.Tensor:
"""rgb_pm1: (1,3,H,W) in [-1,1], mask01: (1,1,H,W) in {0,1}. Returns (1,4,H,W).
Mirrors ``_build_inpaint_cond_image`` in the sd-scripts training / inference
code: the mask channel is rescaled to ``[-1, +1]`` (matches the RGB range),
and if ``masked_input`` is set the RGB is zeroed where ``mask >= 0.5``.
"""
if masked_input:
keep = (mask01 < 0.5).to(rgb_pm1.dtype)
rgb_pm1 = rgb_pm1 * keep
mask_pm1 = mask01.to(rgb_pm1.dtype) * 2.0 - 1.0
return torch.cat([rgb_pm1, mask_pm1], dim=1)
ETNControlNet = io.Custom("ETN_CONTROL_NET")
class ControlLoad(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_control_load",
display_name="Load ControlNet (tooling-nodes)",
description="Loads ControlNet weights. Currently only supports Anima LLLite weights.",
category="external_tooling",
inputs=[
io.Model.Input("model"),
io.Combo.Input("weights", folder_paths.get_filename_list("controlnet")),
],
outputs=[
io.Model.Output("out_model", "model"),
ETNControlNet.Output("control_net"),
],
)
@classmethod
def execute(cls, model: ModelPatcher, weights: str): # type: ignore[override]
weights_path = folder_paths.get_full_path("controlnet", weights)
if weights_path is None or not os.path.isfile(weights_path):
raise FileNotFoundError(f"LLLite weights not found: {weights}")
# Architecture is fully determined by the trained weights — read everything
# from metadata rather than exposing knobs that would just cause load errors.
meta = read_lllite_metadata(weights_path)
if "lllite.version" not in meta:
raise RuntimeError(
"Unrecognized model. This node currently only loads Anima LLLite weights."
)
ce_dim = int(meta.get("lllite.cond_emb_dim", 32))
m_dim = int(meta.get("lllite.mlp_dim", 64))
# v2 records the canonical atomic form under lllite.target_atomics; fall back
# to the legacy preset key, then to the v1 default.
tl = meta.get("lllite.target_atomics", meta.get("lllite.target_layers", "self_attn_q"))
cond_dim = int(meta.get("lllite.cond_dim", 64))
cond_resblocks = int(meta.get("lllite.cond_resblocks", 1))
use_aspp = str(meta.get("lllite.use_aspp", "false")).lower() == "true"
aspp_dilations_meta = meta.get("lllite.aspp_dilations")
if use_aspp and aspp_dilations_meta:
aspp_dilations = tuple(int(d) for d in aspp_dilations_meta.split(",") if d.strip())
else:
aspp_dilations = ASPP_DEFAULT_DILATIONS
cond_in_channels = int(meta.get("lllite.cond_in_channels", 3))
inpaint_masked_input = (
str(meta.get("lllite.inpaint_masked_input", "false")).lower() == "true"
)
lllite = ControlNetLLLiteDiT(
_get_inner_dit(model),
cond_emb_dim=ce_dim,
mlp_dim=m_dim,
target_layers=tl,
multiplier=1.0,
cond_dim=cond_dim,
cond_resblocks=cond_resblocks,
use_aspp=use_aspp,
aspp_dilations=aspp_dilations,
cond_in_channels=cond_in_channels,
inpaint_masked_input=inpaint_masked_input,
)
load_lllite_weights(lllite, weights_path, strict=False)
lllite.eval().requires_grad_(False)
return io.NodeOutput(model, lllite)
class ControlApply(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_control_apply",
display_name="Apply ControlNet (tooling-nodes)",
description="Applies ControlNet conditioning. Currently only supports Anima LLLite weights.",
category="external_tooling",
inputs=[
io.Model.Input("model"),
ETNControlNet.Input("control_net"),
io.Image.Input("image"),
io.Mask.Input("mask", optional=True),
io.Float.Input("strength", default=1.0, min=-10.0, max=10.0, step=0.01),
io.Float.Input("start_percent", default=0.0, min=0.0, max=1.0, step=0.001),
io.Float.Input("end_percent", default=1.0, min=0.0, max=1.0, step=0.001),
],
outputs=[io.Model.Output("model")],
)
@classmethod
def execute( # type: ignore[override]
cls,
model: ModelPatcher,
control_net: ControlNetLLLiteDiT,
image: torch.Tensor,
strength: float,
start_percent: float,
end_percent: float,
mask: torch.Tensor | None = None,
):
dit = _get_inner_dit(model)
patch_spatial = int(getattr(dit, "patch_spatial", 2))
lllite = control_net
lllite.set_multiplier(strength)
# Mask / cond_in_channels consistency: 4ch weights need a MASK, 3ch weights ignore it.
if lllite.cond_in_channels == 4 and mask is None:
raise ValueError("ControlNet weights require a mask input (inpaint mode)")
if lllite.cond_in_channels != 4 and mask is not None:
mask = None
# Convert percent range -> sigma range (start_percent=0 → sigma_max).
model_sampling = model.get_model_object("model_sampling")
sigma_start = float(model_sampling.percent_to_sigma(start_percent))
sigma_end = float(model_sampling.percent_to_sigma(end_percent))
# Capture image / mask tensors (cloned to detach from any upstream caching)
src_image = image.detach().clone()
src_mask = mask.detach().clone() if mask is not None else None
is_inpaint = lllite.cond_in_channels == 4
# Cache for the per-resolution preprocessed cond image (avoids repeat resize)
cache: dict[str, Any] = {"cond_image_pp": None, "key": None, "lllite_loaded_to": None}
# Capture any previously-installed wrapper BEFORE we clone — model_options
# has a single "model_function_wrapper" slot, so without delegation a second
# wrapper-installing node would silently no-op the first. Mirrors the
# ChromaRadianceOptions pattern in comfy_extras/nodes_chroma_radiance.py.
old_wrapper = model.model_options.get("model_function_wrapper")
def _call_next(apply_model, input_x, timestep, c):
if old_wrapper is not None:
return old_wrapper(apply_model, {"input": input_x, "timestep": timestep, "c": c})
return apply_model(input_x, timestep, **c)
def wrapper(apply_model, args):
input_x = args["input"]
timestep = args["timestep"]
c = args["c"]
# Step-range gate: skip LLLite entirely when current sigma is outside
# [sigma_end, sigma_start]. percent_to_sigma maps 0.0 → sigma_max,
# 1.0 → sigma_min, so the active window is sigma_end <= sigma <= sigma_start.
sigma = float(timestep.max().item())
if not (sigma_end <= sigma <= sigma_start):
return _call_next(apply_model, input_x, timestep, c)
# Anima latent shape: (B, C, T, H, W) — take spatial dims from the tail.
latent_h, latent_w = int(input_x.shape[-2]), int(input_x.shape[-1])
device = input_x.device
dtype = input_x.dtype
# Move LLLite to the runtime device/dtype lazily.
tag = (device, dtype)
if cache["lllite_loaded_to"] != tag:
lllite.to(device=device, dtype=dtype)
cache["lllite_loaded_to"] = tag
cache["cond_image_pp"] = None # invalidate
key = (latent_h, latent_w, device, dtype)
if cache["key"] != key or cache["cond_image_pp"] is None:
rgb = _prepare_cond_image(
src_image, latent_h, latent_w, device, dtype, patch_spatial
)
if is_inpaint:
assert src_mask is not None, "Cannot use inpaint control-net without a mask"
mk = _prepare_mask(src_mask, latent_h, latent_w, device, dtype, patch_spatial)
cache["cond_image_pp"] = _build_inpaint_cond_image(
rgb, mk, lllite.inpaint_masked_input
)
else:
cache["cond_image_pp"] = rgb
cache["key"] = key
lllite.set_multiplier(strength)
lllite.set_cond_image(cache["cond_image_pp"])
lllite.apply_to()
try:
return _call_next(apply_model, input_x, timestep, c)
finally:
lllite.restore()
lllite.clear_cond_image()
m = model.clone()
m.set_model_unet_function_wrapper(wrapper)
return (m,)
+3 -5
View File
@@ -32,7 +32,6 @@ function loadImage(base64) {
}
const canvasIcon = loadImage("data:image/webp;base64,UklGRg4KAABXRUJQVlA4WAoAAAAQAAAAYwAAYwAAQUxQSNsDAAARoIRs/yI5+uAHU1kZiOu6u7u7xnObuK9NpBoKCoqi1jd6WonrsD7ROc2e3OJJr28T94aGhobYj+8w9v/X/7+n3UNETAD+b1IW1crLLrMhbYJbRs9Zv2VvuXbqVK28d8v6OaNvCdpIXkaT5InJgJgRABeNbirRYKlp9EUAJB8ZlUq2XAtI1wQIhn1RIUlV1Y5UVUmy0jwsACQPt7CtsloQSBcE6D6jSFKVRlVJFmd0B8QeClSSSn7/ACCdQjBjL0mlRSW5d0YA+4KESpJKnd8dHQsweBeptK7krsGAWIIgppKkkn+NAKSNoO9KUplLJVf2hliCIKWyrZIfXwYIgCdLzHXpSVgXvElt07bcKAAKSs2TUgvWIHiH2p6SX98pi+jgIrEFwUJqO6Ty1HY6qGwObEHwPrU9KqkOUNkS2AKwjNoeqXRSuV7EGtZQO3BVuQhiDV9Q3aIyhFiTDVS3qHwMYgtBK9Utcn932A/20nHlGoglwdOn1DEqn4XYQfA7PfB7AKuCRiqdVzZCbKD+EL14qB4WBTOoPlA2Qsyh7i9f/FUH44LBVHpRORhibr0/ms2hf5XerPaFYcFoqi+UDeaafNJkCsFeenRvnRnBLfSp3mBqNNUfygZTc/wyx9QGv6w3A2zxyxZDwV56da+h3mW/lA1dVPPLKUOXnfIL/8UurvnltKH+Vb9UDXU/4ZeyoaDkl5Ih/ET1h/InU81+aTYjmOOXOaYm06sTzAD3n/JJ7X5T/YtUXyiLfU3JGnp0jZjCS+oPfQnG7ylR/aAs3WmuexO92dTdHKZWfFGZCot3fkX1gbL1dhv1SZVerCb1NvB0K9U9ZevTsNo/OUEPnkj628Hgj33w8WBY7h9up7ql3B72tYUn51ToeHnOk7B+/tSV6pYum3q+PVwbtVDdUbZE1yKPzyY/UV1R/hg/i1zKhKRIdUNZTCZIPtB9Rvo71QXl78mM7sjrxVFSpJPFJLoY+b06jn904ccovhp5vjaOWkjNk5Ithfha5PuKqLCozFyXFxWiK5D3/o1xuiVPW9K4sT/yf35DFi47mpejy8Ks4Xw4+UAcxRvKeaisj6P4Abjaf3xWSDYcbaNmtM3RDUkhG98fDt8+IyvEK3fV2Fa1M6psW9u1Mi5kM26H28EDM7IofPPj7WUaLG//+M0wymY8EMD54M7JaRqF8cKPvy4eqtROq56uVQ4Vv/54YRxGaTr5zgB+vOjpl5IsicIwSrJ35sx5J0uiMIySLHnp6YvgUel/z6iXojTL0jRJ0jTL0uilUff0F/i3/qIb7nn06WHDnn70nhsuqhf8hxQAVlA4IAwGAACwHgCdASpkAGQAPm0wk0akIqGhLRGrUIANiWYA1BHh/t2rC93/Hf2Was/feJ2MrzB6VP6M9gD9Nunp5rP2y9Z70q+gV/Y/+B1kHoAeWv+1Xwa/t5+5PtQXQfhgKtwfp++wwU1Dx6fSXsD+VV7JPQ5/aRrblS2nsaggvO1Mch8UhK6pYtQxLM/VgrswZ0vLV8b6SwudSWaCFSHUiXEUQWX6krc9GtWanHMeaDd9wRYCfO5TwpYkgGAIkaLI4p6taB375EUfaVubYzKMfHSz2KpivsjWF0Vf+YJbACgi8j86d6EiJhQFF31NBJdS+QrGtJ2RUJRbahp1MXso6/J8AAD+/TKL/9q5tf/zOmF5Fe8B0Zn0yX3C0VLv0zxxv2+WH/dbz//rc2RS4TC1UzFVQiXVn5+Y0r+RsfJPsfPNuT02INz8gty7fI7fA/D1Wj2Jv+4RwdpyXs+cRxaT84bme5rMmPf+BH7NDUPKsj7GJ+w/6nBW2vsiPalWPfvBk6AQ3kCHmVecXkcnOgpoZ4ruAF/9Ze93DG5/8Y32x8b/CKPRt1jaXXy2LnoPvSNUT77gbB+/7vI1pfBfUHJsSwheIXY7QSixh7Ya8IliO3wqvI/uIFZAZd9pL8R1gRpYouBoyL5uIuGWQAZC5SKY0SruTf66stUOJVO9hlokeb5lWVzo7FO/Oeb/oj9iK4bqFhNZLCfqsBlH/OeefoP9sFdl7Mq1xmsevmzkfgwyiXg5hxMIP/Wa0JMPVl+XEFqTveAf1M8IBDu/pX/hCEnMn1n15Smyf72eDXKQqBrvp6BugyXXaJ05FDoz8MONUFh4rcjGL7AcijbcZ0SYwJkoeAKBW/I/sjKzTRtTP2E1fLB/8TWnzieHznDAKdlTuY2nSVTwCqFZNcFeFn7boziHOmYBLJin52d874mq1pHmJnulhT96LbKVW4vAT5PnY5F9TzmnDMwIFm4IAuEaA8X8XLE4Hp+AUEG4oswxRbVfOfxNJRyxFO3UB+v+ALgMP8kOf0uK3/3WOq4o/roivfzvW/fXviTC0mx+352hGaO+axx6vFa3eIkFsUEXCdo2LFHIlM8BtPuGUhgvM3oygIMAgvmUKILe0DFYVXhG/QoLi3sYfaoK/f0tX+fNnXhhxwEj1/Ct2Z64g0qWmkgwNkyy8m90EK1HsX0Q10CHVakDZePz5ts37u3GCANwGHQzWB+hNsevqjuU3qT95yGs0jjOtI/IjKsH9JbAmZkjGvNCPOC+FYUkOkwao9sOESY6zCgx9CM7g2LU4/CSHGoe2t0vWV/cMDH1HzI+Wa/yYp9CLDIh7J7iJd/2KnixeJvOhbUvbr9gubyyQU1iO5bnD9T536j++jKDVIk0Fwzk+d+j2eueHsIFJUvdyo2TyxP0kJbWr36R1s3giryqPvrsR5SkXx16+xqDrX4elhqh+1FwzNnSF5Lj5EUT/UC2rJvoAikbnvQ3NtJ9e83++idf3ja4FaLcUDxhoN5Rl5Ziz1LvF9iVeb6Su0QWYoRyBbyZ/pRbgYyhlAU/tonH7Wt+KhPDmXKIo0u4FDbAXM8avbFk4ax6e/dYITOCe+9dVEgcTOnBfhv0Yotd3EzNjZkLz4ksKGtFXcWIZRJ5YAyfzPYsyPex6/6ud9r2Ha9oxhVSIJV418e83qcPOIPlpe+LVGc69W6eC83l/zloqM9D6zQMkfqrjBZNpRkQS0sn8sxSu3s5qzhtH8cvjZk83gMqdfnHnl+1bvA7BI/g4+ePU7HUb9vK3Qw35bVmDcXa8xxWS2NQj8iWMH1cbHLXlboQsaCxIZoo+SeXR6ePUw3k6C/OxgqhjzExMJjLdBjoBeWYt3RPG2foTvx0T0Iz8ukdrCRJMG6HaR+6/f4nG/4xkr/fLhGlqOE/hBDBhuqnANj1CrujVDs2YayTvPuIcqCpNd3i8fOR8DfCq9ytS55F8akKneS6poHfB3bhjWbcIQXPvFzS7S5xLHEVoaixOwp0TL/8cQ8dxriyeddu5kCTyY7KepMQoeR+Pyn04nElkt9qqfYCTqHDtXBriC/UZh9AAAAAAAA=")
const outputIcon = loadImage("data:image/webp;base64,UklGRrIHAABXRUJQVlA4WAoAAAAQAAAAjwAAOwAAQUxQSKoCAAARkMbsnyFJ/2TVySSdzPJs27bxzbbtuznbtm3btm3bt+ikkkonlfxPU1X9n57zh4iYAPifZKQ2T7LIkKDkQ1FPrezCF/hjvq+dx7mwiqM3HnJ058yGXkorEfGCKZV6I6o+quOM4bMsgU5bfE1KOh4LEbGnv2RnUSer50DGFwxJ2qwWGTATERFfZPzBZNR9P1ZXTksgVdaODHgt/H5hCJjP0ME6eiLfCqTLypKBSPYdsmY2OjpWy0KOlF+EkIFk/DvnZ2tIzZA0a0kHUtok0KfWkxieJQQZBQksq3QWiXODEGSnYRsq76mxjJQK0cDdKjY1XpLSWyKYVYFTSyxLqCZSvexKppZnZDCZTNpD6BKTi2iIRbqzZc6gW8x5m1ptOCEuYaJ74B1T6RkhjPQXkugiuC9EBSk38wftRACZsdKLEHGRQjJSCyUgnwichSgdj4jYW64KqQsywN1E1JJqReqtWyEvLtOTFHMtfG8GXMYT6C5zQLIjqXiJm+guNw2ZOqTu+0uJ7sKyg2w+Urv9GdxdGoN0JKnh/qBnIE1LlH6LiHNAMZ5SWQkoJAJHcQ7iTUNlNyE7Vga4a7CsoFqP0AmQjqdm5dTWGJRTvqfTSu4ONe7VNRs0LiTzLK3cNEHsGWhuZ+gom0hlNMgXZ7T4WF16Q6YRuZZPAS7TYrGUoPgFEnZPUC3EKDEf0O6WSGFmMiVoxejw3UDcHC6c21oFNLZjVNhGgxrkHC+cOtQKtBa5aQkCLL4VBGDZsbYzu7uF6AGouPQ1t7mTIv5AMw8EZLmhbx0QizqHgJMer5MQwHkG7NP2YmgzCMqRPYff0cIWDSgOwbrcgOFnhcqLhQNaeSB4h1XxDZh29LX4kXVt5YABrZJBkM/acIBsx5IG/AyGpS1UrkrF4tm98K8uVlA4IOIEAABwHACdASqQADwAPm0uk0ckIiGhLjUJmIANiWgOuBpEsADI+tE9j+ifkl+QHyd1n+77tGXf04+Cpx/549gD9X+lF5t/2d9cn0T+gB/d/8R1gHoAeWx+3nwZfuZ6T+aq/0Dtm/zaCN5VfOk9UfrX8A/Ry9E1elWp+7DqL5PqqU3p7QQgEkexRJ8VwS4d4Xk5cyjjrDbvzKxXEwCzN96k0RfNIpiQ2YsSbcyoRxX5fFzllf+dOy8uCMy/ebkjS0ONwuzkRR55zP3zjA4e+C969ch1Ab9LgcrpUqQ8MOvEXnQmqQkxXM4x4nA1f+jJAAD+8tXq96F285yEhMGbOWPp352/yRnzPmWRyRibmd800tluUOW4IyIZz2Hw1xYA9/xsSgKy0yQS//7BbNldSPJ+MCX/mxBqrttfeQX/mf/+2AcZ7Z1wDrdNoOnt8ISIu33p34GUAqqEyPwtdrhMf55SjsQmwUtm/I/qPiQ0ZOPv5ci6kCP0Ddb2jRr38UOXvi54DlMMmkxTs7j/J3jUQYepc7xEgTVVJFcf+8P//18L/9gA//9fHuK11sTUs+RbYDzZKn0uM2PnTEUpAJAT4wETKSU4KDn3wR3rUGGycaAPVo40AjzO5g7VkynMJuo3M2vclcmfmS3ygBGDqjGHQybd03tnGdkGOCGLLlTEABV+UgJxYxH2YxRX9zUULXFnfnBpPxQcOdq+zs1zi9uI2iAqSKwh8Dhch2Ytz8iaZLW39S3+3pmGdITR49+nlHjcG4xNVSYRLFLRmEj/H/I+7qd90N6AF9aRDuUFH1O7ONRGjEQGvPMEF0Fj5atb5w9tjc1pcKTsaWvT2GbF9NQ31HmpaLAgs8szVbuEC8GHKCESxKmx+Hrh5ZwjrNihG3KL0H1n3/g/WetlSYEFYsYTXQmgyUGCVIILkJYrRIdLB5iVAPrseYWKCT8HgJuCUAhaqRO+6jn1fkplsC0yYCveVI+yyDsVr98kmO5arhQ3u+aqKVUvJ8xZL9as4008lN9DkKcRhvC4BwWdhupsqUYwLQmaQhLxP15875P/r43c8r4NI4sLDiCi7Rzww1dWNTyThiA07x8b/zTaFC9Sz+jtZDpRPoSf3LS+TmvHZQ+yv/N9nSK/CGpimH/qjTJOQRStf5ppvzT0FzGMX2tqNndJbZD8idLxJFXZekFF16KC/6scsX/lTNL+XFfTqsreVXu7bL/wjNVTPeGkJJE7aWcXP2+3qQTMv+LaO9INAsG3cyp5Co/F06O8XoVtYZXBjH3f3r9Y8Wp89/fqq2OfQSD2/Ujo1t0fNnMA14gpYdtm6+/RcRgNQGIPPGxgAaFjsfC4+63CcHr1nczuKyXiQjmoIH7n/0NCmJv3O+v/Lp30d3n/060TaO5ffQGrrx0O7TYUAC6pdQxfOeuX4/EsKJgMKTW18feF5m1SX4ODnH1SWutwnm5T/k0/l4YXbLUi8QRbdtx74QL9DJRtKP8bDT+yyf//8nAEApoAJH7jMHoQv7XKzIUdH1TDS7Phokc3PP5m68+eUTHU17v50avNmnEHCfybI4FC35LTpSaGqRsgNJJliiV37VIfbUlfDfgIqZmmxEHmnCQTSg2zcf5+9LWPblYTxLShx/2U34N3Rf/3Zvie6j8SS/8X+Yo+dXhKIyg1WX040AAAAA==")
function setIconImage(nodeType, image, size, padRows, padCols) {
const onAdded = nodeType.prototype.onAdded
@@ -77,7 +76,8 @@ function defaultParameterType(widgetType, connectedNode, connectedWidget) {
if (connectedNode.comfyClass === "CLIPTextEncode") {
paramType = "prompt (positive)"
}
if (connectedWidget.options?.round === 1) {
const round = connectedWidget.options?.round
if ((paramType == "number" && round === undefined) || round === 1) {
paramType = "number (integer)"
}
return paramType
@@ -96,7 +96,7 @@ function valueMatchesType(value, type, options) {
function optionalWidgetValue(widgets, index, fallback) {
const result = widgets.length > index ? widgets[index].value : null
return result === null || result === 0 ? fallback : result
return result === null || result === -1e10 || result === 1e10 ? fallback : result
}
function changeWidgets(node, type, connectedNode, connectedWidget) {
@@ -206,8 +206,6 @@ app.registerExtension({
beforeRegisterNodeDef(nodeType /*typeof LGraphNode*/, nodeData /*ComfyObjectInfo*/, app) {
if (nodeData.name === "ETN_KritaCanvas") {
setIconImage(nodeType, canvasIcon, [200, 100], 0, 2)
} else if (nodeData.name === "ETN_KritaOutput") {
setIconImage(nodeType, outputIcon, [200, 100], 1, 0)
} else if (nodeData.name === "ETN_Parameter") {
setupParameterNode(nodeType)
} else if (nodeData.name === "ETN_SendText") {
+241 -140
View File
@@ -1,13 +1,16 @@
import sys
import torch
import numpy as np
from enum import Enum
from pathlib import Path
from typing import Any, NamedTuple
import comfy.samplers
import numpy as np
import server
import torch
from comfy.comfy_types.node_typing import IO
from comfy_api.latest import io
from PIL import Image
import server
import comfy.samplers
from comfy.comfy_types.node_typing import IO
from .nodes import SendImageWebSocket
@@ -72,37 +75,74 @@ class _BasicTypes(str):
BasicTypes = _BasicTypes("BASIC")
class KritaOutput:
class OutputBatchMode(Enum):
default = "default"
images = "images"
animation = "animation"
layers = "layers"
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"),
io.Int.Input("x", "offset x", default=0),
io.Int.Input("y", "offset y", default=0),
io.String.Input("name", default=""),
io.Combo.Input(
"batch_mode", OutputBatchMode, "batch mode", default=OutputBatchMode.default
),
io.Boolean.Input("resize_canvas", "resize canvas", default=False),
],
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( # type: ignore
cls,
images: torch.Tensor,
x: int = 0,
y: int = 0,
name="",
batch_mode: OutputBatchMode | str = OutputBatchMode.default,
resize_canvas=False,
):
batch_mode = batch_mode.value if isinstance(batch_mode, OutputBatchMode) else batch_mode
info = {
"name": name,
"offset_x": x,
"offset_y": y,
"batch_mode": batch_mode,
"resize_canvas": resize_canvas,
}
output = SendImageWebSocket.execute(images, "PNG")
assert isinstance(output.ui, dict)
output.ui["info"] = [info]
return output
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): # type: ignore
mime = {
"text": "text/plain",
"markdown": "text/markdown",
@@ -115,72 +155,109 @@ 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"),
io.Mask.Output(display_name="mask"),
],
)
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, **kwargs):
return io.NodeOutput(_placeholder_image(), 512, 512, 0, torch.ones(1, 512, 512))
class KritaImageLayer:
class SelectionContext(Enum):
automatic = "automatic"
entire_image = "entire image"
mask_bounds = "mask bounds"
_selection_context_help = """
Determines the section (crop bounding box) of the image and mask to transmit:
- automatic: area around the selection determined by Krita settings
- entire image: always use the entire canvas area
- mask bounds: tight bounding box of the current selection
This affects the Selection and Canvas nodes. The offset x/y outputs indicate the top-left corner of the context area relative to the full canvas."""
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",
inputs=[
io.Combo.Input(
"context",
options=SelectionContext,
default=SelectionContext.entire_image,
tooltip=_selection_context_help,
),
io.Int.Input("padding", "padding", default=0, min=0),
],
outputs=[
io.Mask.Output("mask", "mask"),
io.Boolean.Output("active", "active"),
io.Int.Output("x", "offset x"),
io.Int.Output("y", "offset y"),
],
)
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, **kwargs):
return io.NodeOutput(torch.ones(1, 512, 512), False, 0, 0)
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): # type: ignore
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): # type: ignore
return io.NodeOutput(torch.ones(1, 512, 512))
_param_types = [
@@ -193,71 +270,95 @@ _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=-1e10, min=-_fmax, max=_fmax, optional=True),
io.Float.Input("max", default=1e10, 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): # type: ignore
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"
@classmethod
def execute(cls, name: str, sampler_preset: str): # type: ignore
raise NotImplementedError("This workflow must be started from Krita!")
def placeholder(self, name: str, sampler_preset: str):
class KritaStyleAndPrompt(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_KritaStyleAndPrompt",
display_name="Krita Style & Prompt",
category="krita",
inputs=[
io.Combo.Input("sampler_preset", options=["auto", "regular", "live"]),
],
outputs=[
io.Model.Output(display_name="model (with loras)"),
io.Clip.Output(display_name="clip"),
io.Vae.Output(display_name="vae"),
io.String.Output(display_name="positive prompt (evaluated)"),
io.String.Output(display_name="negative prompt (evaluated)"),
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"),
],
)
@classmethod
def execute(cls, name: str, sampler_preset: str): # type: ignore
raise NotImplementedError("This workflow must be started from Krita!")
+242 -119
View File
@@ -1,35 +1,43 @@
from __future__ import annotations
from copy import copy
from typing import NamedTuple
from PIL import Image
import numpy as np
import base64
import time
from copy import copy
from dataclasses import dataclass
from io import BytesIO
from typing import NamedTuple
from uuid import uuid4
import numpy as np
import torch
import torch.nn.functional as F
from io import BytesIO
from server import PromptServer, BinaryEventTypes
from comfy.clip_vision import ClipVisionModel
from comfy.sd import StyleModel
from comfy_api.latest import io
from PIL import Image
from server import BinaryEventTypes, PromptServer
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): # type: ignore
_strip_prefix(image, "data:image/png;base64,")
imgdata = base64.b64decode(image)
img = Image.open(BytesIO(imgdata))
if "A" in img.getbands():
mask = np.array(img.getchannel("A")).astype(np.float32) / 255.0
mask = torch.from_numpy(mask)
mask = torch.from_numpy(mask)[None,]
else:
mask = None
@@ -40,16 +48,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): # type: ignore
_strip_prefix(mask, "data:image/png;base64,")
imgdata = base64.b64decode(mask)
img = Image.open(BytesIO(imgdata))
@@ -60,22 +71,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): # type: ignore
results = []
for tensor in images:
array = 255.0 * tensor.cpu().numpy()
@@ -93,43 +104,155 @@ 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:
timeout = 600 # 10 minutes
max_size = 100 * 1024 * 1024 # 100 MB
@dataclass
class Entry:
data: bytes
content_type: str
timestamp: float
retrieved: int
class OldEntry(NamedTuple):
last_used: float
deleted: float
size: int
retrieved: int
def __init__(self):
self.images: dict[str, ImageCache.Entry] = {}
self.old: dict[str, ImageCache.OldEntry] = {}
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:
if old := self.old.get(key):
now = time.time()
print(
f"[comfyui-tooling-nodes] requested image {key} has been deleted ",
f"(last used {now - old.last_used:.0f}s ago, deleted {now - old.deleted:.0f}s ago, "
f"size {old.size / 1024**2:.1f}MB, retrieved {old.retrieved} times)",
)
return None, None
entry.retrieved += 1
if extend:
entry.timestamp = time.time()
self.prune()
return entry.data, entry.content_type
def prune(self):
total_size = sum(len(entry.data) for entry in self.images.values())
if total_size <= self.max_size:
return
# Remove least recently used entries until under max size
sorted_entries = sorted(self.images.items(), key=lambda item: item[1].timestamp)
now = time.time()
for key, entry in sorted_entries:
age = now - entry.timestamp
if age > self.timeout or (age > 60 and entry.retrieved > 0):
self.old[key] = ImageCache.OldEntry(
entry.timestamp, now, len(entry.data), entry.retrieved
)
del self.images[key]
total_size -= len(entry.data)
if total_size <= self.max_size:
break
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): # type: ignore
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))
def crop(self, image, x, y, width, height):
out = image[:, y : y + height, x : x + width, :]
return (out,)
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]
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): # type: ignore
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 +271,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): # type: ignore
out = to_bchw(image)
if out.shape[1] == 3: # Assuming RGB images
out = torch.cat([out, torch.ones_like(out[:, :1, :, :])], dim=1)
@@ -178,7 +302,7 @@ class ApplyMaskToImage:
# Apply each mask in the batch to its corresponding image's alpha channel
for i in range(out.shape[0]):
alpha = mask[i] if is_mask_batch else mask[0]
out[i, 3, :, :] = alpha
out[i, 3, :, :] *= alpha
return (to_bhwc(out),)
@@ -189,28 +313,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( # type: ignore
cls,
image: torch.Tensor,
weight: float,
range_start: float,
@@ -222,24 +344,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( # type: ignore
cls,
conditioning: list[list],
clip_vision: ClipVisionModel,
style_model: StyleModel,
+27 -28
View File
@@ -1,5 +1,4 @@
from __future__ import annotations
from weakref import ref as WeakRef
from pathlib import Path
from tqdm import tqdm
import torch
@@ -8,6 +7,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
@@ -38,6 +38,10 @@ class CLIPSafetyChecker(PreTrainedModel):
self.concept_embeds_weights = nn.Parameter(torch.ones(17), requires_grad=False)
self.special_care_embeds_weights = nn.Parameter(torch.ones(3), requires_grad=False)
# Model requires post_init after transformers v4.57.3
if hasattr(self, "post_init"):
self.post_init()
def forward(self, clip_input, images: Tensor, sensitivity: float):
with torch.no_grad():
image_batch = self.vision_model(clip_input)[1]
@@ -76,7 +80,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 +95,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 +120,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))
+2 -2
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.4.0"
license = { file = "LICENSE" }
[project.urls]
@@ -13,7 +13,7 @@ line-length = 100
preview = true
[tool.ruff.lint]
ignore = ["E741"]
ignore = ["E741", "BLE001"]
[tool.black]
line-length = 100
+299 -74
View File
@@ -1,13 +1,32 @@
# Adapted from https://github.com/pamparamm/ComfyUI-ppm
# Adapted from https://github.com/laksjdjf/cgem156-ComfyUI/blob/main/scripts/attention_couple/node.py
# by @laksjdjf
from __future__ import annotations
from typing import NamedTuple
from functools import partial
from typing import Any, NamedTuple
import torch
import torch.nn.functional as F
import math
from torch import Tensor, Size
import comfy.model_management
import comfy.patcher_extension
from comfy.model_patcher import ModelPatcher
from comfy.model_base import Anima, CosmosPredict2
from comfy.ldm.cosmos.predict2 import Attention as CosmosAttention
from comfy.sampler_helpers import convert_cond
from comfy.samplers import process_conds
from comfy_api.latest import io
COND = 0
UNCOND = 1
ANIMA_COUPLE_WRAPPER_KEY = "etn_attention_mask_anima"
ANIMA_COUPLE_PATCH_KEY = "etn_attention_mask_patch"
CONDS_COUPLE_KEY = "etn_couple_conds"
COND_UNCOND_COUPLE_KEY = "etn_couple_cond_or_uncond"
COUPLE_ACTIVE_KEY = "etn_couple_active"
NUM_TOKENS_COUPLE_KEY = "etn_couple_num_tokens"
def downsample_mask(mask: Tensor, batch: int, target_size: int, original_shape: Size) -> Tensor:
@@ -33,6 +52,12 @@ def downsample_mask(mask: Tensor, batch: int, target_size: int, original_shape:
return result
def reshape_mask(mask: Tensor, size: tuple[int, int], batch: int, target_size: int) -> Tensor:
result = F.interpolate(mask, size=size, mode="nearest")
result = result.view(mask.shape[0], target_size, 1)
return result.repeat_interleave(batch, dim=0)
def lcm(a: int, b: int):
return a * b // math.gcd(a, b)
@@ -65,104 +90,115 @@ 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.region_conds = [r.conditioning for r in region_list]
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())
if _is_anima_couple_model(model):
return patch.apply_anima(model)
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 +209,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 +219,218 @@ 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
def apply_anima(self, model: ModelPatcher):
new_model = model.clone()
_patch_cosmos_attention(new_model)
device = comfy.model_management.get_torch_device()
conds_converted = [convert_cond(cond)[0] for cond in self.region_conds]
new_model.add_wrapper_with_key(
comfy.patcher_extension.WrappersMP.SAMPLER_SAMPLE,
ANIMA_COUPLE_WRAPPER_KEY,
_anima_couple_sample_wrapper(conds_converted, device),
)
new_model.add_wrapper_with_key(
comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL,
ANIMA_COUPLE_WRAPPER_KEY,
_anima_couple_diffusion_wrapper(self),
)
new_model.set_attachments("etn_attention_mask", self)
return new_model
def _is_anima_couple_model(model: ModelPatcher) -> bool:
model_type = type(model.model)
return issubclass(model_type, (Anima, CosmosPredict2))
def _anima_couple_sample_wrapper(conds_converted: list, device):
def sample_wrapper(executor, *args, **kwargs):
if len(conds_converted) > 0:
guider = args[0]
extra_options: dict[str, Any] = args[2]
seed: int = extra_options["seed"]
noise: Tensor = args[4]
latent_image: Tensor = args[5]
denoise_mask: Tensor | None = args[6]
conds_processed = process_conds(
guider.inner_model,
noise,
{"positive": conds_converted},
device,
latent_image,
denoise_mask,
seed,
latent_shapes=[latent_image.shape],
)["positive"]
conds_couple = [cond["model_conds"]["c_crossattn"].cond for cond in conds_processed]
model_options: dict[str, Any] = extra_options["model_options"]
transformer_options: dict[str, Any] = model_options.get("transformer_options", {}).copy()
transformer_options[CONDS_COUPLE_KEY] = conds_couple
transformer_options[NUM_TOKENS_COUPLE_KEY] = [cond.shape[1] for cond in conds_couple]
model_options["transformer_options"] = transformer_options
return executor(*args, **kwargs)
return sample_wrapper
def _anima_couple_diffusion_wrapper(patch: AttentionMaskPatch):
def diffusion_wrapper(executor, *args, **kwargs):
anima_model = executor.class_obj
x: Tensor = args[0]
transformer_options: dict[str, Any] = kwargs.get("transformer_options", {}).copy()
patch_spatial = getattr(anima_model, "patch_spatial", 1)
activations_shape = list(x.shape)
activations_shape[-2] = activations_shape[-2] // patch_spatial
activations_shape[-1] = activations_shape[-1] // patch_spatial
transformer_options["activations_shape"] = activations_shape
transformer_options[ANIMA_COUPLE_PATCH_KEY] = patch
kwargs["transformer_options"] = transformer_options
return executor(*args, **kwargs)
return diffusion_wrapper
def pre_cross_attention(
patch: AttentionMaskPatch,
transformer_options: dict,
x: Tensor,
context: Tensor,
rope_emb: Tensor | None,
) -> tuple[Tensor, Tensor, Tensor | None, dict]:
transformer_options = transformer_options.copy()
if CONDS_COUPLE_KEY not in transformer_options:
transformer_options[COND_UNCOND_COUPLE_KEY] = list(transformer_options["cond_or_uncond"])
transformer_options[COUPLE_ACTIVE_KEY] = False
return x, context, rope_emb, transformer_options
conds: list[Tensor] = transformer_options[CONDS_COUPLE_KEY]
num_tokens_c: list[int] = transformer_options[NUM_TOKENS_COUPLE_KEY]
cond_or_uncond = transformer_options["cond_or_uncond"]
num_chunks = len(cond_or_uncond)
batch = x.shape[0] // num_chunks
x_chunks = x.chunk(num_chunks, dim=0)
c_chunks = context.chunk(num_chunks, dim=0)
lcm_tokens_c = lcm_for_list(num_tokens_c + [context.shape[1]])
conds_c_tensor = torch.cat(
[cond.repeat(batch, lcm_tokens_c // num_tokens_c[i], 1) for i, cond in enumerate(conds)],
dim=0,
)
xs, cs = [], []
cond_or_uncond_couple = []
for i, cond_type in enumerate(cond_or_uncond):
x_target = x_chunks[i]
c_target = c_chunks[i].repeat(1, lcm_tokens_c // context.shape[1], 1)
if cond_type == UNCOND:
xs.append(x_target)
cs.append(c_target)
cond_or_uncond_couple.append(UNCOND)
else:
xs.append(x_target.repeat(patch.num_conds, 1, 1))
cs.append(conds_c_tensor)
cond_or_uncond_couple.extend([COND] * patch.num_conds)
transformer_options[COND_UNCOND_COUPLE_KEY] = cond_or_uncond_couple
transformer_options[COUPLE_ACTIVE_KEY] = True
return torch.cat(xs, dim=0), torch.cat(cs, dim=0), rope_emb, transformer_options
def cross_attention_output(patch: AttentionMaskPatch, transformer_options: dict, out: Tensor):
cond_or_uncond = transformer_options[COND_UNCOND_COUPLE_KEY]
size = tuple(transformer_options["activations_shape"][-2:])
batch = out.shape[0] // len(cond_or_uncond)
mask = patch.mask.to(out.device, dtype=out.dtype)
mask_downsample = reshape_mask(mask, size, batch, out.shape[1])
outputs = []
cond_outputs = []
i_cond = 0
for i, cond_type in enumerate(cond_or_uncond):
pos, next_pos = i * batch, (i + 1) * batch
if cond_type == UNCOND:
outputs.append(out[pos:next_pos])
else:
pos_cond, next_pos_cond = i_cond * batch, (i_cond + 1) * batch
cond_outputs.append(out[pos:next_pos] * mask_downsample[pos_cond:next_pos_cond])
i_cond += 1
if len(cond_outputs) > 0:
outputs.append(torch.stack(cond_outputs).sum(0))
return torch.cat(outputs, dim=0)
def _patch_cosmos_attention(model_patcher: ModelPatcher):
cosmos_model = model_patcher.get_model_object("diffusion_model")
for block_name, block in (
(n, b)
for n, b in cosmos_model.named_modules()
if ("cross_attn" in n or "self_attn" in n) and isinstance(b, CosmosAttention)
):
patch_name = f"diffusion_model.{block_name}.forward"
if patch_name not in model_patcher.object_patches:
model_patcher.add_object_patch(patch_name, partial(_cosmos_attention_forward_patched, block))
def _cosmos_attention_forward_patched(
self,
x: Tensor,
context: Tensor | None = None,
rope_emb: Tensor | None = None,
transformer_options: dict | None = None,
) -> Tensor:
transformer_options = transformer_options if transformer_options is not None else {}
patch: AttentionMaskPatch | None = transformer_options.get(ANIMA_COUPLE_PATCH_KEY)
if context is not None and patch is not None:
x, context, rope_emb, transformer_options = pre_cross_attention(
patch, transformer_options, x, context, rope_emb
)
q, k, v = self.compute_qkv(x, context, rope_emb=rope_emb)
output = self.compute_attention(q, k, v, transformer_options=transformer_options)
if context is not None and patch is not None and transformer_options.get(COUPLE_ACTIVE_KEY, False):
output = cross_attention_output(patch, transformer_options, output)
return output
+102 -93
View File
@@ -3,49 +3,29 @@ 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):
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"
def __init__(
self, image: Tensor, min_tile_size: int, padding: int, blending: int, multiple: int
):
assert all([x % multiple == 0 for x in image.shape[-3:-1]]), (
"Image size must be divisible by multiple"
)
assert min_tile_size % multiple == 0, "Tile size must be divisible by multiple"
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 / multiple) * multiple).astype(int)
def size(self, coord: IntArray):
return self.end(coord) - self.start(coord)
@@ -96,80 +76,109 @@ 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),
io.Int.Input("multiple", default=8, min=1, max=1024, step=1),
],
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, multiple: int):
return io.NodeOutput(TileLayout(image, min_tile_size, padding, blending, multiple))
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)")