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

* Added adaptation acknowledgement for Anima Attention couple implementation

* Address Anima attention cleanup feedback
2026-06-27 13:21:43 +09:00
Acly d82675092e Version 3.2.0 2026-05-31 11:15:36 +02:00
Acly a1e51904de Fix some type errors 2026-05-30 13:50:52 +02:00
Acly ffa130239b Add Anima Control LLLite nodes
* from kohya-ss/ComfyUI-Anima-LLLite
* split nodes into load/apply to cache model loading
2026-05-30 13:50:34 +02:00
Acly 2fd51d0d47 Workaround for import failure in transformers with certain pytorch installs #66
* seems to affects pytorch compiled with DISTRIBUTED=0, eg. Windows ROCm
* should probably be fixed in transformers somehow?
2026-05-30 11:49:19 +02:00
Acly d3c75155b4 Fix already cached image PUT requests (#65)
* Support already cached image PUT requests without aborting the connection
* conditionally send 100 Continue instead
2026-05-11 09:14:53 +02:00
Acly cbaef8d9c5 Version 3.1.4, fix type checks 2026-05-03 18:15:02 +02:00
FeepingCreature b2783d82a6 Add Anima model type 2026-05-01 17:53:32 +02:00
VERIGEN 09759222de Add ERNIE Image model detection 2026-05-01 17:53:02 +02:00
Acly 7fc3df1174 Version 3.1.3 2026-02-21 20:33:28 +01:00
Acly ed99942f86 Add mask output to KritaCanvas 2026-02-20 13:08:25 +01:00
Acly 2d395424ea Fix nsfw filter with transformers>5 2026-02-05 16:28:32 +01:00
Acly 9b9ea62dd8 Version 3.1.2 2026-01-31 18:01:29 +01:00
Acly ad36f89af3 Tiles: make multiple for tile layout configurable
* can now ensure eg. multiple of 16 tiles to be compatible with flux2 latent downsample factor
* default is 8, which matches previous hardcoded value
2026-01-26 17:28:46 +01:00
Alex 7130dcb2df Add KritaStyleAndPrompt node for synced prompts across workspaces
New node ETN_KritaStyleAndPrompt that works like KritaStyle but:
- Prompts and style sync between Generate/Live/Animation/Graph workspaces
- Outputs fully prepared prompts (wildcards evaluated, style merged)
- Model output includes extracted LoRAs from prompts
2026-01-24 13:15:12 +01:00
Acly 77186eda87 Model inspection: detect Flux 2 klein GGUF variants 2026-01-20 17:11:14 +01:00
Acly 24a7bd1a77 Version 3.1.1 2026-01-18 21:02:30 +01:00
Acly 2d14a03ad8 Model inspection: detect variants of Flux 2 (Klein-4B, Klein-9B) 2026-01-16 19:28:58 +01:00
Acly 9d2e03e8d5 Version 3.1.0 2026-01-05 10:18:20 +01:00
Acly c5606f8e8f API: print filename with stack traces when there is an error during inspection 2026-01-05 10:16:54 +01:00
Acly 79e9b6426f Support transmitting partial tiles/crops of the canvas in Krita workflows 2025-12-30 23:47:29 +01:00
Acly ad320a218c Support additional output info for Krita workflows: name, animation, layers 2025-12-29 21:20:57 +01:00
Jax a310f4593b Transmit request to resize the canvas with Krita Output node (#52)
* Added Krita Resize node for plugin
* Registered canvas resize node to the __init__.py
* Removed resizenode and integrated it into the "Krita output"
* Quick update to simplify the node to return the re-sized image to krita instead of a json
2025-12-29 19:35:06 +01:00
Acly 7d957dcfa7 Model inspection: support Z-Image SVDQ (Nunchaku) files 2025-12-22 11:25:09 +01:00
Acly 22cfd71f95 Fix detection of integer widget for Parameter node 2025-12-17 11:05:00 +01:00
Acly 21a2f44d4c Fix Parameter node min/max being reset to default when it's set to 0 #53
* Use a different default than 0 as workaround
* Don't want to change type of min/max as that would break workflows
2025-12-17 10:40:58 +01:00
Acly 0220252912 Version 3.0.1 2025-12-01 09:36:18 +01:00
Acly f447ef70fa Model inspection: support Z-Image GGUFs 2025-11-29 20:26:00 +01:00
Acly fb27a5bda8 Model inspection: support Lumina2, Z-Image, Flux2 2025-11-28 20:00:00 +01:00
Acly aa83259e66 Change image cache to take size into account 2025-11-09 15:01:52 +01:00
Acly 75c632df4b Version 3.0.0 2025-11-03 10:57:01 +01:00
Acly a088a2dde2 API: support pagination for /api/etn/model_info 2025-10-23 14:41:24 +02:00
Acly fbf99f2a08 Add LoadImageCached and SaveImageCached (renamed from SendImageHTTP)
* short-lived in-memory cache for image transfers
* upload images via HTTP to cache and load/reference them in workflows
* save/store images in workflows to cache and download them via HTTP
2025-10-20 14:39:42 +02:00
Acly dfe014ae88 Fix region attention mask not being applied 2025-10-19 12:45:21 +02:00
Acly 6a7ae5ab70 Add SendImageHTTP node
* alternative to SendImageWebSocket
* requires an extra step, but transfers are much faster for large images
* doesn't involve saving files to disk
2025-10-19 12:06:19 +02:00
Acly d4cac6ac95 Change node definitions to "V3" schema, remove CropImage node 2025-10-18 23:59:05 +02:00
Aoi 929fdfcc13 Comment out translation package download print statement
Comment out print statement to prevent encoding errors.
2025-10-18 10:55:56 +02:00
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
13 changed files with 2395 additions and 529 deletions
+2 -1
View File
@@ -11,11 +11,12 @@ jobs:
publish-node: publish-node:
name: Publish Custom Node to registry name: Publish Custom Node to registry
runs-on: ubuntu-latest runs-on: ubuntu-latest
if: ${{ github.repository_owner == 'Acly' }}
steps: steps:
- name: Check out code - name: Check out code
uses: actions/checkout@v4 uses: actions/checkout@v4
- name: Publish Custom Node - name: Publish Custom Node
uses: Comfy-Org/publish-node-action@main uses: Comfy-Org/publish-node-action@v1
with: with:
## Add your own personal access token to your Github Repository secrets and reference it here. ## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }} personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+81 -8
View File
@@ -13,9 +13,9 @@ Provides nodes and API geared towards using ComfyUI as a backend for external to
## <a id="images" href="#toc">Sending and receiving images</a> ## <a id="images" href="#toc">Sending and receiving images</a>
ComfyUI exchanges images via the filesystem. This requires a ComfyUI exchanges images via the filesystem. This requires a
multi-step process (upload images, prompt, download images), is rather multi-step process (upload images, prompt, download images), which
inefficient, and invites a whole class of potential issues. It's also unclear invites a whole class of potential issues you might not want to deal with.
at which point those images will get cleaned up if ComfyUI is used It's also unclear at which point those images will get cleaned up if ComfyUI is used
via external tools. via external tools.
### Load Image (Base64) ### Load Image (Base64)
@@ -33,7 +33,7 @@ Loads a mask (single channel) from a PNG embedded into the prompt as base64 stri
### Send Image (WebSocket) ### Send Image (WebSocket)
Sends an output image over the client WebSocket connection as PNG binary data. Sends an output image over the client WebSocket connection as PNG binary data.
* Inputs: the image (RGB or RGBA) * Inputs: the image (RGB or RGBA), supports batches
This will first send one binary message for each image in the batch via WebSocket: This will first send one binary message for each image in the batch via WebSocket:
``` ```
@@ -44,6 +44,47 @@ That is two 32-bit integers (big endian) with values 1 and 2 followed by the PNG
{'type': 'executed', 'data': {'node': '<node ID>', 'output': {'images': [{'source': 'websocket', 'content-type': 'image/png', 'type': 'output'}, ...]}, 'prompt_id': '<prompt ID>}} {'type': 'executed', 'data': {'node': '<node ID>', 'output': {'images': [{'source': 'websocket', 'content-type': 'image/png', 'type': 'output'}, ...]}, 'prompt_id': '<prompt ID>}}
``` ```
### Load Image from Cache
Loads an image or mask that has been uploaded previously into the workflow.
Uploaded images are temporarily stored in RAM rather than written to disk. This
method has less overhead compared to embedding images as base64 into the prompt,
but is more complex to implement.
* Inputs: id of an image that was uploaded previously
* Outputs: image (RGB) and mask (A of RGBA input, or first channel if no alpha present).
To upload an image, upload the _bytes_ of a PNG via a HTTP PUT request to
`/api/etn/image/{id}`. JPEG or other formats also work. Choose any `id` which
does not clash with other images you upload, and reference it in the node. The
request returns `201` if the image was uploaded and `200` if it was already
cached.
### Save Image to Cache
Stores an output image in RAM temporarily and allows retrieval over HTTP.
This is typically faster than WebSocket, especially for large images.
* Inputs: the image (RGB or RGBA). Batches are supported.
This node will send a JSON message over WebSocket when an image is ready:
```json
{
"type": "executed",
"data": {
"node": "<node ID>",
"output": {
"images": [
{"source": "http", "id": "<image ID>", "content-type": "image/png", "type": "output"}
]
},
"prompt_id": "prompt ID"
}
}
```
To download the images, send a HTTP GET request to `/api/etn/image/{id}` with
the image IDs from the message. Images will be cached for a few minutes.
## <a id="regions" href="#toc">Regions</a> ## <a id="regions" href="#toc">Regions</a>
These nodes implement attention masking for arbitrary number of image regions. Text prompts only apply to the masked area. These nodes implement attention masking for arbitrary number of image regions. Text prompts only apply to the masked area.
@@ -163,7 +204,9 @@ There are various types of models that can be loaded as checkpoint, LoRA, Contro
#### Paramters #### Paramters
* `folder_name`: sub-directory in ComfyUI's models folder. * `folder_name`: sub-directory in ComfyUI's models folder.
Supported model types: `checkpoints`, `diffusion_models` Supported model types: `checkpoints`, `diffusion_models`, `unet`, `unet_gguf`
* `limit=n`: (query parameter, optional) inspect at `n` models
* `offset=i`: (query parameter, optional) start with the `i`th model
#### Output #### Output
Lists available models with additional classification info: Lists available models with additional classification info:
@@ -172,14 +215,38 @@ Lists available models with additional classification info:
"checkpoint_file.safetensors": { "checkpoint_file.safetensors": {
"base_model": "sd15", "base_model": "sd15",
"is_inpaint": false, "is_inpaint": false,
"is_refiner": false "type": "eps"
}, },
... ...
} }
``` ```
Possible values for base model: `sd15, sd20, sd21, sd3, sdxl, ssd1b, svd, cascade-b, cascade-c, aura-flow, hunyuan-dit, flux, flux-schnell` Possible values for base model: `sd15, sd20, sd21, sd3, sdxl, sdxl-refiner, ssd1b, svd, cascade-b, cascade-c, aura-flow, hunyuan-dit, flux, flux-schnell, flux2, lumina2, z-image, chroma, qwen-image`
If base model is `sdxl`, the `type` attribute is set with possible values: `eps, edm, v-prediction, v-prediction-edm`
Detection supports quantized models:
* GGUF: if the `gguf` module is installed, .gguf files are detected and will set the `quant` field
* Nunchaku: SVDQuant models are detected and will set the `quant` field to `svdq`
Returns an entry `{"base_model": "unknown"}` for models with unknown format or which do not match any of the known base models.
#### Pagination
The query parameters limit and offset allow inspecting a subset of models per request.
Usually inspection is quite fast (it only looks at model headers), but it can be slow
in some cases due to anti-virus or slow harddrives.
```
GET /api/etn/model_info/checkpoints?limit=10&offset=20
```
This will return at most 10 models, starting with the 20th model in the list.
It also returns a special `_meta` entry in the output JSON:
```json
{
"checkpoint_20.safetensors": { ... },
"_meta": { "offset": 20, "count": 1, "total": 21 }
}
```
The entry is `{"base_model": "unknown"}` for models which are not in safetensors format or do not match any of the known base models.
### GET /api/etn/languages ### GET /api/etn/languages
@@ -226,3 +293,9 @@ git clone https://github.com/Acly/comfyui-tooling-nodes.git
``` ```
Restart ComfyUI and the nodes are functional. Restart ComfyUI and the nodes are functional.
## Acknowledgements
* Region nodes adapted from [laksjdjf/cgem156-ComfyUI](https://github.com/laksjdjf/cgem156-ComfyUI/blob/main/scripts/attention_couple/node.py)
* Control nodes adapted from [kohya-ss/ComfyUI-Anima-LLLite](https://github.com/kohya-ss/ComfyUI-Anima-LLLite)
+55 -51
View File
@@ -1,53 +1,57 @@
from . import api, nodes, tile, region, nsfw, translation, krita from comfy_api.latest import ComfyExtension, io
from . import api as api
from . import control, krita, nodes, region, tile, translation
class ExternalToolingNodes(ComfyExtension):
async def get_node_list(self) -> list[type[io.ComfyNode]]:
node_list = [
nodes.LoadImageCache,
nodes.SaveImageCache,
nodes.LoadImageBase64,
nodes.LoadMaskBase64,
nodes.SendImageWebSocket,
nodes.ApplyMaskToImage,
nodes.ReferenceImage,
nodes.ApplyReferenceImages,
tile.CreateTileLayout,
tile.ExtractImageTile,
tile.ExtractMaskTile,
tile.GenerateTileMask,
tile.MergeImageTile,
region.BackgroundRegion,
region.DefineRegion,
region.ListRegionMasks,
region.AttentionMask,
translation.Translate,
krita.KritaOutput,
krita.KritaSendText,
krita.KritaCanvas,
krita.KritaSelection,
krita.KritaImageLayer,
krita.KritaMaskLayer,
krita.Parameter,
krita.KritaStyle,
krita.KritaStyleAndPrompt,
control.ControlApply,
control.ControlLoad,
]
try: # see #66
from . import nsfw
node_list.append(nsfw.NSFWFilter)
except (ImportError, ModuleNotFoundError):
import traceback
print("[comfyui-tooling-nodes] WARNING: Could not import all nodes.")
traceback.print_exc()
return node_list
async def comfy_entrypoint():
return ExternalToolingNodes()
NODE_CLASS_MAPPINGS = {
"ETN_LoadImageBase64": nodes.LoadImageBase64,
"ETN_LoadMaskBase64": nodes.LoadMaskBase64,
"ETN_SendImageWebSocket": nodes.SendImageWebSocket,
"ETN_CropImage": nodes.CropImage,
"ETN_ApplyMaskToImage": nodes.ApplyMaskToImage,
"ETN_TileLayout": tile.TileLayout,
"ETN_ExtractImageTile": tile.ExtractImageTile,
"ETN_ExtractMaskTile": tile.ExtractMaskTile,
"ETN_GenerateTileMask": tile.GenerateTileMask,
"ETN_MergeImageTile": tile.MergeImageTile,
"ETN_BackgroundRegion": region.BackgroundRegion,
"ETN_DefineRegion": region.DefineRegion,
"ETN_ListRegionMasks": region.ListRegionMasks,
"ETN_AttentionMask": region.AttentionMask,
"ETN_NSFWFilter": nsfw.NSFWFilter,
"ETN_Translate": translation.Translate,
"ETN_KritaOutput": krita.KritaOutput,
"ETN_KritaCanvas": krita.KritaCanvas,
"ETN_KritaSelection": krita.KritaSelection,
"ETN_KritaImageLayer": krita.KritaImageLayer,
"ETN_KritaMaskLayer": krita.KritaMaskLayer,
"ETN_Parameter": krita.Parameter,
"ETN_KritaStyle": krita.KritaStyle,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ETN_LoadImageBase64": "Load Image (Base64)",
"ETN_LoadMaskBase64": "Load Mask (Base64)",
"ETN_SendImageWebSocket": "Send Image (WebSocket)",
"ETN_CropImage": "Crop Image",
"ETN_ApplyMaskToImage": "Apply Mask to Image",
"ETN_TileLayout": "Create Tile Layout",
"ETN_ExtractImageTile": "Extract Image Tile",
"ETN_ExtractMaskTile": "Extract Mask Tile",
"ETN_MergeImageTile": "Merge Image Tile",
"ETN_GenerateTileMask": "Generate Tile Mask",
"ETN_BackgroundRegion": "Background Region",
"ETN_DefineRegion": "Define Region",
"ETN_ListRegionMasks": "List Region Masks",
"ETN_AttentionMask": "Regions Attention Mask",
"ETN_NSFWFilter": "NSFW Filter",
"ETN_Translate": "Translate Text",
"ETN_KritaOutput": "Krita Output",
"ETN_KritaCanvas": "Krita Canvas",
"ETN_KritaSelection": "Krita Selection",
"ETN_KritaImageLayer": "Krita Image Layer",
"ETN_KritaMaskLayer": "Krita Mask Layer",
"ETN_Parameter": "Parameter",
"ETN_KritaStyle": "Krita Style",
}
WEB_DIRECTORY = "./js" WEB_DIRECTORY = "./js"
+246 -30
View File
@@ -1,19 +1,21 @@
from __future__ import annotations from __future__ import annotations
from aiohttp import web from aiohttp import web
from typing import NamedTuple from typing import Any, NamedTuple
from pathlib import Path from pathlib import Path
import json import json
import traceback import traceback
import re import re
import logging import logging
import itertools
from comfy import model_detection, supported_models from comfy import model_detection
import comfy.utils import comfy.utils
import folder_paths import folder_paths
import server import server
from .translation import available_languages, translate from .translation import available_languages, translate
from .krita import WorkflowExchange from .krita import WorkflowExchange
from .nodes import image_cache
input_block_name = "model.diffusion_model.input_blocks.0.0.weight" input_block_name = "model.diffusion_model.input_blocks.0.0.weight"
@@ -22,7 +24,7 @@ model_names = {
"SD20": "sd20", "SD20": "sd20",
"SD21UnclipL": "sd21", "SD21UnclipL": "sd21",
"SD21UnclipH": "sd21", "SD21UnclipH": "sd21",
"SDXLRefiner": "sdxl", "SDXLRefiner": "sdxl-refiner",
"SDXL": "sdxl", "SDXL": "sdxl",
"SSD1B": "ssd1b", "SSD1B": "ssd1b",
"SVD_img2vid": "svd", "SVD_img2vid": "svd",
@@ -33,7 +35,37 @@ model_names = {
"HunyuanDiT": "hunyuan-dit", "HunyuanDiT": "hunyuan-dit",
"HunyuanDiT1": "hunyuan-dit", "HunyuanDiT1": "hunyuan-dit",
"Flux": "flux", "Flux": "flux",
"FluxInpaint": "flux",
"FluxSchnell": "flux-schnell", "FluxSchnell": "flux-schnell",
"GenmoMochi": "mochi",
"LTXV": "ltxv",
"HunyuanVideo": "hunyuan-video",
"CosmosT2V": "cosmos",
"CosmosI2V": "cosmos",
"CosmosT2IPredict2": "cosmos-predict2",
"CosmosI2VPredict2": "cosmos-predict2",
"ZImage": "z-image",
"Lumina2": "lumina2",
"WAN21_T2V": "wan21",
"WAN21_I2V": "wan21",
"WAN21_FunControl2V": "wan21-fun",
"WAN21_Vace": "wan21-vace",
"WAN21_Camera": "wan21-camera",
"HiDream": "hi-dream",
"Chroma": "chroma",
"ACEStep": "ace-step",
"Omnigen2": "omnigen2",
"QwenImage": "qwen-image",
"QwenImage21": "qwen-image21",
"ErnieImage": "ernie-image",
"Flux2": "flux2",
"Anima": "anima",
"Krea2": "krea2",
}
gguf_architectures = {
"sd1": "sd15",
"qwen_image": "qwen-image",
} }
@@ -48,7 +80,7 @@ class FakeTensor(NamedTuple):
return d return d
def inspect_diffusion_model(filename: str, prefix: str | None, model_type: str): def inspect_safetensors(filename: str, model_type: str, is_checkpoint: bool):
try: try:
# Read header of safetensors file # Read header of safetensors file
path = folder_paths.get_full_path(model_type, filename) path = folder_paths.get_full_path(model_type, filename)
@@ -62,8 +94,10 @@ def inspect_diffusion_model(filename: str, prefix: str | None, model_type: str):
cfg[key] = FakeTensor.from_dict(cfg[key]) cfg[key] = FakeTensor.from_dict(cfg[key])
# Reuse Comfy's model detection # Reuse Comfy's model detection
if prefix is None: prefix = model_detection.unet_prefix_from_state_dict(cfg)
prefix = model_detection.unet_prefix_from_state_dict(cfg) if not is_checkpoint:
cfg = comfy.utils.state_dict_prefix_replace(cfg, {prefix: ""}, filter_keys=False)
prefix = ""
try: # latest ComfyUI takes 2 args try: # latest ComfyUI takes 2 args
unet_config = model_detection.detect_unet_config(cfg, prefix) unet_config = model_detection.detect_unet_config(cfg, prefix)
except TypeError as e: # older ComfyUI versions take 3 args except TypeError as e: # older ComfyUI versions take 3 args
@@ -76,30 +110,161 @@ def inspect_diffusion_model(filename: str, prefix: str | None, model_type: str):
input_count = 4 input_count = 4
# Find a matching base model depending on unet config # Find a matching base model depending on unet config
base_model = model_detection.model_config_from_unet_config(unet_config) base_model = None
if base_model is None: model_type = None
model_quant = None
# Check if it's a Nunchaku SVDQ model by inspecting metadata
raw_name = detect_svdq(cfg)
if raw_name:
model_quant = "svdq"
# Otherwise try ComfyUI's model detection
elif unet_config is not None:
base_model = model_detection.model_config_from_unet_config(unet_config)
if base_model:
raw_name = base_model.__class__.__name__
if raw_name == "SDXL":
model_type = base_model.model_type(cfg).name.lower().replace("_", "-")
if raw_name == "Flux2":
hidden_size = unet_config.get("hidden_size", 0)
model_type = {3072: "klein-4b", 4096: "klein-9b"}.get(hidden_size, "dev")
if not raw_name:
return {"base_model": "unknown"} return {"base_model": "unknown"}
base_model_class = base_model.__class__ base_model_name = model_names.get(raw_name, "unknown")
base_model_name = model_names.get(base_model_class.__name__, "unknown") result: dict[str, Any] = {"base_model": base_model_name}
return { result["is_inpaint"] = (
"base_model": base_model_name, base_model_name in ["sd15", "sdxl"] and input_count > 4
"is_inpaint": base_model_name in ["sd15", "sdxl"] and input_count > 4, ) or raw_name == "FluxInpaint"
"is_refiner": base_model_class is supported_models.SDXLRefiner, if model_quant:
} result["quant"] = model_quant
if model_type:
result["type"] = model_type
elif "T2I" in raw_name:
result["type"] = "t2i"
elif "I2V" in raw_name:
result["type"] = "i2v"
elif "T2V" in raw_name:
result["type"] = "t2v"
elif "Control2V" in raw_name:
result["type"] = "control2v"
return result
return {"base_model": "unknown"} return {"base_model": "unknown"}
except Exception as e:
print("[comfyui-tooling-nodes] Error inspecting file", filename)
traceback.print_exc()
return {"base_model": "unknown", "error": f"Failed to detect base model: {e}"}
def detect_svdq(cfg: dict) -> str | None:
if md := cfg.get("__metadata__"):
if comfy_config := md.get("comfy_config"):
if isinstance(comfy_config, str):
comfy_config = json.loads(comfy_config)
if model_class := comfy_config.get("model_class"):
return model_class
match md.get("model_class"):
case "NunchakuFluxTransformer2dModel":
return "Flux"
case "NunchakuQwenImageTransformer2DModel":
return "QwenImage"
case "NunchakuZImageTransformer2DModel":
return "ZImage"
return None
def inspect_gguf(filename: str, model_type: str):
try:
import gguf
except ImportError:
return {"base_model": "unknown", "error": "GGUF module not found"}
try:
path = folder_paths.get_full_path(model_type, filename)
if path is None:
raise Exception(f"Could not find full path for {model_type}/{filename}")
reader = gguf.GGUFReader(path)
arch_field = reader.get_field("general.architecture")
if arch_field is not None:
if len(arch_field.types) != 1 or arch_field.types[0] != gguf.GGUFValueType.STRING:
raise TypeError(
f"Bad type for GGUF general.architecture key: expected string, got {arch_field.types!r}"
)
arch_str = str(arch_field.parts[arch_field.data[-1]], encoding="utf-8")
else: # stable-diffusion.cpp, requires conversion. not handled for now
return {"base_model": "flux", "is_inpaint": False}
if arch_str == "flux" and any(
t.name.startswith("distilled_guidance_layer")
for t in itertools.islice(reader.tensors, 5)
):
arch_str = "chroma"
# Detect Z-Image (modified Lumina2)
if arch_str == "lumina2":
for t in reader.tensors:
if t.name == "cap_embedder.1.bias" and t.shape[0] == 3840:
arch_str = "z-image"
break
# Detect Flux variants
result_type = None
if arch_str == "flux":
for t in reader.tensors:
if t.name.startswith("distilled_guidance_layer"):
arch_str = "chroma"
break
elif t.name == "double_stream_modulation_img.lin.weight":
arch_str = "flux2"
if t.shape[0] == 3072:
result_type = "klein-4b"
elif t.shape[0] == 4096:
result_type = "klein-9b"
break
result = {
"base_model": gguf_architectures.get(arch_str, arch_str),
"is_inpaint": False,
}
if result_type is not None:
result["type"] = result_type
try:
if file_type := reader.get_field("general.file_type"):
result["quant"] = file_type.contents().lower()
except Exception:
result["quant"] = "gguf"
return result
except Exception as e: except Exception as e:
# traceback.print_exc() # traceback.print_exc()
return {"base_model": "unknown", "error": f"Failed to detect base model: {e}"} return {"base_model": "unknown", "error": f"Failed to detect base model: {e}"}
def inspect_models(model_type: str): def inspect_diffusion_model(filename: str, model_type: str, is_checkpoint: bool):
if filename.endswith(".gguf"):
return inspect_gguf(filename, model_type)
return inspect_safetensors(filename, model_type, is_checkpoint)
def inspect_models(model_type: str, params: dict[str, str]):
try: try:
prefix = "" if model_type in ("unet", "diffusion_models") else None try:
files = folder_paths.get_filename_list(model_type)
except KeyError:
return web.json_response({"error": f"Model folder not found: {model_type}"})
limit = int(params.get("limit", "1000"))
offset = int(params.get("offset", "0"))
files_range = files[offset : offset + limit]
is_checkpoint = model_type == "checkpoints"
info = { info = {
filename: inspect_diffusion_model(filename, prefix, model_type) filename: inspect_diffusion_model(filename, model_type, is_checkpoint)
for filename in folder_paths.get_filename_list(model_type) for filename in files_range
} }
if "limit" in params:
info["_meta"] = dict(offset=offset, count=len(files_range), total=len(files))
return web.json_response(info) return web.json_response(info)
except Exception as e: except Exception as e:
traceback.print_exc() traceback.print_exc()
@@ -128,25 +293,28 @@ def has_invalid_filename(filename: str):
return None return None
async def image_sender(data: bytes):
mem = memoryview(data)
csize = 2**14
for i in range(0, len(mem), csize):
yield mem[i : i + csize]
_server: server.PromptServer | None = getattr(server.PromptServer, "instance", None) _server: server.PromptServer | None = getattr(server.PromptServer, "instance", None)
if _server is not None: if _server is not None:
_workflow_exchange = WorkflowExchange(_server) _workflow_exchange = WorkflowExchange(_server)
@_server.routes.get("/api/etn/model_info/{folder_name}") @_server.routes.get("/api/etn/model_info/{folder_name}")
async def model_info(request: web.Request): async def model_info(request: web.Request):
folder_name = request.match_info.get("folder_name", "checkpoints") folder_name = request.match_info.get("folder_name", "checkpoints")
if error := has_invalid_folder_name(folder_name): error = has_invalid_folder_name(folder_name)
if error is not None:
return error return error
return inspect_models(folder_name) return inspect_models(folder_name, request.rel_url.query)
@_server.routes.get("/api/etn/model_info") @_server.routes.get("/api/etn/model_info")
async def api_model_info(request): async def api_model_info(request):
return inspect_models("checkpoints") return inspect_models("checkpoints", request.rel_url.query)
@_server.routes.get("/etn/model_info")
async def api_model_info(request):
return inspect_models("checkpoints")
@_server.routes.get("/api/etn/languages") @_server.routes.get("/api/etn/languages")
async def languages(request): async def languages(request):
@@ -166,14 +334,60 @@ if _server is not None:
except Exception as e: except Exception as e:
return web.json_response(dict(error=str(e)), status=500) return web.json_response(dict(error=str(e)), status=500)
@_server.routes.get("/api/etn/image/{id}")
async def get_image(request: web.Request):
try:
id = request.match_info.get("id", "")
data, content_type = image_cache.get(id)
if data is None or content_type is None:
return web.json_response(dict(error="Image not found"), status=404)
response = web.Response(
body=image_sender(data),
content_type=content_type,
headers={"Content-Length": str(len(data))},
)
return response
except Exception as e:
return web.json_response(dict(error=str(e)), status=500)
async def put_image(request: web.Request):
try:
id = request.match_info.get("id", "")
if id in image_cache:
await request.release() # Consume and discard the data to avoid connection abort
return web.json_response(dict(status="cached"), status=200)
content_type = request.headers.get("Content-Type", "application/octet-stream")
data = bytearray()
async for chunk, _ in request.content.iter_chunks():
data.extend(chunk)
image_cache.insert(id, bytes(data), content_type)
return web.json_response(dict(status="success"), status=201)
except Exception as e:
return web.json_response(dict(error=str(e)), status=500)
async def _put_image_expect_handler(request: web.Request):
if request.match_info.get("id", "") in image_cache:
# Skip "100 Continue" since we don't need the data, return 200 immediately.
return web.json_response(dict(status="cached"), status=200)
# otherwise run default aiohttp handler
return None
_server.app.router.add_route(
"PUT", "/api/etn/image/{id}", put_image, expect_handler=_put_image_expect_handler
)
@_server.routes.put("/api/etn/upload/{folder_name}/{filename}") @_server.routes.put("/api/etn/upload/{folder_name}/{filename}")
async def upload(request: web.Request): async def upload(request: web.Request):
folder_name = request.match_info.get("folder_name", "") folder_name = request.match_info.get("folder_name", "")
if error := has_invalid_folder_name(folder_name): error = has_invalid_folder_name(folder_name)
if error is not None:
return error return error
filename = request.match_info.get("filename", "") filename = request.match_info.get("filename", "")
if error := has_invalid_filename(filename): error = has_invalid_filename(filename)
if error is not None:
return error return error
try: try:
@@ -182,7 +396,9 @@ if _server is not None:
folder = Path(folder_paths.folder_names_and_paths[folder_name][0][0]) folder = Path(folder_paths.folder_names_and_paths[folder_name][0][0])
total_size = int(request.headers.get("Content-Length", "0")) total_size = int(request.headers.get("Content-Length", "0"))
logging.info(f"Uploading {filename} ({total_size/(1024**2):.1f} MB) to {folder} folder") logging.info(
f"Uploading {filename} ({total_size / (1024**2):.1f} MB) to {folder} folder"
)
with open(folder / filename, "wb") as f: with open(folder / filename, "wb") as f:
async for chunk, _ in request.content.iter_chunks(): async for chunk, _ in request.content.iter_chunks():
+859
View File
@@ -0,0 +1,859 @@
"""ControlNet-LLLite for Anima (DiT) — ComfyUI port (v2 architecture).
Adapted from kohya-ss/ComfyUI-Anima-LLLite
https://github.com/kohya-ss/ComfyUI-Anima-LLLite
Apache-2.0 license
Adapted from kohya-ss/sd-scripts. The on-disk weight format is the v2
named-key format (per-module key prefix = lllite_name, shared encoder under
``lllite_conditioning1.*``, depth embedding split per-module as
``{name}.depth_embed``); legacy ``lllite_modules.*`` files are rejected.
Differences vs. the sd-scripts reference (``networks/control_net_lllite_anima.py``):
* No dependency on ``library.utils`` — uses stdlib logging.
* Module discovery filters the LLM-Adapter sub-tree by class identity in
addition to the path-based check (ComfyUI ships two distinct ``Attention``
classes that share the bare class name).
* ``LLLiteModuleDiT`` keeps a ``restore()`` method (and an idempotent
``apply_to()``); ComfyUI patches/unpatches the original Linear around
every sampler call via ``set_model_unet_function_wrapper``.
* Forward pass casts ``x`` and ``cond_emb`` to the LLLite parameter dtype
so autocast / mixed-precision flows that hand us a different dtype than
the LLLite weights still work.
* CFG batch-size and sequence-length mismatches fall back to identity
instead of asserting, so a slightly-off cond image cannot abort sampling.
* The training-side ``AnimaControlNetLLLiteWrapper`` is omitted; ComfyUI
integrates via ``model_function_wrapper`` in nodes.py instead.
"""
from __future__ import annotations
from copy import copy
import logging
import os
from dataclasses import dataclass
from typing import Any
import folder_paths
import safetensors
import safetensors.torch
import torch
import torch.nn.functional as F
from comfy.model_patcher import ModelPatcher
from comfy_api.latest import io
from torch import nn
logger = logging.getLogger("comfyui-tooling-nodes")
# Class names of the modules that LLLite injects into. The LLM-Adapter uses
# a different ``Attention`` class with the same bare name; we filter it by
# path (``llm_adapter`` in the qualified name) and by the ``is_selfattn``
# attribute presence.
TARGET_ATTENTION_CLASS = "Attention"
TARGET_MLP_CLASS = "GPT2FeedForward"
LLM_ADAPTER_NAME = "llm_adapter"
LLLITE_ARCH_VERSION = "2"
# ----------------------------------------------------------------------------
# target_layers: atomic specifiers and presets
# ----------------------------------------------------------------------------
ATOMIC_SPECIFIERS: tuple[str, ...] = (
"self_attn_q_pre",
"self_attn_kv_pre",
"cross_attn_q_pre",
"mlp_fc1_pre",
)
PRESETS: dict = {
"self_attn_q": ("self_attn_q_pre",),
"self_attn_qkv": ("self_attn_q_pre", "self_attn_kv_pre"),
"self_attn_qkv_cross_q": ("self_attn_q_pre", "self_attn_kv_pre", "cross_attn_q_pre"),
}
def parse_target_layers(spec: str) -> tuple[str, ...]:
"""Resolve a ``target_layers`` spec to a canonical atomic tuple.
Accepts a preset name (``"self_attn_qkv"``) or a comma-separated list of
atomic specifiers (``"self_attn_q_pre,mlp_fc1_pre"``). Returns the atomics
in ``ATOMIC_SPECIFIERS`` order with duplicates removed.
"""
if not isinstance(spec, str):
raise TypeError(f"target_layers must be str, got {type(spec).__name__}")
spec = spec.strip()
if not spec:
raise ValueError("target_layers spec is empty")
if spec in PRESETS:
parts = list(PRESETS[spec])
else:
parts = [p.strip() for p in spec.split(",") if p.strip()]
bad = [p for p in parts if p not in ATOMIC_SPECIFIERS]
if bad:
raise ValueError(
f"unknown target_layers atomic specifier(s): {bad}. "
f"valid atomic={list(ATOMIC_SPECIFIERS)}, presets={list(PRESETS)}"
)
return tuple(a for a in ATOMIC_SPECIFIERS if a in parts)
# ----------------------------------------------------------------------------
# Conditioning1 trunk (v2)
# ----------------------------------------------------------------------------
def _gn(channels: int) -> nn.GroupNorm:
g = 8
while g > 1 and channels % g != 0:
g //= 2
return nn.GroupNorm(g, channels)
class _ResBlock(nn.Module):
def __init__(self, ch: int):
super().__init__()
self.norm1 = _gn(ch)
self.conv1 = nn.Conv2d(ch, ch, kernel_size=3, padding=1)
self.norm2 = _gn(ch)
self.conv2 = nn.Conv2d(ch, ch, kernel_size=3, padding=1)
def forward(self, x: torch.Tensor) -> torch.Tensor:
h = self.conv1(F.silu(self.norm1(x)))
h = self.conv2(F.silu(self.norm2(h)))
return x + h
ASPP_DEFAULT_DILATIONS: tuple[int, ...] = (1, 2, 4, 8)
class _ASPP(nn.Module):
def __init__(self, ch: int, dilations: tuple[int, ...] = ASPP_DEFAULT_DILATIONS):
super().__init__()
assert len(dilations) >= 1, "ASPP needs at least one dilation"
branches = []
for d in dilations:
if d == 1:
conv = nn.Conv2d(ch, ch, kernel_size=1)
else:
conv = nn.Conv2d(ch, ch, kernel_size=3, padding=d, dilation=d)
branches.append(nn.Sequential(conv, _gn(ch), nn.SiLU()))
self.branches = nn.ModuleList(branches)
self.global_pool = nn.AdaptiveAvgPool2d(1)
self.global_conv = nn.Sequential(nn.Conv2d(ch, ch, kernel_size=1), _gn(ch), nn.SiLU())
n_branches = len(dilations) + 1
self.proj = nn.Sequential(nn.Conv2d(ch * n_branches, ch, kernel_size=1), _gn(ch), nn.SiLU())
def forward(self, x: torch.Tensor) -> torch.Tensor:
h, w = x.shape[-2:]
outs = [b(x) for b in self.branches]
g = self.global_conv(self.global_pool(x))
g = F.interpolate(g, size=(h, w), mode="bilinear", align_corners=False)
outs.append(g)
return self.proj(torch.cat(outs, dim=1))
class _Conditioning1(nn.Module):
def __init__(
self,
cond_dim: int,
cond_emb_dim: int,
n_resblocks: int,
use_aspp: bool = False,
aspp_dilations: tuple[int, ...] = ASPP_DEFAULT_DILATIONS,
cond_in_channels: int = 3,
):
super().__init__()
assert cond_dim % 2 == 0, f"cond_dim must be even, got {cond_dim}"
assert cond_in_channels >= 1, f"cond_in_channels must be >= 1, got {cond_in_channels}"
ch_half = cond_dim // 2
self.cond_in_channels = cond_in_channels
self.conv1 = nn.Conv2d(cond_in_channels, ch_half, kernel_size=4, stride=4, padding=0)
self.norm1 = _gn(ch_half)
self.conv2 = nn.Conv2d(ch_half, ch_half, kernel_size=3, stride=1, padding=1)
self.norm2 = _gn(ch_half)
self.conv3 = nn.Conv2d(ch_half, cond_dim, kernel_size=4, stride=4, padding=0)
self.norm3 = _gn(cond_dim)
self.resblocks = nn.ModuleList([_ResBlock(cond_dim) for _ in range(n_resblocks)])
self.aspp = _ASPP(cond_dim, aspp_dilations) if use_aspp else None
self.proj = nn.Conv2d(cond_dim, cond_emb_dim, kernel_size=1)
self.out_norm = nn.LayerNorm(cond_emb_dim)
def forward(self, x: torch.Tensor) -> torch.Tensor:
h = F.silu(self.norm1(self.conv1(x)))
h = F.silu(self.norm2(self.conv2(h)))
h = F.silu(self.norm3(self.conv3(h)))
for rb in self.resblocks:
h = rb(h)
if self.aspp is not None:
h = self.aspp(h)
h = self.proj(h)
b, c, hh, ww = h.shape
h = h.view(b, c, hh * ww).permute(0, 2, 1).contiguous()
h = self.out_norm(h)
return h
# ----------------------------------------------------------------------------
# LLLite module (v2: FiLM + SiLU + 5D path + depth embedding)
# ----------------------------------------------------------------------------
class LLLiteModuleDiT(nn.Module):
def __init__(
self,
name: str,
org_module: nn.Linear,
cond_emb_dim: int,
mlp_dim: int,
dropout: float | None = None,
multiplier: float = 1.0,
):
super().__init__()
self.lllite_name = name
# Wrap in a list so the original Linear is not registered as a submodule
# and its weights stay out of state_dict.
self.org_module = [org_module]
self.cond_emb_dim = cond_emb_dim
self.mlp_dim = mlp_dim
self.dropout = dropout
self.multiplier = multiplier
in_dim = org_module.in_features
self.down = nn.Linear(in_dim, mlp_dim)
self.mid = nn.Linear(mlp_dim + cond_emb_dim, mlp_dim)
# FiLM: cond_local -> (gamma, beta), zero-init for identity at start.
self.cond_to_film = nn.Linear(cond_emb_dim, 2 * mlp_dim)
nn.init.zeros_(self.cond_to_film.weight)
nn.init.zeros_(self.cond_to_film.bias)
self.up = nn.Linear(mlp_dim, in_dim)
nn.init.zeros_(self.up.weight)
nn.init.zeros_(self.up.bias)
self.cond_emb: torch.Tensor | None = None
self.org_forward = None
# Set by the parent ControlNetLLLiteDiT after construction.
self.layer_idx: int = -1
self._depth_embeds_ref: list[nn.Parameter] = []
def apply_to(self):
if self.org_forward is None:
self.org_forward = self.org_module[0].forward
self.org_module[0].forward = self.forward
def restore(self):
if self.org_forward is not None:
self.org_module[0].forward = self.org_forward
self.org_forward = None
def forward(self, x: torch.Tensor) -> torch.Tensor:
# Input layouts:
# self/cross attention q/k/v: (B, S, D) — already flattened in the Anima block
# mlp.layer1: (B, T, H, W, D) — passed un-flattened
# Flatten the 5D case to 3D for the LLLite path and reshape on exit.
if self.multiplier == 0.0 or self.cond_emb is None:
return self.org_forward(x)
orig_shape = x.shape
is_5d = x.dim() == 5
if is_5d:
B, T, H, W, D = orig_shape
x = x.reshape(B, T * H * W, D)
cx = self.cond_emb # (B_c, S, cond_emb_dim)
# Broadcast cond_emb to the runtime batch (CFG cond+uncond, multi-cond).
if x.shape[0] != cx.shape[0]:
if x.shape[0] % cx.shape[0] != 0:
return self.org_forward(x.reshape(orig_shape) if is_5d else x)
cx = cx.repeat(x.shape[0] // cx.shape[0], 1, 1)
if x.shape[1] != cx.shape[1]:
return self.org_forward(x.reshape(orig_shape) if is_5d else x)
# Run the LLLite mini-MLP in its own parameter dtype, then cast the
# correction back to ``x``'s dtype before adding. Robust to autocast
# flows where x and LLLite weights have different dtypes.
param_dtype = self.down.weight.dtype
x_proc = x if x.dtype == param_dtype else x.to(param_dtype)
if cx.dtype != param_dtype or cx.device != x.device:
cx = cx.to(device=x.device, dtype=param_dtype)
# Per-module depth embedding (zero-init so it's a no-op at train start).
if self._depth_embeds_ref:
depth_e = self._depth_embeds_ref[0][self.layer_idx]
if depth_e.dtype != param_dtype or depth_e.device != x.device:
depth_e = depth_e.to(device=x.device, dtype=param_dtype)
cond_local = cx + depth_e
else:
cond_local = cx
h = F.silu(self.down(x_proc))
gb = self.cond_to_film(cond_local)
gamma, beta = gb.chunk(2, dim=-1)
m = self.mid(torch.cat([cond_local, h], dim=-1))
m = m * (1 + gamma) + beta
m = F.silu(m)
if self.dropout is not None and self.training:
m = F.dropout(m, p=self.dropout)
out = self.up(m) * self.multiplier
if out.dtype != x.dtype:
out = out.to(x.dtype)
y = self.org_forward(x + out)
if is_5d:
# org Linear out_features may differ from in_features — recover with -1.
y = y.reshape(orig_shape[0], orig_shape[1], orig_shape[2], orig_shape[3], -1)
return y
# ----------------------------------------------------------------------------
# ControlNetLLLiteDiT
# ----------------------------------------------------------------------------
class ControlNetLLLiteDiT(nn.Module):
def __init__(
self,
dit: nn.Module,
cond_emb_dim: int = 32,
mlp_dim: int = 64,
target_layers: str = "self_attn_q",
dropout: float | None = None,
multiplier: float = 1.0,
cond_dim: int = 64,
cond_resblocks: int = 1,
use_aspp: bool = False,
aspp_dilations: tuple[int, ...] = ASPP_DEFAULT_DILATIONS,
cond_in_channels: int = 3,
inpaint_masked_input: bool = False,
):
super().__init__()
atomics = parse_target_layers(target_layers)
self.cond_emb_dim = cond_emb_dim
self.mlp_dim = mlp_dim
self.target_layers = target_layers
self.target_atomics = atomics
self.dropout = dropout
self.multiplier = multiplier
self.cond_dim = cond_dim
self.cond_resblocks = cond_resblocks
self.use_aspp = use_aspp
self.aspp_dilations = tuple(aspp_dilations) if use_aspp else ()
# 4ch (RGB+mask) inpainting metadata. `inpaint_masked_input` records the training-time
# RGB-masking policy for cond_image preparation; it does not alter the forward pass here.
self.cond_in_channels = cond_in_channels
self.inpaint_masked_input = inpaint_masked_input
self.conditioning1 = _Conditioning1(
cond_dim,
cond_emb_dim,
cond_resblocks,
use_aspp=use_aspp,
aspp_dilations=aspp_dilations,
cond_in_channels=cond_in_channels,
)
modules = self._create_modules(dit, cond_emb_dim, mlp_dim, atomics, dropout, multiplier)
self.lllite_modules = nn.ModuleList(modules)
n = len(self.lllite_modules)
self.depth_embeds = nn.Parameter(torch.zeros(n, cond_emb_dim))
for i, m in enumerate(self.lllite_modules):
m.layer_idx = i
m._depth_embeds_ref = [self.depth_embeds]
aspp_info = f"aspp={'on' + str(list(self.aspp_dilations)) if use_aspp else 'off'}"
inpaint_info = (
f", inpaint=on(masked_input={inpaint_masked_input})" if cond_in_channels != 3 else ""
)
logger.info(
"ControlNet-LLLite (Anima v%s): created %d modules for target=%r "
"(atomics=%s), cond_in_channels=%d, cond_dim=%d, cond_resblocks=%d, %s, "
"cond_emb_dim=%d, mlp_dim=%d%s",
LLLITE_ARCH_VERSION,
n,
target_layers,
list(atomics),
cond_in_channels,
cond_dim,
cond_resblocks,
aspp_info,
cond_emb_dim,
mlp_dim,
inpaint_info,
)
@staticmethod
def _attn_atomic_match(is_self_attn: bool, child_name: str, atomics: tuple[str, ...]) -> bool:
if "output_proj" in child_name:
return False
if is_self_attn:
if child_name == "q_proj":
return "self_attn_q_pre" in atomics
if child_name in ("k_proj", "v_proj"):
return "self_attn_kv_pre" in atomics
return False
else:
if child_name == "q_proj":
return "cross_attn_q_pre" in atomics
return False # cross_attn K,V live in text-embedding space
def _create_modules(
self,
dit: nn.Module,
cond_emb_dim: int,
mlp_dim: int,
atomics: tuple[str, ...],
dropout: float | None,
multiplier: float,
) -> list[LLLiteModuleDiT]:
modules: list[LLLiteModuleDiT] = []
want_mlp_fc1 = "mlp_fc1_pre" in atomics
any_attn = any(
a in atomics for a in ("self_attn_q_pre", "self_attn_kv_pre", "cross_attn_q_pre")
)
for name, module in dit.named_modules():
if LLM_ADAPTER_NAME in name:
continue
cls = module.__class__.__name__
def _is_linear_like(module):
return (
hasattr(module, "in_features")
and hasattr(module, "out_features")
and callable(getattr(module, "forward", None))
)
if any_attn and cls == TARGET_ATTENTION_CLASS:
# The Anima-block Attention exposes is_selfattn; the LLM-Adapter
# Attention does not — skip the latter even if path filter misses.
if not hasattr(module, "is_selfattn"):
continue
is_self_attn = bool(module.is_selfattn)
for child_name, child in module.named_children():
if not _is_linear_like(child):
continue
if not self._attn_atomic_match(is_self_attn, child_name, atomics):
continue
full_name = f"lllite_dit.{name}.{child_name}".replace(".", "_")
modules.append(
LLLiteModuleDiT(
full_name, child, cond_emb_dim, mlp_dim, dropout, multiplier
)
)
elif want_mlp_fc1 and cls == TARGET_MLP_CLASS:
child = getattr(module, "layer1", None)
if not _is_linear_like(child):
continue
full_name = f"lllite_dit.{name}.layer1".replace(".", "_")
modules.append(
LLLiteModuleDiT(full_name, child, cond_emb_dim, mlp_dim, dropout, multiplier)
)
return modules
def set_cond_image(self, cond_image: torch.Tensor | None):
"""cond_image: (B, 3, H*16, W*16) in [-1, 1]; ``None`` clears."""
if cond_image is None:
for m in self.lllite_modules:
m.cond_emb = None
return
cx = self.conditioning1(cond_image) # (B, S, cond_emb_dim)
for m in self.lllite_modules:
m.cond_emb = cx
def clear_cond_image(self):
self.set_cond_image(None)
def set_multiplier(self, multiplier: float):
self.multiplier = multiplier
for m in self.lllite_modules:
m.multiplier = multiplier
def apply_to(self):
for m in self.lllite_modules:
m.apply_to()
def restore(self):
for m in self.lllite_modules:
m.restore()
# ----------------------------------------------------------------------------
# Save / load (named-key format; legacy lllite_modules.* is rejected)
# ----------------------------------------------------------------------------
_INTERNAL_MODULES_PREFIX = "lllite_modules."
_INTERNAL_COND_PREFIX = "conditioning1."
_INTERNAL_DEPTH_KEY = "depth_embeds"
_SAVED_COND_PREFIX = "lllite_conditioning1."
_SAVED_DEPTH_SUFFIX = ".depth_embed"
def _from_saved_state_dict(lllite: ControlNetLLLiteDiT, weights_sd: dict) -> dict:
"""Rewrite a v2 named-key state dict back to the internal layout."""
name_to_idx = {m.lllite_name: i for i, m in enumerate(lllite.lllite_modules)}
n_modules = len(name_to_idx)
out: dict = {}
depth_slices: dict = {}
for k, v in weights_sd.items():
if k.startswith(_SAVED_COND_PREFIX):
out[_INTERNAL_COND_PREFIX + k[len(_SAVED_COND_PREFIX) :]] = v
continue
if k.endswith(_SAVED_DEPTH_SUFFIX):
name = k[: -len(_SAVED_DEPTH_SUFFIX)]
if name in name_to_idx:
depth_slices[name_to_idx[name]] = v
continue
head, dot, tail = k.partition(".")
if dot and head in name_to_idx:
out[f"{_INTERNAL_MODULES_PREFIX}{name_to_idx[head]}.{tail}"] = v
continue
out[k] = v
if depth_slices:
missing = [i for i in range(n_modules) if i not in depth_slices]
if missing:
raise RuntimeError(f"depth_embed slices missing for module idx(es) {missing}")
out[_INTERNAL_DEPTH_KEY] = torch.stack([depth_slices[i] for i in range(n_modules)], dim=0)
return out
def load_lllite_weights(lllite: ControlNetLLLiteDiT, file: str, strict: bool = False):
weights_sd = safetensors.torch.load_file(file)
if any(k.startswith(_INTERNAL_MODULES_PREFIX) for k in weights_sd):
raise RuntimeError(
f"weights at {file} appear to be in a legacy ControlNet-LLLite weight format "
f"(keys starting with '{_INTERNAL_MODULES_PREFIX}'). The current code uses a "
f"named-key format (per-module key prefix = lllite_name, e.g. "
f"'lllite_dit_blocks_0_self_attn_q_proj.down.weight'). Re-train with the current codebase."
)
converted = _from_saved_state_dict(lllite, weights_sd)
info = lllite.load_state_dict(converted, strict=strict)
logger.info("loaded LLLite weights from %s: %s", file, info)
return info
def read_lllite_metadata(file: str) -> dict:
if os.path.splitext(file)[1] != ".safetensors":
raise RuntimeError(f"Must use .safetensors files, got {file}")
with safetensors.safe_open(file, framework="pt") as f:
return f.metadata() or {}
# ----------------------------------------------------------------------------
# ComfyUI nodes for Anima ControlNet-LLLite
# ----------------------------------------------------------------------------
def _get_inner_dit(model) -> torch.nn.Module:
"""Reach the underlying Anima DiT (nn.Module) from a ComfyUI ModelPatcher."""
inner = getattr(model, "model", None)
if inner is None:
raise RuntimeError("Input MODEL has no .model attribute (not a ModelPatcher?)")
dit = getattr(inner, "diffusion_model", None)
if dit is None:
raise RuntimeError("MODEL.model has no .diffusion_model — not a UNet/DiT model?")
return dit
def _target_cond_hw(latent_h: int, latent_w: int, patch_spatial: int = 2) -> tuple[int, int]:
"""Return the (H, W) the cond image / mask must be resized to.
The LLLite ``conditioning1`` Conv has stride 16, so the cond image must be
sized to ``latent_HW * 8`` in input pixel space (= ``token_HW * 16`` after
DiT patchify with patch_spatial=2). The DiT internally pads the latent up
to a multiple of ``patch_spatial`` (see ``MiniTrainDIT.forward`` →
``pad_to_patch_size``), so we mirror that rounding here — otherwise odd
latent dims (e.g. 1032 px → 129 latent) yield a token-count mismatch that
silently bypasses every LLLite module.
"""
padded_h = ((latent_h + patch_spatial - 1) // patch_spatial) * patch_spatial
padded_w = ((latent_w + patch_spatial - 1) // patch_spatial) * patch_spatial
return padded_h * 8, padded_w * 8
def _prepare_cond_image(
image: torch.Tensor,
latent_h: int,
latent_w: int,
device: torch.device,
dtype: torch.dtype,
patch_spatial: int = 2,
) -> torch.Tensor:
"""ComfyUI IMAGE (B,H,W,3) in [0,1] → (1,3,H*8,W*8) in [-1,1]."""
if image.ndim == 4 and image.shape[-1] == 3:
# (B, H, W, 3) -> (B, 3, H, W)
img = image.permute(0, 3, 1, 2).contiguous()
else:
raise ValueError(f"Unexpected cond image shape: {tuple(image.shape)} (expected B,H,W,3)")
img = img[:1] # use first frame only
target_h, target_w = _target_cond_hw(latent_h, latent_w, patch_spatial)
if img.shape[-2] != target_h or img.shape[-1] != target_w:
img = F.interpolate(img, size=(target_h, target_w), mode="bicubic", align_corners=False)
img = img.clamp(0.0, 1.0)
img = img * 2.0 - 1.0
return img.to(device=device, dtype=dtype)
def _prepare_mask(
mask: torch.Tensor,
latent_h: int,
latent_w: int,
device: torch.device,
dtype: torch.dtype,
patch_spatial: int = 2,
) -> torch.Tensor:
"""ComfyUI MASK (B,H,W) in [0,1] → (1,1,H*8,W*8) binarized at 0.5.
Returns the mask in ``{0.0, 1.0}`` (1 = inpaint area, 0 = keep). The caller
is responsible for the ``*2-1`` rescale before concat with RGB.
"""
if mask.ndim == 3:
m = mask.unsqueeze(1) # (B, 1, H, W)
elif mask.ndim == 4 and mask.shape[1] == 1:
m = mask
else:
raise ValueError(f"Unexpected mask shape: {tuple(mask.shape)} (expected B,H,W or B,1,H,W)")
m = m[:1]
target_h, target_w = _target_cond_hw(latent_h, latent_w, patch_spatial)
if m.shape[-2] != target_h or m.shape[-1] != target_w:
m = F.interpolate(m.float(), size=(target_h, target_w), mode="nearest")
m = (m >= 0.5).to(dtype=dtype)
return m.to(device=device)
def _build_inpaint_cond_image(
rgb_pm1: torch.Tensor, mask01: torch.Tensor, masked_input: bool
) -> torch.Tensor:
"""rgb_pm1: (1,3,H,W) in [-1,1], mask01: (1,1,H,W) in {0,1}. Returns (1,4,H,W).
Mirrors ``_build_inpaint_cond_image`` in the sd-scripts training / inference
code: the mask channel is rescaled to ``[-1, +1]`` (matches the RGB range),
and if ``masked_input`` is set the RGB is zeroed where ``mask >= 0.5``.
"""
if masked_input:
keep = (mask01 < 0.5).to(rgb_pm1.dtype)
rgb_pm1 = rgb_pm1 * keep
mask_pm1 = mask01.to(rgb_pm1.dtype) * 2.0 - 1.0
return torch.cat([rgb_pm1, mask_pm1], dim=1)
ETNControlNet = io.Custom("ETN_CONTROL_NET")
class ControlLoad(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_control_load",
display_name="Load ControlNet (tooling-nodes)",
description="Loads ControlNet weights. Currently only supports Anima LLLite weights.",
category="external_tooling",
inputs=[
io.Model.Input("model"),
io.Combo.Input("weights", folder_paths.get_filename_list("controlnet")),
],
outputs=[
io.Model.Output("out_model", "model"),
ETNControlNet.Output("control_net"),
],
)
@classmethod
def execute(cls, model: ModelPatcher, weights: str): # type: ignore[override]
weights_path = folder_paths.get_full_path("controlnet", weights)
if weights_path is None or not os.path.isfile(weights_path):
raise FileNotFoundError(f"LLLite weights not found: {weights}")
# Architecture is fully determined by the trained weights — read everything
# from metadata rather than exposing knobs that would just cause load errors.
meta = read_lllite_metadata(weights_path)
if "lllite.version" not in meta:
raise RuntimeError(
"Unrecognized model. This node currently only loads Anima LLLite weights."
)
ce_dim = int(meta.get("lllite.cond_emb_dim", 32))
m_dim = int(meta.get("lllite.mlp_dim", 64))
# v2 records the canonical atomic form under lllite.target_atomics; fall back
# to the legacy preset key, then to the v1 default.
tl = meta.get("lllite.target_atomics", meta.get("lllite.target_layers", "self_attn_q"))
cond_dim = int(meta.get("lllite.cond_dim", 64))
cond_resblocks = int(meta.get("lllite.cond_resblocks", 1))
use_aspp = str(meta.get("lllite.use_aspp", "false")).lower() == "true"
aspp_dilations_meta = meta.get("lllite.aspp_dilations")
if use_aspp and aspp_dilations_meta:
aspp_dilations = tuple(int(d) for d in aspp_dilations_meta.split(",") if d.strip())
else:
aspp_dilations = ASPP_DEFAULT_DILATIONS
cond_in_channels = int(meta.get("lllite.cond_in_channels", 3))
inpaint_masked_input = (
str(meta.get("lllite.inpaint_masked_input", "false")).lower() == "true"
)
lllite = ControlNetLLLiteDiT(
_get_inner_dit(model),
cond_emb_dim=ce_dim,
mlp_dim=m_dim,
target_layers=tl,
multiplier=1.0,
cond_dim=cond_dim,
cond_resblocks=cond_resblocks,
use_aspp=use_aspp,
aspp_dilations=aspp_dilations,
cond_in_channels=cond_in_channels,
inpaint_masked_input=inpaint_masked_input,
)
load_lllite_weights(lllite, weights_path, strict=False)
lllite.eval().requires_grad_(False)
return io.NodeOutput(model, lllite)
class ControlApply(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_control_apply",
display_name="Apply ControlNet (tooling-nodes)",
description="Applies ControlNet conditioning. Currently only supports Anima LLLite weights.",
category="external_tooling",
inputs=[
io.Model.Input("model"),
ETNControlNet.Input("control_net"),
io.Image.Input("image"),
io.Mask.Input("mask", optional=True),
io.Float.Input("strength", default=1.0, min=-10.0, max=10.0, step=0.01),
io.Float.Input("start_percent", default=0.0, min=0.0, max=1.0, step=0.001),
io.Float.Input("end_percent", default=1.0, min=0.0, max=1.0, step=0.001),
],
outputs=[io.Model.Output("model")],
)
@classmethod
def execute( # type: ignore[override]
cls,
model: ModelPatcher,
control_net: ControlNetLLLiteDiT,
image: torch.Tensor,
strength: float,
start_percent: float,
end_percent: float,
mask: torch.Tensor | None = None,
):
dit = _get_inner_dit(model)
patch_spatial = int(getattr(dit, "patch_spatial", 2))
lllite = control_net
lllite.set_multiplier(strength)
# Mask / cond_in_channels consistency: 4ch weights need a MASK, 3ch weights ignore it.
if lllite.cond_in_channels == 4 and mask is None:
raise ValueError("ControlNet weights require a mask input (inpaint mode)")
if lllite.cond_in_channels != 4 and mask is not None:
mask = None
# Convert percent range -> sigma range (start_percent=0 → sigma_max).
model_sampling = model.get_model_object("model_sampling")
sigma_start = float(model_sampling.percent_to_sigma(start_percent))
sigma_end = float(model_sampling.percent_to_sigma(end_percent))
# Capture image / mask tensors (cloned to detach from any upstream caching)
src_image = image.detach().clone()
src_mask = mask.detach().clone() if mask is not None else None
is_inpaint = lllite.cond_in_channels == 4
# Cache for the per-resolution preprocessed cond image (avoids repeat resize)
cache: dict[str, Any] = {"cond_image_pp": None, "key": None, "lllite_loaded_to": None}
# Capture any previously-installed wrapper BEFORE we clone — model_options
# has a single "model_function_wrapper" slot, so without delegation a second
# wrapper-installing node would silently no-op the first. Mirrors the
# ChromaRadianceOptions pattern in comfy_extras/nodes_chroma_radiance.py.
old_wrapper = model.model_options.get("model_function_wrapper")
def _call_next(apply_model, input_x, timestep, c):
if old_wrapper is not None:
return old_wrapper(apply_model, {"input": input_x, "timestep": timestep, "c": c})
return apply_model(input_x, timestep, **c)
def wrapper(apply_model, args):
input_x = args["input"]
timestep = args["timestep"]
c = args["c"]
# Step-range gate: skip LLLite entirely when current sigma is outside
# [sigma_end, sigma_start]. percent_to_sigma maps 0.0 → sigma_max,
# 1.0 → sigma_min, so the active window is sigma_end <= sigma <= sigma_start.
sigma = float(timestep.max().item())
if not (sigma_end <= sigma <= sigma_start):
return _call_next(apply_model, input_x, timestep, c)
# Anima latent shape: (B, C, T, H, W) — take spatial dims from the tail.
latent_h, latent_w = int(input_x.shape[-2]), int(input_x.shape[-1])
device = input_x.device
dtype = input_x.dtype
# Move LLLite to the runtime device/dtype lazily.
tag = (device, dtype)
if cache["lllite_loaded_to"] != tag:
lllite.to(device=device, dtype=dtype)
cache["lllite_loaded_to"] = tag
cache["cond_image_pp"] = None # invalidate
key = (latent_h, latent_w, device, dtype)
if cache["key"] != key or cache["cond_image_pp"] is None:
rgb = _prepare_cond_image(
src_image, latent_h, latent_w, device, dtype, patch_spatial
)
if is_inpaint:
assert src_mask is not None, "Cannot use inpaint control-net without a mask"
mk = _prepare_mask(src_mask, latent_h, latent_w, device, dtype, patch_spatial)
cache["cond_image_pp"] = _build_inpaint_cond_image(
rgb, mk, lllite.inpaint_masked_input
)
else:
cache["cond_image_pp"] = rgb
cache["key"] = key
lllite.set_multiplier(strength)
lllite.set_cond_image(cache["cond_image_pp"])
lllite.apply_to()
try:
return _call_next(apply_model, input_x, timestep, c)
finally:
lllite.restore()
lllite.clear_cond_image()
m = model.clone()
m.set_model_unet_function_wrapper(wrapper)
return (m,)
+65 -24
View File
@@ -32,7 +32,6 @@ function loadImage(base64) {
} }
const canvasIcon = loadImage("data:image/webp;base64,UklGRg4KAABXRUJQVlA4WAoAAAAQAAAAYwAAYwAAQUxQSNsDAAARoIRs/yI5+uAHU1kZiOu6u7u7xnObuK9NpBoKCoqi1jd6WonrsD7ROc2e3OJJr28T94aGhobYj+8w9v/X/7+n3UNETAD+b1IW1crLLrMhbYJbRs9Zv2VvuXbqVK28d8v6OaNvCdpIXkaT5InJgJgRABeNbirRYKlp9EUAJB8ZlUq2XAtI1wQIhn1RIUlV1Y5UVUmy0jwsACQPt7CtsloQSBcE6D6jSFKVRlVJFmd0B8QeClSSSn7/ACCdQjBjL0mlRSW5d0YA+4KESpJKnd8dHQsweBeptK7krsGAWIIgppKkkn+NAKSNoO9KUplLJVf2hliCIKWyrZIfXwYIgCdLzHXpSVgXvElt07bcKAAKSs2TUgvWIHiH2p6SX98pi+jgIrEFwUJqO6Ty1HY6qGwObEHwPrU9KqkOUNkS2AKwjNoeqXRSuV7EGtZQO3BVuQhiDV9Q3aIyhFiTDVS3qHwMYgtBK9Utcn932A/20nHlGoglwdOn1DEqn4XYQfA7PfB7AKuCRiqdVzZCbKD+EL14qB4WBTOoPlA2Qsyh7i9f/FUH44LBVHpRORhibr0/ms2hf5XerPaFYcFoqi+UDeaafNJkCsFeenRvnRnBLfSp3mBqNNUfygZTc/wyx9QGv6w3A2zxyxZDwV56da+h3mW/lA1dVPPLKUOXnfIL/8UurvnltKH+Vb9UDXU/4ZeyoaDkl5Ih/ET1h/InU81+aTYjmOOXOaYm06sTzAD3n/JJ7X5T/YtUXyiLfU3JGnp0jZjCS+oPfQnG7ylR/aAs3WmuexO92dTdHKZWfFGZCot3fkX1gbL1dhv1SZVerCb1NvB0K9U9ZevTsNo/OUEPnkj628Hgj33w8WBY7h9up7ql3B72tYUn51ToeHnOk7B+/tSV6pYum3q+PVwbtVDdUbZE1yKPzyY/UV1R/hg/i1zKhKRIdUNZTCZIPtB9Rvo71QXl78mM7sjrxVFSpJPFJLoY+b06jn904ccovhp5vjaOWkjNk5Ithfha5PuKqLCozFyXFxWiK5D3/o1xuiVPW9K4sT/yf35DFi47mpejy8Ks4Xw4+UAcxRvKeaisj6P4Abjaf3xWSDYcbaNmtM3RDUkhG98fDt8+IyvEK3fV2Fa1M6psW9u1Mi5kM26H28EDM7IofPPj7WUaLG//+M0wymY8EMD54M7JaRqF8cKPvy4eqtROq56uVQ4Vv/54YRxGaTr5zgB+vOjpl5IsicIwSrJ35sx5J0uiMIySLHnp6YvgUel/z6iXojTL0jRJ0jTL0uilUff0F/i3/qIb7nn06WHDnn70nhsuqhf8hxQAVlA4IAwGAACwHgCdASpkAGQAPm0wk0akIqGhLRGrUIANiWYA1BHh/t2rC93/Hf2Was/feJ2MrzB6VP6M9gD9Nunp5rP2y9Z70q+gV/Y/+B1kHoAeWv+1Xwa/t5+5PtQXQfhgKtwfp++wwU1Dx6fSXsD+VV7JPQ5/aRrblS2nsaggvO1Mch8UhK6pYtQxLM/VgrswZ0vLV8b6SwudSWaCFSHUiXEUQWX6krc9GtWanHMeaDd9wRYCfO5TwpYkgGAIkaLI4p6taB375EUfaVubYzKMfHSz2KpivsjWF0Vf+YJbACgi8j86d6EiJhQFF31NBJdS+QrGtJ2RUJRbahp1MXso6/J8AAD+/TKL/9q5tf/zOmF5Fe8B0Zn0yX3C0VLv0zxxv2+WH/dbz//rc2RS4TC1UzFVQiXVn5+Y0r+RsfJPsfPNuT02INz8gty7fI7fA/D1Wj2Jv+4RwdpyXs+cRxaT84bme5rMmPf+BH7NDUPKsj7GJ+w/6nBW2vsiPalWPfvBk6AQ3kCHmVecXkcnOgpoZ4ruAF/9Ze93DG5/8Y32x8b/CKPRt1jaXXy2LnoPvSNUT77gbB+/7vI1pfBfUHJsSwheIXY7QSixh7Ya8IliO3wqvI/uIFZAZd9pL8R1gRpYouBoyL5uIuGWQAZC5SKY0SruTf66stUOJVO9hlokeb5lWVzo7FO/Oeb/oj9iK4bqFhNZLCfqsBlH/OeefoP9sFdl7Mq1xmsevmzkfgwyiXg5hxMIP/Wa0JMPVl+XEFqTveAf1M8IBDu/pX/hCEnMn1n15Smyf72eDXKQqBrvp6BugyXXaJ05FDoz8MONUFh4rcjGL7AcijbcZ0SYwJkoeAKBW/I/sjKzTRtTP2E1fLB/8TWnzieHznDAKdlTuY2nSVTwCqFZNcFeFn7boziHOmYBLJin52d874mq1pHmJnulhT96LbKVW4vAT5PnY5F9TzmnDMwIFm4IAuEaA8X8XLE4Hp+AUEG4oswxRbVfOfxNJRyxFO3UB+v+ALgMP8kOf0uK3/3WOq4o/roivfzvW/fXviTC0mx+352hGaO+axx6vFa3eIkFsUEXCdo2LFHIlM8BtPuGUhgvM3oygIMAgvmUKILe0DFYVXhG/QoLi3sYfaoK/f0tX+fNnXhhxwEj1/Ct2Z64g0qWmkgwNkyy8m90EK1HsX0Q10CHVakDZePz5ts37u3GCANwGHQzWB+hNsevqjuU3qT95yGs0jjOtI/IjKsH9JbAmZkjGvNCPOC+FYUkOkwao9sOESY6zCgx9CM7g2LU4/CSHGoe2t0vWV/cMDH1HzI+Wa/yYp9CLDIh7J7iJd/2KnixeJvOhbUvbr9gubyyQU1iO5bnD9T536j++jKDVIk0Fwzk+d+j2eueHsIFJUvdyo2TyxP0kJbWr36R1s3giryqPvrsR5SkXx16+xqDrX4elhqh+1FwzNnSF5Lj5EUT/UC2rJvoAikbnvQ3NtJ9e83++idf3ja4FaLcUDxhoN5Rl5Ziz1LvF9iVeb6Su0QWYoRyBbyZ/pRbgYyhlAU/tonH7Wt+KhPDmXKIo0u4FDbAXM8avbFk4ax6e/dYITOCe+9dVEgcTOnBfhv0Yotd3EzNjZkLz4ksKGtFXcWIZRJ5YAyfzPYsyPex6/6ud9r2Ha9oxhVSIJV418e83qcPOIPlpe+LVGc69W6eC83l/zloqM9D6zQMkfqrjBZNpRkQS0sn8sxSu3s5qzhtH8cvjZk83gMqdfnHnl+1bvA7BI/g4+ePU7HUb9vK3Qw35bVmDcXa8xxWS2NQj8iWMH1cbHLXlboQsaCxIZoo+SeXR6ePUw3k6C/OxgqhjzExMJjLdBjoBeWYt3RPG2foTvx0T0Iz8ukdrCRJMG6HaR+6/f4nG/4xkr/fLhGlqOE/hBDBhuqnANj1CrujVDs2YayTvPuIcqCpNd3i8fOR8DfCq9ytS55F8akKneS6poHfB3bhjWbcIQXPvFzS7S5xLHEVoaixOwp0TL/8cQ8dxriyeddu5kCTyY7KepMQoeR+Pyn04nElkt9qqfYCTqHDtXBriC/UZh9AAAAAAAA=") const canvasIcon = loadImage("data:image/webp;base64,UklGRg4KAABXRUJQVlA4WAoAAAAQAAAAYwAAYwAAQUxQSNsDAAARoIRs/yI5+uAHU1kZiOu6u7u7xnObuK9NpBoKCoqi1jd6WonrsD7ROc2e3OJJr28T94aGhobYj+8w9v/X/7+n3UNETAD+b1IW1crLLrMhbYJbRs9Zv2VvuXbqVK28d8v6OaNvCdpIXkaT5InJgJgRABeNbirRYKlp9EUAJB8ZlUq2XAtI1wQIhn1RIUlV1Y5UVUmy0jwsACQPt7CtsloQSBcE6D6jSFKVRlVJFmd0B8QeClSSSn7/ACCdQjBjL0mlRSW5d0YA+4KESpJKnd8dHQsweBeptK7krsGAWIIgppKkkn+NAKSNoO9KUplLJVf2hliCIKWyrZIfXwYIgCdLzHXpSVgXvElt07bcKAAKSs2TUgvWIHiH2p6SX98pi+jgIrEFwUJqO6Ty1HY6qGwObEHwPrU9KqkOUNkS2AKwjNoeqXRSuV7EGtZQO3BVuQhiDV9Q3aIyhFiTDVS3qHwMYgtBK9Utcn932A/20nHlGoglwdOn1DEqn4XYQfA7PfB7AKuCRiqdVzZCbKD+EL14qB4WBTOoPlA2Qsyh7i9f/FUH44LBVHpRORhibr0/ms2hf5XerPaFYcFoqi+UDeaafNJkCsFeenRvnRnBLfSp3mBqNNUfygZTc/wyx9QGv6w3A2zxyxZDwV56da+h3mW/lA1dVPPLKUOXnfIL/8UurvnltKH+Vb9UDXU/4ZeyoaDkl5Ih/ET1h/InU81+aTYjmOOXOaYm06sTzAD3n/JJ7X5T/YtUXyiLfU3JGnp0jZjCS+oPfQnG7ylR/aAs3WmuexO92dTdHKZWfFGZCot3fkX1gbL1dhv1SZVerCb1NvB0K9U9ZevTsNo/OUEPnkj628Hgj33w8WBY7h9up7ql3B72tYUn51ToeHnOk7B+/tSV6pYum3q+PVwbtVDdUbZE1yKPzyY/UV1R/hg/i1zKhKRIdUNZTCZIPtB9Rvo71QXl78mM7sjrxVFSpJPFJLoY+b06jn904ccovhp5vjaOWkjNk5Ithfha5PuKqLCozFyXFxWiK5D3/o1xuiVPW9K4sT/yf35DFi47mpejy8Ks4Xw4+UAcxRvKeaisj6P4Abjaf3xWSDYcbaNmtM3RDUkhG98fDt8+IyvEK3fV2Fa1M6psW9u1Mi5kM26H28EDM7IofPPj7WUaLG//+M0wymY8EMD54M7JaRqF8cKPvy4eqtROq56uVQ4Vv/54YRxGaTr5zgB+vOjpl5IsicIwSrJ35sx5J0uiMIySLHnp6YvgUel/z6iXojTL0jRJ0jTL0uilUff0F/i3/qIb7nn06WHDnn70nhsuqhf8hxQAVlA4IAwGAACwHgCdASpkAGQAPm0wk0akIqGhLRGrUIANiWYA1BHh/t2rC93/Hf2Was/feJ2MrzB6VP6M9gD9Nunp5rP2y9Z70q+gV/Y/+B1kHoAeWv+1Xwa/t5+5PtQXQfhgKtwfp++wwU1Dx6fSXsD+VV7JPQ5/aRrblS2nsaggvO1Mch8UhK6pYtQxLM/VgrswZ0vLV8b6SwudSWaCFSHUiXEUQWX6krc9GtWanHMeaDd9wRYCfO5TwpYkgGAIkaLI4p6taB375EUfaVubYzKMfHSz2KpivsjWF0Vf+YJbACgi8j86d6EiJhQFF31NBJdS+QrGtJ2RUJRbahp1MXso6/J8AAD+/TKL/9q5tf/zOmF5Fe8B0Zn0yX3C0VLv0zxxv2+WH/dbz//rc2RS4TC1UzFVQiXVn5+Y0r+RsfJPsfPNuT02INz8gty7fI7fA/D1Wj2Jv+4RwdpyXs+cRxaT84bme5rMmPf+BH7NDUPKsj7GJ+w/6nBW2vsiPalWPfvBk6AQ3kCHmVecXkcnOgpoZ4ruAF/9Ze93DG5/8Y32x8b/CKPRt1jaXXy2LnoPvSNUT77gbB+/7vI1pfBfUHJsSwheIXY7QSixh7Ya8IliO3wqvI/uIFZAZd9pL8R1gRpYouBoyL5uIuGWQAZC5SKY0SruTf66stUOJVO9hlokeb5lWVzo7FO/Oeb/oj9iK4bqFhNZLCfqsBlH/OeefoP9sFdl7Mq1xmsevmzkfgwyiXg5hxMIP/Wa0JMPVl+XEFqTveAf1M8IBDu/pX/hCEnMn1n15Smyf72eDXKQqBrvp6BugyXXaJ05FDoz8MONUFh4rcjGL7AcijbcZ0SYwJkoeAKBW/I/sjKzTRtTP2E1fLB/8TWnzieHznDAKdlTuY2nSVTwCqFZNcFeFn7boziHOmYBLJin52d874mq1pHmJnulhT96LbKVW4vAT5PnY5F9TzmnDMwIFm4IAuEaA8X8XLE4Hp+AUEG4oswxRbVfOfxNJRyxFO3UB+v+ALgMP8kOf0uK3/3WOq4o/roivfzvW/fXviTC0mx+352hGaO+axx6vFa3eIkFsUEXCdo2LFHIlM8BtPuGUhgvM3oygIMAgvmUKILe0DFYVXhG/QoLi3sYfaoK/f0tX+fNnXhhxwEj1/Ct2Z64g0qWmkgwNkyy8m90EK1HsX0Q10CHVakDZePz5ts37u3GCANwGHQzWB+hNsevqjuU3qT95yGs0jjOtI/IjKsH9JbAmZkjGvNCPOC+FYUkOkwao9sOESY6zCgx9CM7g2LU4/CSHGoe2t0vWV/cMDH1HzI+Wa/yYp9CLDIh7J7iJd/2KnixeJvOhbUvbr9gubyyQU1iO5bnD9T536j++jKDVIk0Fwzk+d+j2eueHsIFJUvdyo2TyxP0kJbWr36R1s3giryqPvrsR5SkXx16+xqDrX4elhqh+1FwzNnSF5Lj5EUT/UC2rJvoAikbnvQ3NtJ9e83++idf3ja4FaLcUDxhoN5Rl5Ziz1LvF9iVeb6Su0QWYoRyBbyZ/pRbgYyhlAU/tonH7Wt+KhPDmXKIo0u4FDbAXM8avbFk4ax6e/dYITOCe+9dVEgcTOnBfhv0Yotd3EzNjZkLz4ksKGtFXcWIZRJ5YAyfzPYsyPex6/6ud9r2Ha9oxhVSIJV418e83qcPOIPlpe+LVGc69W6eC83l/zloqM9D6zQMkfqrjBZNpRkQS0sn8sxSu3s5qzhtH8cvjZk83gMqdfnHnl+1bvA7BI/g4+ePU7HUb9vK3Qw35bVmDcXa8xxWS2NQj8iWMH1cbHLXlboQsaCxIZoo+SeXR6ePUw3k6C/OxgqhjzExMJjLdBjoBeWYt3RPG2foTvx0T0Iz8ukdrCRJMG6HaR+6/f4nG/4xkr/fLhGlqOE/hBDBhuqnANj1CrujVDs2YayTvPuIcqCpNd3i8fOR8DfCq9ytS55F8akKneS6poHfB3bhjWbcIQXPvFzS7S5xLHEVoaixOwp0TL/8cQ8dxriyeddu5kCTyY7KepMQoeR+Pyn04nElkt9qqfYCTqHDtXBriC/UZh9AAAAAAAA=")
const outputIcon = loadImage("data:image/webp;base64,UklGRrIHAABXRUJQVlA4WAoAAAAQAAAAjwAAOwAAQUxQSKoCAAARkMbsnyFJ/2TVySSdzPJs27bxzbbtuznbtm3btm3bt+ikkkonlfxPU1X9n57zh4iYAPifZKQ2T7LIkKDkQ1FPrezCF/hjvq+dx7mwiqM3HnJ058yGXkorEfGCKZV6I6o+quOM4bMsgU5bfE1KOh4LEbGnv2RnUSer50DGFwxJ2qwWGTATERFfZPzBZNR9P1ZXTksgVdaODHgt/H5hCJjP0ME6eiLfCqTLypKBSPYdsmY2OjpWy0KOlF+EkIFk/DvnZ2tIzZA0a0kHUtok0KfWkxieJQQZBQksq3QWiXODEGSnYRsq76mxjJQK0cDdKjY1XpLSWyKYVYFTSyxLqCZSvexKppZnZDCZTNpD6BKTi2iIRbqzZc6gW8x5m1ptOCEuYaJ74B1T6RkhjPQXkugiuC9EBSk38wftRACZsdKLEHGRQjJSCyUgnwichSgdj4jYW64KqQsywN1E1JJqReqtWyEvLtOTFHMtfG8GXMYT6C5zQLIjqXiJm+guNw2ZOqTu+0uJ7sKyg2w+Urv9GdxdGoN0JKnh/qBnIE1LlH6LiHNAMZ5SWQkoJAJHcQ7iTUNlNyE7Vga4a7CsoFqP0AmQjqdm5dTWGJRTvqfTSu4ONe7VNRs0LiTzLK3cNEHsGWhuZ+gom0hlNMgXZ7T4WF16Q6YRuZZPAS7TYrGUoPgFEnZPUC3EKDEf0O6WSGFmMiVoxejw3UDcHC6c21oFNLZjVNhGgxrkHC+cOtQKtBa5aQkCLL4VBGDZsbYzu7uF6AGouPQ1t7mTIv5AMw8EZLmhbx0QizqHgJMer5MQwHkG7NP2YmgzCMqRPYff0cIWDSgOwbrcgOFnhcqLhQNaeSB4h1XxDZh29LX4kXVt5YABrZJBkM/acIBsx5IG/AyGpS1UrkrF4tm98K8uVlA4IOIEAABwHACdASqQADwAPm0uk0ckIiGhLjUJmIANiWgOuBpEsADI+tE9j+ifkl+QHyd1n+77tGXf04+Cpx/549gD9X+lF5t/2d9cn0T+gB/d/8R1gHoAeWx+3nwZfuZ6T+aq/0Dtm/zaCN5VfOk9UfrX8A/Ry9E1elWp+7DqL5PqqU3p7QQgEkexRJ8VwS4d4Xk5cyjjrDbvzKxXEwCzN96k0RfNIpiQ2YsSbcyoRxX5fFzllf+dOy8uCMy/ebkjS0ONwuzkRR55zP3zjA4e+C969ch1Ab9LgcrpUqQ8MOvEXnQmqQkxXM4x4nA1f+jJAAD+8tXq96F285yEhMGbOWPp352/yRnzPmWRyRibmd800tluUOW4IyIZz2Hw1xYA9/xsSgKy0yQS//7BbNldSPJ+MCX/mxBqrttfeQX/mf/+2AcZ7Z1wDrdNoOnt8ISIu33p34GUAqqEyPwtdrhMf55SjsQmwUtm/I/qPiQ0ZOPv5ci6kCP0Ddb2jRr38UOXvi54DlMMmkxTs7j/J3jUQYepc7xEgTVVJFcf+8P//18L/9gA//9fHuK11sTUs+RbYDzZKn0uM2PnTEUpAJAT4wETKSU4KDn3wR3rUGGycaAPVo40AjzO5g7VkynMJuo3M2vclcmfmS3ygBGDqjGHQybd03tnGdkGOCGLLlTEABV+UgJxYxH2YxRX9zUULXFnfnBpPxQcOdq+zs1zi9uI2iAqSKwh8Dhch2Ytz8iaZLW39S3+3pmGdITR49+nlHjcG4xNVSYRLFLRmEj/H/I+7qd90N6AF9aRDuUFH1O7ONRGjEQGvPMEF0Fj5atb5w9tjc1pcKTsaWvT2GbF9NQ31HmpaLAgs8szVbuEC8GHKCESxKmx+Hrh5ZwjrNihG3KL0H1n3/g/WetlSYEFYsYTXQmgyUGCVIILkJYrRIdLB5iVAPrseYWKCT8HgJuCUAhaqRO+6jn1fkplsC0yYCveVI+yyDsVr98kmO5arhQ3u+aqKVUvJ8xZL9as4008lN9DkKcRhvC4BwWdhupsqUYwLQmaQhLxP15875P/r43c8r4NI4sLDiCi7Rzww1dWNTyThiA07x8b/zTaFC9Sz+jtZDpRPoSf3LS+TmvHZQ+yv/N9nSK/CGpimH/qjTJOQRStf5ppvzT0FzGMX2tqNndJbZD8idLxJFXZekFF16KC/6scsX/lTNL+XFfTqsreVXu7bL/wjNVTPeGkJJE7aWcXP2+3qQTMv+LaO9INAsG3cyp5Co/F06O8XoVtYZXBjH3f3r9Y8Wp89/fqq2OfQSD2/Ujo1t0fNnMA14gpYdtm6+/RcRgNQGIPPGxgAaFjsfC4+63CcHr1nczuKyXiQjmoIH7n/0NCmJv3O+v/Lp30d3n/060TaO5ffQGrrx0O7TYUAC6pdQxfOeuX4/EsKJgMKTW18feF5m1SX4ODnH1SWutwnm5T/k0/l4YXbLUi8QRbdtx74QL9DJRtKP8bDT+yyf//8nAEApoAJH7jMHoQv7XKzIUdH1TDS7Phokc3PP5m68+eUTHU17v50avNmnEHCfybI4FC35LTpSaGqRsgNJJliiV37VIfbUlfDfgIqZmmxEHmnCQTSg2zcf5+9LWPblYTxLShx/2U34N3Rf/3Zvie6j8SS/8X+Yo+dXhKIyg1WX040AAAAA==")
function setIconImage(nodeType, image, size, padRows, padCols) { function setIconImage(nodeType, image, size, padRows, padCols) {
const onAdded = nodeType.prototype.onAdded const onAdded = nodeType.prototype.onAdded
@@ -72,36 +71,70 @@ const parameterTypes = {
"text": ["text", "prompt (positive)", "prompt (negative)"], "text": ["text", "prompt (positive)", "prompt (negative)"],
} }
function changeWidget(widget, type, value, options) { function defaultParameterType(widgetType, connectedNode, connectedWidget) {
widget.type = type let paramType = parameterTypes[widgetType][0]
widget.value = value if (connectedNode.comfyClass === "CLIPTextEncode") {
widget.options = options paramType = "prompt (positive)"
}
const round = connectedWidget.options?.round
if ((paramType == "number" && round === undefined) || round === 1) {
paramType = "number (integer)"
}
return paramType
} }
function changeWidgets(node, type, value, options) { function valueMatchesType(value, type, options) {
if (type === "number") {
return typeof value === "number"
} else if (type === "combo") {
return options?.values?.includes(value)
} else if (type === "toggle") {
return typeof value === "boolean"
}
return typeof value === "string"
}
function optionalWidgetValue(widgets, index, fallback) {
const result = widgets.length > index ? widgets[index].value : null
return result === null || result === -1e10 || result === 1e10 ? fallback : result
}
function changeWidgets(node, type, connectedNode, connectedWidget) {
if (type === "customtext") { if (type === "customtext") {
type = "text" type = "text"
} }
const options = connectedWidget.options
node.widgets[1].value = parameterTypes[type][0] const parameterTypeHint = node.widgets[1].value
node.widgets[1].options = {values: parameterTypes[type]} const notSpecialized = node.widgets[1].options.values.includes("auto")
changeWidget(node.widgets[2], type, value, options) const parameterTypeMismatch = !parameterTypes[type].includes(parameterTypeHint)
if (notSpecialized || parameterTypeMismatch) {
node.widgets[1].options = {values: parameterTypes[type]}
}
if (parameterTypeMismatch) {
node.widgets[1].value = defaultParameterType(type, connectedNode, connectedWidget)
}
const oldDefault = node.widgets.length > 2 ? node.widgets[2].value : connectedWidget.value
const oldMin = optionalWidgetValue(node.widgets, 3, options?.min ?? 0)
const oldMax = optionalWidgetValue(node.widgets, 4, options?.max ?? 100)
const isDefaultValid = valueMatchesType(oldDefault, type, connectedWidget.options)
while (node.widgets.length > 2) {
node.widgets.pop()
}
const value = isDefaultValid && oldDefault !== "" ? oldDefault : connectedWidget.value
node.addWidget(type, "default", value, null, options)
if (type === "number") { if (type === "number") {
changeWidget(node.widgets[3], "number", options?.min ?? 0, options) node.addWidget("number", "min", oldMin, null, options)
changeWidget(node.widgets[4], "number", options?.max ?? 100, options) node.addWidget("number", "max", oldMax, null, options)
} else {
changeWidget(node.widgets[3], "number", 0, {min: 0, max: 0})
changeWidget(node.widgets[4], "number", 0, {min: 0, max: 0})
} }
} }
function adaptWidgetsToConnection(node) { function adaptWidgetsToConnection(node) {
if (!node.outputs || node.outputs.length === 0 || !node.outputs[0].links) { if (!node.outputs || node.outputs.length === 0) {
return return
} }
const links = node.outputs[0].links const links = node.outputs[0].links
if (links.length === 1) { if (links && links.length === 1) {
const link = node.graph.links[links[0]] const link = node.graph.links[links[0]]
if (!link) return if (!link) return
@@ -109,7 +142,7 @@ function adaptWidgetsToConnection(node) {
if (!theirNode || !theirNode.inputs) return if (!theirNode || !theirNode.inputs) return
const input = theirNode.inputs[link.target_slot] const input = theirNode.inputs[link.target_slot]
if (!input) return if (!input || !input.widget || theirNode.widgets === undefined) return
node.outputs[0].type = input.type node.outputs[0].type = input.type
@@ -119,11 +152,15 @@ function adaptWidgetsToConnection(node) {
const widgetName = input.widget.name const widgetName = input.widget.name
const theirWidget = theirNode.widgets.find((w) => w.name === widgetName) const theirWidget = theirNode.widgets.find((w) => w.name === widgetName)
const widgetType = theirWidget.origType ?? theirWidget.type if (!theirWidget) return // connected to a custom node that isn't installed
changeWidgets(node, widgetType, theirWidget.value, theirWidget.options)
} else if (links.length === 0) { const widgetType = theirWidget.origType ?? theirWidget.type
changeWidgets(node, widgetType, theirNode, theirWidget)
} else if (!links || links.length === 0) {
node.outputs[0].type = "*" node.outputs[0].type = "*"
node.widgets[1].value = "auto"
node.widgets[1].options = {values: ["auto"]}
} }
} }
@@ -169,10 +206,14 @@ app.registerExtension({
beforeRegisterNodeDef(nodeType /*typeof LGraphNode*/, nodeData /*ComfyObjectInfo*/, app) { beforeRegisterNodeDef(nodeType /*typeof LGraphNode*/, nodeData /*ComfyObjectInfo*/, app) {
if (nodeData.name === "ETN_KritaCanvas") { if (nodeData.name === "ETN_KritaCanvas") {
setIconImage(nodeType, canvasIcon, [200, 100], 0, 2) setIconImage(nodeType, canvasIcon, [200, 100], 0, 2)
} else if (nodeData.name === "ETN_KritaOutput") { } else if (nodeData.name === "ETN_Parameter") {
setIconImage(nodeType, outputIcon, [200, 120], 2, 0)
} else if (nodeData.name == "ETN_Parameter") {
setupParameterNode(nodeType) setupParameterNode(nodeType)
} else if (nodeData.name === "ETN_SendText") {
const onAdded = nodeType.prototype.onAdded
nodeType.prototype.onAdded = function() {
onAdded?.apply(this, arguments)
this.inputs[0].type = "*"
}
} }
}, },
+286 -118
View File
@@ -1,11 +1,16 @@
import torch import sys
import numpy as np from enum import Enum
from pathlib import Path from pathlib import Path
from typing import NamedTuple from typing import Any, NamedTuple
import comfy.samplers
import numpy as np
import server
import torch
from comfy.comfy_types.node_typing import IO
from comfy_api.latest import io
from PIL import Image from PIL import Image
import server
import comfy.samplers
from .nodes import SendImageWebSocket from .nodes import SendImageWebSocket
@@ -34,8 +39,11 @@ class WorkflowExchange:
for publisher in self._publishers.values(): for publisher in self._publishers.values():
await self._notify(client_id, publisher) await self._notify(client_id, publisher)
def unsubscribe(self, client_id: str): async def unsubscribe(self, client_id: str):
self._subscribers.remove(client_id) if client_id in self._subscribers:
self._subscribers.remove(client_id)
else:
raise KeyError("No subscriber found with id " + client_id)
async def _notify(self, client_id: str, publisher: Publisher): async def _notify(self, client_id: str, publisher: Publisher):
data = { data = {
@@ -52,75 +60,204 @@ def _placeholder_image():
return torch.from_numpy(image)[None,] return torch.from_numpy(image)[None,]
class KritaOutput(SendImageWebSocket): class _BasicTypes(str):
RETURN_TYPES = () """Matches IO.PRIMITIVE, but also any list of choices"""
FUNCTION = "send_images"
OUTPUT_NODE = True basic_types = IO.PRIMITIVE.split(",") # STRING, FLOAT, INT, BOOLEAN
CATEGORY = "krita"
def __eq__(self, other):
return other in self.basic_types or isinstance(other, (list, _BasicTypes))
def __ne__(self, other):
return not self.__eq__(other)
class KritaCanvas: BasicTypes = _BasicTypes("BASIC")
class OutputBatchMode(Enum):
default = "default"
images = "images"
animation = "animation"
layers = "layers"
class KritaOutput(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(cls): def define_schema(cls):
return {} return io.Schema(
node_id="ETN_KritaOutput",
display_name="Krita Output",
category="krita",
inputs=[
io.Image.Input("images"),
io.Int.Input("x", "offset x", default=0),
io.Int.Input("y", "offset y", default=0),
io.String.Input("name", default=""),
io.Combo.Input(
"batch_mode", OutputBatchMode, "batch mode", default=OutputBatchMode.default
),
io.Boolean.Input("resize_canvas", "resize canvas", default=False),
],
is_output_node=True,
)
RETURN_TYPES = ("IMAGE", "INT", "INT", "INT")
RETURN_NAMES = ("image", "width", "height", "seed")
FUNCTION = "placeholder"
CATEGORY = "krita"
def placeholder(self):
return (_placeholder_image(), 512, 512, 0)
class KritaSelection:
@classmethod @classmethod
def INPUT_TYPES(cls): def execute( # type: ignore
return {} cls,
images: torch.Tensor,
RETURN_TYPES = ("MASK",) x: int = 0,
RETURN_NAMES = ("mask",) y: int = 0,
FUNCTION = "placeholder" name="",
CATEGORY = "krita" batch_mode: OutputBatchMode | str = OutputBatchMode.default,
resize_canvas=False,
def placeholder(self): ):
return (torch.ones(1, 512, 512),) batch_mode = batch_mode.value if isinstance(batch_mode, OutputBatchMode) else batch_mode
info = {
"name": name,
class KritaImageLayer: "offset_x": x,
@classmethod "offset_y": y,
def INPUT_TYPES(cls): "batch_mode": batch_mode,
return { "resize_canvas": resize_canvas,
"required": {
"name": ("STRING", {"default": "Image"}),
}
} }
output = SendImageWebSocket.execute(images, "PNG")
RETURN_TYPES = ("IMAGE",) assert isinstance(output.ui, dict)
RETURN_NAMES = ("image",) output.ui["info"] = [info]
FUNCTION = "placeholder" return output
CATEGORY = "krita"
def placeholder(self, name: str):
return (_placeholder_image(),)
class KritaMaskLayer: class KritaSendText(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(cls): def define_schema(cls):
return { return io.Schema(
"required": { node_id="ETN_KritaSendText",
"name": ("STRING", {"default": "Mask"}), display_name="Send Text",
} category="krita",
} inputs=[
io.AnyType.Input("value"),
io.String.Input("name", default="Output"),
io.Combo.Input("type", options=["text", "markdown", "html"], default="text"),
],
is_output_node=True,
)
RETURN_TYPES = ("MASK",) @classmethod
RETURN_NAMES = ("mask",) def execute(cls, value: Any, name: str, type: str): # type: ignore
FUNCTION = "placeholder" mime = {
CATEGORY = "krita" "text": "text/plain",
"markdown": "text/markdown",
"html": "text/html",
}[type]
text = "None"
if value is not None:
try:
text = str(value)
except Exception as e:
text = f"Could not convert to text: {e}"
def placeholder(self, name: str): return io.NodeOutput(ui={"text": [{"name": name, "text": text, "content-type": mime}]})
return (torch.ones(1, 512, 512),)
class KritaCanvas(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_KritaCanvas",
display_name="Krita Canvas",
category="krita",
outputs=[
io.Image.Output(display_name="image"),
io.Int.Output(display_name="width"),
io.Int.Output(display_name="height"),
io.Int.Output(display_name="seed"),
io.Mask.Output(display_name="mask"),
],
)
@classmethod
def execute(cls, **kwargs):
return io.NodeOutput(_placeholder_image(), 512, 512, 0, torch.ones(1, 512, 512))
class SelectionContext(Enum):
automatic = "automatic"
entire_image = "entire image"
mask_bounds = "mask bounds"
_selection_context_help = """
Determines the section (crop bounding box) of the image and mask to transmit:
- automatic: area around the selection determined by Krita settings
- entire image: always use the entire canvas area
- mask bounds: tight bounding box of the current selection
This affects the Selection and Canvas nodes. The offset x/y outputs indicate the top-left corner of the context area relative to the full canvas."""
class KritaSelection(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_KritaSelection",
display_name="Krita Selection",
category="krita",
inputs=[
io.Combo.Input(
"context",
options=SelectionContext,
default=SelectionContext.entire_image,
tooltip=_selection_context_help,
),
io.Int.Input("padding", "padding", default=0, min=0),
],
outputs=[
io.Mask.Output("mask", "mask"),
io.Boolean.Output("active", "active"),
io.Int.Output("x", "offset x"),
io.Int.Output("y", "offset y"),
],
)
@classmethod
def execute(cls, **kwargs):
return io.NodeOutput(torch.ones(1, 512, 512), False, 0, 0)
class KritaImageLayer(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_KritaImageLayer",
display_name="Krita Image Layer",
category="krita",
inputs=[io.String.Input("name", default="Image")],
outputs=[
io.Image.Output(display_name="image"),
io.Mask.Output(display_name="mask"),
],
)
@classmethod
def execute(cls, name: str): # type: ignore
return io.NodeOutput(_placeholder_image(), torch.ones(1, 512, 512))
class KritaMaskLayer(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_KritaMaskLayer",
display_name="Krita Mask Layer",
category="krita",
inputs=[io.String.Input("name", default="Mask")],
outputs=[
io.Mask.Output(display_name="mask"),
],
)
@classmethod
def execute(cls, name: str): # type: ignore
return io.NodeOutput(torch.ones(1, 512, 512))
_param_types = [ _param_types = [
@@ -133,64 +270,95 @@ _param_types = [
"prompt (positive)", "prompt (positive)",
"prompt (negative)", "prompt (negative)",
] ]
_fmax = sys.float_info.max
class Parameter: class Parameter(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(cls): def define_schema(cls):
return { return io.Schema(
"required": { node_id="ETN_Parameter",
"name": ("STRING", {"default": "Parameter"}), display_name="Parameter",
"type": (_param_types, {"default": "auto"}), category="krita",
"default": ("STRING", {"default": ""}), inputs=[
"min": ("FLOAT", {"default": 0.0}), io.String.Input("name", default="Parameter"),
"max": ("FLOAT", {"default": 1.0}), io.Combo.Input("type", options=_param_types, default="auto"),
} io.String.Input("default", default=""),
} io.Float.Input("min", default=-1e10, min=-_fmax, max=_fmax, optional=True),
io.Float.Input("max", default=1e10, min=-_fmax, max=_fmax, optional=True),
],
outputs=[io.AnyType.Output(display_name="value")],
)
RETURN_TYPES = ("*",)
RETURN_NAMES = ("value",)
FUNCTION = "placeholder"
CATEGORY = "krita"
def placeholder(self, name: str, type: str, default, min, max):
return (default,)
class KritaStyle:
@classmethod @classmethod
def INPUT_TYPES(cls): def execute(cls, name: str, type: str, default, min=0.0, max=1.0): # type: ignore
return { if type == "number":
"required": { return io.NodeOutput(float(default))
"name": ("STRING", {"default": "Style"}), elif type == "number (integer)":
"sampler_preset": (["auto", "regular", "live"],), return io.NodeOutput(int(default))
} return io.NodeOutput(default)
}
RETURN_TYPES = (
"MODEL",
"CLIP",
"VAE",
"STRING",
"STRING",
comfy.samplers.KSampler.SAMPLERS,
comfy.samplers.KSampler.SCHEDULERS,
"INT",
"FLOAT",
)
RETURN_NAMES = (
"model",
"clip",
"vae",
"positive prompt",
"negative prompt",
"sampler name",
"scheduler",
"steps",
"guidance",
)
FUNCTION = "placeholder"
CATEGORY = "krita"
def placeholder(self, name: str, sampler_preset: str): class KritaStyle(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_KritaStyle",
display_name="Krita Style",
category="krita",
inputs=[
io.String.Input("name", default="Style"),
io.Combo.Input("sampler_preset", options=["auto", "regular", "live"]),
],
outputs=[
io.Model.Output(display_name="model"),
io.Clip.Output(display_name="clip"),
io.Vae.Output(display_name="vae"),
io.String.Output(display_name="positive prompt"),
io.String.Output(display_name="negative prompt"),
io.Combo.Output(
display_name="sampler name", options=comfy.samplers.KSampler.SAMPLERS
),
io.Combo.Output(
display_name="scheduler", options=comfy.samplers.KSampler.SCHEDULERS
),
io.Int.Output(display_name="steps"),
io.Float.Output(display_name="guidance"),
],
)
@classmethod
def execute(cls, name: str, sampler_preset: str): # type: ignore
raise NotImplementedError("This workflow must be started from Krita!")
class KritaStyleAndPrompt(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_KritaStyleAndPrompt",
display_name="Krita Style & Prompt",
category="krita",
inputs=[
io.Combo.Input("sampler_preset", options=["auto", "regular", "live"]),
],
outputs=[
io.Model.Output(display_name="model (with loras)"),
io.Clip.Output(display_name="clip"),
io.Vae.Output(display_name="vae"),
io.String.Output(display_name="positive prompt (evaluated)"),
io.String.Output(display_name="negative prompt (evaluated)"),
io.Combo.Output(
display_name="sampler name", options=comfy.samplers.KSampler.SAMPLERS
),
io.Combo.Output(
display_name="scheduler", options=comfy.samplers.KSampler.SCHEDULERS
),
io.Int.Output(display_name="steps"),
io.Float.Output(display_name="guidance"),
],
)
@classmethod
def execute(cls, name: str, sampler_preset: str): # type: ignore
raise NotImplementedError("This workflow must be started from Krita!") raise NotImplementedError("This workflow must be started from Krita!")
+348 -88
View File
@@ -1,30 +1,45 @@
from __future__ import annotations from __future__ import annotations
from PIL import Image
import numpy as np
import base64 import base64
import torch import time
from copy import copy
from dataclasses import dataclass
from io import BytesIO from io import BytesIO
from server import PromptServer, BinaryEventTypes from typing import NamedTuple
from uuid import uuid4
import numpy as np
import torch
import torch.nn.functional as F
from comfy.clip_vision import ClipVisionModel
from comfy.sd import StyleModel
from comfy_api.latest import io
from PIL import Image
from server import BinaryEventTypes, PromptServer
class LoadImageBase64: class LoadImageBase64(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(s): def define_schema(cls):
return {"required": {"image": ("STRING", {"multiline": False})}} return io.Schema(
node_id="ETN_LoadImageBase64",
display_name="Load Image (Base64)",
category="external_tooling",
inputs=[io.String.Input("image", multiline=False)],
outputs=[io.Image.Output(display_name="image"), io.Mask.Output(display_name="mask")],
)
RETURN_TYPES = ("IMAGE", "MASK") @classmethod
CATEGORY = "external_tooling" def execute(cls, image: str): # type: ignore
FUNCTION = "load_image" _strip_prefix(image, "data:image/png;base64,")
def load_image(self, image):
imgdata = base64.b64decode(image) imgdata = base64.b64decode(image)
img = Image.open(BytesIO(imgdata)) img = Image.open(BytesIO(imgdata))
if "A" in img.getbands(): if "A" in img.getbands():
mask = np.array(img.getchannel("A")).astype(np.float32) / 255.0 mask = np.array(img.getchannel("A")).astype(np.float32) / 255.0
mask = 1.0 - torch.from_numpy(mask) mask = torch.from_numpy(mask)[None,]
else: else:
mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu") mask = None
img = img.convert("RGB") img = img.convert("RGB")
img = np.array(img).astype(np.float32) / 255.0 img = np.array(img).astype(np.float32) / 255.0
@@ -33,16 +48,20 @@ class LoadImageBase64:
return (img, mask) return (img, mask)
class LoadMaskBase64: class LoadMaskBase64(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(s): def define_schema(cls):
return {"required": {"mask": ("STRING", {"multiline": False})}} return io.Schema(
node_id="ETN_LoadMaskBase64",
display_name="Load Mask (Base64)",
category="external_tooling",
inputs=[io.String.Input("mask", multiline=False)],
outputs=[io.Mask.Output(display_name="mask")],
)
RETURN_TYPES = ("MASK",) @classmethod
CATEGORY = "external_tooling" def execute(cls, mask: str): # type: ignore
FUNCTION = "load_mask" _strip_prefix(mask, "data:image/png;base64,")
def load_mask(self, mask):
imgdata = base64.b64decode(mask) imgdata = base64.b64decode(mask)
img = Image.open(BytesIO(imgdata)) img = Image.open(BytesIO(imgdata))
img = np.array(img).astype(np.float32) / 255.0 img = np.array(img).astype(np.float32) / 255.0
@@ -50,23 +69,24 @@ class LoadMaskBase64:
if img.dim() == 3: # RGB(A) input, use red channel if img.dim() == 3: # RGB(A) input, use red channel
img = img[:, :, 0] img = img[:, :, 0]
return (img.unsqueeze(0),) return (img.unsqueeze(0),)
class SendImageWebSocket:
class SendImageWebSocket(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(s): def define_schema(cls):
return { return io.Schema(
"required": { node_id="ETN_SendImageWebSocket",
"images": ("IMAGE",), display_name="Send Image (WebSocket)",
"format": (["PNG", "JPEG"], {"default": "PNG"}), category="external_tooling",
} inputs=[
} io.Image.Input("images"),
io.Combo.Input("format", options=["PNG", "JPEG"], default="PNG"),
],
is_output_node=True,
)
RETURN_TYPES = () @classmethod
FUNCTION = "send_images" def execute(cls, images: torch.Tensor, format: str): # type: ignore
OUTPUT_NODE = True
CATEGORY = "external_tooling"
def send_images(self, images, format):
results = [] results = []
for tensor in images: for tensor in images:
array = 255.0 * tensor.cpu().numpy() array = 255.0 * tensor.cpu().numpy()
@@ -78,46 +98,161 @@ class SendImageWebSocket:
[format, image, None], [format, image, None],
server.client_id, server.client_id,
) )
results.append( results.append({
{"source": "websocket", "content-type": f"image/{format.lower()}", "type": "output"} "source": "websocket",
) "content-type": f"image/{format.lower()}",
"type": "output",
})
return {"ui": {"images": results}} return io.NodeOutput(ui={"images": results})
class CropImage:
"""Deprecated, ComfyUI has an ImageCrop node now which does the same.""" class ImageCache:
timeout = 600 # 10 minutes
max_size = 100 * 1024 * 1024 # 100 MB
@dataclass
class Entry:
data: bytes
content_type: str
timestamp: float
retrieved: int
class OldEntry(NamedTuple):
last_used: float
deleted: float
size: int
retrieved: int
def __init__(self):
self.images: dict[str, ImageCache.Entry] = {}
self.old: dict[str, ImageCache.OldEntry] = {}
def add(self, image: Image.Image, format: str):
key = uuid4().hex
with BytesIO() as output:
image.save(output, format=format, quality=95, compress_level=1)
image_data = output.getvalue()
self.insert(key, image_data, f"image/{format.lower()}")
return key
def insert(self, key: str, data: bytes, content_type: str):
self.images[key] = ImageCache.Entry(
data=data,
content_type=content_type,
timestamp=time.time(),
retrieved=0,
)
def get(self, key: str, extend: bool = False):
entry = self.images.get(key)
if entry is None:
if old := self.old.get(key):
now = time.time()
print(
f"[comfyui-tooling-nodes] requested image {key} has been deleted ",
f"(last used {now - old.last_used:.0f}s ago, deleted {now - old.deleted:.0f}s ago, "
f"size {old.size / 1024**2:.1f}MB, retrieved {old.retrieved} times)",
)
return None, None
entry.retrieved += 1
if extend:
entry.timestamp = time.time()
self.prune()
return entry.data, entry.content_type
def prune(self):
total_size = sum(len(entry.data) for entry in self.images.values())
if total_size <= self.max_size:
return
# Remove least recently used entries until under max size
sorted_entries = sorted(self.images.items(), key=lambda item: item[1].timestamp)
now = time.time()
for key, entry in sorted_entries:
age = now - entry.timestamp
if age > self.timeout or (age > 60 and entry.retrieved > 0):
self.old[key] = ImageCache.OldEntry(
entry.timestamp, now, len(entry.data), entry.retrieved
)
del self.images[key]
total_size -= len(entry.data)
if total_size <= self.max_size:
break
def __contains__(self, key: str):
return key in self.images
image_cache = ImageCache()
class LoadImageCache(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_LoadImageCache",
display_name="Load Image from Cache",
category="external_tooling",
inputs=[io.String.Input("id", multiline=False)],
outputs=[io.Image.Output(display_name="image"), io.Mask.Output(display_name="mask")],
)
@classmethod @classmethod
def INPUT_TYPES(cls): def execute(cls, id: str): # type: ignore
return { image_data, content_type = image_cache.get(id, extend=True)
"required": { if image_data is None:
"image": ("IMAGE",), raise ValueError(f"Image with ID {id} not found in cache.")
"x": (
"INT",
{"default": 0, "min": 0, "max": 8192, "step": 1},
),
"y": (
"INT",
{"default": 0, "min": 0, "max": 8192, "step": 1},
),
"width": (
"INT",
{"default": 512, "min": 1, "max": 8192, "step": 1},
),
"height": (
"INT",
{"default": 512, "min": 1, "max": 8192, "step": 1},
),
}
}
CATEGORY = "external_tooling" img = Image.open(BytesIO(image_data))
RETURN_TYPES = ("IMAGE",)
FUNCTION = "crop"
def crop(self, image, x, y, width, height): w, h = img.size
out = image[:, y : y + height, x : x + width, :] c = len(img.getbands())
return (out,) normalized = np.array(img).astype(np.float32) / 255.0
tensor = torch.from_numpy(normalized).reshape(1, h, w, c)
match c:
case 1:
image = tensor.expand(1, h, w, 3)
mask = tensor.reshape(1, h, w)
case 3:
image = tensor
mask = tensor[..., 0]
case 4:
image = tensor[..., :3]
mask = tensor[..., 3]
return io.NodeOutput(image, mask)
class SaveImageCache(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_SaveImageCache",
display_name="Save Image to Cache",
category="external_tooling",
inputs=[
io.Image.Input("images"),
io.Combo.Input("format", options=["PNG", "JPEG"], default="PNG"),
],
is_output_node=True,
)
@classmethod
def execute(cls, images: torch.Tensor, format: str): # type: ignore
results = []
for tensor in images:
array = 255.0 * tensor.cpu().numpy()
image = Image.fromarray(np.clip(array, 0, 255).astype(np.uint8))
key = image_cache.add(image, format)
results.append({
"source": "http",
"id": key,
"content-type": f"image/{format.lower()}",
"type": "output",
})
return io.NodeOutput(ui={"images": results})
def to_bchw(image: torch.Tensor): def to_bchw(image: torch.Tensor):
@@ -136,21 +271,22 @@ def mask_batch(mask: torch.Tensor):
return mask return mask
class ApplyMaskToImage: class ApplyMaskToImage(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(cls): def define_schema(cls):
return { return io.Schema(
"required": { node_id="ETN_ApplyMaskToImage",
"image": ("IMAGE",), display_name="Apply Mask to Image",
"mask": ("MASK",), category="external_tooling",
} inputs=[
} io.Image.Input("image"),
io.Mask.Input("mask"),
],
outputs=[io.Image.Output(display_name="masked")],
)
CATEGORY = "external_tooling" @classmethod
RETURN_TYPES = ("IMAGE",) def execute(cls, image: torch.Tensor, mask: torch.Tensor): # type: ignore
FUNCTION = "apply_mask"
def apply_mask(self, image: torch.Tensor, mask: torch.Tensor):
out = to_bchw(image) out = to_bchw(image)
if out.shape[1] == 3: # Assuming RGB images if out.shape[1] == 3: # Assuming RGB images
out = torch.cat([out, torch.ones_like(out[:, :1, :, :])], dim=1) out = torch.cat([out, torch.ones_like(out[:, :1, :, :])], dim=1)
@@ -158,14 +294,138 @@ class ApplyMaskToImage:
assert mask.ndim == 3, f"Mask should have shape [B, H, W]. {mask.shape}" assert mask.ndim == 3, f"Mask should have shape [B, H, W]. {mask.shape}"
assert out.ndim == 4, f"Image should have shape [B, C, H, W]. {out.shape}" assert out.ndim == 4, f"Image should have shape [B, C, H, W]. {out.shape}"
assert ( assert out.shape[-2:] == mask.shape[-2:], (
out.shape[-2:] == mask.shape[-2:] f"Image size {out.shape[-2:]} must match mask size {mask.shape[-2:]}"
), f"Image size {out.shape[-2:]} must match mask size {mask.shape[-2:]}" )
is_mask_batch = mask.shape[0] == out.shape[0] is_mask_batch = mask.shape[0] == out.shape[0]
# Apply each mask in the batch to its corresponding image's alpha channel # Apply each mask in the batch to its corresponding image's alpha channel
for i in range(out.shape[0]): for i in range(out.shape[0]):
alpha = mask[i] if is_mask_batch else mask[0] alpha = mask[i] if is_mask_batch else mask[0]
out[i, 3, :, :] = alpha out[i, 3, :, :] *= alpha
return (to_bhwc(out),) return (to_bhwc(out),)
class _ReferenceImageData(NamedTuple):
image: torch.Tensor
weight: float
range: tuple[float, float]
class ReferenceImage(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_ReferenceImage",
display_name="Reference Image",
category="external_tooling",
inputs=[
io.Image.Input("image"),
io.Float.Input("weight", default=1.0, min=0.0, max=10.0),
io.Float.Input("range_start", default=0.0, min=0.0, max=1.0),
io.Float.Input("range_end", default=1.0, min=0.0, max=1.0),
io.Custom("ReferenceImage").Input("reference_images", optional=True),
],
outputs=[io.Custom("ReferenceImage").Output(display_name="reference_images")],
)
@classmethod
def execute( # type: ignore
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( # type: ignore
cls,
conditioning: list[list],
clip_vision: ClipVisionModel,
style_model: StyleModel,
references: list[_ReferenceImageData],
):
delimiters = {0.0, 1.0}
delimiters |= set(r.range[0] for r in references)
delimiters |= set(r.range[1] for r in references)
delimiters = sorted(delimiters)
ranges = [(delimiters[i], delimiters[i + 1]) for i in range(len(delimiters) - 1)]
embeds = [_encode_image(r.image, clip_vision, style_model, r.weight) for r in references]
base = conditioning[0][0]
result = []
for start, end in ranges:
e = [
embeds[i]
for i, r in enumerate(references)
if r.range[0] <= start and r.range[1] >= end
]
options = conditioning[0][1].copy()
options["start_percent"] = start
options["end_percent"] = end
result.append((torch.cat([base] + e, dim=1), options))
return (result,)
def _encode_image(
image: torch.Tensor, clip_vision: ClipVisionModel, style_model: StyleModel, weight: float
):
e = clip_vision.encode_image(image)
e = style_model.get_cond(e).flatten(start_dim=0, end_dim=1).unsqueeze(dim=0)
e = _downsample_image_cond(e, weight)
return e
def _downsample_image_cond(cond: torch.Tensor, weight: float):
if weight >= 1.0:
return cond
elif weight <= 0.0:
return torch.zeros_like(cond)
elif weight >= 0.6:
factor = 2
elif weight >= 0.3:
factor = 3
else:
factor = 4
# Downsample the clip vision embedding to make it smaller, resulting in less impact
# compared to other conditioning.
# See https://github.com/kaibioinfo/ComfyUI_AdvancedRefluxControl
(b, t, h) = cond.shape
m = int(np.sqrt(t))
cond = F.interpolate(
cond.view(b, m, m, h).transpose(1, -1),
size=(m // factor, m // factor),
mode="area",
)
return cond.transpose(1, -1).reshape(b, -1, h)
def _strip_prefix(s: str, prefix: str) -> str:
if s.startswith(prefix):
return s[len(prefix) :]
return s
+27 -28
View File
@@ -1,5 +1,4 @@
from __future__ import annotations from __future__ import annotations
from weakref import ref as WeakRef
from pathlib import Path from pathlib import Path
from tqdm import tqdm from tqdm import tqdm
import torch import torch
@@ -8,6 +7,7 @@ import torch.nn.functional as F
from torch import Tensor from torch import Tensor
from transformers import CLIPImageProcessor, CLIPConfig, CLIPVisionModel, PreTrainedModel from transformers import CLIPImageProcessor, CLIPConfig, CLIPVisionModel, PreTrainedModel
from kornia.filters import box_blur from kornia.filters import box_blur
from comfy_api.latest import io
from .nodes import to_bchw, to_bhwc from .nodes import to_bchw, to_bhwc
@@ -38,6 +38,10 @@ class CLIPSafetyChecker(PreTrainedModel):
self.concept_embeds_weights = nn.Parameter(torch.ones(17), requires_grad=False) self.concept_embeds_weights = nn.Parameter(torch.ones(17), requires_grad=False)
self.special_care_embeds_weights = nn.Parameter(torch.ones(3), requires_grad=False) self.special_care_embeds_weights = nn.Parameter(torch.ones(3), requires_grad=False)
# Model requires post_init after transformers v4.57.3
if hasattr(self, "post_init"):
self.post_init()
def forward(self, clip_input, images: Tensor, sensitivity: float): def forward(self, clip_input, images: Tensor, sensitivity: float):
with torch.no_grad(): with torch.no_grad():
image_batch = self.vision_model(clip_input)[1] image_batch = self.vision_model(clip_input)[1]
@@ -76,7 +80,7 @@ class CLIPSafetyChecker(PreTrainedModel):
class CachedModels: class CachedModels:
_instance: WeakRef | None = None _instance: CachedModels | None = None
def __init__(self): def __init__(self):
model_dir = Path(__file__).parent / "safetychecker" model_dir = Path(__file__).parent / "safetychecker"
@@ -91,11 +95,9 @@ class CachedModels:
@classmethod @classmethod
def load(cls): def load(cls):
models = cls._instance and cls._instance() if cls._instance is None:
if models is None: cls._instance = CachedModels()
models = cls() return cls._instance
cls._instance = WeakRef(models)
return models
def download(self, url: str, target: Path): def download(self, url: str, target: Path):
import requests import requests
@@ -118,29 +120,26 @@ class CachedModels:
) from e ) from e
class NSFWFilter: class NSFWFilter(io.ComfyNode):
models: CachedModels @classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_NSFWFilter",
display_name="NSFW Filter",
category="external_tooling",
inputs=[
io.Image.Input("image"),
io.Float.Input("sensitivity", default=0.5, min=0.0, max=1.0, step=0.1),
],
outputs=[io.Image.Output(display_name="image")],
)
@classmethod @classmethod
def INPUT_TYPES(cls): def execute(cls, image: Tensor, sensitivity: float):
return { models = CachedModels.load()
"required": {
"image": ("IMAGE",),
"sensitivity": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.10}),
},
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "check"
CATEGORY = "external_tooling"
def __init__(self):
self.models = CachedModels.load()
def check(self, image, sensitivity):
image = to_bchw(image) image = to_bchw(image)
input = self.models.feature_extractor(image, do_rescale=False, return_tensors="pt") input = models.feature_extractor(image, do_rescale=False, return_tensors="pt")
filtered = self.models.safety_checker( filtered = models.safety_checker(
images=image, clip_input=input.pixel_values, sensitivity=sensitivity images=image, clip_input=input.pixel_values, sensitivity=sensitivity
) )
return (to_bhwc(filtered),) return io.NodeOutput(to_bhwc(filtered))
+9 -1
View File
@@ -1,12 +1,20 @@
[project] [project]
name = "comfyui-tooling-nodes" name = "comfyui-tooling-nodes"
description = "Provides nodes and server API extensions geared towards using ComfyUI as a backend for external tools." description = "Provides nodes and server API extensions geared towards using ComfyUI as a backend for external tools."
version = "1.5.0" version = "3.4.0"
license = { file = "LICENSE" } license = { file = "LICENSE" }
[project.urls] [project.urls]
Repository = "https://github.com/Acly/comfyui-tooling-nodes" Repository = "https://github.com/Acly/comfyui-tooling-nodes"
[tool.ruff]
target-version = "py311"
line-length = 100
preview = true
[tool.ruff.lint]
ignore = ["E741", "BLE001"]
[tool.black] [tool.black]
line-length = 100 line-length = 100
preview = true preview = true
+299 -75
View File
@@ -1,13 +1,32 @@
# Adapted from https://github.com/pamparamm/ComfyUI-ppm
# Adapted from https://github.com/laksjdjf/cgem156-ComfyUI/blob/main/scripts/attention_couple/node.py # Adapted from https://github.com/laksjdjf/cgem156-ComfyUI/blob/main/scripts/attention_couple/node.py
# by @laksjdjf # by @laksjdjf
from __future__ import annotations from __future__ import annotations
from typing import NamedTuple from functools import partial
from typing import Any, NamedTuple
import torch import torch
import torch.nn.functional as F import torch.nn.functional as F
import math import math
from torch import Tensor, Size from torch import Tensor, Size
import comfy.model_management
import comfy.patcher_extension
from comfy.model_patcher import ModelPatcher from comfy.model_patcher import ModelPatcher
from comfy.model_base import Anima, CosmosPredict2
from comfy.ldm.cosmos.predict2 import Attention as CosmosAttention
from comfy.sampler_helpers import convert_cond
from comfy.samplers import process_conds
from comfy_api.latest import io
COND = 0
UNCOND = 1
ANIMA_COUPLE_WRAPPER_KEY = "etn_attention_mask_anima"
ANIMA_COUPLE_PATCH_KEY = "etn_attention_mask_patch"
CONDS_COUPLE_KEY = "etn_couple_conds"
COND_UNCOND_COUPLE_KEY = "etn_couple_cond_or_uncond"
COUPLE_ACTIVE_KEY = "etn_couple_active"
NUM_TOKENS_COUPLE_KEY = "etn_couple_num_tokens"
def downsample_mask(mask: Tensor, batch: int, target_size: int, original_shape: Size) -> Tensor: def downsample_mask(mask: Tensor, batch: int, target_size: int, original_shape: Size) -> Tensor:
@@ -33,6 +52,12 @@ def downsample_mask(mask: Tensor, batch: int, target_size: int, original_shape:
return result return result
def reshape_mask(mask: Tensor, size: tuple[int, int], batch: int, target_size: int) -> Tensor:
result = F.interpolate(mask, size=size, mode="nearest")
result = result.view(mask.shape[0], target_size, 1)
return result.repeat_interleave(batch, dim=0)
def lcm(a: int, b: int): def lcm(a: int, b: int):
return a * b // math.gcd(a, b) return a * b // math.gcd(a, b)
@@ -65,105 +90,115 @@ class Region(NamedTuple):
return result return result
class BackgroundRegion: Regions = io.Custom("Regions")
class BackgroundRegion(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(cls): def define_schema(cls):
return {"required": {"conditioning": ("CONDITIONING",)}} return io.Schema(
node_id="ETN_BackgroundRegion",
display_name="Background Region",
category="external_tooling/regions",
inputs=[io.Conditioning.Input("conditioning")],
outputs=[Regions.Output(display_name="regions")],
)
CATEGORY = "external_tooling/regions" @classmethod
RETURN_TYPES = ("REGIONS",) def execute(cls, conditioning: list):
FUNCTION = "define"
def define(self, conditioning: list):
return (Region(None, None, conditioning),) return (Region(None, None, conditioning),)
class DefineRegion: class DefineRegion(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(cls): def define_schema(cls):
return { return io.Schema(
"required": { node_id="ETN_DefineRegion",
"mask": ("MASK",), display_name="Define Region",
"conditioning": ("CONDITIONING",), category="external_tooling/regions",
}, inputs=[
"optional": { io.Mask.Input("mask"),
"regions": ("REGIONS",), io.Conditioning.Input("conditioning"),
}, Regions.Input("regions", optional=True),
} ],
outputs=[Regions.Output(display_name="regions")],
)
CATEGORY = "external_tooling/regions" @classmethod
RETURN_TYPES = ("REGIONS",) def execute(cls, mask: Tensor, conditioning: list, regions: Region | None = None):
FUNCTION = "define"
def define(self, mask: Tensor, conditioning: list, regions: Region | None = None):
if mask.dim() < 3: if mask.dim() < 3:
mask = mask.unsqueeze(0) mask = mask.unsqueeze(0)
return (Region(regions, mask, conditioning),) return io.NodeOutput(Region(regions, mask, conditioning))
class ListRegionMasks: class ListRegionMasks(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(cls): def define_schema(cls):
return {"required": {"regions": ("REGIONS",)}} return io.Schema(
node_id="ETN_ListRegionMasks",
CATEGORY = "external_tooling/regions" display_name="List Region Masks",
RETURN_TYPES = ("MASK",) category="external_tooling/regions",
FUNCTION = "get_masks" inputs=[Regions.Input("regions")],
outputs=[io.Mask.Output(display_name="masks")],
def get_masks(self, regions: Region): )
return (torch.stack([r.mask for r in regions.preprocess()], dim=0),)
class AttentionMask:
@classmethod @classmethod
def INPUT_TYPES(s): def execute(cls, regions: Region):
return { return io.NodeOutput(torch.stack([r.mask for r in regions.preprocess()], dim=0))
"required": {
"model": ("MODEL",),
"regions": ("REGIONS",),
}
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "attention_mask"
CATEGORY = "external_tooling/regions"
mask: Tensor class AttentionMask(io.ComfyNode):
conds: list[Tensor] @classmethod
batch_size: int def define_schema(cls):
return io.Schema(
node_id="ETN_AttentionMask",
display_name="Regions Attention Mask",
category="external_tooling/regions",
inputs=[io.Model.Input("model"), Regions.Input("regions")],
outputs=[io.Model.Output(display_name="model")],
)
def attention_mask(self, model: ModelPatcher, regions: Region): @classmethod
new_model = model.clone() def execute(cls, model: ModelPatcher, regions: Region):
region_list = regions.preprocess() return io.NodeOutput(AttentionMaskPatch.apply(model, regions))
num_conds = len(region_list)
class AttentionMaskPatch:
def __init__(self, region_list: list[Region]):
mask = torch.stack([r.mask for r in region_list], dim=0) mask = torch.stack([r.mask for r in region_list], dim=0)
mask_sum = mask.sum(dim=0, keepdim=True) mask_sum = mask.sum(dim=0, keepdim=True)
assert mask_sum.sum() > 0, "There are areas that are zero in all masks." assert mask_sum.sum() > 0, "There are areas that are zero in all masks."
self.mask = mask / mask_sum self.mask = mask / mask_sum
self.region_conds = [r.conditioning for r in region_list]
self.conds = [r.conditioning[0][0] for r in region_list] self.conds = [r.conditioning[0][0] for r in region_list]
num_tokens = [cond.shape[1] for cond in self.conds] self.num_tokens = [cond.shape[1] for cond in self.conds]
self.num_conds = len(region_list)
self.batch_size = 0
@staticmethod
def apply(model: ModelPatcher, regions: Region):
patch = AttentionMaskPatch(regions.preprocess())
if _is_anima_couple_model(model):
return patch.apply_anima(model)
def attn2_patch(q: Tensor, k: Tensor, v: Tensor, extra_options: dict): def attn2_patch(q: Tensor, k: Tensor, v: Tensor, extra_options: dict):
assert k.mean() == v.mean(), "k and v must be the same." assert k.mean() == v.mean(), "k and v must be the same."
device, dtype = q.device, q.dtype device, dtype = q.device, q.dtype
if self.conds[0].device != device or self.conds[0].dtype != dtype: if patch.conds[0].device != device or patch.conds[0].dtype != dtype:
self.conds = [cond.to(device, dtype=dtype) for cond in self.conds] patch.conds = [cond.to(device, dtype=dtype) for cond in patch.conds]
if self.mask.device != device or self.mask.dtype != dtype: if patch.mask.device != device or patch.mask.dtype != dtype:
self.mask = self.mask.to(device, dtype=dtype) patch.mask = patch.mask.to(device, dtype=dtype)
cond_or_unconds = extra_options["cond_or_uncond"] cond_or_unconds = extra_options["cond_or_uncond"]
num_chunks = len(cond_or_unconds) num_chunks = len(cond_or_unconds)
self.batch_size = q.shape[0] // num_chunks patch.batch_size = q.shape[0] // num_chunks
q_chunks = q.chunk(num_chunks, dim=0) q_chunks = q.chunk(num_chunks, dim=0)
k_chunks = k.chunk(num_chunks, dim=0) k_chunks = k.chunk(num_chunks, dim=0)
lcm_tokens = lcm_for_list(num_tokens + [k.shape[1]]) lcm_tokens = lcm_for_list(patch.num_tokens + [k.shape[1]])
conds_tensor = [ conds_tensor = [
cond.repeat(self.batch_size, lcm_tokens // num_tokens[i], 1) cond.repeat(patch.batch_size, lcm_tokens // patch.num_tokens[i], 1)
for i, cond in enumerate(self.conds) for i, cond in enumerate(patch.conds)
] ]
conds_tensor = torch.cat(conds_tensor, dim=0) conds_tensor = torch.cat(conds_tensor, dim=0)
@@ -174,9 +209,9 @@ class AttentionMask:
qs.insert(0, q_chunks[i]) qs.insert(0, q_chunks[i])
ks.insert(0, k_target) ks.insert(0, k_target)
else: else:
qs.insert(0, q_chunks[i].repeat(num_conds, 1, 1)) qs.insert(0, q_chunks[i].repeat(patch.num_conds, 1, 1))
ks.insert(0, conds_tensor) ks.insert(0, conds_tensor)
for _ in range(num_conds - 1): for _ in range(patch.num_conds - 1):
cond_or_unconds.insert(i, 0) cond_or_unconds.insert(i, 0)
qs = torch.cat(qs, dim=0) qs = torch.cat(qs, dim=0)
@@ -184,29 +219,218 @@ class AttentionMask:
return qs, ks, ks return qs, ks, ks
def attn2_output_patch(out: Tensor, extra_options: dict): def attn2_output_patch(out: Tensor, extra_options: dict):
num_conds = patch.num_conds
cond_or_unconds = extra_options["cond_or_uncond"] cond_or_unconds = extra_options["cond_or_uncond"]
mask_downsample = downsample_mask( mask_downsample = downsample_mask(
self.mask, self.batch_size, out.shape[1], extra_options["original_shape"] patch.mask, patch.batch_size, out.shape[1], extra_options["original_shape"]
) )
outputs: list[Tensor] = [] outputs: list[Tensor] = []
pos = 0 pos = 0
i = 0 i = 0
while i < len(cond_or_unconds): while i < len(cond_or_unconds):
if cond_or_unconds[i] == 1: # uncond if cond_or_unconds[i] == 1: # uncond
outputs.append(out[pos : pos + self.batch_size]) outputs.append(out[pos : pos + patch.batch_size])
pos += self.batch_size pos += patch.batch_size
else: else:
masked = out[pos : pos + num_conds * self.batch_size] * mask_downsample masked = out[pos : pos + num_conds * patch.batch_size] * mask_downsample
masked = masked.view(num_conds, self.batch_size, out.shape[1], out.shape[2]) masked = masked.view(num_conds, patch.batch_size, out.shape[1], out.shape[2])
masked = masked.sum(dim=0) masked = masked.sum(dim=0)
outputs.append(masked) outputs.append(masked)
pos += num_conds * self.batch_size pos += num_conds * patch.batch_size
for _ in range(num_conds - 1): for _ in range(num_conds - 1):
cond_or_unconds.pop(i) cond_or_unconds.pop(i)
i += 1 i += 1
return torch.cat(outputs, dim=0) return torch.cat(outputs, dim=0)
new_model = model.clone()
new_model.set_model_attn2_patch(attn2_patch) new_model.set_model_attn2_patch(attn2_patch)
new_model.set_model_attn2_output_patch(attn2_output_patch) new_model.set_model_attn2_output_patch(attn2_output_patch)
return (new_model,) new_model.set_attachments("etn_attention_mask", patch)
return new_model
def apply_anima(self, model: ModelPatcher):
new_model = model.clone()
_patch_cosmos_attention(new_model)
device = comfy.model_management.get_torch_device()
conds_converted = [convert_cond(cond)[0] for cond in self.region_conds]
new_model.add_wrapper_with_key(
comfy.patcher_extension.WrappersMP.SAMPLER_SAMPLE,
ANIMA_COUPLE_WRAPPER_KEY,
_anima_couple_sample_wrapper(conds_converted, device),
)
new_model.add_wrapper_with_key(
comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL,
ANIMA_COUPLE_WRAPPER_KEY,
_anima_couple_diffusion_wrapper(self),
)
new_model.set_attachments("etn_attention_mask", self)
return new_model
def _is_anima_couple_model(model: ModelPatcher) -> bool:
model_type = type(model.model)
return issubclass(model_type, (Anima, CosmosPredict2))
def _anima_couple_sample_wrapper(conds_converted: list, device):
def sample_wrapper(executor, *args, **kwargs):
if len(conds_converted) > 0:
guider = args[0]
extra_options: dict[str, Any] = args[2]
seed: int = extra_options["seed"]
noise: Tensor = args[4]
latent_image: Tensor = args[5]
denoise_mask: Tensor | None = args[6]
conds_processed = process_conds(
guider.inner_model,
noise,
{"positive": conds_converted},
device,
latent_image,
denoise_mask,
seed,
latent_shapes=[latent_image.shape],
)["positive"]
conds_couple = [cond["model_conds"]["c_crossattn"].cond for cond in conds_processed]
model_options: dict[str, Any] = extra_options["model_options"]
transformer_options: dict[str, Any] = model_options.get("transformer_options", {}).copy()
transformer_options[CONDS_COUPLE_KEY] = conds_couple
transformer_options[NUM_TOKENS_COUPLE_KEY] = [cond.shape[1] for cond in conds_couple]
model_options["transformer_options"] = transformer_options
return executor(*args, **kwargs)
return sample_wrapper
def _anima_couple_diffusion_wrapper(patch: AttentionMaskPatch):
def diffusion_wrapper(executor, *args, **kwargs):
anima_model = executor.class_obj
x: Tensor = args[0]
transformer_options: dict[str, Any] = kwargs.get("transformer_options", {}).copy()
patch_spatial = getattr(anima_model, "patch_spatial", 1)
activations_shape = list(x.shape)
activations_shape[-2] = activations_shape[-2] // patch_spatial
activations_shape[-1] = activations_shape[-1] // patch_spatial
transformer_options["activations_shape"] = activations_shape
transformer_options[ANIMA_COUPLE_PATCH_KEY] = patch
kwargs["transformer_options"] = transformer_options
return executor(*args, **kwargs)
return diffusion_wrapper
def pre_cross_attention(
patch: AttentionMaskPatch,
transformer_options: dict,
x: Tensor,
context: Tensor,
rope_emb: Tensor | None,
) -> tuple[Tensor, Tensor, Tensor | None, dict]:
transformer_options = transformer_options.copy()
if CONDS_COUPLE_KEY not in transformer_options:
transformer_options[COND_UNCOND_COUPLE_KEY] = list(transformer_options["cond_or_uncond"])
transformer_options[COUPLE_ACTIVE_KEY] = False
return x, context, rope_emb, transformer_options
conds: list[Tensor] = transformer_options[CONDS_COUPLE_KEY]
num_tokens_c: list[int] = transformer_options[NUM_TOKENS_COUPLE_KEY]
cond_or_uncond = transformer_options["cond_or_uncond"]
num_chunks = len(cond_or_uncond)
batch = x.shape[0] // num_chunks
x_chunks = x.chunk(num_chunks, dim=0)
c_chunks = context.chunk(num_chunks, dim=0)
lcm_tokens_c = lcm_for_list(num_tokens_c + [context.shape[1]])
conds_c_tensor = torch.cat(
[cond.repeat(batch, lcm_tokens_c // num_tokens_c[i], 1) for i, cond in enumerate(conds)],
dim=0,
)
xs, cs = [], []
cond_or_uncond_couple = []
for i, cond_type in enumerate(cond_or_uncond):
x_target = x_chunks[i]
c_target = c_chunks[i].repeat(1, lcm_tokens_c // context.shape[1], 1)
if cond_type == UNCOND:
xs.append(x_target)
cs.append(c_target)
cond_or_uncond_couple.append(UNCOND)
else:
xs.append(x_target.repeat(patch.num_conds, 1, 1))
cs.append(conds_c_tensor)
cond_or_uncond_couple.extend([COND] * patch.num_conds)
transformer_options[COND_UNCOND_COUPLE_KEY] = cond_or_uncond_couple
transformer_options[COUPLE_ACTIVE_KEY] = True
return torch.cat(xs, dim=0), torch.cat(cs, dim=0), rope_emb, transformer_options
def cross_attention_output(patch: AttentionMaskPatch, transformer_options: dict, out: Tensor):
cond_or_uncond = transformer_options[COND_UNCOND_COUPLE_KEY]
size = tuple(transformer_options["activations_shape"][-2:])
batch = out.shape[0] // len(cond_or_uncond)
mask = patch.mask.to(out.device, dtype=out.dtype)
mask_downsample = reshape_mask(mask, size, batch, out.shape[1])
outputs = []
cond_outputs = []
i_cond = 0
for i, cond_type in enumerate(cond_or_uncond):
pos, next_pos = i * batch, (i + 1) * batch
if cond_type == UNCOND:
outputs.append(out[pos:next_pos])
else:
pos_cond, next_pos_cond = i_cond * batch, (i_cond + 1) * batch
cond_outputs.append(out[pos:next_pos] * mask_downsample[pos_cond:next_pos_cond])
i_cond += 1
if len(cond_outputs) > 0:
outputs.append(torch.stack(cond_outputs).sum(0))
return torch.cat(outputs, dim=0)
def _patch_cosmos_attention(model_patcher: ModelPatcher):
cosmos_model = model_patcher.get_model_object("diffusion_model")
for block_name, block in (
(n, b)
for n, b in cosmos_model.named_modules()
if ("cross_attn" in n or "self_attn" in n) and isinstance(b, CosmosAttention)
):
patch_name = f"diffusion_model.{block_name}.forward"
if patch_name not in model_patcher.object_patches:
model_patcher.add_object_patch(patch_name, partial(_cosmos_attention_forward_patched, block))
def _cosmos_attention_forward_patched(
self,
x: Tensor,
context: Tensor | None = None,
rope_emb: Tensor | None = None,
transformer_options: dict | None = None,
) -> Tensor:
transformer_options = transformer_options if transformer_options is not None else {}
patch: AttentionMaskPatch | None = transformer_options.get(ANIMA_COUPLE_PATCH_KEY)
if context is not None and patch is not None:
x, context, rope_emb, transformer_options = pre_cross_attention(
patch, transformer_options, x, context, rope_emb
)
q, k, v = self.compute_qkv(x, context, rope_emb=rope_emb)
output = self.compute_attention(q, k, v, transformer_options=transformer_options)
if context is not None and patch is not None and transformer_options.get(COUPLE_ACTIVE_KEY, False):
output = cross_attention_output(patch, transformer_options, output)
return output
+103 -94
View File
@@ -3,49 +3,29 @@ import numpy as np
import numpy.typing as npt import numpy.typing as npt
import torch import torch
from torch import Tensor from torch import Tensor
from comfy_api.latest import io
IntArray = npt.NDArray[np.int_] IntArray = npt.NDArray[np.int_]
class TileLayout: class TileLayout:
@classmethod def __init__(
def INPUT_TYPES(cls): self, image: Tensor, min_tile_size: int, padding: int, blending: int, multiple: int
return { ):
"required": { assert all([x % multiple == 0 for x in image.shape[-3:-1]]), (
"image": ("IMAGE",), "Image size must be divisible by multiple"
"min_tile_size": ("INT", {"default": 512, "min": 64, "max": 8192, "step": 8}), )
"padding": ("INT", {"default": 32, "min": 0, "max": 8192, "step": 8}), assert min_tile_size % multiple == 0, "Tile size must be divisible by multiple"
"blending": ("INT", {"default": 8, "min": 0, "max": 256, "step": 8}), assert blending <= padding, "Blending must be smaller than padding"
}
}
CATEGORY = "external_tooling/tiles" self.image_size: IntArray = np.array(image.shape[-3:-1])
RETURN_TYPES = ("TILE_LAYOUT",) self.padding: int = padding
FUNCTION = "node" self.blending: int = blending
self.tile_count: IntArray = np.maximum(1, self.image_size // (min_tile_size - 2 * padding))
image_size: IntArray
tile_size: IntArray
padding: int
blending: int
tile_count: IntArray
def node(self, image: Tensor, min_tile_size: int, padding: int, blending: int):
self.init(image, min_tile_size, padding, blending)
return (self,)
def init(self, image: Tensor, min_tile_size: int, padding: int, blending: int):
assert all([x % 8 == 0 for x in image.shape[-3:-1]]), "Image size must be divisible by 8"
assert min_tile_size % 8 == 0, "Tile size must be divisible by 8"
assert blending < padding, "Blending must be smaller than padding"
self.image_size = np.array(image.shape[-3:-1])
self.padding = padding
self.blending = blending
self.tile_count = np.maximum(1, self.image_size // (min_tile_size - 2 * padding))
image_size_with_overlap = self.image_size + (self.tile_count - 1) * 2 * padding image_size_with_overlap = self.image_size + (self.tile_count - 1) * 2 * padding
tile_size = np.ceil(image_size_with_overlap / self.tile_count) tile_size = np.ceil(image_size_with_overlap / self.tile_count)
self.tile_size = (np.ceil(tile_size / 8) * 8).astype(int) self.tile_size: IntArray = (np.ceil(tile_size / multiple) * multiple).astype(int)
def size(self, coord: IntArray): def size(self, coord: IntArray):
return self.end(coord) - self.start(coord) return self.end(coord) - self.start(coord)
@@ -96,80 +76,109 @@ class TileLayout:
image[rect] = (1 - mask) * image[rect] + mask * tile image[rect] = (1 - mask) * image[rect] + mask * tile
class ExtractImageTile: class CreateTileLayout(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(cls): def define_schema(cls):
return { return io.Schema(
"required": { node_id="ETN_TileLayout",
"image": ("IMAGE",), display_name="Create Tile Layout",
"layout": ("TILE_LAYOUT",), category="external_tooling/tiles",
"index": ("INT", {"min": 0}), inputs=[
} io.Image.Input("image"),
} io.Int.Input("min_tile_size", default=512, min=64, max=8192, step=8),
io.Int.Input("padding", default=32, min=0, max=8192, step=8),
io.Int.Input("blending", default=8, min=0, max=256, step=8),
io.Int.Input("multiple", default=8, min=1, max=1024, step=1),
],
outputs=[io.Custom("TileLayout").Output(display_name="layout")],
)
CATEGORY = "external_tooling/tiles"
RETURN_TYPES = ("IMAGE",)
FUNCTION = "slice"
def slice(self, image: Tensor, layout: TileLayout, index: int):
return (layout.tile(image, index),)
class ExtractMaskTile:
@classmethod @classmethod
def INPUT_TYPES(cls): def execute(cls, image: Tensor, min_tile_size: int, padding: int, blending: int, multiple: int):
return { return io.NodeOutput(TileLayout(image, min_tile_size, padding, blending, multiple))
"required": {
"mask": ("MASK",),
"layout": ("TILE_LAYOUT",),
"index": ("INT", {"min": 0}),
}
}
CATEGORY = "external_tooling/tiles"
RETURN_TYPES = ("MASK",)
FUNCTION = "slice"
def slice(self, mask: Tensor, layout: TileLayout, index: int): class ExtractImageTile(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_ExtractImageTile",
display_name="Extract Image Tile",
category="external_tooling/tiles",
inputs=[
io.Image.Input("image"),
io.Custom("TileLayout").Input("layout"),
io.Int.Input("index", default=0, min=0),
],
outputs=[io.Image.Output(display_name="tile")],
)
@classmethod
def execute(cls, image: Tensor, layout: TileLayout, index: int):
return io.NodeOutput(layout.tile(image, index))
class ExtractMaskTile(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_ExtractMaskTile",
display_name="Extract Mask Tile",
category="external_tooling/tiles",
inputs=[
io.Mask.Input("mask"),
io.Custom("TileLayout").Input("layout"),
io.Int.Input("index", default=0, min=0),
],
outputs=[io.Mask.Output(display_name="tile")],
)
@classmethod
def execute(cls, mask: Tensor, layout: TileLayout, index: int):
tile = layout.tile(mask.unsqueeze(3), index) tile = layout.tile(mask.unsqueeze(3), index)
return (tile.squeeze(3),) return io.NodeOutput(tile.squeeze(3))
class GenerateTileMask: class GenerateTileMask(io.ComfyNode):
@classmethod @classmethod
def INPUT_TYPES(cls): def define_schema(cls):
return { return io.Schema(
"required": {"layout": ("TILE_LAYOUT",), "index": ("INT", {"min": 0})}, node_id="ETN_GenerateTileMask",
"optional": {"blend": ("BOOLEAN",)}, display_name="Generate Tile Mask",
} category="external_tooling/tiles",
inputs=[
io.Custom("TileLayout").Input("layout"),
io.Int.Input("index", default=0, min=0),
io.Boolean.Input("blend", default=False, optional=True),
],
outputs=[io.Mask.Output(display_name="mask")],
)
CATEGORY = "external_tooling/tiles"
RETURN_TYPES = ("MASK",)
FUNCTION = "generate"
def generate(self, layout: TileLayout, index: int, blend: bool = False):
return (layout.mask(layout.coord(index), blend=blend),)
class MergeImageTile:
@classmethod @classmethod
def INPUT_TYPES(cls): def execute(cls, layout: TileLayout, index: int, blend: bool = False):
return { return io.NodeOutput(layout.mask(layout.coord(index), blend=blend))
"required": {
"image": ("IMAGE",),
"layout": ("TILE_LAYOUT",),
"index": ("INT", {"min": 0}),
"tile": ("IMAGE",),
}
}
CATEGORY = "external_tooling/tiles"
RETURN_TYPES = ("IMAGE",)
FUNCTION = "merge"
def merge(self, image: Tensor, layout: TileLayout, index: int, tile: Tensor): class MergeImageTile(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="ETN_MergeImageTile",
display_name="Merge Image Tile",
category="external_tooling/tiles",
inputs=[
io.Image.Input("image"),
io.Custom("TileLayout").Input("layout"),
io.Int.Input("index", default=0, min=0),
io.Image.Input("tile"),
],
outputs=[io.Image.Output(display_name="image")],
)
@classmethod
def execute(cls, image: Tensor, layout: TileLayout, index: int, tile: Tensor):
assert index < layout.total_count, f"Index {index} out of range" assert index < layout.total_count, f"Index {index} out of range"
if index == 0: if index == 0:
image = image.clone() image = image.clone()
layout.merge(image, index, tile) layout.merge(image, index, tile)
return (image,) return io.NodeOutput(image)
+15 -11
View File
@@ -10,6 +10,7 @@ from __future__ import annotations
import re import re
from functools import cache from functools import cache
from typing import NamedTuple from typing import NamedTuple
from comfy_api.latest import io
@cache @cache
@@ -43,7 +44,7 @@ def translate_chunk(text: str, language: str):
(p for p in available if p.from_code == language and p.to_code == target), None (p for p in available if p.from_code == language and p.to_code == target), None
) )
assert pkg, f"Couldn't find package for translation from {language}" assert pkg, f"Couldn't find package for translation from {language}"
print("Downloading and installing translation package", pkg) # print("Downloading and installing translation package", pkg) # this will cause encoding errors
pkg.install() pkg.install()
text, embeddings = _extract_embeddings(text) text, embeddings = _extract_embeddings(text)
@@ -61,17 +62,20 @@ def translate(text: str):
return " ".join(translate_chunk(c.text, c.lang) for c in chunks) return " ".join(translate_chunk(c.text, c.lang) for c in chunks)
class Translate: class Translate(io.ComfyNode):
@staticmethod @classmethod
def INPUT_TYPES(): def define_schema(cls):
return {"required": {"text": ("STRING", {"multiline": True})}} return io.Schema(
node_id="ETN_Translate",
display_name="Translate Text",
category="external_tooling",
inputs=[io.String.Input("text", multiline=True)],
outputs=[io.String.Output(display_name="translation")],
)
CATEGORY = "external_tooling" @classmethod
RETURN_TYPES = ("STRING",) def execute(cls, text: str):
FUNCTION = "translate" return io.NodeOutput(translate(text))
def translate(self, text: str):
return (translate(text),)
_lang_regex = re.compile(r"(lang:\w\w)") _lang_regex = re.compile(r"(lang:\w\w)")