97 Commits
Author SHA1 Message Date
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
Acly 9a9cbe78a5 High level krita parameters (style, control-net, ip-adapter) 2024-10-11 09:30:12 +02:00
Acly 9a8d90dd95 Replace typed parameter nodes with a universal parameter node which adapts to the first widget it is connected to 2024-10-06 23:58:25 +02:00
Acly 24b7aabf8b Add placeholder image when running external nodes from web ui 2024-10-04 21:27:19 +02:00
Acly 81f944f119 Add custom icons to some of the krita interop nodes 2024-10-04 20:46:56 +02:00
Acly 327b2a1fe3 More parameter nodes for shared workflows 2024-10-04 11:04:35 +02:00
Acly fb847a5225 Publish workflows only if there's a related sink node in the graph 2024-10-02 17:56:19 +02:00
Acly 5a45172d02 API to exchange workflows between multiple connected clients
Placeholder nodes to parametrize and run custom workflows from Krita
2024-10-02 11:07:02 +02:00
Acly 29e24ec52c Initial API for workflow exchange between ComfyUI clients 2024-09-23 20:59:57 +02:00
Marco Tundo e5e62a4a79 Added filetype selector to SendImageWebSocket 2024-09-23 20:59:19 +02:00
Acly 1a24975f99 Don't print trace if a model can't be detected (leads to more confusion than it helps) 2024-09-20 09:53:12 +02:00
Acly 61fa161c34 Regions: also check if dtype matches 2024-09-12 10:17:37 +02:00
Acly f986f6a442 Document upload api 2024-08-30 12:10:16 +02:00
Acly e0d0c3cc2c Add model upload API endpoint
- folder must match existing model folder
- file must be safetensors
2024-08-26 23:05:17 +02:00
Acly f5ec9d830c Make /api/etn/model_info work with diffusion_models folder (formerly unet)
- endpoint is now `/api/etn/model_info/{folder_name}`
- old endpoints are still available
- also works with unet folder (deprecated)
2024-08-20 16:17:38 +02:00
Acly d1dcf12f10 Fix detection for HunyuanDit #17 2024-08-09 18:54:05 +02:00
Acly cb92e547c6 Version 1.4.0, fix toml license directive 2024-08-09 10:08:30 +02:00
Acly b5fec4a062 Model info api: add aura-flow, hunyuan-dit, flux 2024-08-05 00:25:26 +02:00
Acly 42965013f9 Add __future__ imports for older python 2024-07-28 13:14:29 +02:00
Acly d20615fb48 Support language directives in translate api, add documentation 2024-07-27 15:37:20 +02:00
Acly 5bad00f72f Remove debug prints 2024-07-24 14:38:51 +02:00
Acly df54344077 Translate: parse language directives included in the text 2024-07-24 12:00:07 +02:00
Acly f42c0f29b6 Don't translate embeddings 2024-07-22 18:48:02 +02:00
Acly 547c3d5c97 Add text translation node & API 2024-07-22 16:53:09 +02:00
Acly 73babbd00e Add NSFWFilter node 2024-07-21 20:47:18 +02:00
Acly cac32fe37c Move image channel permutation to separate functions 2024-07-21 18:31:42 +02:00
Acly 5620b5c6e2 Document tiling nodes 2024-07-20 23:52:09 +02:00
Acly 3d4a960982 Document region nodes 2024-07-20 23:12:46 +02:00
Acly 9d533984c2 Version 1.2.0 2024-06-24 17:55:55 +02:00
Acly 715a41e04f Prefix server API route with api/ 2024-06-20 16:39:15 +02:00
Acly aff32e8da6 Tiles: fix div by zero when image is smaller than tile size 2024-06-20 11:26:00 +02:00
Acly e46123612d Model info API: support SD3 2024-06-12 16:56:25 +02:00
Acly 6e7b2445db Remove seperable=True for box_blur
- not supported by older kornia versions, and probably no actual speed up at typical tile sizes
2024-06-12 09:51:03 +02:00
Acly 2f39365248 Bump version to 1.1.0 2024-06-11 12:14:23 +02:00
Acly c324f6741d Expand mask batch dimension if it doesn't exist 2024-06-08 09:31:29 +02:00
Acly c27b662fd8 Don't unpack tuple within index operation (not supported by older Python) 2024-06-07 17:25:52 +02:00
21 changed files with 3923 additions and 355 deletions
+2 -1
View File
@@ -11,11 +11,12 @@ jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
if: ${{ github.repository_owner == 'Acly' }}
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@main
uses: Comfy-Org/publish-node-action@v1
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+3 -1
View File
@@ -1,4 +1,6 @@
.vscode
.env
.dev
__pycache__
__pycache__
safetychecker/*.safetensors
+239 -26
View File
@@ -2,12 +2,20 @@
Provides nodes and API geared towards using ComfyUI as a backend for external tools.
## Nodes for sending and receiving images
* <a href="#images">Sending and receiving images</a>
* <a href="#regions">Regions (Attention Masking)
* <a href="#tiles">Tiled image processing
* <a href="#misc">Miscellanious nodes
* <a href="#api">Http API extensions (Model inspection)
* <a href="#installation">⭳ Installation</a>
## <a id="images" href="#toc">Sending and receiving images</a>
ComfyUI exchanges images via the filesystem. This requires a
multi-step process (upload images, prompt, download images), is rather
inefficient, and invites a whole class of potential issues. It's also unclear
at which point those images will get cleaned up if ComfyUI is used
multi-step process (upload images, prompt, download images), which
invites a whole class of potential issues you might not want to deal with.
It's also unclear at which point those images will get cleaned up if ComfyUI is used
via external tools.
### Load Image (Base64)
@@ -25,7 +33,7 @@ Loads a mask (single channel) from a PNG embedded into the prompt as base64 stri
### Send Image (WebSocket)
Sends an output image over the client WebSocket connection as PNG binary data.
* Inputs: the image (RGB or RGBA)
* Inputs: the image (RGB or RGBA), supports batches
This will first send one binary message for each image in the batch via WebSocket:
```
@@ -36,9 +44,74 @@ 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>}}
```
## Nodes for working on regions
### Load Image from Cache
When integrating ComfyUI into tools which use layers and compose them on the fly, it is useful to only receive relevant masked regions.
Loads an image or mask that has been uploaded previously into the workflow.
Uploaded images are temporarily stored in RAM rather than written to disk. This
method has less overhead compared to embedding images as base64 into the prompt,
but is more complex to implement.
* Inputs: id of an image that was uploaded previously
* Outputs: image (RGB) and mask (A of RGBA input, or first channel if no alpha present).
To upload an image, upload the _bytes_ of a PNG via a HTTP PUT request to
`/api/etn/image/{id}`. JPEG or other formats also work. Choose any `id` which
does not clash with other images you upload, and reference it in the node. The
request returns `201` if the image was uploaded and `200` if it was already
cached.
### Save Image to Cache
Stores an output image in RAM temporarily and allows retrieval over HTTP.
This is typically faster than WebSocket, especially for large images.
* Inputs: the image (RGB or RGBA). Batches are supported.
This node will send a JSON message over WebSocket when an image is ready:
```json
{
"type": "executed",
"data": {
"node": "<node ID>",
"output": {
"images": [
{"source": "http", "id": "<image ID>", "content-type": "image/png", "type": "output"}
]
},
"prompt_id": "prompt ID"
}
}
```
To download the images, send a HTTP GET request to `/api/etn/image/{id}` with
the image IDs from the message. Images will be cached for a few minutes.
## <a id="regions" href="#toc">Regions</a>
These nodes implement attention masking for arbitrary number of image regions. Text prompts only apply to the masked area.
In contrast to condition masking, this method is less "forceful", but leads to more natural image compositions.
![Regions Attention Mask](workflows/region_attention_mask.png)
[Workflow: region_attention_mask.json](workflows/region_attention_mask.json)
### Background Region
This node starts a list of regions. It takes a prompt, but no mask. The prompt is assigned to all image areas which are _not_
covered by another region mask in the list.
### Define Region
Appends a new region to a region list (or starts a new list). Takes a prompt, and mask which defines the area in the image
the prompt will apply to. Masks must be the same size as the image _or_ the latent (which is factor 8 smaller).
### List Region Masks
This node takes a list of regions and outputs all their masks. It can be useful for inspection, debugging or to reuse the
computed background mask.
### Regions Attention Mask
Patches the model to use the provided list of regions. This replaces the positive text conditioning which is provided
to the sampler. It's still possible to pass ControlNet and other conditioning to the sampler.
### Apply Mask to Image
@@ -46,30 +119,170 @@ Copies a mask into the alpha channel of an image.
* Inputs: image and mask
* Outputs: RGBA image with mask used as transparency
## API for model inspection
There are various types of models that can be loaded as checkpoint, LoRA, ControlNet, etc. which cannot be used interchangeably. The following API helps to categorize and filter them.
## <a id="tiles" href="#toc">Tiles</a>
### /etn/model_info
Splitting an image into tiles to be processed individually is a useful method to speed up
diffusion and save VRAM. There are various nodes out there which provide a fixed pipeline.
In contrast, the following nodes only provide a way to split an image into tiles and merge
it back together. With tools and scripts it is feasible to generate individual workflows
for each tile. This allows maximum flexibility (different prompts, regions, control, etc.).
Lists available models with additional classification info.
* Paramters: _none_
* Output: list of model files
```
{
"checkpoint_file.safetensors": {
"base_model": "sd15"|"sd20"|"sd21"|"sdxl"|"ssd1b"|"svd"|"cascade-b"|"cascade-c",
"is_inpaint": true|false,
"is_refiner": true|false
},
...
}
```
The entry is `{"base_model": "unknown"}` for models which are not in safetensors format or do not match any of the known base models.
![Image tiles](workflows/image_tiles.png)
[Workflow: image_tiles.json](workflows/image_tiles.json)
_Note: currently only supports checkpoints. May add other models in the future._
### Create Tile Layout
## Installation
This node defines the tiling parameters:
* **min_tile_size**: Minimum resolution of each tile in pixels. Tiles may be larger to fit the image size evenly.
* **padding**: Padding around each tile in pixels. Overlaps with neighbour tiles. There is no padding at the image borders.
* **blending**: The part of the padding area which is used for smooth blending to avoid seams. Affects masks which are generated from this layout.
The number of tiles is: `image_size // (min_tile_size + 2 * padding)`
### Extract Image Tile
Splits out part of an image. Tile indices range from 0 to number of tiles and are column-major
(tile 1 is usually below tile 0).
### Extract Mask Tile
Same as "Extract Image Tile" but for masks.
### Merge Image Tile
Merges a tile into a full image, usually after sampling. Uses a smooth transition overlap
between neighbouring tiles depending on padding and blending values.
### Generate Tile Mask
Creates a coverage mask for a certain tile. The size of the mask matches the image tile size.
The image area will be white (1) and the padding area black (0), with a smooth transition
depending on the chosen blend size.
This mask is used internally by "Merge Image Tile", but it can also be useful as input for "Set Latent Noise Mask" in upscale workflows.
## <a id="misc" href="#toc">Miscellaneous Nodes</a>
<a id="node-translate"></a>
### Translate Text
Node which translates a string into English. The language to translate from is indicated with a
_language directive_ of the form `lang:xx` where xx is a 2-letter language code. Multiple
directives are allowed and change language for any text that comes after, until the next
directive. `lang:en` (the default) passes through text fragments untouched. Useful
for keywords, tags and such.
Examples:
| Input | Output |
|:-|:-|
| lang:de eine modische handtasche aus grünem kunstleder | a fashionable handbag made of green suede |
| origami paperwork, lang:zh 狐狸和鹤, lang:en mountain view | origami paperwork, Fox and crane, mountain view |
Translation happens entirely local, powered by [argosopentech/argos-translate](https://github.com/argosopentech/argos-translate):
* Install with `pip install argostranslate` or `pip install -r requirements.txt`
* Models are automatically downloaded on first use.
There is also a [translation API](#api-translation) for immediate feedback in tool UI.
### NSFW Filter
Checks images for NSFW content using [Safety-Checker](https://huggingface.co/CompVis/stable-diffusion-safety-checker). Images which don't pass the check are blurred to
obfuscate contents. Model is downloaded on first use.
Inputs: image and sensitivity (0.5 for explicit content only, 0.7+ to include partial nudity).
**Important:** the filter isn't perfect. Some explicit content may slip through.
## <a id="api" href="#toc">API extensions</a>
### GET /api/etn/model_info/{folder_name}
There are various types of models that can be loaded as checkpoint, LoRA, ControlNet, etc. which cannot be used interchangeably. This endpoint helps to categorize and filter them.
#### Paramters
* `folder_name`: sub-directory in ComfyUI's models folder.
Supported model types: `checkpoints`, `diffusion_models`, `unet`, `unet_gguf`
* `limit=n`: (query parameter, optional) inspect at `n` models
* `offset=i`: (query parameter, optional) start with the `i`th model
#### Output
Lists available models with additional classification info:
```json
{
"checkpoint_file.safetensors": {
"base_model": "sd15",
"is_inpaint": false,
"type": "eps"
},
...
}
```
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 }
}
```
### GET /api/etn/languages
Returns a list of available languages for translation.
```json
[
{ "name": "English", "code": "en" },
{ ... }
]
```
<a id="api-translation"></a>
### GET /api/etn/translate/{lang}/{text}
Translates `text` into English. `lang` is a 2-letter code indicating the language to translate
from. `text` may also contain _language directives_ to only translate some fragments.
See the [node documentation](#node-translate) for details.
* Output: JSON string
* Example: `/api/etn/translate/de/eine%20modische%20Handtasche` -> `"a fashionable handbag"`
### PUT /api/etn/upload/{folder_name}/{filename}
Uploads a model to ComfyUI's local model folder.
#### Parameters
* `folder_name`: the model type. Must match one of the existing folders in ComfyUI's models folder.
* `filename`: target filename for the model. Must not contain any (absolute or relative) path. Extension must be .safetensors.
#### Output
* Code `201` and `{ "status": "success" }` after successful upload.
* Code `200` and `{ "status": "cached" }` if the file already exists.
* Code `400` and `{ "error": "..." }` if the parameters are invalid.
## <a id="installation" href="#toc">Installation</a>
Download the repository and unpack into the `custom_nodes` folder in the ComfyUI installation directory.
+41 -35
View File
@@ -1,36 +1,42 @@
from . import api, nodes, tile, region
from comfy_api.latest import ComfyExtension, io
from . import api as api, nodes, tile, region, nsfw, translation, krita
NODE_CLASS_MAPPINGS = {
"ETN_LoadImageBase64": nodes.LoadImageBase64,
"ETN_LoadMaskBase64": nodes.LoadMaskBase64,
"ETN_SendImageWebSocket": nodes.SendImageWebSocket,
"ETN_CropImage": nodes.CropImage,
"ETN_ApplyMaskToImage": nodes.ApplyMaskToImage,
"ETN_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,
}
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_ListAppend": "List 🢒 Append",
"ETN_ListElement": "List 🢒 Get Element",
"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",
}
class ExternalToolingNodes(ComfyExtension):
async def get_node_list(self) -> list[type[io.ComfyNode]]:
return [
nodes.LoadImageCache,
nodes.SaveImageCache,
nodes.LoadImageBase64,
nodes.LoadMaskBase64,
nodes.SendImageWebSocket,
nodes.ApplyMaskToImage,
nodes.ReferenceImage,
nodes.ApplyReferenceImages,
tile.CreateTileLayout,
tile.ExtractImageTile,
tile.ExtractMaskTile,
tile.GenerateTileMask,
tile.MergeImageTile,
region.BackgroundRegion,
region.DefineRegion,
region.ListRegionMasks,
region.AttentionMask,
nsfw.NSFWFilter,
translation.Translate,
krita.KritaOutput,
krita.KritaSendText,
krita.KritaCanvas,
krita.KritaSelection,
krita.KritaImageLayer,
krita.KritaMaskLayer,
krita.Parameter,
krita.KritaStyle,
]
async def comfy_entrypoint():
return ExternalToolingNodes()
WEB_DIRECTORY = "./js"
+334 -26
View File
@@ -1,13 +1,22 @@
from __future__ import annotations
from aiohttp import web
from typing import NamedTuple
from typing import Any, NamedTuple
from pathlib import Path
import json
import traceback
import re
import logging
import itertools
import comfy.utils
from comfy import supported_models
from comfy import model_detection
import comfy.utils
import folder_paths
import server
from .translation import available_languages, translate
from .krita import WorkflowExchange
from .nodes import image_cache
input_block_name = "model.diffusion_model.input_blocks.0.0.weight"
model_names = {
@@ -15,12 +24,44 @@ model_names = {
"SD20": "sd20",
"SD21UnclipL": "sd21",
"SD21UnclipH": "sd21",
"SDXLRefiner": "sdxl",
"SDXLRefiner": "sdxl-refiner",
"SDXL": "sdxl",
"SSD1B": "ssd1b",
"SVD_img2vid": "svd",
"Stable_Cascade_B": "cascade-b",
"Stable_Cascade_C": "cascade-c",
"SD3": "sd3",
"AuraFlow": "aura-flow",
"HunyuanDiT": "hunyuan-dit",
"HunyuanDiT1": "hunyuan-dit",
"Flux": "flux",
"FluxInpaint": "flux",
"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",
"Flux2": "flux2",
}
gguf_architectures = {
"sd1": "sd15",
"qwen_image": "qwen-image",
}
@@ -35,10 +76,10 @@ class FakeTensor(NamedTuple):
return d
def inspect_checkpoint(filename):
def inspect_safetensors(filename: str, model_type: str, is_checkpoint: bool):
try:
# Read header of safetensors file
path = folder_paths.get_full_path("checkpoints", filename)
path = folder_paths.get_full_path(model_type, filename)
header = comfy.utils.safetensors_header(path)
if header:
cfg = json.loads(header.decode("utf-8"))
@@ -49,11 +90,14 @@ def inspect_checkpoint(filename):
cfg[key] = FakeTensor.from_dict(cfg[key])
# Reuse Comfy's model detection
unet_args = [cfg, "model.diffusion_model.", "F32"]
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
unet_config = model_detection.detect_unet_config(*unet_args[:-1])
unet_config = model_detection.detect_unet_config(cfg, prefix)
except TypeError as e: # older ComfyUI versions take 3 args
unet_config = model_detection.detect_unet_config(*unet_args)
raise TypeError(f"{e} when calling detect_unet_config - old version of ComfyUI?")
# Get input count to detect inpaint models
if input_block := cfg.get(input_block_name, None):
@@ -62,31 +106,295 @@ def inspect_checkpoint(filename):
input_count = 4
# Find a matching base model depending on unet config
base_model = model_detection.model_config_from_unet_config(unet_config)
if base_model is None:
base_model = 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 not raw_name:
return {"base_model": "unknown"}
base_model_class = base_model.__class__
base_model_name = model_names.get(base_model_class.__name__, "unknown")
return {
"base_model": base_model_name,
"is_inpaint": base_model_name in ["sd15", "sdxl"] and input_count > 4,
"is_refiner": base_model_class is supported_models.SDXLRefiner,
}
base_model_name = model_names.get(raw_name, "unknown")
result: dict[str, Any] = {"base_model": base_model_name}
result["is_inpaint"] = (
base_model_name in ["sd15", "sdxl"] and input_count > 4
) or raw_name == "FluxInpaint"
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"}
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}"}
if _server := getattr(server.PromptServer, "instance", None):
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
@_server.routes.get("/etn/model_info")
async def model_info(request):
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}
# Detect Chroma (modified Flux)
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
result = {
"base_model": gguf_architectures.get(arch_str, arch_str),
"is_inpaint": False,
}
try:
info = {
filename: inspect_checkpoint(filename)
for filename in folder_paths.get_filename_list("checkpoints")
}
return web.json_response(info)
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:
# traceback.print_exc()
return {"base_model": "unknown", "error": f"Failed to detect base model: {e}"}
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:
files = folder_paths.get_filename_list(model_type)
except KeyError:
return web.json_response({"error": f"Model folder not found: {model_type}"})
limit = int(params.get("limit", "1000"))
offset = int(params.get("offset", "0"))
files_range = files[offset : offset + limit]
is_checkpoint = model_type == "checkpoints"
info = {
filename: inspect_diffusion_model(filename, model_type, is_checkpoint)
for filename in files_range
}
if "limit" in params:
info["_meta"] = dict(offset=offset, count=len(files_range), total=len(files))
return web.json_response(info)
except Exception as e:
traceback.print_exc()
return web.json_response(dict(error=str(e)), status=500)
def has_invalid_folder_name(folder_name: str):
valid_names = list(folder_paths.folder_names_and_paths.keys())
if folder_name not in valid_names:
return web.json_response(
dict(error=f"Invalid folder path, must be one of {', '.join(valid_names)}"),
status=400,
)
return None
def has_invalid_filename(filename: str):
if not filename.lower().endswith((".sft", ".safetensors")):
return web.json_response(dict(error="File extension must be .safetensors"), status=400)
if not filename or not filename.strip() or len(filename) > 255:
return web.json_response(dict(error="Invalid filename"), status=400)
if any(char in filename for char in ["..", "/", "\\", "\n", "\r", "\t", "\0"]):
return web.json_response(dict(error="Invalid filename"), status=400)
if filename.startswith(".") or not re.match(r"^[a-zA-Z0-9_\-. ]+$", filename):
return web.json_response(dict(error="Invalid filename"), status=400)
return None
async def image_sender(data: bytes):
mem = memoryview(data)
csize = 2**14
for i in range(0, len(mem), csize):
yield mem[i : i + csize]
_server: server.PromptServer | None = getattr(server.PromptServer, "instance", None)
if _server is not None:
_workflow_exchange = WorkflowExchange(_server)
@_server.routes.get("/api/etn/model_info/{folder_name}")
async def model_info(request: web.Request):
folder_name = request.match_info.get("folder_name", "checkpoints")
error = has_invalid_folder_name(folder_name)
if error is not None:
return error
return inspect_models(folder_name, request.rel_url.query)
@_server.routes.get("/api/etn/model_info")
async def api_model_info(request):
return inspect_models("checkpoints", request.rel_url.query)
@_server.routes.get("/api/etn/languages")
async def languages(request):
try:
result = [dict(name=name, code=code) for code, name in available_languages()]
return web.json_response(result)
except Exception as e:
return web.json_response(dict(error=str(e)), status=500)
@_server.routes.get("/api/etn/translate/{lang}/{text}")
async def translate_text(request):
try:
language = request.match_info.get("lang", "en")
text = request.match_info.get("text", "")
result = translate(f"lang:{language} {text}")
return web.json_response(result)
except Exception as e:
return web.json_response(dict(error=str(e)), status=500)
@_server.routes.get("/api/etn/image/{id}")
async def get_image(request: web.Request):
try:
id = request.match_info.get("id", "")
data, content_type = image_cache.get(id)
if data is None or content_type is None:
return web.json_response(dict(error="Image not found"), status=404)
response = web.Response(
body=image_sender(data),
content_type=content_type,
headers={"Content-Length": str(len(data))},
)
return response
except Exception as e:
return web.json_response(dict(error=str(e)), status=500)
@_server.routes.put("/api/etn/image/{id}")
async def put_image(request: web.Request):
try:
id = request.match_info.get("id", "")
if id in image_cache:
return web.json_response(dict(status="cached"), status=200)
content_type = request.headers.get("Content-Type", "application/octet-stream")
data = bytearray()
async for chunk, _ in request.content.iter_chunks():
data.extend(chunk)
image_cache.insert(id, bytes(data), content_type)
return web.json_response(dict(status="success"), status=201)
except Exception as e:
return web.json_response(dict(error=str(e)), status=500)
@_server.routes.put("/api/etn/upload/{folder_name}/{filename}")
async def upload(request: web.Request):
folder_name = request.match_info.get("folder_name", "")
error = has_invalid_folder_name(folder_name)
if error is not None:
return error
filename = request.match_info.get("filename", "")
error = has_invalid_filename(filename)
if error is not None:
return error
try:
if folder_paths.get_full_path(folder_name, filename) is not None:
return web.json_response(dict(status="cached"), status=200)
folder = Path(folder_paths.folder_names_and_paths[folder_name][0][0])
total_size = int(request.headers.get("Content-Length", "0"))
logging.info(
f"Uploading {filename} ({total_size / (1024**2):.1f} MB) to {folder} folder"
)
with open(folder / filename, "wb") as f:
async for chunk, _ in request.content.iter_chunks():
f.write(chunk)
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 _handle_workflow_request(request: web.Request, handler, *arg_keys):
try:
data = await request.json()
args = [data[key] for key in arg_keys]
await handler(*args)
return web.json_response(dict(status="success"), status=200)
except KeyError as e:
return web.json_response(dict(error=str(e)), status=400)
except Exception as e:
return web.json_response(dict(error=str(e)), status=500)
@_server.routes.post("/api/etn/workflow/publish")
async def publish_workflow(request: web.Request):
return await _handle_workflow_request(
request, _workflow_exchange.publish, "name", "client_id", "workflow"
)
@_server.routes.post("/api/etn/workflow/subscribe")
async def subscribe_workflow(request: web.Request):
return await _handle_workflow_request(request, _workflow_exchange.subscribe, "client_id")
@_server.routes.post("/api/etn/workflow/unsubscribe")
async def unsubscribe_workflow(request: web.Request):
return await _handle_workflow_request(request, _workflow_exchange.unsubscribe, "client_id")
Binary file not shown.

After

Width:  |  Height:  |  Size: 26 KiB

+235
View File
@@ -0,0 +1,235 @@
import { app } from "/scripts/app.js"
import { api } from "/scripts/api.js"
(function() {
// Workflow publishing
// - done whenever the graph changes, as long as there is a KritaOutput node
let publisherRegistered = false
async function publishWorkflow(e) {
const prompt = await app.graphToPrompt()
await api.fetchApi("/api/etn/workflow/publish", {
method: "POST",
body: JSON.stringify({
name: "ComfyUI Web",
client_id: api.clientId,
workflow: prompt["output"]
}, null, 2)
})
}
// Image background for nodes
// - this is just for visuals
function loadImage(base64) {
const image = new Image()
image.src = base64
// image.onerror = () => console.error("Failed to load image");
return image
}
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=")
function setIconImage(nodeType, image, size, padRows, padCols) {
const onAdded = nodeType.prototype.onAdded
nodeType.prototype.onAdded = function () {
onAdded?.apply(this, arguments)
this.size = size
}
const onDrawBackground = nodeType.prototype.onDrawBackground
nodeType.prototype.onDrawBackground = function(ctx) {
onDrawBackground?.apply(this, arguments)
const pad = [padCols * 20, LiteGraph.NODE_SLOT_HEIGHT * padRows + 8];
if(this.flags.collapsed || pad[1] + 32 > this.size[1] || image.width === 0) {
return
}
const avail = [this.size[0] - pad[0], this.size[1] - pad[1]]
const scale = Math.min(1.0, avail[0] / image.width, avail[1] / image.height)
const size = [Math.floor(image.width * scale), Math.floor(image.height * scale)]
const offset = [Math.max(0, (avail[0] - size[0]) / 2), Math.max(0, (avail[1] - size[1]) / 2)]
ctx.drawImage(image, offset[0], pad[1] + offset[1], size[0], size[1])
}
}
// Parameter node
// - represents a customizable parameter that should be exposed in external tools
// - adapts to whichever node it is connected to, similar to the built-in "Primitive" node
// - can only be connected to slots which are converted widgets
const replaceableWidgets = ["INT", "FLOAT", "BOOLEAN", "STRING", "COMBO", "INT:seed"]
const parameterTypes = {
"combo": ["choice"],
"number": ["number", "number (integer)"],
"toggle": ["toggle"],
"text": ["text", "prompt (positive)", "prompt (negative)"],
}
function defaultParameterType(widgetType, connectedNode, connectedWidget) {
let paramType = parameterTypes[widgetType][0]
if (connectedNode.comfyClass === "CLIPTextEncode") {
paramType = "prompt (positive)"
}
const round = connectedWidget.options?.round
if ((paramType == "number" && round === undefined) || round === 1) {
paramType = "number (integer)"
}
return paramType
}
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") {
type = "text"
}
const options = connectedWidget.options
const parameterTypeHint = node.widgets[1].value
const notSpecialized = node.widgets[1].options.values.includes("auto")
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") {
node.addWidget("number", "min", oldMin, null, options)
node.addWidget("number", "max", oldMax, null, options)
}
}
function adaptWidgetsToConnection(node) {
if (!node.outputs || node.outputs.length === 0) {
return
}
const links = node.outputs[0].links
if (links && links.length === 1) {
const link = node.graph.links[links[0]]
if (!link) return
const theirNode = node.graph.getNodeById(link.target_id)
if (!theirNode || !theirNode.inputs) return
const input = theirNode.inputs[link.target_slot]
if (!input || !input.widget || theirNode.widgets === undefined) return
node.outputs[0].type = input.type
if (node.widgets[0].value === "Parameter") {
node.widgets[0].value = input.name
}
const widgetName = input.widget.name
const theirWidget = theirNode.widgets.find((w) => w.name === widgetName)
if (!theirWidget) return // connected to a custom node that isn't installed
const widgetType = theirWidget.origType ?? theirWidget.type
changeWidgets(node, widgetType, theirNode, theirWidget)
} else if (!links || links.length === 0) {
node.outputs[0].type = "*"
node.widgets[1].value = "auto"
node.widgets[1].options = {values: ["auto"]}
}
}
function setupParameterNode(nodeType) {
const onAdded = nodeType.prototype.onAdded
nodeType.prototype.onAdded = function() {
onAdded?.apply(this, arguments)
adaptWidgetsToConnection(this)
}
const onAfterGraphConfigured = nodeType.prototype.onAfterGraphConfigured
nodeType.prototype.onAfterGraphConfigured = function() {
onAfterGraphConfigured?.apply(this, arguments)
adaptWidgetsToConnection(this)
}
const onConnectOutput = nodeType.prototype.onConnectOutput
nodeType.prototype.onConnectOutput = function(slot, type, input, target_node, target_slot) {
if (!input.widget && !(input.type in replaceableWidgets)) {
return false
} else if (onConnectOutput) {
result = onConnectOutput.apply(this, arguments)
return result
}
return true
}
const onConnectionsChange = nodeType.prototype.onConnectionsChange
nodeType.prototype.onConnectionsChange = function(_, index, connected) {
if (!app.configuringGraph) {
adaptWidgetsToConnection(this)
}
onConnectionsChange?.apply(this, arguments)
}
}
// Register the extension
app.registerExtension({
name: "external_tooling_nodes",
beforeRegisterNodeDef(nodeType /*typeof LGraphNode*/, nodeData /*ComfyObjectInfo*/, app) {
if (nodeData.name === "ETN_KritaCanvas") {
setIconImage(nodeType, canvasIcon, [200, 100], 0, 2)
} else if (nodeData.name === "ETN_Parameter") {
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 = "*"
}
}
},
nodeCreated(node /*ComfyNode*/, app) {
if (publisherRegistered || node.comfyClass !== "ETN_KritaOutput") {
return
}
api.addEventListener('graphChanged', publishWorkflow)
publisherRegistered = true
},
setup(app) {
if (publisherRegistered) {
publishWorkflow(null)
}
},
});
})();
+330
View File
@@ -0,0 +1,330 @@
import sys
import torch
import numpy as np
from enum import Enum
from pathlib import Path
from typing import Any, NamedTuple
from PIL import Image
import server
import comfy.samplers
from comfy.comfy_types.node_typing import IO
from comfy_api.latest import io
from .nodes import SendImageWebSocket
class Publisher(NamedTuple):
name: str
id: str
workflow: dict
class WorkflowExchange:
def __init__(self, server: server.PromptServer):
self._server = server
self._publishers: dict[str, Publisher] = {}
self._subscribers: list[str] = []
async def publish(self, publisher_name: str, publisher_id: str, workflow: dict):
publisher = Publisher(publisher_name, publisher_id, workflow)
for client_id in self._subscribers:
await self._notify(client_id, publisher)
self._publishers[publisher_id] = publisher
async def subscribe(self, client_id: str):
if client_id in self._subscribers:
raise KeyError("Already subscribed")
self._subscribers.append(client_id)
for publisher in self._publishers.values():
await self._notify(client_id, publisher)
async def unsubscribe(self, client_id: str):
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):
data = {
"publisher": {"name": publisher.name, "id": publisher.id},
"workflow": publisher.workflow,
}
await self._server.send_json("etn_workflow_published", data, client_id)
def _placeholder_image():
path = Path(__file__).parent / "data" / "external-image-placeholder.webp"
image = Image.open(path).convert("RGB")
image = np.array(image).astype(np.float32) / 255.0
return torch.from_numpy(image)[None,]
class _BasicTypes(str):
"""Matches IO.PRIMITIVE, but also any list of choices"""
basic_types = IO.PRIMITIVE.split(",") # STRING, FLOAT, INT, BOOLEAN
def __eq__(self, other):
return other in self.basic_types or isinstance(other, (list, _BasicTypes))
def __ne__(self, other):
return not self.__eq__(other)
BasicTypes = _BasicTypes("BASIC")
class OutputBatchMode(Enum):
default = "default"
images = "images"
animation = "animation"
layers = "layers"
class KritaOutput(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_KritaOutput",
display_name="Krita Output",
category="krita",
inputs=[
io.Image.Input("images"),
io.Int.Input("x", "offset x", default=0),
io.Int.Input("y", "offset y", default=0),
io.String.Input("name", default=""),
io.Combo.Input(
"batch_mode", OutputBatchMode, "batch mode", default=OutputBatchMode.default
),
io.Boolean.Input("resize_canvas", "resize canvas", default=False),
],
is_output_node=True,
)
@classmethod
def execute(
cls,
images: torch.Tensor,
x: int = 0,
y: int = 0,
name="",
batch_mode: OutputBatchMode | str = OutputBatchMode.default,
resize_canvas=False,
):
batch_mode = batch_mode.value if isinstance(batch_mode, OutputBatchMode) else batch_mode
info = {
"name": name,
"offset_x": x,
"offset_y": y,
"batch_mode": batch_mode,
"resize_canvas": resize_canvas,
}
output = SendImageWebSocket.execute(images, "PNG")
assert isinstance(output.ui, dict)
output.ui["info"] = [info]
return output
class KritaSendText(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_KritaSendText",
display_name="Send Text",
category="krita",
inputs=[
io.AnyType.Input("value"),
io.String.Input("name", default="Output"),
io.Combo.Input("type", options=["text", "markdown", "html"], default="text"),
],
is_output_node=True,
)
@classmethod
def execute(cls, value: Any, name: str, type: str):
mime = {
"text": "text/plain",
"markdown": "text/markdown",
"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}"
return io.NodeOutput(ui={"text": [{"name": name, "text": text, "content-type": mime}]})
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"),
],
)
@classmethod
def execute(cls):
return io.NodeOutput(_placeholder_image(), 512, 512, 0)
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):
return io.NodeOutput(_placeholder_image(), torch.ones(1, 512, 512))
class KritaMaskLayer(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_KritaMaskLayer",
display_name="Krita Mask Layer",
category="krita",
inputs=[io.String.Input("name", default="Mask")],
outputs=[
io.Mask.Output(display_name="mask"),
],
)
@classmethod
def execute(cls, name: str):
return io.NodeOutput(torch.ones(1, 512, 512))
_param_types = [
"auto",
"number",
"number (integer)",
"toggle",
"choice",
"text",
"prompt (positive)",
"prompt (negative)",
]
_fmax = sys.float_info.max
class Parameter(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_Parameter",
display_name="Parameter",
category="krita",
inputs=[
io.String.Input("name", default="Parameter"),
io.Combo.Input("type", options=_param_types, default="auto"),
io.String.Input("default", default=""),
io.Float.Input("min", default=-1e10, min=-_fmax, max=_fmax, optional=True),
io.Float.Input("max", default=1e10, min=-_fmax, max=_fmax, optional=True),
],
outputs=[io.AnyType.Output(display_name="value")],
)
@classmethod
def execute(cls, name: str, type: str, default, min=0.0, max=1.0):
if type == "number":
return io.NodeOutput(float(default))
elif type == "number (integer)":
return io.NodeOutput(int(default))
return io.NodeOutput(default)
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):
raise NotImplementedError("This workflow must be started from Krita!")
+360 -92
View File
@@ -1,30 +1,44 @@
from __future__ import annotations
from copy import copy
from dataclasses import dataclass
import time
from typing import NamedTuple
from uuid import uuid4
from PIL import Image
import numpy as np
import base64
import torch
import torch.nn.functional as F
from io import BytesIO
from server import PromptServer, BinaryEventTypes
from comfy.clip_vision import ClipVisionModel
from comfy.sd import StyleModel
from comfy_api.latest import io
class LoadImageBase64:
class LoadImageBase64(io.ComfyNode):
@classmethod
def INPUT_TYPES(s):
return {"required": {"image": ("STRING", {"multiline": False})}}
def define_schema(cls):
return io.Schema(
node_id="ETN_LoadImageBase64",
display_name="Load Image (Base64)",
category="external_tooling",
inputs=[io.String.Input("image", multiline=False)],
outputs=[io.Image.Output(display_name="image"), io.Mask.Output(display_name="mask")],
)
RETURN_TYPES = ("IMAGE", "MASK")
CATEGORY = "external_tooling"
FUNCTION = "load_image"
def load_image(self, image):
@classmethod
def execute(cls, image: str):
_strip_prefix(image, "data:image/png;base64,")
imgdata = base64.b64decode(image)
img = Image.open(BytesIO(imgdata))
if "A" in img.getbands():
mask = np.array(img.getchannel("A")).astype(np.float32) / 255.0
mask = 1.0 - torch.from_numpy(mask)
mask = torch.from_numpy(mask)
else:
mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu")
mask = None
img = img.convert("RGB")
img = np.array(img).astype(np.float32) / 255.0
@@ -33,16 +47,20 @@ class LoadImageBase64:
return (img, mask)
class LoadMaskBase64:
class LoadMaskBase64(io.ComfyNode):
@classmethod
def INPUT_TYPES(s):
return {"required": {"mask": ("STRING", {"multiline": False})}}
def define_schema(cls):
return io.Schema(
node_id="ETN_LoadMaskBase64",
display_name="Load Mask (Base64)",
category="external_tooling",
inputs=[io.String.Input("mask", multiline=False)],
outputs=[io.Mask.Output(display_name="mask")],
)
RETURN_TYPES = ("MASK",)
CATEGORY = "external_tooling"
FUNCTION = "load_mask"
def load_mask(self, mask):
@classmethod
def execute(cls, mask: str):
_strip_prefix(mask, "data:image/png;base64,")
imgdata = base64.b64decode(mask)
img = Image.open(BytesIO(imgdata))
img = np.array(img).astype(np.float32) / 255.0
@@ -52,17 +70,22 @@ class LoadMaskBase64:
return (img.unsqueeze(0),)
class SendImageWebSocket:
class SendImageWebSocket(io.ComfyNode):
@classmethod
def INPUT_TYPES(s):
return {"required": {"images": ("IMAGE",)}}
def define_schema(cls):
return io.Schema(
node_id="ETN_SendImageWebSocket",
display_name="Send Image (WebSocket)",
category="external_tooling",
inputs=[
io.Image.Input("images"),
io.Combo.Input("format", options=["PNG", "JPEG"], default="PNG"),
],
is_output_node=True,
)
RETURN_TYPES = ()
FUNCTION = "send_images"
OUTPUT_NODE = True
CATEGORY = "external_tooling"
def send_images(self, images):
@classmethod
def execute(cls, images: torch.Tensor, format: str):
results = []
for tensor in images:
array = 255.0 * tensor.cpu().numpy()
@@ -71,91 +94,336 @@ class SendImageWebSocket:
server = PromptServer.instance
server.send_sync(
BinaryEventTypes.UNENCODED_PREVIEW_IMAGE,
["PNG", image, None],
[format, image, None],
server.client_id,
)
results.append(
# Could put some kind of ID here, but for now just match them by index
{"source": "websocket", "content-type": "image/png", "type": "output"}
)
results.append({
"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
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"x": (
"INT",
{"default": 0, "min": 0, "max": 8192, "step": 1},
),
"y": (
"INT",
{"default": 0, "min": 0, "max": 8192, "step": 1},
),
"width": (
"INT",
{"default": 512, "min": 1, "max": 8192, "step": 1},
),
"height": (
"INT",
{"default": 512, "min": 1, "max": 8192, "step": 1},
),
}
}
def execute(cls, id: str):
image_data, content_type = image_cache.get(id, extend=True)
if image_data is None:
raise ValueError(f"Image with ID {id} not found in cache.")
CATEGORY = "external_tooling"
RETURN_TYPES = ("IMAGE",)
FUNCTION = "crop"
img = Image.open(BytesIO(image_data))
w, h = img.size
c = len(img.getbands())
normalized = np.array(img).astype(np.float32) / 255.0
tensor = torch.from_numpy(normalized).reshape(1, h, w, c)
match c:
case 1:
image = tensor.expand(1, h, w, 3)
mask = tensor.reshape(1, h, w)
case 3:
image = tensor
mask = tensor[..., 0]
case 4:
image = tensor[..., :3]
mask = tensor[..., 3]
def crop(self, image, x, y, width, height):
out = image[:, y : y + height, x : x + width, :]
return (out,)
return io.NodeOutput(image, mask)
class ApplyMaskToImage:
class SaveImageCache(io.ComfyNode):
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"mask": ("MASK",),
}
}
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,
)
CATEGORY = "external_tooling"
RETURN_TYPES = ("IMAGE",)
FUNCTION = "apply_mask"
@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)
def apply_mask(self, image: torch.Tensor, mask: torch.Tensor):
# Move the channel to the second dimension for processing
out = image.movedim(-1, 1)
results.append({
"source": "http",
"id": key,
"content-type": f"image/{format.lower()}",
"type": "output",
})
return io.NodeOutput(ui={"images": results})
# Check if the images are RGB, and if so, add an alpha channel initialized to 1
def to_bchw(image: torch.Tensor):
if image.ndim == 3:
image = image.unsqueeze(0)
return image.movedim(-1, 1)
def to_bhwc(image: torch.Tensor):
return image.movedim(1, -1)
def mask_batch(mask: torch.Tensor):
if mask.ndim == 2:
mask = mask.unsqueeze(0)
return mask
class ApplyMaskToImage(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_ApplyMaskToImage",
display_name="Apply Mask to Image",
category="external_tooling",
inputs=[
io.Image.Input("image"),
io.Mask.Input("mask"),
],
outputs=[io.Image.Output(display_name="masked")],
)
@classmethod
def execute(cls, image: torch.Tensor, mask: torch.Tensor):
out = to_bchw(image)
if out.shape[1] == 3: # Assuming RGB images
out = torch.cat([out, torch.ones_like(out[:, :1, :, :])], dim=1)
# Ensure masks are unsqueezed to match the alpha channel dimension if needed
if mask.ndim == 2:
mask = mask.unsqueeze(0) # Add a batch dimension to masks
# For single mask, expand it to match size of image batch size.
if mask.shape[0] == 1:
mask = mask.repeat(out.shape[0], 1, 1)
mask = mask_batch(mask)
assert mask.ndim == 3, f"Mask should have shape [B, H, W]. {mask.shape}"
assert out.ndim == 4, f"Image should have shsape [B, C, H, W]. {out.shape}"
assert out.shape[-2:] == mask.shape[-2:], f"{out.shape[-2:]} != {mask.shape[-2:]}"
assert out.shape[0] == mask.shape[0], f"{out.shape[0]} != {mask.shape[0]}"
assert out.ndim == 4, f"Image should have shape [B, C, H, W]. {out.shape}"
assert out.shape[-2:] == 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]
# Apply each mask in the batch to its corresponding image's alpha channel
for i in range(out.shape[0]):
out[i, 3, :, :] = mask[i]
alpha = mask[i] if is_mask_batch else mask[0]
out[i, 3, :, :] = alpha
# Move the channel back to its original dimension
out = out.movedim(1, -1)
return (to_bhwc(out),)
return (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
+142
View File
@@ -0,0 +1,142 @@
from __future__ import annotations
from weakref import ref as WeakRef
from pathlib import Path
from tqdm import tqdm
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch import Tensor
from transformers import CLIPImageProcessor, CLIPConfig, CLIPVisionModel, PreTrainedModel
from kornia.filters import box_blur
from comfy_api.latest import io
from .nodes import to_bchw, to_bhwc
def cosine_similarity(image_embeds: Tensor, text_embeds: Tensor):
if image_embeds.dim() == 2 and text_embeds.dim() == 2:
image_embeds = image_embeds.unsqueeze(1)
return F.cosine_similarity(image_embeds, text_embeds, dim=-1)
class CLIPSafetyChecker(PreTrainedModel):
# https://huggingface.co/CompVis/stable-diffusion-safety-checker
# Adapted from:
# https://github.com/huggingface/diffusers/blob/main/src/diffusers/pipelines/stable_diffusion/safety_checker.py
config_class = CLIPConfig
_no_split_modules = ["CLIPEncoderLayer"]
def __init__(self, config: CLIPConfig):
super().__init__(config)
projdim = config.projection_dim
self.vision_model = CLIPVisionModel(config.vision_config)
self.visual_projection = nn.Linear(config.vision_config.hidden_size, projdim, bias=False)
self.concept_embeds = nn.Parameter(torch.ones(17, projdim), requires_grad=False)
self.special_care_embeds = nn.Parameter(torch.ones(3, projdim), 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)
def forward(self, clip_input, images: Tensor, sensitivity: float):
with torch.no_grad():
image_batch = self.vision_model(clip_input)[1]
image_embeds = self.visual_projection(image_batch)
sensitivity = -0.1 + 0.14 * sensitivity
special_cos_dist = cosine_similarity(image_embeds, self.special_care_embeds)
special_scores_threshold = self.special_care_embeds_weights.unsqueeze(0)
special_scores = special_cos_dist - special_scores_threshold + sensitivity
if torch.any(special_scores > 0):
sensitivity = sensitivity + 0.01
cos_dist = cosine_similarity(image_embeds, self.concept_embeds)
concept_threshold = self.concept_embeds_weights.unsqueeze(0)
concept_scores = cos_dist - concept_threshold + sensitivity
is_nsfw = [torch.any(concept_scores[i] > 0) for i in range(concept_scores.shape[0])]
is_nsfw = [x.item() for x in is_nsfw]
return self.filter_images(images, is_nsfw)
def filter_images(self, images: Tensor, is_nsfw: list[bool]):
if not any(is_nsfw):
return images
images = images.clone()
images_to_filter = (i for i, nsfw in enumerate(is_nsfw) if nsfw)
orig_size = images.shape[-2:]
for idx in images_to_filter:
filtered = images[idx].unsqueeze(0)
filtered = F.interpolate(filtered, size=64, mode="nearest")
filtered = box_blur(filtered, 11, separable=True)
filtered = F.interpolate(filtered, size=orig_size, mode="bilinear")
images[idx] = filtered.squeeze(0)
return images
class CachedModels:
_instance: CachedModels | None = None
def __init__(self):
model_dir = Path(__file__).parent / "safetychecker"
model_file = model_dir / "model.safetensors"
if not model_file.exists():
self.download(
"https://huggingface.co/CompVis/stable-diffusion-safety-checker/resolve/refs%2Fpr%2F41/model.safetensors",
target=model_file,
)
self.feature_extractor = CLIPImageProcessor.from_pretrained(model_dir)
self.safety_checker = CLIPSafetyChecker.from_pretrained(model_dir)
@classmethod
def load(cls):
if cls._instance is None:
cls._instance = CachedModels()
return cls._instance
def download(self, url: str, target: Path):
import requests
try:
target_temp = target.with_suffix(".download")
with requests.get(url, stream=True) as response:
text = "NSFWFilter model download"
total = int(response.headers.get("content-length", 0))
pbar = tqdm(None, total=total, unit="b", unit_scale=True, desc=text)
with open(target_temp, "wb") as f:
for chunk in response.iter_content(chunk_size=8192):
f.write(chunk)
pbar.update(len(chunk))
pbar.close()
target_temp.rename(target)
except Exception as e:
raise RuntimeError(
f"NSFWFilter: Failed to download safety-checker model from {url} to target location {target}: {e}"
) from e
class NSFWFilter(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_NSFWFilter",
display_name="NSFW Filter",
category="external_tooling",
inputs=[
io.Image.Input("image"),
io.Float.Input("sensitivity", default=0.5, min=0.0, max=1.0, step=0.1),
],
outputs=[io.Image.Output(display_name="image")],
)
@classmethod
def execute(cls, image: Tensor, sensitivity: float):
models = CachedModels.load()
image = to_bchw(image)
input = models.feature_extractor(image, do_rescale=False, return_tensors="pt")
filtered = models.safety_checker(
images=image, clip_input=input.pixel_values, sensitivity=sensitivity
)
return io.NodeOutput(to_bhwc(filtered))
+10 -2
View File
@@ -1,12 +1,20 @@
[project]
name = "comfyui-tooling-nodes"
description = "Provides nodes and server API extensions geared towards using ComfyUI as a backend for external tools."
version = "1.0.0"
license = "LICENSE"
version = "3.1.0"
license = { file = "LICENSE" }
[project.urls]
Repository = "https://github.com/Acly/comfyui-tooling-nodes"
[tool.ruff]
target-version = "py311"
line-length = 100
preview = true
[tool.ruff.lint]
ignore = ["E741"]
[tool.black]
line-length = 100
preview = true
+91 -78
View File
@@ -8,6 +8,7 @@ import torch.nn.functional as F
import math
from torch import Tensor, Size
from comfy.model_patcher import ModelPatcher
from comfy_api.latest import io
def downsample_mask(mask: Tensor, batch: int, target_size: int, original_shape: Size) -> Tensor:
@@ -65,103 +66,112 @@ class Region(NamedTuple):
return result
class BackgroundRegion:
Regions = io.Custom("Regions")
class BackgroundRegion(io.ComfyNode):
@classmethod
def INPUT_TYPES(cls):
return {"required": {"conditioning": ("CONDITIONING",)}}
def define_schema(cls):
return io.Schema(
node_id="ETN_BackgroundRegion",
display_name="Background Region",
category="external_tooling/regions",
inputs=[io.Conditioning.Input("conditioning")],
outputs=[Regions.Output(display_name="regions")],
)
CATEGORY = "external_tooling/regions"
RETURN_TYPES = ("REGIONS",)
FUNCTION = "define"
def define(self, conditioning: list):
@classmethod
def execute(cls, conditioning: list):
return (Region(None, None, conditioning),)
class DefineRegion:
class DefineRegion(io.ComfyNode):
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"mask": ("MASK",),
"conditioning": ("CONDITIONING",),
},
"optional": {
"regions": ("REGIONS",),
},
}
CATEGORY = "external_tooling/regions"
RETURN_TYPES = ("REGIONS",)
FUNCTION = "define"
def define(self, mask: Tensor, conditioning: list, regions: Region | None = None):
return (Region(regions, mask, conditioning),)
class ListRegionMasks:
@classmethod
def INPUT_TYPES(cls):
return {"required": {"regions": ("REGIONS",)}}
CATEGORY = "external_tooling/regions"
RETURN_TYPES = ("MASK",)
FUNCTION = "get_masks"
def get_masks(self, regions: Region):
return (torch.stack([r.mask for r in regions.preprocess()], dim=0),)
class AttentionMask:
def define_schema(cls):
return io.Schema(
node_id="ETN_DefineRegion",
display_name="Define Region",
category="external_tooling/regions",
inputs=[
io.Mask.Input("mask"),
io.Conditioning.Input("conditioning"),
Regions.Input("regions", optional=True),
],
outputs=[Regions.Output(display_name="regions")],
)
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("MODEL",),
"regions": ("REGIONS",),
}
}
def execute(cls, mask: Tensor, conditioning: list, regions: Region | None = None):
if mask.dim() < 3:
mask = mask.unsqueeze(0)
return io.NodeOutput(Region(regions, mask, conditioning))
RETURN_TYPES = ("MODEL",)
FUNCTION = "attention_mask"
CATEGORY = "external_tooling/regions"
mask: Tensor
conds: list[Tensor]
batch_size: int
class ListRegionMasks(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_ListRegionMasks",
display_name="List Region Masks",
category="external_tooling/regions",
inputs=[Regions.Input("regions")],
outputs=[io.Mask.Output(display_name="masks")],
)
def attention_mask(self, model: ModelPatcher, regions: Region):
new_model = model.clone()
region_list = regions.preprocess()
num_conds = len(region_list)
@classmethod
def execute(cls, regions: Region):
return io.NodeOutput(torch.stack([r.mask for r in regions.preprocess()], dim=0))
class AttentionMask(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_AttentionMask",
display_name="Regions Attention Mask",
category="external_tooling/regions",
inputs=[io.Model.Input("model"), Regions.Input("regions")],
outputs=[io.Model.Output(display_name="model")],
)
@classmethod
def execute(cls, model: ModelPatcher, regions: Region):
return io.NodeOutput(AttentionMaskPatch.apply(model, regions))
class AttentionMaskPatch:
def __init__(self, region_list: list[Region]):
mask = torch.stack([r.mask for r in region_list], dim=0)
mask_sum = mask.sum(dim=0, keepdim=True)
assert mask_sum.sum() > 0, "There are areas that are zero in all masks."
self.mask = mask / mask_sum
self.conds = [r.conditioning[0][0] for r in region_list]
num_tokens = [cond.shape[1] for cond in self.conds]
self.num_tokens = [cond.shape[1] for cond in self.conds]
self.num_conds = len(region_list)
self.batch_size = 0
@staticmethod
def apply(model: ModelPatcher, regions: Region):
patch = AttentionMaskPatch(regions.preprocess())
def attn2_patch(q: Tensor, k: Tensor, v: Tensor, extra_options: dict):
assert k.mean() == v.mean(), "k and v must be the same."
device, dtype = q.device, q.dtype
if self.conds[0].device != device:
self.conds = [cond.to(device, dtype=dtype) for cond in self.conds]
if self.mask.device != device:
self.mask = self.mask.to(device, dtype=dtype)
if patch.conds[0].device != device or patch.conds[0].dtype != dtype:
patch.conds = [cond.to(device, dtype=dtype) for cond in patch.conds]
if patch.mask.device != device or patch.mask.dtype != dtype:
patch.mask = patch.mask.to(device, dtype=dtype)
cond_or_unconds = extra_options["cond_or_uncond"]
num_chunks = len(cond_or_unconds)
self.batch_size = q.shape[0] // num_chunks
patch.batch_size = q.shape[0] // num_chunks
q_chunks = q.chunk(num_chunks, dim=0)
k_chunks = k.chunk(num_chunks, dim=0)
lcm_tokens = lcm_for_list(num_tokens + [k.shape[1]])
lcm_tokens = lcm_for_list(patch.num_tokens + [k.shape[1]])
conds_tensor = [
cond.repeat(self.batch_size, lcm_tokens // num_tokens[i], 1)
for i, cond in enumerate(self.conds)
cond.repeat(patch.batch_size, lcm_tokens // patch.num_tokens[i], 1)
for i, cond in enumerate(patch.conds)
]
conds_tensor = torch.cat(conds_tensor, dim=0)
@@ -172,9 +182,9 @@ class AttentionMask:
qs.insert(0, q_chunks[i])
ks.insert(0, k_target)
else:
qs.insert(0, q_chunks[i].repeat(num_conds, 1, 1))
qs.insert(0, q_chunks[i].repeat(patch.num_conds, 1, 1))
ks.insert(0, conds_tensor)
for _ in range(num_conds - 1):
for _ in range(patch.num_conds - 1):
cond_or_unconds.insert(i, 0)
qs = torch.cat(qs, dim=0)
@@ -182,29 +192,32 @@ class AttentionMask:
return qs, ks, ks
def attn2_output_patch(out: Tensor, extra_options: dict):
num_conds = patch.num_conds
cond_or_unconds = extra_options["cond_or_uncond"]
mask_downsample = downsample_mask(
self.mask, self.batch_size, out.shape[1], extra_options["original_shape"]
patch.mask, patch.batch_size, out.shape[1], extra_options["original_shape"]
)
outputs: list[Tensor] = []
pos = 0
i = 0
while i < len(cond_or_unconds):
if cond_or_unconds[i] == 1: # uncond
outputs.append(out[pos : pos + self.batch_size])
pos += self.batch_size
outputs.append(out[pos : pos + patch.batch_size])
pos += patch.batch_size
else:
masked = out[pos : pos + num_conds * self.batch_size] * mask_downsample
masked = masked.view(num_conds, self.batch_size, out.shape[1], out.shape[2])
masked = out[pos : pos + num_conds * patch.batch_size] * mask_downsample
masked = masked.view(num_conds, patch.batch_size, out.shape[1], out.shape[2])
masked = masked.sum(dim=0)
outputs.append(masked)
pos += num_conds * self.batch_size
pos += num_conds * patch.batch_size
for _ in range(num_conds - 1):
cond_or_unconds.pop(i)
i += 1
return torch.cat(outputs, dim=0)
new_model = model.clone()
new_model.set_model_attn2_patch(attn2_patch)
new_model.set_model_attn2_output_patch(attn2_output_patch)
return (new_model,)
new_model.set_attachments("etn_attention_mask", patch)
return new_model
+2
View File
@@ -0,0 +1,2 @@
# Optional, only required for Translate node:
argostranslate
+171
View File
@@ -0,0 +1,171 @@
{
"_name_or_path": "clip-vit-large-patch14/",
"architectures": [
"SafetyChecker"
],
"initializer_factor": 1.0,
"logit_scale_init_value": 2.6592,
"model_type": "clip",
"projection_dim": 768,
"text_config": {
"_name_or_path": "",
"add_cross_attention": false,
"architectures": null,
"attention_dropout": 0.0,
"bad_words_ids": null,
"bos_token_id": 0,
"chunk_size_feed_forward": 0,
"cross_attention_hidden_size": null,
"decoder_start_token_id": null,
"diversity_penalty": 0.0,
"do_sample": false,
"dropout": 0.0,
"early_stopping": false,
"encoder_no_repeat_ngram_size": 0,
"eos_token_id": 2,
"exponential_decay_length_penalty": null,
"finetuning_task": null,
"forced_bos_token_id": null,
"forced_eos_token_id": null,
"hidden_act": "quick_gelu",
"hidden_size": 768,
"id2label": {
"0": "LABEL_0",
"1": "LABEL_1"
},
"initializer_factor": 1.0,
"initializer_range": 0.02,
"intermediate_size": 3072,
"is_decoder": false,
"is_encoder_decoder": false,
"label2id": {
"LABEL_0": 0,
"LABEL_1": 1
},
"layer_norm_eps": 1e-05,
"length_penalty": 1.0,
"max_length": 20,
"max_position_embeddings": 77,
"min_length": 0,
"model_type": "clip_text_model",
"no_repeat_ngram_size": 0,
"num_attention_heads": 12,
"num_beam_groups": 1,
"num_beams": 1,
"num_hidden_layers": 12,
"num_return_sequences": 1,
"output_attentions": false,
"output_hidden_states": false,
"output_scores": false,
"pad_token_id": 1,
"prefix": null,
"problem_type": null,
"pruned_heads": {},
"remove_invalid_values": false,
"repetition_penalty": 1.0,
"return_dict": true,
"return_dict_in_generate": false,
"sep_token_id": null,
"task_specific_params": null,
"temperature": 1.0,
"tie_encoder_decoder": false,
"tie_word_embeddings": true,
"tokenizer_class": null,
"top_k": 50,
"top_p": 1.0,
"torch_dtype": null,
"torchscript": false,
"transformers_version": "4.21.0.dev0",
"typical_p": 1.0,
"use_bfloat16": false,
"vocab_size": 49408
},
"text_config_dict": {
"hidden_size": 768,
"intermediate_size": 3072,
"num_attention_heads": 12,
"num_hidden_layers": 12
},
"torch_dtype": "float32",
"transformers_version": null,
"vision_config": {
"_name_or_path": "",
"add_cross_attention": false,
"architectures": null,
"attention_dropout": 0.0,
"bad_words_ids": null,
"bos_token_id": null,
"chunk_size_feed_forward": 0,
"cross_attention_hidden_size": null,
"decoder_start_token_id": null,
"diversity_penalty": 0.0,
"do_sample": false,
"dropout": 0.0,
"early_stopping": false,
"encoder_no_repeat_ngram_size": 0,
"eos_token_id": null,
"exponential_decay_length_penalty": null,
"finetuning_task": null,
"forced_bos_token_id": null,
"forced_eos_token_id": null,
"hidden_act": "quick_gelu",
"hidden_size": 1024,
"id2label": {
"0": "LABEL_0",
"1": "LABEL_1"
},
"image_size": 224,
"initializer_factor": 1.0,
"initializer_range": 0.02,
"intermediate_size": 4096,
"is_decoder": false,
"is_encoder_decoder": false,
"label2id": {
"LABEL_0": 0,
"LABEL_1": 1
},
"layer_norm_eps": 1e-05,
"length_penalty": 1.0,
"max_length": 20,
"min_length": 0,
"model_type": "clip_vision_model",
"no_repeat_ngram_size": 0,
"num_attention_heads": 16,
"num_beam_groups": 1,
"num_beams": 1,
"num_hidden_layers": 24,
"num_return_sequences": 1,
"output_attentions": false,
"output_hidden_states": false,
"output_scores": false,
"pad_token_id": null,
"patch_size": 14,
"prefix": null,
"problem_type": null,
"pruned_heads": {},
"remove_invalid_values": false,
"repetition_penalty": 1.0,
"return_dict": true,
"return_dict_in_generate": false,
"sep_token_id": null,
"task_specific_params": null,
"temperature": 1.0,
"tie_encoder_decoder": false,
"tie_word_embeddings": true,
"tokenizer_class": null,
"top_k": 50,
"top_p": 1.0,
"torch_dtype": null,
"torchscript": false,
"transformers_version": "4.21.0.dev0",
"typical_p": 1.0,
"use_bfloat16": false
},
"vision_config_dict": {
"hidden_size": 1024,
"intermediate_size": 4096,
"num_attention_heads": 16,
"num_hidden_layers": 24,
"patch_size": 14
}
}
+20
View File
@@ -0,0 +1,20 @@
{
"crop_size": 224,
"do_center_crop": true,
"do_convert_rgb": true,
"do_normalize": true,
"do_resize": true,
"feature_extractor_type": "CLIPFeatureExtractor",
"image_mean": [
0.48145466,
0.4578275,
0.40821073
],
"image_std": [
0.26862954,
0.26130258,
0.27577711
],
"resample": 3,
"size": 224
}
+98 -94
View File
@@ -3,49 +3,25 @@ import numpy as np
import numpy.typing as npt
import torch
from torch import Tensor
from comfy_api.latest import io
IntArray = npt.NDArray[np.int_]
class TileLayout:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"min_tile_size": ("INT", {"default": 512, "min": 64, "max": 8192, "step": 8}),
"padding": ("INT", {"default": 32, "min": 0, "max": 8192, "step": 8}),
"blending": ("INT", {"default": 8, "min": 0, "max": 256, "step": 8}),
}
}
CATEGORY = "external_tooling/tiles"
RETURN_TYPES = ("TILE_LAYOUT",)
FUNCTION = "node"
image_size: IntArray
tile_size: IntArray
padding: int
blending: int
tile_count: IntArray
def node(self, image: Tensor, min_tile_size: int, padding: int, blending: int):
self.init(image, min_tile_size, padding, blending)
return (self,)
def init(self, image: Tensor, min_tile_size: int, padding: int, blending: int):
def __init__(self, image: Tensor, min_tile_size: int, padding: int, blending: int):
assert all([x % 8 == 0 for x in image.shape[-3:-1]]), "Image size must be divisible by 8"
assert min_tile_size % 8 == 0, "Tile size must be divisible by 8"
assert blending < padding, "Blending must be smaller than padding"
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 = self.image_size // (min_tile_size - 2 * padding)
self.image_size: IntArray = np.array(image.shape[-3:-1])
self.padding: int = padding
self.blending: int = blending
self.tile_count: IntArray = np.maximum(1, self.image_size // (min_tile_size - 2 * padding))
image_size_with_overlap = self.image_size + (self.tile_count - 1) * 2 * padding
tile_size = np.ceil(image_size_with_overlap / self.tile_count)
self.tile_size = (np.ceil(tile_size / 8) * 8).astype(int)
self.tile_size: IntArray = (np.ceil(tile_size / 8) * 8).astype(int)
def size(self, coord: IntArray):
return self.end(coord) - self.start(coord)
@@ -85,7 +61,7 @@ class TileLayout:
mask = torch.zeros((1, 1, size[0], size[1]), dtype=torch.float)
mask[:, :, s[0] : e[0], s[1] : e[1]] = 1.0
if blend and self.blending > 0:
mask = box_blur(mask, (self.blending, self.blending), separable=True)
mask = box_blur(mask, (self.blending, self.blending))
return mask.squeeze(0)
def merge(self, image: Tensor, index: int, tile: Tensor):
@@ -93,83 +69,111 @@ class TileLayout:
rect = self.rect(coord)
mask = self.mask(coord, blend=True)
mask = mask.reshape(*mask.shape, 1).repeat(1, 1, 1, image.shape[-1])
image[*rect] = (1 - mask) * image[*rect] + mask * tile
image[rect] = (1 - mask) * image[rect] + mask * tile
class ExtractImageTile:
class CreateTileLayout(io.ComfyNode):
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"layout": ("TILE_LAYOUT",),
"index": ("INT", {"min": 0}),
}
}
def define_schema(cls):
return io.Schema(
node_id="ETN_TileLayout",
display_name="Create Tile Layout",
category="external_tooling/tiles",
inputs=[
io.Image.Input("image"),
io.Int.Input("min_tile_size", default=512, min=64, max=8192, step=8),
io.Int.Input("padding", default=32, min=0, max=8192, step=8),
io.Int.Input("blending", default=8, min=0, max=256, step=8),
],
outputs=[io.Custom("TileLayout").Output(display_name="layout")],
)
CATEGORY = "external_tooling/tiles"
RETURN_TYPES = ("IMAGE",)
FUNCTION = "slice"
def slice(self, image: Tensor, layout: TileLayout, index: int):
return (layout.tile(image, index),)
class ExtractMaskTile:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"mask": ("MASK",),
"layout": ("TILE_LAYOUT",),
"index": ("INT", {"min": 0}),
}
}
def execute(cls, image: Tensor, min_tile_size: int, padding: int, blending: int):
return io.NodeOutput(TileLayout(image, min_tile_size, padding, blending))
CATEGORY = "external_tooling/tiles"
RETURN_TYPES = ("MASK",)
FUNCTION = "slice"
def slice(self, mask: Tensor, layout: TileLayout, index: int):
class ExtractImageTile(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_ExtractImageTile",
display_name="Extract Image Tile",
category="external_tooling/tiles",
inputs=[
io.Image.Input("image"),
io.Custom("TileLayout").Input("layout"),
io.Int.Input("index", default=0, min=0),
],
outputs=[io.Image.Output(display_name="tile")],
)
@classmethod
def execute(cls, image: Tensor, layout: TileLayout, index: int):
return io.NodeOutput(layout.tile(image, index))
class ExtractMaskTile(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_ExtractMaskTile",
display_name="Extract Mask Tile",
category="external_tooling/tiles",
inputs=[
io.Mask.Input("mask"),
io.Custom("TileLayout").Input("layout"),
io.Int.Input("index", default=0, min=0),
],
outputs=[io.Mask.Output(display_name="tile")],
)
@classmethod
def execute(cls, mask: Tensor, layout: TileLayout, index: int):
tile = layout.tile(mask.unsqueeze(3), index)
return (tile.squeeze(3),)
return io.NodeOutput(tile.squeeze(3))
class GenerateTileMask:
class GenerateTileMask(io.ComfyNode):
@classmethod
def INPUT_TYPES(cls):
return {
"required": {"layout": ("TILE_LAYOUT",), "index": ("INT", {"min": 0})},
"optional": {"blend": ("BOOLEAN",)},
}
def define_schema(cls):
return io.Schema(
node_id="ETN_GenerateTileMask",
display_name="Generate Tile Mask",
category="external_tooling/tiles",
inputs=[
io.Custom("TileLayout").Input("layout"),
io.Int.Input("index", default=0, min=0),
io.Boolean.Input("blend", default=False, optional=True),
],
outputs=[io.Mask.Output(display_name="mask")],
)
CATEGORY = "external_tooling/tiles"
RETURN_TYPES = ("MASK",)
FUNCTION = "generate"
def generate(self, layout: TileLayout, index: int, blend: bool = False):
return (layout.mask(layout.coord(index), blend=blend),)
class MergeImageTile:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"layout": ("TILE_LAYOUT",),
"index": ("INT", {"min": 0}),
"tile": ("IMAGE",),
}
}
def execute(cls, layout: TileLayout, index: int, blend: bool = False):
return io.NodeOutput(layout.mask(layout.coord(index), blend=blend))
CATEGORY = "external_tooling/tiles"
RETURN_TYPES = ("IMAGE",)
FUNCTION = "merge"
def merge(self, image: Tensor, layout: TileLayout, index: int, tile: Tensor):
class MergeImageTile(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_MergeImageTile",
display_name="Merge Image Tile",
category="external_tooling/tiles",
inputs=[
io.Image.Input("image"),
io.Custom("TileLayout").Input("layout"),
io.Int.Input("index", default=0, min=0),
io.Image.Input("tile"),
],
outputs=[io.Image.Output(display_name="image")],
)
@classmethod
def execute(cls, image: Tensor, layout: TileLayout, index: int, tile: Tensor):
assert index < layout.total_count, f"Index {index} out of range"
if index == 0:
image = image.clone()
layout.merge(image, index, tile)
return (image,)
return io.NodeOutput(image)
+119
View File
@@ -0,0 +1,119 @@
"""Text translation using Argos Translate.
The node takes text input and translates it to English. The text may contain any
number of language directives in the form `lang:xx` where `xx` is a two-letter
language code. Text fragments after a language directives are translated.
If the language is `en` text is passed through unmodified.
"""
from __future__ import annotations
import re
from functools import cache
from typing import NamedTuple
from comfy_api.latest import io
@cache
def available_languages():
try:
from argostranslate.package import update_package_index, get_available_packages
update_package_index()
list = get_available_packages()
return [(l.from_code, l.from_name) for l in list if l.to_code == "en"]
except ImportError:
return [("NOT INSTALLED", "NOT INSTALLED")]
def translate_chunk(text: str, language: str):
if text.strip() == "":
return text
target = "en"
if language == target:
return text
try:
from argostranslate.package import get_installed_packages, get_available_packages
from argostranslate.translate import translate
installed = get_installed_packages()
if not any(p.from_code == language and p.to_code == target for p in installed):
available = get_available_packages()
pkg = next(
(p for p in available if p.from_code == language and p.to_code == target), None
)
assert pkg, f"Couldn't find package for translation from {language}"
# print("Downloading and installing translation package", pkg) # this will cause encoding errors
pkg.install()
text, embeddings = _extract_embeddings(text)
translation = translate(text, language, target)
return embeddings + translation
except ImportError:
raise ImportError(
"Argos Translate is not installed. Please install it with `pip install argostranslate`"
)
def translate(text: str):
chunks = Chunk.parse(text)
return " ".join(translate_chunk(c.text, c.lang) for c in chunks)
class Translate(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_Translate",
display_name="Translate Text",
category="external_tooling",
inputs=[io.String.Input("text", multiline=True)],
outputs=[io.String.Output(display_name="translation")],
)
@classmethod
def execute(cls, text: str):
return io.NodeOutput(translate(text))
_lang_regex = re.compile(r"(lang:\w\w)")
class Chunk(NamedTuple):
text: str
lang: str
@staticmethod
def parse(text: str):
languages = [code for code, name in available_languages()] + ["en"]
chunks: list[Chunk] = []
lang = "en"
last = 0
for m in _lang_regex.finditer(text):
if m.start() > 0:
chunks.append(Chunk(text[last : m.start()].strip(), lang))
last = m.end()
lang = m.group(0)[5:]
if lang not in languages:
raise ValueError(
f"Invalid language directive {m.group(0)} - {lang} is not a known language code."
f" Available languages: {', '.join(languages)}"
)
if last < len(text):
chunks.append(Chunk(text[last:].strip(), lang))
return [c for c in chunks if c.text != ""]
_embedding_regex = re.compile(r"(embedding:[^\s,]+)")
def _extract_embeddings(text: str):
matches = _embedding_regex.findall(text)
embeddings = " ".join(matches)
if matches:
embeddings += " "
for m in matches:
text = text.replace(m, "")
return text, embeddings
+771
View File
@@ -0,0 +1,771 @@
{
"last_node_id": 78,
"last_link_id": 126,
"nodes": [
{
"id": 60,
"type": "LoadImage",
"pos": [
-692,
-758
],
"size": {
"0": 310.5925598144531,
"1": 335.9309997558594
},
"flags": {},
"order": 0,
"mode": 0,
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
101,
110,
111,
112,
113
],
"shape": 3,
"slot_index": 0
},
{
"name": "MASK",
"type": "MASK",
"links": null,
"shape": 3
}
],
"properties": {
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"photo.jpg",
"image"
]
},
{
"id": 74,
"type": "ETN_MergeImageTile",
"pos": [
-326,
-250
],
"size": {
"0": 315,
"1": 98
},
"flags": {},
"order": 12,
"mode": 0,
"inputs": [
{
"name": "image",
"type": "IMAGE",
"link": 120
},
{
"name": "layout",
"type": "TILE_LAYOUT",
"link": 122,
"slot_index": 1
},
{
"name": "tile",
"type": "IMAGE",
"link": 121
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
123
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "ETN_MergeImageTile"
},
"widgets_values": [
3
],
"color": "#232",
"bgcolor": "#353"
},
{
"id": 61,
"type": "ETN_TileLayout",
"pos": [
-333,
-639
],
"size": {
"0": 315,
"1": 106
},
"flags": {},
"order": 2,
"mode": 0,
"inputs": [
{
"name": "image",
"type": "IMAGE",
"link": 101
}
],
"outputs": [
{
"name": "TILE_LAYOUT",
"type": "TILE_LAYOUT",
"links": [
102,
104,
105,
106,
122,
124
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "ETN_TileLayout"
},
"widgets_values": [
880,
48,
16
]
},
{
"id": 63,
"type": "PreviewImage",
"pos": [
380,
-700
],
"size": {
"0": 221.69317626953125,
"1": 191.44602966308594
},
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 103
}
],
"properties": {
"Node name for S&R": "PreviewImage"
}
},
{
"id": 69,
"type": "PreviewImage",
"pos": [
620,
-470
],
"size": {
"0": 231.39317321777344,
"1": 200.04603576660156
},
"flags": {},
"order": 11,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 109
}
],
"properties": {
"Node name for S&R": "PreviewImage"
}
},
{
"id": 75,
"type": "PreviewImage",
"pos": [
25,
-220
],
"size": {
"0": 321.5931701660156,
"1": 266.5460205078125
},
"flags": {},
"order": 14,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 123
}
],
"properties": {
"Node name for S&R": "PreviewImage"
}
},
{
"id": 76,
"type": "ETN_GenerateTileMask",
"pos": [
376,
-207
],
"size": {
"0": 210,
"1": 85.74603271484375
},
"flags": {},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "layout",
"type": "TILE_LAYOUT",
"link": 124,
"slot_index": 0
}
],
"outputs": [
{
"name": "MASK",
"type": "MASK",
"links": [
125
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "ETN_GenerateTileMask"
},
"widgets_values": [
3,
true
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 77,
"type": "MaskToImage",
"pos": [
620,
-210
],
"size": {
"0": 210,
"1": 26
},
"flags": {},
"order": 13,
"mode": 0,
"inputs": [
{
"name": "mask",
"type": "MASK",
"link": 125
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
126
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "MaskToImage"
}
},
{
"id": 78,
"type": "PreviewImage",
"pos": [
629,
-140
],
"size": {
"0": 227.79318237304688,
"1": 201.64602661132812
},
"flags": {},
"order": 15,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 126
}
],
"properties": {
"Node name for S&R": "PreviewImage"
}
},
{
"id": 71,
"type": "EmptyImage",
"pos": [
-681,
-247
],
"size": {
"0": 315,
"1": 130
},
"flags": {},
"order": 1,
"mode": 0,
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
120
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "EmptyImage"
},
"widgets_values": [
2304,
1728,
1,
0
]
},
{
"id": 67,
"type": "PreviewImage",
"pos": [
380,
-470
],
"size": {
"0": 222.79318237304688,
"1": 197.94602966308594
},
"flags": {},
"order": 9,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 107
}
],
"properties": {
"Node name for S&R": "PreviewImage"
}
},
{
"id": 68,
"type": "PreviewImage",
"pos": [
620,
-703
],
"size": {
"0": 225.89317321777344,
"1": 194.74603271484375
},
"flags": {},
"order": 10,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 108
}
],
"properties": {
"Node name for S&R": "PreviewImage"
}
},
{
"id": 65,
"type": "ETN_ExtractImageTile",
"pos": [
30,
-760
],
"size": {
"0": 278.19317626953125,
"1": 78
},
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "image",
"type": "IMAGE",
"link": 110
},
{
"name": "layout",
"type": "TILE_LAYOUT",
"link": 105
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
108
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "ETN_ExtractImageTile"
},
"widgets_values": [
2
],
"color": "#2a363b",
"bgcolor": "#3f5159"
},
{
"id": 62,
"type": "ETN_ExtractImageTile",
"pos": [
40,
-630
],
"size": {
"0": 274.19317626953125,
"1": 78
},
"flags": {},
"order": 3,
"mode": 0,
"inputs": [
{
"name": "image",
"type": "IMAGE",
"link": 113,
"slot_index": 0
},
{
"name": "layout",
"type": "TILE_LAYOUT",
"link": 102
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
103
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "ETN_ExtractImageTile"
},
"widgets_values": [
0
],
"color": "#2a363b",
"bgcolor": "#3f5159"
},
{
"id": 64,
"type": "ETN_ExtractImageTile",
"pos": [
40,
-510
],
"size": {
"0": 273.2931823730469,
"1": 78
},
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "image",
"type": "IMAGE",
"link": 111
},
{
"name": "layout",
"type": "TILE_LAYOUT",
"link": 104
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
107
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "ETN_ExtractImageTile"
},
"widgets_values": [
1
],
"color": "#2a363b",
"bgcolor": "#3f5159"
},
{
"id": 66,
"type": "ETN_ExtractImageTile",
"pos": [
40,
-370
],
"size": {
"0": 275.4931945800781,
"1": 78
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "image",
"type": "IMAGE",
"link": 112
},
{
"name": "layout",
"type": "TILE_LAYOUT",
"link": 106
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
109,
121
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "ETN_ExtractImageTile"
},
"widgets_values": [
3
],
"color": "#2a363b",
"bgcolor": "#3f5159"
}
],
"links": [
[
101,
60,
0,
61,
0,
"IMAGE"
],
[
102,
61,
0,
62,
1,
"TILE_LAYOUT"
],
[
103,
62,
0,
63,
0,
"IMAGE"
],
[
104,
61,
0,
64,
1,
"TILE_LAYOUT"
],
[
105,
61,
0,
65,
1,
"TILE_LAYOUT"
],
[
106,
61,
0,
66,
1,
"TILE_LAYOUT"
],
[
107,
64,
0,
67,
0,
"IMAGE"
],
[
108,
65,
0,
68,
0,
"IMAGE"
],
[
109,
66,
0,
69,
0,
"IMAGE"
],
[
110,
60,
0,
65,
0,
"IMAGE"
],
[
111,
60,
0,
64,
0,
"IMAGE"
],
[
112,
60,
0,
66,
0,
"IMAGE"
],
[
113,
60,
0,
62,
0,
"IMAGE"
],
[
120,
71,
0,
74,
0,
"IMAGE"
],
[
121,
66,
0,
74,
2,
"IMAGE"
],
[
122,
61,
0,
74,
1,
"TILE_LAYOUT"
],
[
123,
74,
0,
75,
0,
"IMAGE"
],
[
124,
61,
0,
76,
0,
"TILE_LAYOUT"
],
[
125,
76,
0,
77,
0,
"MASK"
],
[
126,
77,
0,
78,
0,
"IMAGE"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 0.9090909090909091,
"offset": [
899.4068198252606,
850.8539704011341
]
}
},
"version": 0.4
}
Binary file not shown.

After

Width:  |  Height:  |  Size: 668 KiB

File diff suppressed because one or more lines are too long
Binary file not shown.

After

Width:  |  Height:  |  Size: 418 KiB