78 Commits
Author SHA1 Message Date
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 3348 additions and 93 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
+177 -25
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)
@@ -36,9 +44,33 @@ 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
## <a id="regions" href="#toc">Regions</a>
When integrating ComfyUI into tools which use layers and compose them on the fly, it is useful to only receive relevant masked regions.
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 +78,150 @@ 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`
#### 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, lumina2, 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.
### 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.
+26 -3
View File
@@ -1,4 +1,4 @@
from . import api, nodes, tile, region
from . import api as api, nodes, tile, region, nsfw, translation, krita
NODE_CLASS_MAPPINGS = {
"ETN_LoadImageBase64": nodes.LoadImageBase64,
@@ -6,6 +6,8 @@ NODE_CLASS_MAPPINGS = {
"ETN_SendImageWebSocket": nodes.SendImageWebSocket,
"ETN_CropImage": nodes.CropImage,
"ETN_ApplyMaskToImage": nodes.ApplyMaskToImage,
"ETN_ReferenceImage": nodes.ReferenceImage,
"ETN_ApplyReferenceImages": nodes.ApplyReferenceImages,
"ETN_TileLayout": tile.TileLayout,
"ETN_ExtractImageTile": tile.ExtractImageTile,
"ETN_ExtractMaskTile": tile.ExtractMaskTile,
@@ -15,6 +17,16 @@ NODE_CLASS_MAPPINGS = {
"ETN_DefineRegion": region.DefineRegion,
"ETN_ListRegionMasks": region.ListRegionMasks,
"ETN_AttentionMask": region.AttentionMask,
"ETN_NSFWFilter": nsfw.NSFWFilter,
"ETN_Translate": translation.Translate,
"ETN_KritaOutput": krita.KritaOutput,
"ETN_KritaSendText": krita.KritaSendText,
"ETN_KritaCanvas": krita.KritaCanvas,
"ETN_KritaSelection": krita.KritaSelection,
"ETN_KritaImageLayer": krita.KritaImageLayer,
"ETN_KritaMaskLayer": krita.KritaMaskLayer,
"ETN_Parameter": krita.Parameter,
"ETN_KritaStyle": krita.KritaStyle,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ETN_LoadImageBase64": "Load Image (Base64)",
@@ -22,8 +34,8 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"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_ReferenceImage": "Reference Image",
"ETN_ApplyReferenceImages": "Apply Reference Images",
"ETN_TileLayout": "Create Tile Layout",
"ETN_ExtractImageTile": "Extract Image Tile",
"ETN_ExtractMaskTile": "Extract Mask Tile",
@@ -33,4 +45,15 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"ETN_DefineRegion": "Define Region",
"ETN_ListRegionMasks": "List Region Masks",
"ETN_AttentionMask": "Regions Attention Mask",
"ETN_NSFWFilter": "NSFW Filter",
"ETN_Translate": "Translate Text",
"ETN_KritaOutput": "Krita Output",
"ETN_KritaSendText": "Send Text",
"ETN_KritaCanvas": "Krita Canvas",
"ETN_KritaSelection": "Krita Selection",
"ETN_KritaImageLayer": "Krita Image Layer",
"ETN_KritaMaskLayer": "Krita Mask Layer",
"ETN_Parameter": "Parameter",
"ETN_KritaStyle": "Krita Style",
}
WEB_DIRECTORY = "./js"
+266 -25
View File
@@ -1,13 +1,21 @@
from __future__ import annotations
from aiohttp import web
from typing import 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
input_block_name = "model.diffusion_model.input_blocks.0.0.weight"
model_names = {
@@ -15,12 +23,41 @@ 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",
"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",
}
gguf_architectures = {
"sd1": "sd15",
"qwen_image": "qwen-image",
}
@@ -35,10 +72,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 +86,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 +102,232 @@ 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 = {"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:
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)
return comfy_config.get("model_class")
model_class = md.get("model_class")
if model_class == "NunchakuFluxTransformer2dModel":
return "Flux"
if model_class == "NunchakuQwenImageTransformer2DModel":
return "QwenImage"
return None
@_server.routes.get("/etn/model_info")
async def model_info(request):
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)
reader = gguf.GGUFReader(path)
arch_field = reader.get_field("general.architecture")
if arch_field is not None:
if len(arch_field.types) != 1 or arch_field.types[0] != gguf.GGUFValueType.STRING:
raise TypeError(
f"Bad type for GGUF general.architecture key: expected string, got {arch_field.types!r}"
)
arch_str = str(arch_field.parts[arch_field.data[-1]], encoding="utf-8")
else: # stable-diffusion.cpp, requires conversion. not handled for now
return {"base_model": "flux", "is_inpaint": False}
if arch_str == "flux" and any(
t.name.startswith("distilled_guidance_layer")
for t in itertools.islice(reader.tensors, 5)
):
arch_str = "chroma"
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)
result["quant"] = reader.get_field("general.file_type").lower()
except Exception as e:
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):
try:
try:
files = folder_paths.get_filename_list(model_type)
except KeyError:
return web.json_response({"error": f"Model folder not found: {model_type}"})
is_checkpoint = model_type == "checkpoints"
info = {
filename: inspect_diffusion_model(filename, model_type, is_checkpoint)
for filename in 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
_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)
@_server.routes.get("/api/etn/model_info")
async def api_model_info(request):
return inspect_models("checkpoints")
@_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.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

+237
View File
@@ -0,0 +1,237 @@
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=")
const outputIcon = loadImage("data:image/webp;base64,UklGRrIHAABXRUJQVlA4WAoAAAAQAAAAjwAAOwAAQUxQSKoCAAARkMbsnyFJ/2TVySSdzPJs27bxzbbtuznbtm3btm3bt+ikkkonlfxPU1X9n57zh4iYAPifZKQ2T7LIkKDkQ1FPrezCF/hjvq+dx7mwiqM3HnJ058yGXkorEfGCKZV6I6o+quOM4bMsgU5bfE1KOh4LEbGnv2RnUSer50DGFwxJ2qwWGTATERFfZPzBZNR9P1ZXTksgVdaODHgt/H5hCJjP0ME6eiLfCqTLypKBSPYdsmY2OjpWy0KOlF+EkIFk/DvnZ2tIzZA0a0kHUtok0KfWkxieJQQZBQksq3QWiXODEGSnYRsq76mxjJQK0cDdKjY1XpLSWyKYVYFTSyxLqCZSvexKppZnZDCZTNpD6BKTi2iIRbqzZc6gW8x5m1ptOCEuYaJ74B1T6RkhjPQXkugiuC9EBSk38wftRACZsdKLEHGRQjJSCyUgnwichSgdj4jYW64KqQsywN1E1JJqReqtWyEvLtOTFHMtfG8GXMYT6C5zQLIjqXiJm+guNw2ZOqTu+0uJ7sKyg2w+Urv9GdxdGoN0JKnh/qBnIE1LlH6LiHNAMZ5SWQkoJAJHcQ7iTUNlNyE7Vga4a7CsoFqP0AmQjqdm5dTWGJRTvqfTSu4ONe7VNRs0LiTzLK3cNEHsGWhuZ+gom0hlNMgXZ7T4WF16Q6YRuZZPAS7TYrGUoPgFEnZPUC3EKDEf0O6WSGFmMiVoxejw3UDcHC6c21oFNLZjVNhGgxrkHC+cOtQKtBa5aQkCLL4VBGDZsbYzu7uF6AGouPQ1t7mTIv5AMw8EZLmhbx0QizqHgJMer5MQwHkG7NP2YmgzCMqRPYff0cIWDSgOwbrcgOFnhcqLhQNaeSB4h1XxDZh29LX4kXVt5YABrZJBkM/acIBsx5IG/AyGpS1UrkrF4tm98K8uVlA4IOIEAABwHACdASqQADwAPm0uk0ckIiGhLjUJmIANiWgOuBpEsADI+tE9j+ifkl+QHyd1n+77tGXf04+Cpx/549gD9X+lF5t/2d9cn0T+gB/d/8R1gHoAeWx+3nwZfuZ6T+aq/0Dtm/zaCN5VfOk9UfrX8A/Ry9E1elWp+7DqL5PqqU3p7QQgEkexRJ8VwS4d4Xk5cyjjrDbvzKxXEwCzN96k0RfNIpiQ2YsSbcyoRxX5fFzllf+dOy8uCMy/ebkjS0ONwuzkRR55zP3zjA4e+C969ch1Ab9LgcrpUqQ8MOvEXnQmqQkxXM4x4nA1f+jJAAD+8tXq96F285yEhMGbOWPp352/yRnzPmWRyRibmd800tluUOW4IyIZz2Hw1xYA9/xsSgKy0yQS//7BbNldSPJ+MCX/mxBqrttfeQX/mf/+2AcZ7Z1wDrdNoOnt8ISIu33p34GUAqqEyPwtdrhMf55SjsQmwUtm/I/qPiQ0ZOPv5ci6kCP0Ddb2jRr38UOXvi54DlMMmkxTs7j/J3jUQYepc7xEgTVVJFcf+8P//18L/9gA//9fHuK11sTUs+RbYDzZKn0uM2PnTEUpAJAT4wETKSU4KDn3wR3rUGGycaAPVo40AjzO5g7VkynMJuo3M2vclcmfmS3ygBGDqjGHQybd03tnGdkGOCGLLlTEABV+UgJxYxH2YxRX9zUULXFnfnBpPxQcOdq+zs1zi9uI2iAqSKwh8Dhch2Ytz8iaZLW39S3+3pmGdITR49+nlHjcG4xNVSYRLFLRmEj/H/I+7qd90N6AF9aRDuUFH1O7ONRGjEQGvPMEF0Fj5atb5w9tjc1pcKTsaWvT2GbF9NQ31HmpaLAgs8szVbuEC8GHKCESxKmx+Hrh5ZwjrNihG3KL0H1n3/g/WetlSYEFYsYTXQmgyUGCVIILkJYrRIdLB5iVAPrseYWKCT8HgJuCUAhaqRO+6jn1fkplsC0yYCveVI+yyDsVr98kmO5arhQ3u+aqKVUvJ8xZL9as4008lN9DkKcRhvC4BwWdhupsqUYwLQmaQhLxP15875P/r43c8r4NI4sLDiCi7Rzww1dWNTyThiA07x8b/zTaFC9Sz+jtZDpRPoSf3LS+TmvHZQ+yv/N9nSK/CGpimH/qjTJOQRStf5ppvzT0FzGMX2tqNndJbZD8idLxJFXZekFF16KC/6scsX/lTNL+XFfTqsreVXu7bL/wjNVTPeGkJJE7aWcXP2+3qQTMv+LaO9INAsG3cyp5Co/F06O8XoVtYZXBjH3f3r9Y8Wp89/fqq2OfQSD2/Ujo1t0fNnMA14gpYdtm6+/RcRgNQGIPPGxgAaFjsfC4+63CcHr1nczuKyXiQjmoIH7n/0NCmJv3O+v/Lp30d3n/060TaO5ffQGrrx0O7TYUAC6pdQxfOeuX4/EsKJgMKTW18feF5m1SX4ODnH1SWutwnm5T/k0/l4YXbLUi8QRbdtx74QL9DJRtKP8bDT+yyf//8nAEApoAJH7jMHoQv7XKzIUdH1TDS7Phokc3PP5m68+eUTHU17v50avNmnEHCfybI4FC35LTpSaGqRsgNJJliiV37VIfbUlfDfgIqZmmxEHmnCQTSg2zcf5+9LWPblYTxLShx/2U34N3Rf/3Zvie6j8SS/8X+Yo+dXhKIyg1WX040AAAAA==")
function setIconImage(nodeType, image, size, padRows, padCols) {
const onAdded = nodeType.prototype.onAdded
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)"
}
if (connectedWidget.options?.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 === 0 ? 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_KritaOutput") {
setIconImage(nodeType, outputIcon, [200, 100], 1, 0)
} 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)
}
},
});
})();
+263
View File
@@ -0,0 +1,263 @@
import sys
import torch
import numpy as np
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 .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 KritaOutput:
@classmethod
def INPUT_TYPES(s):
return {"required": {"images": ("IMAGE",)}}
RETURN_TYPES = ()
FUNCTION = "send_images"
OUTPUT_NODE = True
CATEGORY = "krita"
def send_images(self, images):
return SendImageWebSocket().send_images(images, "PNG")
class KritaSendText:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"value": (IO.ANY, {}),
"name": ("STRING", {"default": "Output"}),
"type": (["text", "markdown", "html"], {"default": "text"}),
}
}
RETURN_TYPES = ()
FUNCTION = "send"
OUTPUT_NODE = True
CATEGORY = "krita"
def send(self, 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}"
print(f"Sending text: {name} = {text}")
return {"ui": {"text": [{"name": name, "text": text, "content-type": mime}]}}
class KritaCanvas:
@classmethod
def INPUT_TYPES(cls):
return {}
RETURN_TYPES = ("IMAGE", "INT", "INT", "INT")
RETURN_NAMES = ("image", "width", "height", "seed")
FUNCTION = "placeholder"
CATEGORY = "krita"
def placeholder(self):
return (_placeholder_image(), 512, 512, 0)
class KritaSelection:
@classmethod
def INPUT_TYPES(cls):
return {}
RETURN_TYPES = (IO.MASK, IO.BOOLEAN)
RETURN_NAMES = ("mask", "active")
FUNCTION = "placeholder"
CATEGORY = "krita"
def placeholder(self):
return (torch.ones(1, 512, 512), False)
class KritaImageLayer:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"name": ("STRING", {"default": "Image"}),
}
}
RETURN_TYPES = ("IMAGE", "MASK")
RETURN_NAMES = ("image", "mask")
FUNCTION = "placeholder"
CATEGORY = "krita"
def placeholder(self, name: str):
return (_placeholder_image(), torch.ones(1, 512, 512))
class KritaMaskLayer:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"name": ("STRING", {"default": "Mask"}),
}
}
RETURN_TYPES = ("MASK",)
RETURN_NAMES = ("mask",)
FUNCTION = "placeholder"
CATEGORY = "krita"
def placeholder(self, name: str):
return (torch.ones(1, 512, 512),)
_param_types = [
"auto",
"number",
"number (integer)",
"toggle",
"choice",
"text",
"prompt (positive)",
"prompt (negative)",
]
_any_float = {"default": 0.0, "min": -sys.float_info.max, "max": sys.float_info.max}
class Parameter:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"name": ("STRING", {"default": "Parameter"}),
"type": (_param_types, {"default": "auto"}),
"default": ("STRING", {"default": ""}),
},
"optional": {
"min": ("FLOAT", _any_float),
"max": ("FLOAT", _any_float),
},
}
RETURN_TYPES = (BasicTypes,)
RETURN_NAMES = ("value",)
FUNCTION = "placeholder"
CATEGORY = "krita"
def placeholder(self, name: str, type: str, default, min=0.0, max=1.0):
if type == "number":
return (float(default),)
elif type == "number (integer)":
return (int(default),)
return (default,)
class KritaStyle:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"name": ("STRING", {"default": "Style"}),
"sampler_preset": (["auto", "regular", "live"],),
}
}
RETURN_TYPES = (
"MODEL",
"CLIP",
"VAE",
"STRING",
"STRING",
comfy.samplers.KSampler.SAMPLERS,
comfy.samplers.KSampler.SCHEDULERS,
"INT",
"FLOAT",
)
RETURN_NAMES = (
"model",
"clip",
"vae",
"positive prompt",
"negative prompt",
"sampler name",
"scheduler",
"steps",
"guidance",
)
FUNCTION = "placeholder"
CATEGORY = "krita"
def placeholder(self, name: str, sampler_preset: str):
raise NotImplementedError("This workflow must be started from Krita!")
+176 -29
View File
@@ -1,11 +1,17 @@
from __future__ import annotations
from copy import copy
from typing import NamedTuple
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
class LoadImageBase64:
@classmethod
@@ -16,15 +22,16 @@ class LoadImageBase64:
CATEGORY = "external_tooling"
FUNCTION = "load_image"
def load_image(self, image):
def load_image(self, 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
@@ -42,7 +49,8 @@ class LoadMaskBase64:
CATEGORY = "external_tooling"
FUNCTION = "load_mask"
def load_mask(self, mask):
def load_mask(self, 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
@@ -55,14 +63,19 @@ class LoadMaskBase64:
class SendImageWebSocket:
@classmethod
def INPUT_TYPES(s):
return {"required": {"images": ("IMAGE",)}}
return {
"required": {
"images": ("IMAGE",),
"format": (["PNG", "JPEG"], {"default": "PNG"}),
}
}
RETURN_TYPES = ()
FUNCTION = "send_images"
OUTPUT_NODE = True
CATEGORY = "external_tooling"
def send_images(self, images):
def send_images(self, images, format):
results = []
for tensor in images:
array = 255.0 * tensor.cpu().numpy()
@@ -71,13 +84,14 @@ 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}}
@@ -118,6 +132,22 @@ class CropImage:
return (out,)
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:
@classmethod
def INPUT_TYPES(cls):
@@ -133,29 +163,146 @@ class ApplyMaskToImage:
FUNCTION = "apply_mask"
def apply_mask(self, image: torch.Tensor, mask: torch.Tensor):
# Move the channel to the second dimension for processing
out = image.movedim(-1, 1)
# Check if the images are RGB, and if so, add an alpha channel initialized to 1
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:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"weight": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0}),
"range_start": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0}),
"range_end": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0}),
},
"optional": {
"reference_images": ("REFERENCE_IMAGE",),
},
}
CATEGORY = "external_tooling"
RETURN_TYPES = ("REFERENCE_IMAGE",)
RETURN_NAMES = ("reference_images",)
FUNCTION = "append"
def append(
self,
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:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"conditioning": ("CONDITIONING",),
"clip_vision": ("CLIP_VISION",),
"style_model": ("STYLE_MODEL",),
"references": ("REFERENCE_IMAGE",),
}
}
CATEGORY = "external_tooling"
RETURN_TYPES = ("CONDITIONING",)
FUNCTION = "apply"
def apply(
self,
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
+146
View File
@@ -0,0 +1,146 @@
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 .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: WeakRef | 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):
models = cls._instance and cls._instance()
if models is None:
models = cls()
cls._instance = WeakRef(models)
return models
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:
models: CachedModels
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"sensitivity": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.10}),
},
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "check"
CATEGORY = "external_tooling"
def __init__(self):
self.models = CachedModels.load()
def check(self, image, sensitivity):
image = to_bchw(image)
input = self.models.feature_extractor(image, do_rescale=False, return_tensors="pt")
filtered = self.models.safety_checker(
images=image, clip_input=input.pixel_values, sensitivity=sensitivity
)
return (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 = "2.0.6"
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
+4 -3
View File
@@ -96,6 +96,8 @@ class DefineRegion:
FUNCTION = "define"
def define(self, mask: Tensor, conditioning: list, regions: Region | None = None):
if mask.dim() < 3:
mask = mask.unsqueeze(0)
return (Region(regions, mask, conditioning),)
@@ -113,7 +115,6 @@ class ListRegionMasks:
class AttentionMask:
@classmethod
def INPUT_TYPES(s):
return {
@@ -148,9 +149,9 @@ class AttentionMask:
assert k.mean() == v.mean(), "k and v must be the same."
device, dtype = q.device, q.dtype
if self.conds[0].device != device:
if self.conds[0].device != device or self.conds[0].dtype != dtype:
self.conds = [cond.to(device, dtype=dtype) for cond in self.conds]
if self.mask.device != device:
if self.mask.device != device or self.mask.dtype != dtype:
self.mask = self.mask.to(device, dtype=dtype)
cond_or_unconds = extra_options["cond_or_uncond"]
+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
}
+4 -4
View File
@@ -36,12 +36,12 @@ class TileLayout:
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.tile_count = np.maximum(1, self.image_size // (min_tile_size - 2 * padding))
image_size_with_overlap = self.image_size + (self.tile_count - 1) * 2 * padding
tile_size = np.ceil(image_size_with_overlap / self.tile_count)
@@ -85,7 +85,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,7 +93,7 @@ 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:
+115
View File
@@ -0,0 +1,115 @@
"""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
@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)
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:
@staticmethod
def INPUT_TYPES():
return {"required": {"text": ("STRING", {"multiline": True})}}
CATEGORY = "external_tooling"
RETURN_TYPES = ("STRING",)
FUNCTION = "translate"
def translate(self, text: str):
return (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