Author SHA1 Message Date
Acly 32b7e1e301 cont 2026-05-10 16:00:42 +02:00
Acly 552c89bec9 Support already cached image PUT requests without aborting the connection
* conditionally send 100 Continue instead
2026-05-10 13:17:11 +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
Acly 20f8d8ecc9 Version 2.0.6 2025-10-11 19:17:02 +02:00
Acly f555efb71b API: fix qwen svdq models not having the quant field set 2025-10-05 12:54:04 +02:00
Acly db4f296533 API: map GGUF "qwen_image" arch to "qwen-image" to be consistent with safetensors models 2025-10-05 12:41:06 +02:00
Acly 17c36ebc70 Parameter node: avoid unhandled exception when input is not a widget 2025-10-05 11:15:12 +02:00
Acly 0697a3ac1f Add active output to KritaSelection node (true if there is a selection, false otherwise) 2025-09-06 23:33:07 +02:00
Acly fa46b93329 Model inspection: support Qwen and Nunchaku quants 2025-08-20 21:30:16 +02:00
Acly fa84eec8fc Model inspection: support more base models 2025-08-09 10:40:43 +02:00
Acly 5ef2fddc1b Version 2.0.3 2025-06-15 12:32:46 +02:00
Acly bff22b8351 Model inspection: don't filter keys for diffusion models (strips things like vpred and zsnr keys)
- this fixes sdxl-vpred diffusion models being detected as eps
2025-05-28 20:45:58 +02:00
Acly ca2b59248e Fix error when importing workflows that contain Parameter nodes connected to nodes that aren't installed 2025-05-22 15:46:42 +02:00
Acly 696899a5fc Allow to run workflows with parameter nodes
- fix validation error due to missing min/max
- fix errors due to numbers being passed as string
2025-05-12 15:51:42 +02:00
Acly 5f4373d71a Version 2.0.2 2025-04-28 10:11:42 +02:00
Acly 61d2a19120 Remove data:image/png;base64, prefix in base64 strings if present #39 2025-04-27 19:44:17 +02:00
Acly a6af76ac39 Fix default values of Parameter node not being editable #38 2025-04-27 17:21:13 +02:00
Acly 6a5c8e02e5 Fix error response when a model folder cannot be found 2025-04-27 09:51:59 +02:00
Acly c2308a0762 Don't run publish action on forks, close #114 2025-03-31 11:49:17 +02:00
Acly b8e4659a10 Add optional mask output for Krita Image Layer node 2025-03-01 20:20:52 +01:00
Acly 93e1932456 Don't return alpha channel from LoadImageBase64 inverted
... why was it ever inverted?
2025-03-01 20:20:24 +01:00
Acly ea755151fe Add Lumina 2 to known base models 2025-02-13 16:48:23 +01:00
Acly facd65995a Version 2.0.1, remove some unused code and imports 2025-02-02 20:17:51 +01:00
Acly d7b18203a4 Use comfy built-in any type (*) matching 2025-02-01 00:03:33 +01:00
Acly 1839c099ad Fix type matching BOOL -> BOOLEAN 2025-01-31 23:44:51 +01:00
Acly bed8b36705 Fix assertion when using tiling with padding=blending=0 2025-01-26 11:57:11 +01:00
Acly 8ed5591574 API breaking: removed is_refiner attribute from model inspection
- sdxl refiner is reported with base model "sdxl-refiner"
- added type attribute for sdxl model, allows to detect eps/v-prediction
2025-01-12 17:53:47 +01:00
Acly fe39d22eb9 Return a more specific error when inspect model folder doesn't exist 2024-12-07 19:49:00 +01:00
Acly 50d3479fba Code compatibility (match was no longer useful anyway) #31 2024-11-30 00:29:00 +01:00
Acly d7d421baaa Model detection: add support for (some) GGUF and Flux Inpaint models
- GGUF detection only works for converted models
2024-11-29 09:46:45 +01:00
Acly e10daee9ed Nodes for stacking and weighting reference images with flux redux model 2024-11-24 22:51:42 +01:00
Acly 50c3ffdf64 Parameter node: Fix type reset to default for connected widget on reload #29 2024-11-15 15:52:50 +01:00
Acly 517790d1d6 Parameter node: fix not being able to enter negative numbers for min/max 2024-11-11 11:47:53 +01:00
Acly e2bd09d7e9 Parameter node: Restrict initial type choice to avoid mismatch between type and default value before connecting the output 2024-10-29 13:03:53 +01:00
Acly 035c68c629 Parameter node: keep configured default values when reloading #25
- make sure default is changed if the node is reconnected to a non-matching type
2024-10-29 12:13:42 +01:00
Acly e86973fedf Remove image format parameter from Krita Output node 2024-10-28 14:56:01 +01:00
Acly 19337dcc0e Fix parameter node not being connectable #23 2024-10-28 13:11:33 +01:00
Acly 1d4ffe14bb Add a Send Text node
- converts any input to string and sends it as output (websocket message)
2024-10-27 10:43:43 +01:00
Acly 20c8039a98 Detect some diffusion models which have prefix like checkpoints 2024-10-25 13:08:41 +02:00
Acly fcf678735c New package version 2024-10-23 13:05:24 +02:00
Acly 63ab33800e Fix Parameter node type comparison for workflow validation 2024-10-21 15:11:01 +02:00
Acly ef5ccfa98f Fix Parameter node widget values being reset when switching or reloading workflows 2024-10-18 16:38:01 +02:00
Acly 0b01696f5b Fix workflow/unsubscribe endpoint 2024-10-14 13:20:41 +02:00
Acly 2fce4c56d5 Fix "Object of type _BasicTypes is not JSON serializable" 2024-10-12 15:37:38 +02:00
Acly a7f77032ec Make parameter nodes validate when executed 2024-10-11 21:46:49 +02:00
Acly 72335898cb Detect better parameter type defaults 2024-10-11 21:20:31 +02:00
12 changed files with 1293 additions and 523 deletions
+2 -1
View File
@@ -11,11 +11,12 @@ jobs:
publish-node: publish-node:
name: Publish Custom Node to registry name: Publish Custom Node to registry
runs-on: ubuntu-latest runs-on: ubuntu-latest
if: ${{ github.repository_owner == 'Acly' }}
steps: steps:
- name: Check out code - name: Check out code
uses: actions/checkout@v4 uses: actions/checkout@v4
- name: Publish Custom Node - name: Publish Custom Node
uses: Comfy-Org/publish-node-action@main uses: Comfy-Org/publish-node-action@v1
with: with:
## Add your own personal access token to your Github Repository secrets and reference it here. ## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }} personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+75 -8
View File
@@ -13,9 +13,9 @@ Provides nodes and API geared towards using ComfyUI as a backend for external to
## <a id="images" href="#toc">Sending and receiving images</a> ## <a id="images" href="#toc">Sending and receiving images</a>
ComfyUI exchanges images via the filesystem. This requires a ComfyUI exchanges images via the filesystem. This requires a
multi-step process (upload images, prompt, download images), is rather multi-step process (upload images, prompt, download images), which
inefficient, and invites a whole class of potential issues. It's also unclear invites a whole class of potential issues you might not want to deal with.
at which point those images will get cleaned up if ComfyUI is used It's also unclear at which point those images will get cleaned up if ComfyUI is used
via external tools. via external tools.
### Load Image (Base64) ### Load Image (Base64)
@@ -33,7 +33,7 @@ Loads a mask (single channel) from a PNG embedded into the prompt as base64 stri
### Send Image (WebSocket) ### Send Image (WebSocket)
Sends an output image over the client WebSocket connection as PNG binary data. 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: 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>}} {'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> ## <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. These nodes implement attention masking for arbitrary number of image regions. Text prompts only apply to the masked area.
@@ -163,7 +204,9 @@ There are various types of models that can be loaded as checkpoint, LoRA, Contro
#### Paramters #### Paramters
* `folder_name`: sub-directory in ComfyUI's models folder. * `folder_name`: sub-directory in ComfyUI's models folder.
Supported model types: `checkpoints`, `diffusion_models` 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 #### Output
Lists available models with additional classification info: Lists available models with additional classification info:
@@ -172,14 +215,38 @@ Lists available models with additional classification info:
"checkpoint_file.safetensors": { "checkpoint_file.safetensors": {
"base_model": "sd15", "base_model": "sd15",
"is_inpaint": false, "is_inpaint": false,
"is_refiner": false "type": "eps"
}, },
... ...
} }
``` ```
Possible values for base model: `sd15, sd20, sd21, sd3, sdxl, ssd1b, svd, cascade-b, cascade-c, aura-flow, hunyuan-dit, flux, flux-schnell` 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`
Detection supports quantized models:
* GGUF: if the `gguf` module is installed, .gguf files are detected and will set the `quant` field
* Nunchaku: SVDQuant models are detected and will set the `quant` field to `svdq`
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 }
}
```
The entry is `{"base_model": "unknown"}` for models which are not in safetensors format or do not match any of the known base models.
### GET /api/etn/languages ### GET /api/etn/languages
+41 -51
View File
@@ -1,53 +1,43 @@
from . import api, nodes, tile, region, nsfw, translation, krita from comfy_api.latest import ComfyExtension, io
from . import api as api, nodes, tile, region, nsfw, translation, krita
class ExternalToolingNodes(ComfyExtension):
async def get_node_list(self) -> list[type[io.ComfyNode]]:
return [
nodes.LoadImageCache,
nodes.SaveImageCache,
nodes.LoadImageBase64,
nodes.LoadMaskBase64,
nodes.SendImageWebSocket,
nodes.ApplyMaskToImage,
nodes.ReferenceImage,
nodes.ApplyReferenceImages,
tile.CreateTileLayout,
tile.ExtractImageTile,
tile.ExtractMaskTile,
tile.GenerateTileMask,
tile.MergeImageTile,
region.BackgroundRegion,
region.DefineRegion,
region.ListRegionMasks,
region.AttentionMask,
nsfw.NSFWFilter,
translation.Translate,
krita.KritaOutput,
krita.KritaSendText,
krita.KritaCanvas,
krita.KritaSelection,
krita.KritaImageLayer,
krita.KritaMaskLayer,
krita.Parameter,
krita.KritaStyle,
krita.KritaStyleAndPrompt,
]
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_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_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_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_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" WEB_DIRECTORY = "./js"
+244 -30
View File
@@ -1,19 +1,21 @@
from __future__ import annotations from __future__ import annotations
from aiohttp import web from aiohttp import web
from typing import NamedTuple from typing import Any, NamedTuple
from pathlib import Path from pathlib import Path
import json import json
import traceback import traceback
import re import re
import logging import logging
import itertools
from comfy import model_detection, supported_models from comfy import model_detection
import comfy.utils import comfy.utils
import folder_paths import folder_paths
import server import server
from .translation import available_languages, translate from .translation import available_languages, translate
from .krita import WorkflowExchange from .krita import WorkflowExchange
from .nodes import image_cache
input_block_name = "model.diffusion_model.input_blocks.0.0.weight" input_block_name = "model.diffusion_model.input_blocks.0.0.weight"
@@ -22,7 +24,7 @@ model_names = {
"SD20": "sd20", "SD20": "sd20",
"SD21UnclipL": "sd21", "SD21UnclipL": "sd21",
"SD21UnclipH": "sd21", "SD21UnclipH": "sd21",
"SDXLRefiner": "sdxl", "SDXLRefiner": "sdxl-refiner",
"SDXL": "sdxl", "SDXL": "sdxl",
"SSD1B": "ssd1b", "SSD1B": "ssd1b",
"SVD_img2vid": "svd", "SVD_img2vid": "svd",
@@ -33,7 +35,35 @@ model_names = {
"HunyuanDiT": "hunyuan-dit", "HunyuanDiT": "hunyuan-dit",
"HunyuanDiT1": "hunyuan-dit", "HunyuanDiT1": "hunyuan-dit",
"Flux": "flux", "Flux": "flux",
"FluxInpaint": "flux",
"FluxSchnell": "flux-schnell", "FluxSchnell": "flux-schnell",
"GenmoMochi": "mochi",
"LTXV": "ltxv",
"HunyuanVideo": "hunyuan-video",
"CosmosT2V": "cosmos",
"CosmosI2V": "cosmos",
"CosmosT2IPredict2": "cosmos-predict2",
"CosmosI2VPredict2": "cosmos-predict2",
"ZImage": "z-image",
"Lumina2": "lumina2",
"WAN21_T2V": "wan21",
"WAN21_I2V": "wan21",
"WAN21_FunControl2V": "wan21-fun",
"WAN21_Vace": "wan21-vace",
"WAN21_Camera": "wan21-camera",
"HiDream": "hi-dream",
"Chroma": "chroma",
"ACEStep": "ace-step",
"Omnigen2": "omnigen2",
"QwenImage": "qwen-image",
"ErnieImage": "ernie-image",
"Flux2": "flux2",
"Anima": "anima",
}
gguf_architectures = {
"sd1": "sd15",
"qwen_image": "qwen-image",
} }
@@ -48,7 +78,7 @@ class FakeTensor(NamedTuple):
return d return d
def inspect_diffusion_model(filename: str, prefix: str | None, model_type: str): def inspect_safetensors(filename: str, model_type: str, is_checkpoint: bool):
try: try:
# Read header of safetensors file # Read header of safetensors file
path = folder_paths.get_full_path(model_type, filename) path = folder_paths.get_full_path(model_type, filename)
@@ -62,8 +92,10 @@ def inspect_diffusion_model(filename: str, prefix: str | None, model_type: str):
cfg[key] = FakeTensor.from_dict(cfg[key]) cfg[key] = FakeTensor.from_dict(cfg[key])
# Reuse Comfy's model detection # Reuse Comfy's model detection
if prefix is None: prefix = model_detection.unet_prefix_from_state_dict(cfg)
prefix = model_detection.unet_prefix_from_state_dict(cfg) if not is_checkpoint:
cfg = comfy.utils.state_dict_prefix_replace(cfg, {prefix: ""}, filter_keys=False)
prefix = ""
try: # latest ComfyUI takes 2 args try: # latest ComfyUI takes 2 args
unet_config = model_detection.detect_unet_config(cfg, prefix) unet_config = model_detection.detect_unet_config(cfg, prefix)
except TypeError as e: # older ComfyUI versions take 3 args except TypeError as e: # older ComfyUI versions take 3 args
@@ -76,30 +108,161 @@ def inspect_diffusion_model(filename: str, prefix: str | None, model_type: str):
input_count = 4 input_count = 4
# Find a matching base model depending on unet config # Find a matching base model depending on unet config
base_model = model_detection.model_config_from_unet_config(unet_config) base_model = None
if base_model is None: model_type = None
model_quant = None
# Check if it's a Nunchaku SVDQ model by inspecting metadata
raw_name = detect_svdq(cfg)
if raw_name:
model_quant = "svdq"
# Otherwise try ComfyUI's model detection
elif unet_config is not None:
base_model = model_detection.model_config_from_unet_config(unet_config)
if base_model:
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"} return {"base_model": "unknown"}
base_model_class = base_model.__class__ base_model_name = model_names.get(raw_name, "unknown")
base_model_name = model_names.get(base_model_class.__name__, "unknown") result: dict[str, Any] = {"base_model": base_model_name}
return { result["is_inpaint"] = (
"base_model": base_model_name, base_model_name in ["sd15", "sdxl"] and input_count > 4
"is_inpaint": base_model_name in ["sd15", "sdxl"] and input_count > 4, ) or raw_name == "FluxInpaint"
"is_refiner": base_model_class is supported_models.SDXLRefiner, if model_quant:
} result["quant"] = model_quant
if model_type:
result["type"] = model_type
elif "T2I" in raw_name:
result["type"] = "t2i"
elif "I2V" in raw_name:
result["type"] = "i2v"
elif "T2V" in raw_name:
result["type"] = "t2v"
elif "Control2V" in raw_name:
result["type"] = "control2v"
return result
return {"base_model": "unknown"} 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}"}
def detect_svdq(cfg: dict) -> str | None:
if md := cfg.get("__metadata__"):
if comfy_config := md.get("comfy_config"):
if isinstance(comfy_config, str):
comfy_config = json.loads(comfy_config)
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
def inspect_gguf(filename: str, model_type: str):
try:
import gguf
except ImportError:
return {"base_model": "unknown", "error": "GGUF module not found"}
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:
if len(arch_field.types) != 1 or arch_field.types[0] != gguf.GGUFValueType.STRING:
raise TypeError(
f"Bad type for GGUF general.architecture key: expected string, got {arch_field.types!r}"
)
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:
if file_type := reader.get_field("general.file_type"):
result["quant"] = file_type.contents().lower()
except Exception:
result["quant"] = "gguf"
return result
except Exception as e: except Exception as e:
# traceback.print_exc() # traceback.print_exc()
return {"base_model": "unknown", "error": f"Failed to detect base model: {e}"} return {"base_model": "unknown", "error": f"Failed to detect base model: {e}"}
def inspect_models(model_type: str): def inspect_diffusion_model(filename: str, model_type: str, is_checkpoint: bool):
if filename.endswith(".gguf"):
return inspect_gguf(filename, model_type)
return inspect_safetensors(filename, model_type, is_checkpoint)
def inspect_models(model_type: str, params: dict[str, str]):
try: try:
prefix = "" if model_type in ("unet", "diffusion_models") else None 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 = { info = {
filename: inspect_diffusion_model(filename, prefix, model_type) filename: inspect_diffusion_model(filename, model_type, is_checkpoint)
for filename in folder_paths.get_filename_list(model_type) 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) return web.json_response(info)
except Exception as e: except Exception as e:
traceback.print_exc() traceback.print_exc()
@@ -128,25 +291,28 @@ def has_invalid_filename(filename: str):
return None 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) _server: server.PromptServer | None = getattr(server.PromptServer, "instance", None)
if _server is not None: if _server is not None:
_workflow_exchange = WorkflowExchange(_server) _workflow_exchange = WorkflowExchange(_server)
@_server.routes.get("/api/etn/model_info/{folder_name}") @_server.routes.get("/api/etn/model_info/{folder_name}")
async def model_info(request: web.Request): async def model_info(request: web.Request):
folder_name = request.match_info.get("folder_name", "checkpoints") folder_name = request.match_info.get("folder_name", "checkpoints")
if error := has_invalid_folder_name(folder_name): error = has_invalid_folder_name(folder_name)
if error is not None:
return error return error
return inspect_models(folder_name) return inspect_models(folder_name, request.rel_url.query)
@_server.routes.get("/api/etn/model_info") @_server.routes.get("/api/etn/model_info")
async def api_model_info(request): async def api_model_info(request):
return inspect_models("checkpoints") return inspect_models("checkpoints", request.rel_url.query)
@_server.routes.get("/etn/model_info")
async def api_model_info(request):
return inspect_models("checkpoints")
@_server.routes.get("/api/etn/languages") @_server.routes.get("/api/etn/languages")
async def languages(request): async def languages(request):
@@ -166,14 +332,60 @@ if _server is not None:
except Exception as e: except Exception as e:
return web.json_response(dict(error=str(e)), status=500) 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}") @_server.routes.put("/api/etn/upload/{folder_name}/{filename}")
async def upload(request: web.Request): async def upload(request: web.Request):
folder_name = request.match_info.get("folder_name", "") folder_name = request.match_info.get("folder_name", "")
if error := has_invalid_folder_name(folder_name): error = has_invalid_folder_name(folder_name)
if error is not None:
return error return error
filename = request.match_info.get("filename", "") filename = request.match_info.get("filename", "")
if error := has_invalid_filename(filename): error = has_invalid_filename(filename)
if error is not None:
return error return error
try: try:
@@ -182,7 +394,9 @@ if _server is not None:
folder = Path(folder_paths.folder_names_and_paths[folder_name][0][0]) folder = Path(folder_paths.folder_names_and_paths[folder_name][0][0])
total_size = int(request.headers.get("Content-Length", "0")) total_size = int(request.headers.get("Content-Length", "0"))
logging.info(f"Uploading {filename} ({total_size/(1024**2):.1f} MB) to {folder} folder") logging.info(
f"Uploading {filename} ({total_size / (1024**2):.1f} MB) to {folder} folder"
)
with open(folder / filename, "wb") as f: with open(folder / filename, "wb") as f:
async for chunk, _ in request.content.iter_chunks(): async for chunk, _ in request.content.iter_chunks():
+65 -24
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 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) { function setIconImage(nodeType, image, size, padRows, padCols) {
const onAdded = nodeType.prototype.onAdded const onAdded = nodeType.prototype.onAdded
@@ -72,36 +71,70 @@ const parameterTypes = {
"text": ["text", "prompt (positive)", "prompt (negative)"], "text": ["text", "prompt (positive)", "prompt (negative)"],
} }
function changeWidget(widget, type, value, options) { function defaultParameterType(widgetType, connectedNode, connectedWidget) {
widget.type = type let paramType = parameterTypes[widgetType][0]
widget.value = value if (connectedNode.comfyClass === "CLIPTextEncode") {
widget.options = options paramType = "prompt (positive)"
}
const round = connectedWidget.options?.round
if ((paramType == "number" && round === undefined) || round === 1) {
paramType = "number (integer)"
}
return paramType
} }
function changeWidgets(node, type, value, options) { function valueMatchesType(value, type, options) {
if (type === "number") {
return typeof value === "number"
} else if (type === "combo") {
return options?.values?.includes(value)
} else if (type === "toggle") {
return typeof value === "boolean"
}
return typeof value === "string"
}
function optionalWidgetValue(widgets, index, fallback) {
const result = widgets.length > index ? widgets[index].value : null
return result === null || result === -1e10 || result === 1e10 ? fallback : result
}
function changeWidgets(node, type, connectedNode, connectedWidget) {
if (type === "customtext") { if (type === "customtext") {
type = "text" type = "text"
} }
const options = connectedWidget.options
node.widgets[1].value = parameterTypes[type][0] const parameterTypeHint = node.widgets[1].value
node.widgets[1].options = {values: parameterTypes[type]} const notSpecialized = node.widgets[1].options.values.includes("auto")
changeWidget(node.widgets[2], type, value, options) const parameterTypeMismatch = !parameterTypes[type].includes(parameterTypeHint)
if (notSpecialized || parameterTypeMismatch) {
node.widgets[1].options = {values: parameterTypes[type]}
}
if (parameterTypeMismatch) {
node.widgets[1].value = defaultParameterType(type, connectedNode, connectedWidget)
}
const oldDefault = node.widgets.length > 2 ? node.widgets[2].value : connectedWidget.value
const oldMin = optionalWidgetValue(node.widgets, 3, options?.min ?? 0)
const oldMax = optionalWidgetValue(node.widgets, 4, options?.max ?? 100)
const isDefaultValid = valueMatchesType(oldDefault, type, connectedWidget.options)
while (node.widgets.length > 2) {
node.widgets.pop()
}
const value = isDefaultValid && oldDefault !== "" ? oldDefault : connectedWidget.value
node.addWidget(type, "default", value, null, options)
if (type === "number") { if (type === "number") {
changeWidget(node.widgets[3], "number", options?.min ?? 0, options) node.addWidget("number", "min", oldMin, null, options)
changeWidget(node.widgets[4], "number", options?.max ?? 100, options) node.addWidget("number", "max", oldMax, null, options)
} else {
changeWidget(node.widgets[3], "number", 0, {min: 0, max: 0})
changeWidget(node.widgets[4], "number", 0, {min: 0, max: 0})
} }
} }
function adaptWidgetsToConnection(node) { function adaptWidgetsToConnection(node) {
if (!node.outputs || node.outputs.length === 0 || !node.outputs[0].links) { if (!node.outputs || node.outputs.length === 0) {
return return
} }
const links = node.outputs[0].links const links = node.outputs[0].links
if (links.length === 1) { if (links && links.length === 1) {
const link = node.graph.links[links[0]] const link = node.graph.links[links[0]]
if (!link) return if (!link) return
@@ -109,7 +142,7 @@ function adaptWidgetsToConnection(node) {
if (!theirNode || !theirNode.inputs) return if (!theirNode || !theirNode.inputs) return
const input = theirNode.inputs[link.target_slot] const input = theirNode.inputs[link.target_slot]
if (!input) return if (!input || !input.widget || theirNode.widgets === undefined) return
node.outputs[0].type = input.type node.outputs[0].type = input.type
@@ -119,11 +152,15 @@ function adaptWidgetsToConnection(node) {
const widgetName = input.widget.name const widgetName = input.widget.name
const theirWidget = theirNode.widgets.find((w) => w.name === widgetName) const theirWidget = theirNode.widgets.find((w) => w.name === widgetName)
const widgetType = theirWidget.origType ?? theirWidget.type if (!theirWidget) return // connected to a custom node that isn't installed
changeWidgets(node, widgetType, theirWidget.value, theirWidget.options)
} else if (links.length === 0) { const widgetType = theirWidget.origType ?? theirWidget.type
changeWidgets(node, widgetType, theirNode, theirWidget)
} else if (!links || links.length === 0) {
node.outputs[0].type = "*" node.outputs[0].type = "*"
node.widgets[1].value = "auto"
node.widgets[1].options = {values: ["auto"]}
} }
} }
@@ -169,10 +206,14 @@ app.registerExtension({
beforeRegisterNodeDef(nodeType /*typeof LGraphNode*/, nodeData /*ComfyObjectInfo*/, app) { beforeRegisterNodeDef(nodeType /*typeof LGraphNode*/, nodeData /*ComfyObjectInfo*/, app) {
if (nodeData.name === "ETN_KritaCanvas") { if (nodeData.name === "ETN_KritaCanvas") {
setIconImage(nodeType, canvasIcon, [200, 100], 0, 2) setIconImage(nodeType, canvasIcon, [200, 100], 0, 2)
} else if (nodeData.name === "ETN_KritaOutput") { } else if (nodeData.name === "ETN_Parameter") {
setIconImage(nodeType, outputIcon, [200, 120], 2, 0)
} else if (nodeData.name == "ETN_Parameter") {
setupParameterNode(nodeType) setupParameterNode(nodeType)
} else if (nodeData.name === "ETN_SendText") {
const onAdded = nodeType.prototype.onAdded
nodeType.prototype.onAdded = function() {
onAdded?.apply(this, arguments)
this.inputs[0].type = "*"
}
} }
}, },
+286 -118
View File
@@ -1,11 +1,16 @@
import torch import sys
import numpy as np from enum import Enum
from pathlib import Path from pathlib import Path
from typing import NamedTuple 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 from PIL import Image
import server
import comfy.samplers
from .nodes import SendImageWebSocket from .nodes import SendImageWebSocket
@@ -34,8 +39,11 @@ class WorkflowExchange:
for publisher in self._publishers.values(): for publisher in self._publishers.values():
await self._notify(client_id, publisher) await self._notify(client_id, publisher)
def unsubscribe(self, client_id: str): async def unsubscribe(self, client_id: str):
self._subscribers.remove(client_id) if client_id in self._subscribers:
self._subscribers.remove(client_id)
else:
raise KeyError("No subscriber found with id " + client_id)
async def _notify(self, client_id: str, publisher: Publisher): async def _notify(self, client_id: str, publisher: Publisher):
data = { data = {
@@ -52,75 +60,204 @@ def _placeholder_image():
return torch.from_numpy(image)[None,] return torch.from_numpy(image)[None,]
class KritaOutput(SendImageWebSocket): class _BasicTypes(str):
RETURN_TYPES = () """Matches IO.PRIMITIVE, but also any list of choices"""
FUNCTION = "send_images"
OUTPUT_NODE = True basic_types = IO.PRIMITIVE.split(",") # STRING, FLOAT, INT, BOOLEAN
CATEGORY = "krita"
def __eq__(self, other):
return other in self.basic_types or isinstance(other, (list, _BasicTypes))
def __ne__(self, other):
return not self.__eq__(other)
class KritaCanvas: BasicTypes = _BasicTypes("BASIC")
class OutputBatchMode(Enum):
default = "default"
images = "images"
animation = "animation"
layers = "layers"
class KritaOutput(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(cls): def define_schema(cls):
return {} 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 = ("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 @classmethod
def INPUT_TYPES(cls): def execute( # type: ignore
return {} cls,
images: torch.Tensor,
RETURN_TYPES = ("MASK",) x: int = 0,
RETURN_NAMES = ("mask",) y: int = 0,
FUNCTION = "placeholder" name="",
CATEGORY = "krita" batch_mode: OutputBatchMode | str = OutputBatchMode.default,
resize_canvas=False,
def placeholder(self): ):
return (torch.ones(1, 512, 512),) batch_mode = batch_mode.value if isinstance(batch_mode, OutputBatchMode) else batch_mode
info = {
"name": name,
class KritaImageLayer: "offset_x": x,
@classmethod "offset_y": y,
def INPUT_TYPES(cls): "batch_mode": batch_mode,
return { "resize_canvas": resize_canvas,
"required": {
"name": ("STRING", {"default": "Image"}),
}
} }
output = SendImageWebSocket.execute(images, "PNG")
RETURN_TYPES = ("IMAGE",) assert isinstance(output.ui, dict)
RETURN_NAMES = ("image",) output.ui["info"] = [info]
FUNCTION = "placeholder" return output
CATEGORY = "krita"
def placeholder(self, name: str):
return (_placeholder_image(),)
class KritaMaskLayer: class KritaSendText(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(cls): def define_schema(cls):
return { return io.Schema(
"required": { node_id="ETN_KritaSendText",
"name": ("STRING", {"default": "Mask"}), 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,
)
RETURN_TYPES = ("MASK",) @classmethod
RETURN_NAMES = ("mask",) def execute(cls, value: Any, name: str, type: str): # type: ignore
FUNCTION = "placeholder" mime = {
CATEGORY = "krita" "text": "text/plain",
"markdown": "text/markdown",
"html": "text/html",
}[type]
text = "None"
if value is not None:
try:
text = str(value)
except Exception as e:
text = f"Could not convert to text: {e}"
def placeholder(self, name: str): return io.NodeOutput(ui={"text": [{"name": name, "text": text, "content-type": mime}]})
return (torch.ones(1, 512, 512),)
class KritaCanvas(io.ComfyNode):
@classmethod
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"),
],
)
@classmethod
def execute(cls, **kwargs):
return io.NodeOutput(_placeholder_image(), 512, 512, 0, torch.ones(1, 512, 512))
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 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"),
],
)
@classmethod
def execute(cls, **kwargs):
return io.NodeOutput(torch.ones(1, 512, 512), False, 0, 0)
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 = [ _param_types = [
@@ -133,64 +270,95 @@ _param_types = [
"prompt (positive)", "prompt (positive)",
"prompt (negative)", "prompt (negative)",
] ]
_fmax = sys.float_info.max
class Parameter: class Parameter(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(cls): def define_schema(cls):
return { return io.Schema(
"required": { node_id="ETN_Parameter",
"name": ("STRING", {"default": "Parameter"}), display_name="Parameter",
"type": (_param_types, {"default": "auto"}), category="krita",
"default": ("STRING", {"default": ""}), inputs=[
"min": ("FLOAT", {"default": 0.0}), io.String.Input("name", default="Parameter"),
"max": ("FLOAT", {"default": 1.0}), 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 = ("*",)
RETURN_NAMES = ("value",)
FUNCTION = "placeholder"
CATEGORY = "krita"
def placeholder(self, name: str, type: str, default, min, max):
return (default,)
class KritaStyle:
@classmethod @classmethod
def INPUT_TYPES(cls): def execute(cls, name: str, type: str, default, min=0.0, max=1.0): # type: ignore
return { if type == "number":
"required": { return io.NodeOutput(float(default))
"name": ("STRING", {"default": "Style"}), elif type == "number (integer)":
"sampler_preset": (["auto", "regular", "live"],), return io.NodeOutput(int(default))
} return io.NodeOutput(default)
}
RETURN_TYPES = (
"MODEL",
"CLIP",
"VAE",
"STRING",
"STRING",
comfy.samplers.KSampler.SAMPLERS,
comfy.samplers.KSampler.SCHEDULERS,
"INT",
"FLOAT",
)
RETURN_NAMES = (
"model",
"clip",
"vae",
"positive prompt",
"negative prompt",
"sampler name",
"scheduler",
"steps",
"guidance",
)
FUNCTION = "placeholder"
CATEGORY = "krita"
def placeholder(self, name: str, sampler_preset: str): class KritaStyle(io.ComfyNode):
@classmethod
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"),
],
)
@classmethod
def execute(cls, name: str, sampler_preset: str): # type: ignore
raise NotImplementedError("This workflow must be started from Krita!")
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!") raise NotImplementedError("This workflow must be started from Krita!")
+341 -83
View File
@@ -1,30 +1,44 @@
from __future__ import annotations from __future__ import annotations
from copy import copy
from dataclasses import dataclass
import time
from typing import NamedTuple
from uuid import uuid4
from PIL import Image from PIL import Image
import numpy as np import numpy as np
import base64 import base64
import torch import torch
import torch.nn.functional as F
from io import BytesIO from io import BytesIO
from server import PromptServer, BinaryEventTypes from server import PromptServer, BinaryEventTypes
from comfy.clip_vision import ClipVisionModel
from comfy.sd import StyleModel
from comfy_api.latest import io
class LoadImageBase64:
class LoadImageBase64(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(s): def define_schema(cls):
return {"required": {"image": ("STRING", {"multiline": False})}} 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") @classmethod
CATEGORY = "external_tooling" def execute(cls, image: str):
FUNCTION = "load_image" _strip_prefix(image, "data:image/png;base64,")
def load_image(self, image):
imgdata = base64.b64decode(image) imgdata = base64.b64decode(image)
img = Image.open(BytesIO(imgdata)) img = Image.open(BytesIO(imgdata))
if "A" in img.getbands(): if "A" in img.getbands():
mask = np.array(img.getchannel("A")).astype(np.float32) / 255.0 mask = np.array(img.getchannel("A")).astype(np.float32) / 255.0
mask = 1.0 - torch.from_numpy(mask) mask = torch.from_numpy(mask)
else: else:
mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu") mask = None
img = img.convert("RGB") img = img.convert("RGB")
img = np.array(img).astype(np.float32) / 255.0 img = np.array(img).astype(np.float32) / 255.0
@@ -33,16 +47,20 @@ class LoadImageBase64:
return (img, mask) return (img, mask)
class LoadMaskBase64: class LoadMaskBase64(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(s): def define_schema(cls):
return {"required": {"mask": ("STRING", {"multiline": False})}} 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",) @classmethod
CATEGORY = "external_tooling" def execute(cls, mask: str):
FUNCTION = "load_mask" _strip_prefix(mask, "data:image/png;base64,")
def load_mask(self, mask):
imgdata = base64.b64decode(mask) imgdata = base64.b64decode(mask)
img = Image.open(BytesIO(imgdata)) img = Image.open(BytesIO(imgdata))
img = np.array(img).astype(np.float32) / 255.0 img = np.array(img).astype(np.float32) / 255.0
@@ -50,23 +68,24 @@ class LoadMaskBase64:
if img.dim() == 3: # RGB(A) input, use red channel if img.dim() == 3: # RGB(A) input, use red channel
img = img[:, :, 0] img = img[:, :, 0]
return (img.unsqueeze(0),) return (img.unsqueeze(0),)
class SendImageWebSocket:
class SendImageWebSocket(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(s): def define_schema(cls):
return { return io.Schema(
"required": { node_id="ETN_SendImageWebSocket",
"images": ("IMAGE",), display_name="Send Image (WebSocket)",
"format": (["PNG", "JPEG"], {"default": "PNG"}), category="external_tooling",
} inputs=[
} io.Image.Input("images"),
io.Combo.Input("format", options=["PNG", "JPEG"], default="PNG"),
],
is_output_node=True,
)
RETURN_TYPES = () @classmethod
FUNCTION = "send_images" def execute(cls, images: torch.Tensor, format: str):
OUTPUT_NODE = True
CATEGORY = "external_tooling"
def send_images(self, images, format):
results = [] results = []
for tensor in images: for tensor in images:
array = 255.0 * tensor.cpu().numpy() array = 255.0 * tensor.cpu().numpy()
@@ -78,46 +97,160 @@ class SendImageWebSocket:
[format, image, None], [format, image, None],
server.client_id, server.client_id,
) )
results.append( results.append({
{"source": "websocket", "content-type": f"image/{format.lower()}", "type": "output"} "source": "websocket",
) "content-type": f"image/{format.lower()}",
"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 @classmethod
def INPUT_TYPES(cls): def execute(cls, id: str):
return { image_data, content_type = image_cache.get(id, extend=True)
"required": { if image_data is None:
"image": ("IMAGE",), raise ValueError(f"Image with ID {id} not found in cache.")
"x": (
"INT",
{"default": 0, "min": 0, "max": 8192, "step": 1},
),
"y": (
"INT",
{"default": 0, "min": 0, "max": 8192, "step": 1},
),
"width": (
"INT",
{"default": 512, "min": 1, "max": 8192, "step": 1},
),
"height": (
"INT",
{"default": 512, "min": 1, "max": 8192, "step": 1},
),
}
}
CATEGORY = "external_tooling" img = Image.open(BytesIO(image_data))
RETURN_TYPES = ("IMAGE",) w, h = img.size
FUNCTION = "crop" c = len(img.getbands())
normalized = np.array(img).astype(np.float32) / 255.0
tensor = torch.from_numpy(normalized).reshape(1, h, w, c)
match c:
case 1:
image = tensor.expand(1, h, w, 3)
mask = tensor.reshape(1, h, w)
case 3:
image = tensor
mask = tensor[..., 0]
case 4:
image = tensor[..., :3]
mask = tensor[..., 3]
def crop(self, image, x, y, width, height): return io.NodeOutput(image, mask)
out = image[:, y : y + height, x : x + width, :]
return (out,)
class SaveImageCache(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_SaveImageCache",
display_name="Save Image to Cache",
category="external_tooling",
inputs=[
io.Image.Input("images"),
io.Combo.Input("format", options=["PNG", "JPEG"], default="PNG"),
],
is_output_node=True,
)
@classmethod
def execute(cls, images: torch.Tensor, format: str):
results = []
for tensor in images:
array = 255.0 * tensor.cpu().numpy()
image = Image.fromarray(np.clip(array, 0, 255).astype(np.uint8))
key = image_cache.add(image, format)
results.append({
"source": "http",
"id": key,
"content-type": f"image/{format.lower()}",
"type": "output",
})
return io.NodeOutput(ui={"images": results})
def to_bchw(image: torch.Tensor): def to_bchw(image: torch.Tensor):
@@ -136,21 +269,22 @@ def mask_batch(mask: torch.Tensor):
return mask return mask
class ApplyMaskToImage: class ApplyMaskToImage(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(cls): def define_schema(cls):
return { return io.Schema(
"required": { node_id="ETN_ApplyMaskToImage",
"image": ("IMAGE",), display_name="Apply Mask to Image",
"mask": ("MASK",), category="external_tooling",
} inputs=[
} io.Image.Input("image"),
io.Mask.Input("mask"),
],
outputs=[io.Image.Output(display_name="masked")],
)
CATEGORY = "external_tooling" @classmethod
RETURN_TYPES = ("IMAGE",) def execute(cls, image: torch.Tensor, mask: torch.Tensor):
FUNCTION = "apply_mask"
def apply_mask(self, image: torch.Tensor, mask: torch.Tensor):
out = to_bchw(image) out = to_bchw(image)
if out.shape[1] == 3: # Assuming RGB images if out.shape[1] == 3: # Assuming RGB images
out = torch.cat([out, torch.ones_like(out[:, :1, :, :])], dim=1) out = torch.cat([out, torch.ones_like(out[:, :1, :, :])], dim=1)
@@ -158,9 +292,9 @@ class ApplyMaskToImage:
assert mask.ndim == 3, f"Mask should have shape [B, H, W]. {mask.shape}" assert mask.ndim == 3, f"Mask should have shape [B, H, W]. {mask.shape}"
assert out.ndim == 4, f"Image should have shape [B, C, H, W]. {out.shape}" assert out.ndim == 4, f"Image should have shape [B, C, H, W]. {out.shape}"
assert ( assert out.shape[-2:] == mask.shape[-2:], (
out.shape[-2:] == mask.shape[-2:] f"Image size {out.shape[-2:]} must match mask size {mask.shape[-2:]}"
), f"Image size {out.shape[-2:]} must match mask size {mask.shape[-2:]}" )
is_mask_batch = mask.shape[0] == out.shape[0] is_mask_batch = mask.shape[0] == out.shape[0]
# Apply each mask in the batch to its corresponding image's alpha channel # Apply each mask in the batch to its corresponding image's alpha channel
@@ -169,3 +303,127 @@ class ApplyMaskToImage:
out[i, 3, :, :] = alpha out[i, 3, :, :] = alpha
return (to_bhwc(out),) return (to_bhwc(out),)
class _ReferenceImageData(NamedTuple):
image: torch.Tensor
weight: float
range: tuple[float, float]
class ReferenceImage(io.ComfyNode):
@classmethod
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")],
)
@classmethod
def execute(
cls,
image: torch.Tensor,
weight: float,
range_start: float,
range_end: float,
reference_images: list[_ReferenceImageData] | None = None,
):
imgs = copy(reference_images) if reference_images is not None else []
imgs.append(_ReferenceImageData(image, weight, (range_start, range_end)))
return (imgs,)
class ApplyReferenceImages(io.ComfyNode):
@classmethod
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")],
)
@classmethod
def execute(
cls,
conditioning: list[list],
clip_vision: ClipVisionModel,
style_model: StyleModel,
references: list[_ReferenceImageData],
):
delimiters = {0.0, 1.0}
delimiters |= set(r.range[0] for r in references)
delimiters |= set(r.range[1] for r in references)
delimiters = sorted(delimiters)
ranges = [(delimiters[i], delimiters[i + 1]) for i in range(len(delimiters) - 1)]
embeds = [_encode_image(r.image, clip_vision, style_model, r.weight) for r in references]
base = conditioning[0][0]
result = []
for start, end in ranges:
e = [
embeds[i]
for i, r in enumerate(references)
if r.range[0] <= start and r.range[1] >= end
]
options = conditioning[0][1].copy()
options["start_percent"] = start
options["end_percent"] = end
result.append((torch.cat([base] + e, dim=1), options))
return (result,)
def _encode_image(
image: torch.Tensor, clip_vision: ClipVisionModel, style_model: StyleModel, weight: float
):
e = clip_vision.encode_image(image)
e = style_model.get_cond(e).flatten(start_dim=0, end_dim=1).unsqueeze(dim=0)
e = _downsample_image_cond(e, weight)
return e
def _downsample_image_cond(cond: torch.Tensor, weight: float):
if weight >= 1.0:
return cond
elif weight <= 0.0:
return torch.zeros_like(cond)
elif weight >= 0.6:
factor = 2
elif weight >= 0.3:
factor = 3
else:
factor = 4
# Downsample the clip vision embedding to make it smaller, resulting in less impact
# compared to other conditioning.
# See https://github.com/kaibioinfo/ComfyUI_AdvancedRefluxControl
(b, t, h) = cond.shape
m = int(np.sqrt(t))
cond = F.interpolate(
cond.view(b, m, m, h).transpose(1, -1),
size=(m // factor, m // factor),
mode="area",
)
return cond.transpose(1, -1).reshape(b, -1, h)
def _strip_prefix(s: str, prefix: str) -> str:
if s.startswith(prefix):
return s[len(prefix) :]
return s
+27 -28
View File
@@ -1,5 +1,4 @@
from __future__ import annotations from __future__ import annotations
from weakref import ref as WeakRef
from pathlib import Path from pathlib import Path
from tqdm import tqdm from tqdm import tqdm
import torch import torch
@@ -8,6 +7,7 @@ import torch.nn.functional as F
from torch import Tensor from torch import Tensor
from transformers import CLIPImageProcessor, CLIPConfig, CLIPVisionModel, PreTrainedModel from transformers import CLIPImageProcessor, CLIPConfig, CLIPVisionModel, PreTrainedModel
from kornia.filters import box_blur from kornia.filters import box_blur
from comfy_api.latest import io
from .nodes import to_bchw, to_bhwc 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.concept_embeds_weights = nn.Parameter(torch.ones(17), requires_grad=False)
self.special_care_embeds_weights = nn.Parameter(torch.ones(3), 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): def forward(self, clip_input, images: Tensor, sensitivity: float):
with torch.no_grad(): with torch.no_grad():
image_batch = self.vision_model(clip_input)[1] image_batch = self.vision_model(clip_input)[1]
@@ -76,7 +80,7 @@ class CLIPSafetyChecker(PreTrainedModel):
class CachedModels: class CachedModels:
_instance: WeakRef | None = None _instance: CachedModels | None = None
def __init__(self): def __init__(self):
model_dir = Path(__file__).parent / "safetychecker" model_dir = Path(__file__).parent / "safetychecker"
@@ -91,11 +95,9 @@ class CachedModels:
@classmethod @classmethod
def load(cls): def load(cls):
models = cls._instance and cls._instance() if cls._instance is None:
if models is None: cls._instance = CachedModels()
models = cls() return cls._instance
cls._instance = WeakRef(models)
return models
def download(self, url: str, target: Path): def download(self, url: str, target: Path):
import requests import requests
@@ -118,29 +120,26 @@ class CachedModels:
) from e ) from e
class NSFWFilter: class NSFWFilter(io.ComfyNode):
models: CachedModels @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 @classmethod
def INPUT_TYPES(cls): def execute(cls, image: Tensor, sensitivity: float):
return { models = CachedModels.load()
"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):
image = to_bchw(image) image = to_bchw(image)
input = self.models.feature_extractor(image, do_rescale=False, return_tensors="pt") input = models.feature_extractor(image, do_rescale=False, return_tensors="pt")
filtered = self.models.safety_checker( filtered = models.safety_checker(
images=image, clip_input=input.pixel_values, sensitivity=sensitivity images=image, clip_input=input.pixel_values, sensitivity=sensitivity
) )
return (to_bhwc(filtered),) return io.NodeOutput(to_bhwc(filtered))
+9 -1
View File
@@ -1,12 +1,20 @@
[project] [project]
name = "comfyui-tooling-nodes" name = "comfyui-tooling-nodes"
description = "Provides nodes and server API extensions geared towards using ComfyUI as a backend for external tools." description = "Provides nodes and server API extensions geared towards using ComfyUI as a backend for external tools."
version = "1.5.0" version = "3.1.4"
license = { file = "LICENSE" } license = { file = "LICENSE" }
[project.urls] [project.urls]
Repository = "https://github.com/Acly/comfyui-tooling-nodes" Repository = "https://github.com/Acly/comfyui-tooling-nodes"
[tool.ruff]
target-version = "py311"
line-length = 100
preview = true
[tool.ruff.lint]
ignore = ["E741", "BLE001"]
[tool.black] [tool.black]
line-length = 100 line-length = 100
preview = true preview = true
+85 -74
View File
@@ -8,6 +8,7 @@ import torch.nn.functional as F
import math import math
from torch import Tensor, Size from torch import Tensor, Size
from comfy.model_patcher import ModelPatcher from comfy.model_patcher import ModelPatcher
from comfy_api.latest import io
def downsample_mask(mask: Tensor, batch: int, target_size: int, original_shape: Size) -> Tensor: def downsample_mask(mask: Tensor, batch: int, target_size: int, original_shape: Size) -> Tensor:
@@ -65,105 +66,112 @@ class Region(NamedTuple):
return result return result
class BackgroundRegion: Regions = io.Custom("Regions")
class BackgroundRegion(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(cls): def define_schema(cls):
return {"required": {"conditioning": ("CONDITIONING",)}} 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" @classmethod
RETURN_TYPES = ("REGIONS",) def execute(cls, conditioning: list):
FUNCTION = "define"
def define(self, conditioning: list):
return (Region(None, None, conditioning),) return (Region(None, None, conditioning),)
class DefineRegion: class DefineRegion(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(cls): def define_schema(cls):
return { return io.Schema(
"required": { node_id="ETN_DefineRegion",
"mask": ("MASK",), display_name="Define Region",
"conditioning": ("CONDITIONING",), category="external_tooling/regions",
}, inputs=[
"optional": { io.Mask.Input("mask"),
"regions": ("REGIONS",), io.Conditioning.Input("conditioning"),
}, Regions.Input("regions", optional=True),
} ],
outputs=[Regions.Output(display_name="regions")],
)
CATEGORY = "external_tooling/regions" @classmethod
RETURN_TYPES = ("REGIONS",) def execute(cls, mask: Tensor, conditioning: list, regions: Region | None = None):
FUNCTION = "define"
def define(self, mask: Tensor, conditioning: list, regions: Region | None = None):
if mask.dim() < 3: if mask.dim() < 3:
mask = mask.unsqueeze(0) mask = mask.unsqueeze(0)
return (Region(regions, mask, conditioning),) return io.NodeOutput(Region(regions, mask, conditioning))
class ListRegionMasks: class ListRegionMasks(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(cls): def define_schema(cls):
return {"required": {"regions": ("REGIONS",)}} return io.Schema(
node_id="ETN_ListRegionMasks",
CATEGORY = "external_tooling/regions" display_name="List Region Masks",
RETURN_TYPES = ("MASK",) category="external_tooling/regions",
FUNCTION = "get_masks" inputs=[Regions.Input("regions")],
outputs=[io.Mask.Output(display_name="masks")],
def get_masks(self, regions: Region): )
return (torch.stack([r.mask for r in regions.preprocess()], dim=0),)
class AttentionMask:
@classmethod @classmethod
def INPUT_TYPES(s): def execute(cls, regions: Region):
return { return io.NodeOutput(torch.stack([r.mask for r in regions.preprocess()], dim=0))
"required": {
"model": ("MODEL",),
"regions": ("REGIONS",),
}
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "attention_mask"
CATEGORY = "external_tooling/regions"
mask: Tensor class AttentionMask(io.ComfyNode):
conds: list[Tensor] @classmethod
batch_size: int def define_schema(cls):
return io.Schema(
node_id="ETN_AttentionMask",
display_name="Regions Attention Mask",
category="external_tooling/regions",
inputs=[io.Model.Input("model"), Regions.Input("regions")],
outputs=[io.Model.Output(display_name="model")],
)
def attention_mask(self, model: ModelPatcher, regions: Region): @classmethod
new_model = model.clone() def execute(cls, model: ModelPatcher, regions: Region):
region_list = regions.preprocess() return io.NodeOutput(AttentionMaskPatch.apply(model, regions))
num_conds = len(region_list)
class AttentionMaskPatch:
def __init__(self, region_list: list[Region]):
mask = torch.stack([r.mask for r in region_list], dim=0) mask = torch.stack([r.mask for r in region_list], dim=0)
mask_sum = mask.sum(dim=0, keepdim=True) mask_sum = mask.sum(dim=0, keepdim=True)
assert mask_sum.sum() > 0, "There are areas that are zero in all masks." assert mask_sum.sum() > 0, "There are areas that are zero in all masks."
self.mask = mask / mask_sum self.mask = mask / mask_sum
self.conds = [r.conditioning[0][0] 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())
def attn2_patch(q: Tensor, k: Tensor, v: Tensor, extra_options: dict): def attn2_patch(q: Tensor, k: Tensor, v: Tensor, extra_options: dict):
assert k.mean() == v.mean(), "k and v must be the same." assert k.mean() == v.mean(), "k and v must be the same."
device, dtype = q.device, q.dtype device, dtype = q.device, q.dtype
if self.conds[0].device != device or self.conds[0].dtype != dtype: if patch.conds[0].device != device or patch.conds[0].dtype != dtype:
self.conds = [cond.to(device, dtype=dtype) for cond in self.conds] patch.conds = [cond.to(device, dtype=dtype) for cond in patch.conds]
if self.mask.device != device or self.mask.dtype != dtype: if patch.mask.device != device or patch.mask.dtype != dtype:
self.mask = self.mask.to(device, dtype=dtype) patch.mask = patch.mask.to(device, dtype=dtype)
cond_or_unconds = extra_options["cond_or_uncond"] cond_or_unconds = extra_options["cond_or_uncond"]
num_chunks = len(cond_or_unconds) 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) q_chunks = q.chunk(num_chunks, dim=0)
k_chunks = k.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 = [ conds_tensor = [
cond.repeat(self.batch_size, lcm_tokens // num_tokens[i], 1) cond.repeat(patch.batch_size, lcm_tokens // patch.num_tokens[i], 1)
for i, cond in enumerate(self.conds) for i, cond in enumerate(patch.conds)
] ]
conds_tensor = torch.cat(conds_tensor, dim=0) conds_tensor = torch.cat(conds_tensor, dim=0)
@@ -174,9 +182,9 @@ class AttentionMask:
qs.insert(0, q_chunks[i]) qs.insert(0, q_chunks[i])
ks.insert(0, k_target) ks.insert(0, k_target)
else: 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) ks.insert(0, conds_tensor)
for _ in range(num_conds - 1): for _ in range(patch.num_conds - 1):
cond_or_unconds.insert(i, 0) cond_or_unconds.insert(i, 0)
qs = torch.cat(qs, dim=0) qs = torch.cat(qs, dim=0)
@@ -184,29 +192,32 @@ class AttentionMask:
return qs, ks, ks return qs, ks, ks
def attn2_output_patch(out: Tensor, extra_options: dict): def attn2_output_patch(out: Tensor, extra_options: dict):
num_conds = patch.num_conds
cond_or_unconds = extra_options["cond_or_uncond"] cond_or_unconds = extra_options["cond_or_uncond"]
mask_downsample = downsample_mask( 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] = [] outputs: list[Tensor] = []
pos = 0 pos = 0
i = 0 i = 0
while i < len(cond_or_unconds): while i < len(cond_or_unconds):
if cond_or_unconds[i] == 1: # uncond if cond_or_unconds[i] == 1: # uncond
outputs.append(out[pos : pos + self.batch_size]) outputs.append(out[pos : pos + patch.batch_size])
pos += self.batch_size pos += patch.batch_size
else: else:
masked = out[pos : pos + num_conds * self.batch_size] * mask_downsample masked = out[pos : pos + num_conds * patch.batch_size] * mask_downsample
masked = masked.view(num_conds, self.batch_size, out.shape[1], out.shape[2]) masked = masked.view(num_conds, patch.batch_size, out.shape[1], out.shape[2])
masked = masked.sum(dim=0) masked = masked.sum(dim=0)
outputs.append(masked) outputs.append(masked)
pos += num_conds * self.batch_size pos += num_conds * patch.batch_size
for _ in range(num_conds - 1): for _ in range(num_conds - 1):
cond_or_unconds.pop(i) cond_or_unconds.pop(i)
i += 1 i += 1
return torch.cat(outputs, dim=0) return torch.cat(outputs, dim=0)
new_model = model.clone()
new_model.set_model_attn2_patch(attn2_patch) new_model.set_model_attn2_patch(attn2_patch)
new_model.set_model_attn2_output_patch(attn2_output_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
+103 -94
View File
@@ -3,49 +3,29 @@ import numpy as np
import numpy.typing as npt import numpy.typing as npt
import torch import torch
from torch import Tensor from torch import Tensor
from comfy_api.latest import io
IntArray = npt.NDArray[np.int_] IntArray = npt.NDArray[np.int_]
class TileLayout: class TileLayout:
@classmethod def __init__(
def INPUT_TYPES(cls): self, image: Tensor, min_tile_size: int, padding: int, blending: int, multiple: int
return { ):
"required": { assert all([x % multiple == 0 for x in image.shape[-3:-1]]), (
"image": ("IMAGE",), "Image size must be divisible by multiple"
"min_tile_size": ("INT", {"default": 512, "min": 64, "max": 8192, "step": 8}), )
"padding": ("INT", {"default": 32, "min": 0, "max": 8192, "step": 8}), assert min_tile_size % multiple == 0, "Tile size must be divisible by multiple"
"blending": ("INT", {"default": 8, "min": 0, "max": 256, "step": 8}), assert blending <= padding, "Blending must be smaller than padding"
}
}
CATEGORY = "external_tooling/tiles" self.image_size: IntArray = np.array(image.shape[-3:-1])
RETURN_TYPES = ("TILE_LAYOUT",) self.padding: int = padding
FUNCTION = "node" self.blending: int = blending
self.tile_count: IntArray = np.maximum(1, self.image_size // (min_tile_size - 2 * padding))
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"
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))
image_size_with_overlap = self.image_size + (self.tile_count - 1) * 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) 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): def size(self, coord: IntArray):
return self.end(coord) - self.start(coord) return self.end(coord) - self.start(coord)
@@ -96,80 +76,109 @@ class TileLayout:
image[rect] = (1 - mask) * image[rect] + mask * tile image[rect] = (1 - mask) * image[rect] + mask * tile
class ExtractImageTile: class CreateTileLayout(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(cls): def define_schema(cls):
return { return io.Schema(
"required": { node_id="ETN_TileLayout",
"image": ("IMAGE",), display_name="Create Tile Layout",
"layout": ("TILE_LAYOUT",), category="external_tooling/tiles",
"index": ("INT", {"min": 0}), 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 @classmethod
def INPUT_TYPES(cls): def execute(cls, image: Tensor, min_tile_size: int, padding: int, blending: int, multiple: int):
return { return io.NodeOutput(TileLayout(image, min_tile_size, padding, blending, multiple))
"required": {
"mask": ("MASK",),
"layout": ("TILE_LAYOUT",),
"index": ("INT", {"min": 0}),
}
}
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) tile = layout.tile(mask.unsqueeze(3), index)
return (tile.squeeze(3),) return io.NodeOutput(tile.squeeze(3))
class GenerateTileMask: class GenerateTileMask(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(cls): def define_schema(cls):
return { return io.Schema(
"required": {"layout": ("TILE_LAYOUT",), "index": ("INT", {"min": 0})}, node_id="ETN_GenerateTileMask",
"optional": {"blend": ("BOOLEAN",)}, 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 @classmethod
def INPUT_TYPES(cls): def execute(cls, layout: TileLayout, index: int, blend: bool = False):
return { return io.NodeOutput(layout.mask(layout.coord(index), blend=blend))
"required": {
"image": ("IMAGE",),
"layout": ("TILE_LAYOUT",),
"index": ("INT", {"min": 0}),
"tile": ("IMAGE",),
}
}
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" assert index < layout.total_count, f"Index {index} out of range"
if index == 0: if index == 0:
image = image.clone() image = image.clone()
layout.merge(image, index, tile) 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 import re
from functools import cache from functools import cache
from typing import NamedTuple from typing import NamedTuple
from comfy_api.latest import io
@cache @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 (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}" 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() pkg.install()
text, embeddings = _extract_embeddings(text) 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) return " ".join(translate_chunk(c.text, c.lang) for c in chunks)
class Translate: class Translate(io.ComfyNode):
@staticmethod @classmethod
def INPUT_TYPES(): def define_schema(cls):
return {"required": {"text": ("STRING", {"multiline": True})}} 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" @classmethod
RETURN_TYPES = ("STRING",) def execute(cls, text: str):
FUNCTION = "translate" return io.NodeOutput(translate(text))
def translate(self, text: str):
return (translate(text),)
_lang_regex = re.compile(r"(lang:\w\w)") _lang_regex = re.compile(r"(lang:\w\w)")