Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0abe742480 | ||
|
|
b3ae4aa2d9 | ||
|
|
1b8d81ce5a | ||
|
|
ca01116495 | ||
|
|
5d3194f4d4 | ||
|
|
d5812f900b | ||
|
|
7064288fbe | ||
|
|
d82675092e | ||
|
|
a1e51904de | ||
|
|
ffa130239b | ||
|
|
2fd51d0d47 | ||
|
|
d3c75155b4 | ||
|
|
cbaef8d9c5 | ||
|
|
b2783d82a6 | ||
|
|
09759222de | ||
|
|
7fc3df1174 | ||
|
|
ed99942f86 | ||
|
|
2d395424ea | ||
|
|
9b9ea62dd8 | ||
|
|
ad36f89af3 | ||
|
|
7130dcb2df | ||
|
|
77186eda87 | ||
|
|
24a7bd1a77 | ||
|
|
2d14a03ad8 | ||
|
|
9d2e03e8d5 | ||
|
|
c5606f8e8f | ||
|
|
79e9b6426f | ||
|
|
ad320a218c | ||
|
|
a310f4593b | ||
|
|
7d957dcfa7 | ||
|
|
22cfd71f95 | ||
|
|
21a2f44d4c | ||
|
|
0220252912 | ||
|
|
f447ef70fa | ||
|
|
fb27a5bda8 | ||
|
|
aa83259e66 | ||
|
|
75c632df4b | ||
|
|
a088a2dde2 | ||
|
|
fbf99f2a08 | ||
|
|
dfe014ae88 | ||
|
|
6a7ae5ab70 | ||
|
|
d4cac6ac95 | ||
|
|
929fdfcc13 | ||
|
|
20f8d8ecc9 | ||
|
|
f555efb71b | ||
|
|
db4f296533 | ||
|
|
17c36ebc70 | ||
|
|
0697a3ac1f | ||
|
|
fa46b93329 | ||
|
|
fa84eec8fc | ||
|
|
5ef2fddc1b | ||
|
|
bff22b8351 | ||
|
|
ca2b59248e | ||
|
|
696899a5fc | ||
|
|
5f4373d71a | ||
|
|
61d2a19120 | ||
|
|
a6af76ac39 | ||
|
|
6a5c8e02e5 | ||
|
|
c2308a0762 | ||
|
|
b8e4659a10 | ||
|
|
93e1932456 | ||
|
|
ea755151fe | ||
|
|
facd65995a | ||
|
|
d7b18203a4 | ||
|
|
1839c099ad | ||
|
|
bed8b36705 | ||
|
|
8ed5591574 | ||
|
|
fe39d22eb9 | ||
|
|
50d3479fba | ||
|
|
d7d421baaa | ||
|
|
e10daee9ed | ||
|
|
50c3ffdf64 | ||
|
|
517790d1d6 | ||
|
|
e2bd09d7e9 | ||
|
|
035c68c629 | ||
|
|
e86973fedf | ||
|
|
19337dcc0e | ||
|
|
1d4ffe14bb | ||
|
|
20c8039a98 | ||
|
|
fcf678735c | ||
|
|
63ab33800e | ||
|
|
ef5ccfa98f | ||
|
|
0b01696f5b | ||
|
|
2fce4c56d5 | ||
|
|
a7f77032ec | ||
|
|
72335898cb | ||
|
|
9a9cbe78a5 | ||
|
|
9a8d90dd95 | ||
|
|
24b7aabf8b | ||
|
|
81f944f119 | ||
|
|
327b2a1fe3 | ||
|
|
fb847a5225 | ||
|
|
5a45172d02 | ||
|
|
29e24ec52c | ||
|
|
e5e62a4a79 | ||
|
|
1a24975f99 | ||
|
|
61fa161c34 | ||
|
|
f986f6a442 | ||
|
|
e0d0c3cc2c | ||
|
|
f5ec9d830c | ||
|
|
d1dcf12f10 | ||
|
|
cb92e547c6 | ||
|
|
b5fec4a062 | ||
|
|
42965013f9 | ||
|
|
d20615fb48 | ||
|
|
5bad00f72f | ||
|
|
df54344077 | ||
|
|
f42c0f29b6 | ||
|
|
547c3d5c97 | ||
|
|
73babbd00e | ||
|
|
cac32fe37c | ||
|
|
5620b5c6e2 | ||
|
|
3d4a960982 | ||
|
|
9d533984c2 | ||
|
|
715a41e04f | ||
|
|
aff32e8da6 | ||
|
|
e46123612d | ||
|
|
6e7b2445db | ||
|
|
2f39365248 | ||
|
|
c324f6741d | ||
|
|
c27b662fd8 |
@@ -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 }}
|
||||||
+3
-1
@@ -1,4 +1,6 @@
|
|||||||
.vscode
|
.vscode
|
||||||
.env
|
.env
|
||||||
.dev
|
.dev
|
||||||
__pycache__
|
__pycache__
|
||||||
|
|
||||||
|
safetychecker/*.safetensors
|
||||||
@@ -2,12 +2,20 @@
|
|||||||
|
|
||||||
Provides nodes and API geared towards using ComfyUI as a backend for external tools.
|
Provides nodes and API geared towards using ComfyUI as a backend for external tools.
|
||||||
|
|
||||||
## Nodes for sending and receiving images
|
* <a href="#images">Sending and receiving images</a>
|
||||||
|
* <a href="#regions">Regions (Attention Masking)
|
||||||
|
* <a href="#tiles">Tiled image processing
|
||||||
|
* <a href="#misc">Miscellanious nodes
|
||||||
|
* <a href="#api">Http API extensions (Model inspection)
|
||||||
|
* <a href="#installation">⭳ Installation</a>
|
||||||
|
|
||||||
|
|
||||||
|
## <a id="images" href="#toc">Sending and receiving images</a>
|
||||||
|
|
||||||
ComfyUI exchanges images via the filesystem. This requires a
|
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)
|
||||||
@@ -25,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:
|
||||||
```
|
```
|
||||||
@@ -36,9 +44,74 @@ That is two 32-bit integers (big endian) with values 1 and 2 followed by the PNG
|
|||||||
{'type': 'executed', 'data': {'node': '<node ID>', 'output': {'images': [{'source': 'websocket', 'content-type': 'image/png', 'type': 'output'}, ...]}, 'prompt_id': '<prompt ID>}}
|
{'type': 'executed', 'data': {'node': '<node ID>', 'output': {'images': [{'source': 'websocket', 'content-type': 'image/png', 'type': 'output'}, ...]}, 'prompt_id': '<prompt ID>}}
|
||||||
```
|
```
|
||||||
|
|
||||||
## Nodes for working on regions
|
### Load Image from Cache
|
||||||
|
|
||||||
When integrating ComfyUI into tools which use layers and compose them on the fly, it is useful to only receive relevant masked regions.
|
Loads an image or mask that has been uploaded previously into the workflow.
|
||||||
|
Uploaded images are temporarily stored in RAM rather than written to disk. This
|
||||||
|
method has less overhead compared to embedding images as base64 into the prompt,
|
||||||
|
but is more complex to implement.
|
||||||
|
* Inputs: id of an image that was uploaded previously
|
||||||
|
* Outputs: image (RGB) and mask (A of RGBA input, or first channel if no alpha present).
|
||||||
|
|
||||||
|
To upload an image, upload the _bytes_ of a PNG via a HTTP PUT request to
|
||||||
|
`/api/etn/image/{id}`. JPEG or other formats also work. Choose any `id` which
|
||||||
|
does not clash with other images you upload, and reference it in the node. The
|
||||||
|
request returns `201` if the image was uploaded and `200` if it was already
|
||||||
|
cached.
|
||||||
|
|
||||||
|
### Save Image to Cache
|
||||||
|
|
||||||
|
Stores an output image in RAM temporarily and allows retrieval over HTTP.
|
||||||
|
This is typically faster than WebSocket, especially for large images.
|
||||||
|
* Inputs: the image (RGB or RGBA). Batches are supported.
|
||||||
|
|
||||||
|
This node will send a JSON message over WebSocket when an image is ready:
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"type": "executed",
|
||||||
|
"data": {
|
||||||
|
"node": "<node ID>",
|
||||||
|
"output": {
|
||||||
|
"images": [
|
||||||
|
{"source": "http", "id": "<image ID>", "content-type": "image/png", "type": "output"}
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"prompt_id": "prompt ID"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
To download the images, send a HTTP GET request to `/api/etn/image/{id}` with
|
||||||
|
the image IDs from the message. Images will be cached for a few minutes.
|
||||||
|
|
||||||
|
|
||||||
|
## <a id="regions" href="#toc">Regions</a>
|
||||||
|
|
||||||
|
These nodes implement attention masking for arbitrary number of image regions. Text prompts only apply to the masked area.
|
||||||
|
In contrast to condition masking, this method is less "forceful", but leads to more natural image compositions.
|
||||||
|
|
||||||
|

|
||||||
|
[Workflow: region_attention_mask.json](workflows/region_attention_mask.json)
|
||||||
|
|
||||||
|
### Background Region
|
||||||
|
|
||||||
|
This node starts a list of regions. It takes a prompt, but no mask. The prompt is assigned to all image areas which are _not_
|
||||||
|
covered by another region mask in the list.
|
||||||
|
|
||||||
|
### Define Region
|
||||||
|
|
||||||
|
Appends a new region to a region list (or starts a new list). Takes a prompt, and mask which defines the area in the image
|
||||||
|
the prompt will apply to. Masks must be the same size as the image _or_ the latent (which is factor 8 smaller).
|
||||||
|
|
||||||
|
### List Region Masks
|
||||||
|
|
||||||
|
This node takes a list of regions and outputs all their masks. It can be useful for inspection, debugging or to reuse the
|
||||||
|
computed background mask.
|
||||||
|
|
||||||
|
### Regions Attention Mask
|
||||||
|
|
||||||
|
Patches the model to use the provided list of regions. This replaces the positive text conditioning which is provided
|
||||||
|
to the sampler. It's still possible to pass ControlNet and other conditioning to the sampler.
|
||||||
|
|
||||||
### Apply Mask to Image
|
### Apply Mask to Image
|
||||||
|
|
||||||
@@ -46,30 +119,170 @@ Copies a mask into the alpha channel of an image.
|
|||||||
* Inputs: image and mask
|
* Inputs: image and mask
|
||||||
* Outputs: RGBA image with mask used as transparency
|
* Outputs: RGBA image with mask used as transparency
|
||||||
|
|
||||||
## API for model inspection
|
|
||||||
|
|
||||||
There are various types of models that can be loaded as checkpoint, LoRA, ControlNet, etc. which cannot be used interchangeably. The following API helps to categorize and filter them.
|
## <a id="tiles" href="#toc">Tiles</a>
|
||||||
|
|
||||||
### /etn/model_info
|
Splitting an image into tiles to be processed individually is a useful method to speed up
|
||||||
|
diffusion and save VRAM. There are various nodes out there which provide a fixed pipeline.
|
||||||
|
In contrast, the following nodes only provide a way to split an image into tiles and merge
|
||||||
|
it back together. With tools and scripts it is feasible to generate individual workflows
|
||||||
|
for each tile. This allows maximum flexibility (different prompts, regions, control, etc.).
|
||||||
|
|
||||||
Lists available models with additional classification info.
|

|
||||||
* Paramters: _none_
|
[Workflow: image_tiles.json](workflows/image_tiles.json)
|
||||||
* Output: list of model files
|
|
||||||
```
|
|
||||||
{
|
|
||||||
"checkpoint_file.safetensors": {
|
|
||||||
"base_model": "sd15"|"sd20"|"sd21"|"sdxl"|"ssd1b"|"svd"|"cascade-b"|"cascade-c",
|
|
||||||
"is_inpaint": true|false,
|
|
||||||
"is_refiner": true|false
|
|
||||||
},
|
|
||||||
...
|
|
||||||
}
|
|
||||||
```
|
|
||||||
The entry is `{"base_model": "unknown"}` for models which are not in safetensors format or do not match any of the known base models.
|
|
||||||
|
|
||||||
_Note: currently only supports checkpoints. May add other models in the future._
|
### Create Tile Layout
|
||||||
|
|
||||||
## Installation
|
This node defines the tiling parameters:
|
||||||
|
* **min_tile_size**: Minimum resolution of each tile in pixels. Tiles may be larger to fit the image size evenly.
|
||||||
|
* **padding**: Padding around each tile in pixels. Overlaps with neighbour tiles. There is no padding at the image borders.
|
||||||
|
* **blending**: The part of the padding area which is used for smooth blending to avoid seams. Affects masks which are generated from this layout.
|
||||||
|
|
||||||
|
The number of tiles is: `image_size // (min_tile_size + 2 * padding)`
|
||||||
|
|
||||||
|
### Extract Image Tile
|
||||||
|
|
||||||
|
Splits out part of an image. Tile indices range from 0 to number of tiles and are column-major
|
||||||
|
(tile 1 is usually below tile 0).
|
||||||
|
|
||||||
|
### Extract Mask Tile
|
||||||
|
|
||||||
|
Same as "Extract Image Tile" but for masks.
|
||||||
|
|
||||||
|
### Merge Image Tile
|
||||||
|
|
||||||
|
Merges a tile into a full image, usually after sampling. Uses a smooth transition overlap
|
||||||
|
between neighbouring tiles depending on padding and blending values.
|
||||||
|
|
||||||
|
### Generate Tile Mask
|
||||||
|
|
||||||
|
Creates a coverage mask for a certain tile. The size of the mask matches the image tile size.
|
||||||
|
The image area will be white (1) and the padding area black (0), with a smooth transition
|
||||||
|
depending on the chosen blend size.
|
||||||
|
|
||||||
|
This mask is used internally by "Merge Image Tile", but it can also be useful as input for "Set Latent Noise Mask" in upscale workflows.
|
||||||
|
|
||||||
|
|
||||||
|
## <a id="misc" href="#toc">Miscellaneous Nodes</a>
|
||||||
|
|
||||||
|
<a id="node-translate"></a>
|
||||||
|
### Translate Text
|
||||||
|
|
||||||
|
Node which translates a string into English. The language to translate from is indicated with a
|
||||||
|
_language directive_ of the form `lang:xx` where xx is a 2-letter language code. Multiple
|
||||||
|
directives are allowed and change language for any text that comes after, until the next
|
||||||
|
directive. `lang:en` (the default) passes through text fragments untouched. Useful
|
||||||
|
for keywords, tags and such.
|
||||||
|
|
||||||
|
Examples:
|
||||||
|
| Input | Output |
|
||||||
|
|:-|:-|
|
||||||
|
| lang:de eine modische handtasche aus grünem kunstleder | a fashionable handbag made of green suede |
|
||||||
|
| origami paperwork, lang:zh 狐狸和鹤, lang:en mountain view | origami paperwork, Fox and crane, mountain view |
|
||||||
|
|
||||||
|
Translation happens entirely local, powered by [argosopentech/argos-translate](https://github.com/argosopentech/argos-translate):
|
||||||
|
* Install with `pip install argostranslate` or `pip install -r requirements.txt`
|
||||||
|
* Models are automatically downloaded on first use.
|
||||||
|
|
||||||
|
There is also a [translation API](#api-translation) for immediate feedback in tool UI.
|
||||||
|
|
||||||
|
### NSFW Filter
|
||||||
|
|
||||||
|
Checks images for NSFW content using [Safety-Checker](https://huggingface.co/CompVis/stable-diffusion-safety-checker). Images which don't pass the check are blurred to
|
||||||
|
obfuscate contents. Model is downloaded on first use.
|
||||||
|
|
||||||
|
Inputs: image and sensitivity (0.5 for explicit content only, 0.7+ to include partial nudity).
|
||||||
|
|
||||||
|
**Important:** the filter isn't perfect. Some explicit content may slip through.
|
||||||
|
|
||||||
|
|
||||||
|
## <a id="api" href="#toc">API extensions</a>
|
||||||
|
|
||||||
|
### GET /api/etn/model_info/{folder_name}
|
||||||
|
|
||||||
|
There are various types of models that can be loaded as checkpoint, LoRA, ControlNet, etc. which cannot be used interchangeably. This endpoint helps to categorize and filter them.
|
||||||
|
|
||||||
|
#### Paramters
|
||||||
|
* `folder_name`: sub-directory in ComfyUI's models folder.
|
||||||
|
Supported model types: `checkpoints`, `diffusion_models`, `unet`, `unet_gguf`
|
||||||
|
* `limit=n`: (query parameter, optional) inspect at `n` models
|
||||||
|
* `offset=i`: (query parameter, optional) start with the `i`th model
|
||||||
|
|
||||||
|
#### Output
|
||||||
|
Lists available models with additional classification info:
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"checkpoint_file.safetensors": {
|
||||||
|
"base_model": "sd15",
|
||||||
|
"is_inpaint": false,
|
||||||
|
"type": "eps"
|
||||||
|
},
|
||||||
|
...
|
||||||
|
}
|
||||||
|
```
|
||||||
|
Possible values for base model: `sd15, sd20, sd21, sd3, sdxl, sdxl-refiner, ssd1b, svd, cascade-b, cascade-c, aura-flow, hunyuan-dit, flux, flux-schnell, flux2, lumina2, z-image, chroma, qwen-image`
|
||||||
|
|
||||||
|
If base model is `sdxl`, the `type` attribute is set with possible values: `eps, edm, v-prediction, v-prediction-edm`
|
||||||
|
|
||||||
|
Detection supports quantized models:
|
||||||
|
* GGUF: if the `gguf` module is installed, .gguf files are detected and will set the `quant` field
|
||||||
|
* Nunchaku: SVDQuant models are detected and will set the `quant` field to `svdq`
|
||||||
|
|
||||||
|
Returns an entry `{"base_model": "unknown"}` for models with unknown format or which do not match any of the known base models.
|
||||||
|
|
||||||
|
|
||||||
|
#### Pagination
|
||||||
|
The query parameters limit and offset allow inspecting a subset of models per request.
|
||||||
|
Usually inspection is quite fast (it only looks at model headers), but it can be slow
|
||||||
|
in some cases due to anti-virus or slow harddrives.
|
||||||
|
```
|
||||||
|
GET /api/etn/model_info/checkpoints?limit=10&offset=20
|
||||||
|
```
|
||||||
|
This will return at most 10 models, starting with the 20th model in the list.
|
||||||
|
It also returns a special `_meta` entry in the output JSON:
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"checkpoint_20.safetensors": { ... },
|
||||||
|
"_meta": { "offset": 20, "count": 1, "total": 21 }
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
|
||||||
|
### GET /api/etn/languages
|
||||||
|
|
||||||
|
Returns a list of available languages for translation.
|
||||||
|
```json
|
||||||
|
[
|
||||||
|
{ "name": "English", "code": "en" },
|
||||||
|
{ ... }
|
||||||
|
]
|
||||||
|
```
|
||||||
|
|
||||||
|
<a id="api-translation"></a>
|
||||||
|
### GET /api/etn/translate/{lang}/{text}
|
||||||
|
|
||||||
|
Translates `text` into English. `lang` is a 2-letter code indicating the language to translate
|
||||||
|
from. `text` may also contain _language directives_ to only translate some fragments.
|
||||||
|
See the [node documentation](#node-translate) for details.
|
||||||
|
|
||||||
|
* Output: JSON string
|
||||||
|
* Example: `/api/etn/translate/de/eine%20modische%20Handtasche` -> `"a fashionable handbag"`
|
||||||
|
|
||||||
|
### PUT /api/etn/upload/{folder_name}/{filename}
|
||||||
|
|
||||||
|
Uploads a model to ComfyUI's local model folder.
|
||||||
|
|
||||||
|
#### Parameters
|
||||||
|
* `folder_name`: the model type. Must match one of the existing folders in ComfyUI's models folder.
|
||||||
|
* `filename`: target filename for the model. Must not contain any (absolute or relative) path. Extension must be .safetensors.
|
||||||
|
|
||||||
|
#### Output
|
||||||
|
* Code `201` and `{ "status": "success" }` after successful upload.
|
||||||
|
* Code `200` and `{ "status": "cached" }` if the file already exists.
|
||||||
|
* Code `400` and `{ "error": "..." }` if the parameters are invalid.
|
||||||
|
|
||||||
|
|
||||||
|
## <a id="installation" href="#toc">Installation</a>
|
||||||
|
|
||||||
Download the repository and unpack into the `custom_nodes` folder in the ComfyUI installation directory.
|
Download the repository and unpack into the `custom_nodes` folder in the ComfyUI installation directory.
|
||||||
|
|
||||||
@@ -80,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)
|
||||||
|
|||||||
+56
-35
@@ -1,36 +1,57 @@
|
|||||||
from . import api, nodes, tile, region
|
from comfy_api.latest import ComfyExtension, io
|
||||||
|
|
||||||
NODE_CLASS_MAPPINGS = {
|
from . import api as api
|
||||||
"ETN_LoadImageBase64": nodes.LoadImageBase64,
|
from . import control, krita, nodes, region, tile, translation
|
||||||
"ETN_LoadMaskBase64": nodes.LoadMaskBase64,
|
|
||||||
"ETN_SendImageWebSocket": nodes.SendImageWebSocket,
|
|
||||||
"ETN_CropImage": nodes.CropImage,
|
class ExternalToolingNodes(ComfyExtension):
|
||||||
"ETN_ApplyMaskToImage": nodes.ApplyMaskToImage,
|
async def get_node_list(self) -> list[type[io.ComfyNode]]:
|
||||||
"ETN_TileLayout": tile.TileLayout,
|
node_list = [
|
||||||
"ETN_ExtractImageTile": tile.ExtractImageTile,
|
nodes.LoadImageCache,
|
||||||
"ETN_ExtractMaskTile": tile.ExtractMaskTile,
|
nodes.SaveImageCache,
|
||||||
"ETN_GenerateTileMask": tile.GenerateTileMask,
|
nodes.LoadImageBase64,
|
||||||
"ETN_MergeImageTile": tile.MergeImageTile,
|
nodes.LoadMaskBase64,
|
||||||
"ETN_BackgroundRegion": region.BackgroundRegion,
|
nodes.SendImageWebSocket,
|
||||||
"ETN_DefineRegion": region.DefineRegion,
|
nodes.ApplyMaskToImage,
|
||||||
"ETN_ListRegionMasks": region.ListRegionMasks,
|
nodes.ReferenceImage,
|
||||||
"ETN_AttentionMask": region.AttentionMask,
|
nodes.ApplyReferenceImages,
|
||||||
}
|
tile.CreateTileLayout,
|
||||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
tile.ExtractImageTile,
|
||||||
"ETN_LoadImageBase64": "Load Image (Base64)",
|
tile.ExtractMaskTile,
|
||||||
"ETN_LoadMaskBase64": "Load Mask (Base64)",
|
tile.GenerateTileMask,
|
||||||
"ETN_SendImageWebSocket": "Send Image (WebSocket)",
|
tile.MergeImageTile,
|
||||||
"ETN_CropImage": "Crop Image",
|
region.BackgroundRegion,
|
||||||
"ETN_ApplyMaskToImage": "Apply Mask to Image",
|
region.DefineRegion,
|
||||||
"ETN_ListAppend": "List 🢒 Append",
|
region.ListRegionMasks,
|
||||||
"ETN_ListElement": "List 🢒 Get Element",
|
region.AttentionMask,
|
||||||
"ETN_TileLayout": "Create Tile Layout",
|
translation.Translate,
|
||||||
"ETN_ExtractImageTile": "Extract Image Tile",
|
krita.KritaOutput,
|
||||||
"ETN_ExtractMaskTile": "Extract Mask Tile",
|
krita.KritaSendText,
|
||||||
"ETN_MergeImageTile": "Merge Image Tile",
|
krita.KritaCanvas,
|
||||||
"ETN_GenerateTileMask": "Generate Tile Mask",
|
krita.KritaSelection,
|
||||||
"ETN_BackgroundRegion": "Background Region",
|
krita.KritaImageLayer,
|
||||||
"ETN_DefineRegion": "Define Region",
|
krita.KritaMaskLayer,
|
||||||
"ETN_ListRegionMasks": "List Region Masks",
|
krita.Parameter,
|
||||||
"ETN_AttentionMask": "Regions Attention Mask",
|
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()
|
||||||
|
|
||||||
|
|
||||||
|
WEB_DIRECTORY = "./js"
|
||||||
|
|||||||
@@ -1,13 +1,22 @@
|
|||||||
|
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
|
||||||
import json
|
import json
|
||||||
|
import traceback
|
||||||
|
import re
|
||||||
|
import logging
|
||||||
|
import itertools
|
||||||
|
|
||||||
import comfy.utils
|
|
||||||
from comfy import supported_models
|
|
||||||
from comfy import model_detection
|
from comfy import model_detection
|
||||||
|
import comfy.utils
|
||||||
import folder_paths
|
import folder_paths
|
||||||
import server
|
import server
|
||||||
|
|
||||||
|
from .translation import available_languages, translate
|
||||||
|
from .krita import WorkflowExchange
|
||||||
|
from .nodes import image_cache
|
||||||
|
|
||||||
input_block_name = "model.diffusion_model.input_blocks.0.0.weight"
|
input_block_name = "model.diffusion_model.input_blocks.0.0.weight"
|
||||||
|
|
||||||
model_names = {
|
model_names = {
|
||||||
@@ -15,12 +24,48 @@ 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",
|
||||||
"Stable_Cascade_B": "cascade-b",
|
"Stable_Cascade_B": "cascade-b",
|
||||||
"Stable_Cascade_C": "cascade-c",
|
"Stable_Cascade_C": "cascade-c",
|
||||||
|
"SD3": "sd3",
|
||||||
|
"AuraFlow": "aura-flow",
|
||||||
|
"HunyuanDiT": "hunyuan-dit",
|
||||||
|
"HunyuanDiT1": "hunyuan-dit",
|
||||||
|
"Flux": "flux",
|
||||||
|
"FluxInpaint": "flux",
|
||||||
|
"FluxSchnell": "flux-schnell",
|
||||||
|
"GenmoMochi": "mochi",
|
||||||
|
"LTXV": "ltxv",
|
||||||
|
"HunyuanVideo": "hunyuan-video",
|
||||||
|
"CosmosT2V": "cosmos",
|
||||||
|
"CosmosI2V": "cosmos",
|
||||||
|
"CosmosT2IPredict2": "cosmos-predict2",
|
||||||
|
"CosmosI2VPredict2": "cosmos-predict2",
|
||||||
|
"ZImage": "z-image",
|
||||||
|
"Lumina2": "lumina2",
|
||||||
|
"WAN21_T2V": "wan21",
|
||||||
|
"WAN21_I2V": "wan21",
|
||||||
|
"WAN21_FunControl2V": "wan21-fun",
|
||||||
|
"WAN21_Vace": "wan21-vace",
|
||||||
|
"WAN21_Camera": "wan21-camera",
|
||||||
|
"HiDream": "hi-dream",
|
||||||
|
"Chroma": "chroma",
|
||||||
|
"ACEStep": "ace-step",
|
||||||
|
"Omnigen2": "omnigen2",
|
||||||
|
"QwenImage": "qwen-image",
|
||||||
|
"QwenImage21": "qwen-image21",
|
||||||
|
"ErnieImage": "ernie-image",
|
||||||
|
"Flux2": "flux2",
|
||||||
|
"Anima": "anima",
|
||||||
|
"Krea2": "krea2",
|
||||||
|
}
|
||||||
|
|
||||||
|
gguf_architectures = {
|
||||||
|
"sd1": "sd15",
|
||||||
|
"qwen_image": "qwen-image",
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -35,10 +80,10 @@ class FakeTensor(NamedTuple):
|
|||||||
return d
|
return d
|
||||||
|
|
||||||
|
|
||||||
def inspect_checkpoint(filename):
|
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("checkpoints", filename)
|
path = folder_paths.get_full_path(model_type, filename)
|
||||||
header = comfy.utils.safetensors_header(path)
|
header = comfy.utils.safetensors_header(path)
|
||||||
if header:
|
if header:
|
||||||
cfg = json.loads(header.decode("utf-8"))
|
cfg = json.loads(header.decode("utf-8"))
|
||||||
@@ -49,11 +94,14 @@ def inspect_checkpoint(filename):
|
|||||||
cfg[key] = FakeTensor.from_dict(cfg[key])
|
cfg[key] = FakeTensor.from_dict(cfg[key])
|
||||||
|
|
||||||
# Reuse Comfy's model detection
|
# Reuse Comfy's model detection
|
||||||
unet_args = [cfg, "model.diffusion_model.", "F32"]
|
prefix = model_detection.unet_prefix_from_state_dict(cfg)
|
||||||
|
if not is_checkpoint:
|
||||||
|
cfg = comfy.utils.state_dict_prefix_replace(cfg, {prefix: ""}, filter_keys=False)
|
||||||
|
prefix = ""
|
||||||
try: # latest ComfyUI takes 2 args
|
try: # latest ComfyUI takes 2 args
|
||||||
unet_config = model_detection.detect_unet_config(*unet_args[:-1])
|
unet_config = model_detection.detect_unet_config(cfg, prefix)
|
||||||
except TypeError as e: # older ComfyUI versions take 3 args
|
except TypeError as e: # older ComfyUI versions take 3 args
|
||||||
unet_config = model_detection.detect_unet_config(*unet_args)
|
raise TypeError(f"{e} when calling detect_unet_config - old version of ComfyUI?")
|
||||||
|
|
||||||
# Get input count to detect inpaint models
|
# Get input count to detect inpaint models
|
||||||
if input_block := cfg.get(input_block_name, None):
|
if input_block := cfg.get(input_block_name, None):
|
||||||
@@ -62,31 +110,325 @@ def inspect_checkpoint(filename):
|
|||||||
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:
|
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}"}
|
return {"base_model": "unknown", "error": f"Failed to detect base model: {e}"}
|
||||||
|
|
||||||
|
|
||||||
if _server := getattr(server.PromptServer, "instance", None):
|
def detect_svdq(cfg: dict) -> str | None:
|
||||||
|
if md := cfg.get("__metadata__"):
|
||||||
|
if comfy_config := md.get("comfy_config"):
|
||||||
|
if isinstance(comfy_config, str):
|
||||||
|
comfy_config = json.loads(comfy_config)
|
||||||
|
if model_class := comfy_config.get("model_class"):
|
||||||
|
return model_class
|
||||||
|
|
||||||
@_server.routes.get("/etn/model_info")
|
match md.get("model_class"):
|
||||||
async def model_info(request):
|
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:
|
try:
|
||||||
info = {
|
if file_type := reader.get_field("general.file_type"):
|
||||||
filename: inspect_checkpoint(filename)
|
result["quant"] = file_type.contents().lower()
|
||||||
for filename in folder_paths.get_filename_list("checkpoints")
|
except Exception:
|
||||||
}
|
result["quant"] = "gguf"
|
||||||
return web.json_response(info)
|
return result
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
# traceback.print_exc()
|
||||||
|
return {"base_model": "unknown", "error": f"Failed to detect base model: {e}"}
|
||||||
|
|
||||||
|
|
||||||
|
def inspect_diffusion_model(filename: str, model_type: str, is_checkpoint: bool):
|
||||||
|
if filename.endswith(".gguf"):
|
||||||
|
return inspect_gguf(filename, model_type)
|
||||||
|
return inspect_safetensors(filename, model_type, is_checkpoint)
|
||||||
|
|
||||||
|
|
||||||
|
def inspect_models(model_type: str, params: dict[str, str]):
|
||||||
|
try:
|
||||||
|
try:
|
||||||
|
files = folder_paths.get_filename_list(model_type)
|
||||||
|
except KeyError:
|
||||||
|
return web.json_response({"error": f"Model folder not found: {model_type}"})
|
||||||
|
limit = int(params.get("limit", "1000"))
|
||||||
|
offset = int(params.get("offset", "0"))
|
||||||
|
files_range = files[offset : offset + limit]
|
||||||
|
is_checkpoint = model_type == "checkpoints"
|
||||||
|
info = {
|
||||||
|
filename: inspect_diffusion_model(filename, model_type, is_checkpoint)
|
||||||
|
for filename in files_range
|
||||||
|
}
|
||||||
|
if "limit" in params:
|
||||||
|
info["_meta"] = dict(offset=offset, count=len(files_range), total=len(files))
|
||||||
|
return web.json_response(info)
|
||||||
|
except Exception as e:
|
||||||
|
traceback.print_exc()
|
||||||
|
return web.json_response(dict(error=str(e)), status=500)
|
||||||
|
|
||||||
|
|
||||||
|
def has_invalid_folder_name(folder_name: str):
|
||||||
|
valid_names = list(folder_paths.folder_names_and_paths.keys())
|
||||||
|
if folder_name not in valid_names:
|
||||||
|
return web.json_response(
|
||||||
|
dict(error=f"Invalid folder path, must be one of {', '.join(valid_names)}"),
|
||||||
|
status=400,
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def has_invalid_filename(filename: str):
|
||||||
|
if not filename.lower().endswith((".sft", ".safetensors")):
|
||||||
|
return web.json_response(dict(error="File extension must be .safetensors"), status=400)
|
||||||
|
if not filename or not filename.strip() or len(filename) > 255:
|
||||||
|
return web.json_response(dict(error="Invalid filename"), status=400)
|
||||||
|
if any(char in filename for char in ["..", "/", "\\", "\n", "\r", "\t", "\0"]):
|
||||||
|
return web.json_response(dict(error="Invalid filename"), status=400)
|
||||||
|
if filename.startswith(".") or not re.match(r"^[a-zA-Z0-9_\-. ]+$", filename):
|
||||||
|
return web.json_response(dict(error="Invalid filename"), status=400)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
async def image_sender(data: bytes):
|
||||||
|
mem = memoryview(data)
|
||||||
|
csize = 2**14
|
||||||
|
for i in range(0, len(mem), csize):
|
||||||
|
yield mem[i : i + csize]
|
||||||
|
|
||||||
|
|
||||||
|
_server: server.PromptServer | None = getattr(server.PromptServer, "instance", None)
|
||||||
|
if _server is not None:
|
||||||
|
_workflow_exchange = WorkflowExchange(_server)
|
||||||
|
|
||||||
|
@_server.routes.get("/api/etn/model_info/{folder_name}")
|
||||||
|
async def model_info(request: web.Request):
|
||||||
|
folder_name = request.match_info.get("folder_name", "checkpoints")
|
||||||
|
error = has_invalid_folder_name(folder_name)
|
||||||
|
if error is not None:
|
||||||
|
return error
|
||||||
|
return inspect_models(folder_name, request.rel_url.query)
|
||||||
|
|
||||||
|
@_server.routes.get("/api/etn/model_info")
|
||||||
|
async def api_model_info(request):
|
||||||
|
return inspect_models("checkpoints", request.rel_url.query)
|
||||||
|
|
||||||
|
@_server.routes.get("/api/etn/languages")
|
||||||
|
async def languages(request):
|
||||||
|
try:
|
||||||
|
result = [dict(name=name, code=code) for code, name in available_languages()]
|
||||||
|
return web.json_response(result)
|
||||||
except Exception as e:
|
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/translate/{lang}/{text}")
|
||||||
|
async def translate_text(request):
|
||||||
|
try:
|
||||||
|
language = request.match_info.get("lang", "en")
|
||||||
|
text = request.match_info.get("text", "")
|
||||||
|
result = translate(f"lang:{language} {text}")
|
||||||
|
return web.json_response(result)
|
||||||
|
except Exception as e:
|
||||||
|
return web.json_response(dict(error=str(e)), status=500)
|
||||||
|
|
||||||
|
@_server.routes.get("/api/etn/image/{id}")
|
||||||
|
async def get_image(request: web.Request):
|
||||||
|
try:
|
||||||
|
id = request.match_info.get("id", "")
|
||||||
|
data, content_type = image_cache.get(id)
|
||||||
|
if data is None or content_type is None:
|
||||||
|
return web.json_response(dict(error="Image not found"), status=404)
|
||||||
|
response = web.Response(
|
||||||
|
body=image_sender(data),
|
||||||
|
content_type=content_type,
|
||||||
|
headers={"Content-Length": str(len(data))},
|
||||||
|
)
|
||||||
|
return response
|
||||||
|
except Exception as e:
|
||||||
|
return web.json_response(dict(error=str(e)), status=500)
|
||||||
|
|
||||||
|
async def put_image(request: web.Request):
|
||||||
|
try:
|
||||||
|
id = request.match_info.get("id", "")
|
||||||
|
if id in image_cache:
|
||||||
|
await request.release() # Consume and discard the data to avoid connection abort
|
||||||
|
return web.json_response(dict(status="cached"), status=200)
|
||||||
|
|
||||||
|
content_type = request.headers.get("Content-Type", "application/octet-stream")
|
||||||
|
data = bytearray()
|
||||||
|
async for chunk, _ in request.content.iter_chunks():
|
||||||
|
data.extend(chunk)
|
||||||
|
|
||||||
|
image_cache.insert(id, bytes(data), content_type)
|
||||||
|
return web.json_response(dict(status="success"), status=201)
|
||||||
|
except Exception as e:
|
||||||
|
return web.json_response(dict(error=str(e)), status=500)
|
||||||
|
|
||||||
|
async def _put_image_expect_handler(request: web.Request):
|
||||||
|
if request.match_info.get("id", "") in image_cache:
|
||||||
|
# Skip "100 Continue" since we don't need the data, return 200 immediately.
|
||||||
|
return web.json_response(dict(status="cached"), status=200)
|
||||||
|
# otherwise run default aiohttp handler
|
||||||
|
return None
|
||||||
|
|
||||||
|
_server.app.router.add_route(
|
||||||
|
"PUT", "/api/etn/image/{id}", put_image, expect_handler=_put_image_expect_handler
|
||||||
|
)
|
||||||
|
|
||||||
|
@_server.routes.put("/api/etn/upload/{folder_name}/{filename}")
|
||||||
|
async def upload(request: web.Request):
|
||||||
|
folder_name = request.match_info.get("folder_name", "")
|
||||||
|
error = has_invalid_folder_name(folder_name)
|
||||||
|
if error is not None:
|
||||||
|
return error
|
||||||
|
|
||||||
|
filename = request.match_info.get("filename", "")
|
||||||
|
error = has_invalid_filename(filename)
|
||||||
|
if error is not None:
|
||||||
|
return error
|
||||||
|
|
||||||
|
try:
|
||||||
|
if folder_paths.get_full_path(folder_name, filename) is not None:
|
||||||
|
return web.json_response(dict(status="cached"), status=200)
|
||||||
|
|
||||||
|
folder = Path(folder_paths.folder_names_and_paths[folder_name][0][0])
|
||||||
|
total_size = int(request.headers.get("Content-Length", "0"))
|
||||||
|
logging.info(
|
||||||
|
f"Uploading {filename} ({total_size / (1024**2):.1f} MB) to {folder} folder"
|
||||||
|
)
|
||||||
|
|
||||||
|
with open(folder / filename, "wb") as f:
|
||||||
|
async for chunk, _ in request.content.iter_chunks():
|
||||||
|
f.write(chunk)
|
||||||
|
|
||||||
|
return web.json_response(dict(status="success"), status=201)
|
||||||
|
except Exception as e:
|
||||||
|
return web.json_response(dict(error=str(e)), status=500)
|
||||||
|
|
||||||
|
async def _handle_workflow_request(request: web.Request, handler, *arg_keys):
|
||||||
|
try:
|
||||||
|
data = await request.json()
|
||||||
|
args = [data[key] for key in arg_keys]
|
||||||
|
await handler(*args)
|
||||||
|
return web.json_response(dict(status="success"), status=200)
|
||||||
|
except KeyError as e:
|
||||||
|
return web.json_response(dict(error=str(e)), status=400)
|
||||||
|
except Exception as e:
|
||||||
|
return web.json_response(dict(error=str(e)), status=500)
|
||||||
|
|
||||||
|
@_server.routes.post("/api/etn/workflow/publish")
|
||||||
|
async def publish_workflow(request: web.Request):
|
||||||
|
return await _handle_workflow_request(
|
||||||
|
request, _workflow_exchange.publish, "name", "client_id", "workflow"
|
||||||
|
)
|
||||||
|
|
||||||
|
@_server.routes.post("/api/etn/workflow/subscribe")
|
||||||
|
async def subscribe_workflow(request: web.Request):
|
||||||
|
return await _handle_workflow_request(request, _workflow_exchange.subscribe, "client_id")
|
||||||
|
|
||||||
|
@_server.routes.post("/api/etn/workflow/unsubscribe")
|
||||||
|
async def unsubscribe_workflow(request: web.Request):
|
||||||
|
return await _handle_workflow_request(request, _workflow_exchange.unsubscribe, "client_id")
|
||||||
|
|||||||
+859
@@ -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,)
|
||||||
Binary file not shown.
|
After Width: | Height: | Size: 26 KiB |
@@ -0,0 +1,235 @@
|
|||||||
|
import { app } from "/scripts/app.js"
|
||||||
|
import { api } from "/scripts/api.js"
|
||||||
|
|
||||||
|
(function() {
|
||||||
|
|
||||||
|
// Workflow publishing
|
||||||
|
// - done whenever the graph changes, as long as there is a KritaOutput node
|
||||||
|
|
||||||
|
let publisherRegistered = false
|
||||||
|
|
||||||
|
async function publishWorkflow(e) {
|
||||||
|
const prompt = await app.graphToPrompt()
|
||||||
|
await api.fetchApi("/api/etn/workflow/publish", {
|
||||||
|
method: "POST",
|
||||||
|
body: JSON.stringify({
|
||||||
|
name: "ComfyUI Web",
|
||||||
|
client_id: api.clientId,
|
||||||
|
workflow: prompt["output"]
|
||||||
|
}, null, 2)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
// Image background for nodes
|
||||||
|
// - this is just for visuals
|
||||||
|
|
||||||
|
function loadImage(base64) {
|
||||||
|
const image = new Image()
|
||||||
|
image.src = base64
|
||||||
|
// image.onerror = () => console.error("Failed to load image");
|
||||||
|
return image
|
||||||
|
}
|
||||||
|
|
||||||
|
const canvasIcon = loadImage("data:image/webp;base64,UklGRg4KAABXRUJQVlA4WAoAAAAQAAAAYwAAYwAAQUxQSNsDAAARoIRs/yI5+uAHU1kZiOu6u7u7xnObuK9NpBoKCoqi1jd6WonrsD7ROc2e3OJJr28T94aGhobYj+8w9v/X/7+n3UNETAD+b1IW1crLLrMhbYJbRs9Zv2VvuXbqVK28d8v6OaNvCdpIXkaT5InJgJgRABeNbirRYKlp9EUAJB8ZlUq2XAtI1wQIhn1RIUlV1Y5UVUmy0jwsACQPt7CtsloQSBcE6D6jSFKVRlVJFmd0B8QeClSSSn7/ACCdQjBjL0mlRSW5d0YA+4KESpJKnd8dHQsweBeptK7krsGAWIIgppKkkn+NAKSNoO9KUplLJVf2hliCIKWyrZIfXwYIgCdLzHXpSVgXvElt07bcKAAKSs2TUgvWIHiH2p6SX98pi+jgIrEFwUJqO6Ty1HY6qGwObEHwPrU9KqkOUNkS2AKwjNoeqXRSuV7EGtZQO3BVuQhiDV9Q3aIyhFiTDVS3qHwMYgtBK9Utcn932A/20nHlGoglwdOn1DEqn4XYQfA7PfB7AKuCRiqdVzZCbKD+EL14qB4WBTOoPlA2Qsyh7i9f/FUH44LBVHpRORhibr0/ms2hf5XerPaFYcFoqi+UDeaafNJkCsFeenRvnRnBLfSp3mBqNNUfygZTc/wyx9QGv6w3A2zxyxZDwV56da+h3mW/lA1dVPPLKUOXnfIL/8UurvnltKH+Vb9UDXU/4ZeyoaDkl5Ih/ET1h/InU81+aTYjmOOXOaYm06sTzAD3n/JJ7X5T/YtUXyiLfU3JGnp0jZjCS+oPfQnG7ylR/aAs3WmuexO92dTdHKZWfFGZCot3fkX1gbL1dhv1SZVerCb1NvB0K9U9ZevTsNo/OUEPnkj628Hgj33w8WBY7h9up7ql3B72tYUn51ToeHnOk7B+/tSV6pYum3q+PVwbtVDdUbZE1yKPzyY/UV1R/hg/i1zKhKRIdUNZTCZIPtB9Rvo71QXl78mM7sjrxVFSpJPFJLoY+b06jn904ccovhp5vjaOWkjNk5Ithfha5PuKqLCozFyXFxWiK5D3/o1xuiVPW9K4sT/yf35DFi47mpejy8Ks4Xw4+UAcxRvKeaisj6P4Abjaf3xWSDYcbaNmtM3RDUkhG98fDt8+IyvEK3fV2Fa1M6psW9u1Mi5kM26H28EDM7IofPPj7WUaLG//+M0wymY8EMD54M7JaRqF8cKPvy4eqtROq56uVQ4Vv/54YRxGaTr5zgB+vOjpl5IsicIwSrJ35sx5J0uiMIySLHnp6YvgUel/z6iXojTL0jRJ0jTL0uilUff0F/i3/qIb7nn06WHDnn70nhsuqhf8hxQAVlA4IAwGAACwHgCdASpkAGQAPm0wk0akIqGhLRGrUIANiWYA1BHh/t2rC93/Hf2Was/feJ2MrzB6VP6M9gD9Nunp5rP2y9Z70q+gV/Y/+B1kHoAeWv+1Xwa/t5+5PtQXQfhgKtwfp++wwU1Dx6fSXsD+VV7JPQ5/aRrblS2nsaggvO1Mch8UhK6pYtQxLM/VgrswZ0vLV8b6SwudSWaCFSHUiXEUQWX6krc9GtWanHMeaDd9wRYCfO5TwpYkgGAIkaLI4p6taB375EUfaVubYzKMfHSz2KpivsjWF0Vf+YJbACgi8j86d6EiJhQFF31NBJdS+QrGtJ2RUJRbahp1MXso6/J8AAD+/TKL/9q5tf/zOmF5Fe8B0Zn0yX3C0VLv0zxxv2+WH/dbz//rc2RS4TC1UzFVQiXVn5+Y0r+RsfJPsfPNuT02INz8gty7fI7fA/D1Wj2Jv+4RwdpyXs+cRxaT84bme5rMmPf+BH7NDUPKsj7GJ+w/6nBW2vsiPalWPfvBk6AQ3kCHmVecXkcnOgpoZ4ruAF/9Ze93DG5/8Y32x8b/CKPRt1jaXXy2LnoPvSNUT77gbB+/7vI1pfBfUHJsSwheIXY7QSixh7Ya8IliO3wqvI/uIFZAZd9pL8R1gRpYouBoyL5uIuGWQAZC5SKY0SruTf66stUOJVO9hlokeb5lWVzo7FO/Oeb/oj9iK4bqFhNZLCfqsBlH/OeefoP9sFdl7Mq1xmsevmzkfgwyiXg5hxMIP/Wa0JMPVl+XEFqTveAf1M8IBDu/pX/hCEnMn1n15Smyf72eDXKQqBrvp6BugyXXaJ05FDoz8MONUFh4rcjGL7AcijbcZ0SYwJkoeAKBW/I/sjKzTRtTP2E1fLB/8TWnzieHznDAKdlTuY2nSVTwCqFZNcFeFn7boziHOmYBLJin52d874mq1pHmJnulhT96LbKVW4vAT5PnY5F9TzmnDMwIFm4IAuEaA8X8XLE4Hp+AUEG4oswxRbVfOfxNJRyxFO3UB+v+ALgMP8kOf0uK3/3WOq4o/roivfzvW/fXviTC0mx+352hGaO+axx6vFa3eIkFsUEXCdo2LFHIlM8BtPuGUhgvM3oygIMAgvmUKILe0DFYVXhG/QoLi3sYfaoK/f0tX+fNnXhhxwEj1/Ct2Z64g0qWmkgwNkyy8m90EK1HsX0Q10CHVakDZePz5ts37u3GCANwGHQzWB+hNsevqjuU3qT95yGs0jjOtI/IjKsH9JbAmZkjGvNCPOC+FYUkOkwao9sOESY6zCgx9CM7g2LU4/CSHGoe2t0vWV/cMDH1HzI+Wa/yYp9CLDIh7J7iJd/2KnixeJvOhbUvbr9gubyyQU1iO5bnD9T536j++jKDVIk0Fwzk+d+j2eueHsIFJUvdyo2TyxP0kJbWr36R1s3giryqPvrsR5SkXx16+xqDrX4elhqh+1FwzNnSF5Lj5EUT/UC2rJvoAikbnvQ3NtJ9e83++idf3ja4FaLcUDxhoN5Rl5Ziz1LvF9iVeb6Su0QWYoRyBbyZ/pRbgYyhlAU/tonH7Wt+KhPDmXKIo0u4FDbAXM8avbFk4ax6e/dYITOCe+9dVEgcTOnBfhv0Yotd3EzNjZkLz4ksKGtFXcWIZRJ5YAyfzPYsyPex6/6ud9r2Ha9oxhVSIJV418e83qcPOIPlpe+LVGc69W6eC83l/zloqM9D6zQMkfqrjBZNpRkQS0sn8sxSu3s5qzhtH8cvjZk83gMqdfnHnl+1bvA7BI/g4+ePU7HUb9vK3Qw35bVmDcXa8xxWS2NQj8iWMH1cbHLXlboQsaCxIZoo+SeXR6ePUw3k6C/OxgqhjzExMJjLdBjoBeWYt3RPG2foTvx0T0Iz8ukdrCRJMG6HaR+6/f4nG/4xkr/fLhGlqOE/hBDBhuqnANj1CrujVDs2YayTvPuIcqCpNd3i8fOR8DfCq9ytS55F8akKneS6poHfB3bhjWbcIQXPvFzS7S5xLHEVoaixOwp0TL/8cQ8dxriyeddu5kCTyY7KepMQoeR+Pyn04nElkt9qqfYCTqHDtXBriC/UZh9AAAAAAAA=")
|
||||||
|
|
||||||
|
function setIconImage(nodeType, image, size, padRows, padCols) {
|
||||||
|
const onAdded = nodeType.prototype.onAdded
|
||||||
|
nodeType.prototype.onAdded = function () {
|
||||||
|
onAdded?.apply(this, arguments)
|
||||||
|
this.size = size
|
||||||
|
}
|
||||||
|
|
||||||
|
const onDrawBackground = nodeType.prototype.onDrawBackground
|
||||||
|
nodeType.prototype.onDrawBackground = function(ctx) {
|
||||||
|
onDrawBackground?.apply(this, arguments)
|
||||||
|
|
||||||
|
const pad = [padCols * 20, LiteGraph.NODE_SLOT_HEIGHT * padRows + 8];
|
||||||
|
if(this.flags.collapsed || pad[1] + 32 > this.size[1] || image.width === 0) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
const avail = [this.size[0] - pad[0], this.size[1] - pad[1]]
|
||||||
|
const scale = Math.min(1.0, avail[0] / image.width, avail[1] / image.height)
|
||||||
|
const size = [Math.floor(image.width * scale), Math.floor(image.height * scale)]
|
||||||
|
const offset = [Math.max(0, (avail[0] - size[0]) / 2), Math.max(0, (avail[1] - size[1]) / 2)]
|
||||||
|
ctx.drawImage(image, offset[0], pad[1] + offset[1], size[0], size[1])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
// Parameter node
|
||||||
|
// - represents a customizable parameter that should be exposed in external tools
|
||||||
|
// - adapts to whichever node it is connected to, similar to the built-in "Primitive" node
|
||||||
|
// - can only be connected to slots which are converted widgets
|
||||||
|
|
||||||
|
const replaceableWidgets = ["INT", "FLOAT", "BOOLEAN", "STRING", "COMBO", "INT:seed"]
|
||||||
|
|
||||||
|
const parameterTypes = {
|
||||||
|
"combo": ["choice"],
|
||||||
|
"number": ["number", "number (integer)"],
|
||||||
|
"toggle": ["toggle"],
|
||||||
|
"text": ["text", "prompt (positive)", "prompt (negative)"],
|
||||||
|
}
|
||||||
|
|
||||||
|
function defaultParameterType(widgetType, connectedNode, connectedWidget) {
|
||||||
|
let paramType = parameterTypes[widgetType][0]
|
||||||
|
if (connectedNode.comfyClass === "CLIPTextEncode") {
|
||||||
|
paramType = "prompt (positive)"
|
||||||
|
}
|
||||||
|
const round = connectedWidget.options?.round
|
||||||
|
if ((paramType == "number" && round === undefined) || round === 1) {
|
||||||
|
paramType = "number (integer)"
|
||||||
|
}
|
||||||
|
return paramType
|
||||||
|
}
|
||||||
|
|
||||||
|
function valueMatchesType(value, type, options) {
|
||||||
|
if (type === "number") {
|
||||||
|
return typeof value === "number"
|
||||||
|
} else if (type === "combo") {
|
||||||
|
return options?.values?.includes(value)
|
||||||
|
} else if (type === "toggle") {
|
||||||
|
return typeof value === "boolean"
|
||||||
|
}
|
||||||
|
return typeof value === "string"
|
||||||
|
}
|
||||||
|
|
||||||
|
function optionalWidgetValue(widgets, index, fallback) {
|
||||||
|
const result = widgets.length > index ? widgets[index].value : null
|
||||||
|
return result === null || result === -1e10 || result === 1e10 ? fallback : result
|
||||||
|
}
|
||||||
|
|
||||||
|
function changeWidgets(node, type, connectedNode, connectedWidget) {
|
||||||
|
if (type === "customtext") {
|
||||||
|
type = "text"
|
||||||
|
}
|
||||||
|
const options = connectedWidget.options
|
||||||
|
|
||||||
|
const parameterTypeHint = node.widgets[1].value
|
||||||
|
const notSpecialized = node.widgets[1].options.values.includes("auto")
|
||||||
|
const parameterTypeMismatch = !parameterTypes[type].includes(parameterTypeHint)
|
||||||
|
if (notSpecialized || parameterTypeMismatch) {
|
||||||
|
node.widgets[1].options = {values: parameterTypes[type]}
|
||||||
|
}
|
||||||
|
if (parameterTypeMismatch) {
|
||||||
|
node.widgets[1].value = defaultParameterType(type, connectedNode, connectedWidget)
|
||||||
|
}
|
||||||
|
const oldDefault = node.widgets.length > 2 ? node.widgets[2].value : connectedWidget.value
|
||||||
|
const oldMin = optionalWidgetValue(node.widgets, 3, options?.min ?? 0)
|
||||||
|
const oldMax = optionalWidgetValue(node.widgets, 4, options?.max ?? 100)
|
||||||
|
const isDefaultValid = valueMatchesType(oldDefault, type, connectedWidget.options)
|
||||||
|
while (node.widgets.length > 2) {
|
||||||
|
node.widgets.pop()
|
||||||
|
}
|
||||||
|
const value = isDefaultValid && oldDefault !== "" ? oldDefault : connectedWidget.value
|
||||||
|
node.addWidget(type, "default", value, null, options)
|
||||||
|
if (type === "number") {
|
||||||
|
node.addWidget("number", "min", oldMin, null, options)
|
||||||
|
node.addWidget("number", "max", oldMax, null, options)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
function adaptWidgetsToConnection(node) {
|
||||||
|
if (!node.outputs || node.outputs.length === 0) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
const links = node.outputs[0].links
|
||||||
|
if (links && links.length === 1) {
|
||||||
|
const link = node.graph.links[links[0]]
|
||||||
|
if (!link) return
|
||||||
|
|
||||||
|
const theirNode = node.graph.getNodeById(link.target_id)
|
||||||
|
if (!theirNode || !theirNode.inputs) return
|
||||||
|
|
||||||
|
const input = theirNode.inputs[link.target_slot]
|
||||||
|
if (!input || !input.widget || theirNode.widgets === undefined) return
|
||||||
|
|
||||||
|
node.outputs[0].type = input.type
|
||||||
|
|
||||||
|
if (node.widgets[0].value === "Parameter") {
|
||||||
|
node.widgets[0].value = input.name
|
||||||
|
}
|
||||||
|
|
||||||
|
const widgetName = input.widget.name
|
||||||
|
const theirWidget = theirNode.widgets.find((w) => w.name === widgetName)
|
||||||
|
if (!theirWidget) return // connected to a custom node that isn't installed
|
||||||
|
|
||||||
|
const widgetType = theirWidget.origType ?? theirWidget.type
|
||||||
|
changeWidgets(node, widgetType, theirNode, theirWidget)
|
||||||
|
|
||||||
|
} else if (!links || links.length === 0) {
|
||||||
|
node.outputs[0].type = "*"
|
||||||
|
node.widgets[1].value = "auto"
|
||||||
|
node.widgets[1].options = {values: ["auto"]}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
function setupParameterNode(nodeType) {
|
||||||
|
const onAdded = nodeType.prototype.onAdded
|
||||||
|
nodeType.prototype.onAdded = function() {
|
||||||
|
onAdded?.apply(this, arguments)
|
||||||
|
adaptWidgetsToConnection(this)
|
||||||
|
}
|
||||||
|
|
||||||
|
const onAfterGraphConfigured = nodeType.prototype.onAfterGraphConfigured
|
||||||
|
nodeType.prototype.onAfterGraphConfigured = function() {
|
||||||
|
onAfterGraphConfigured?.apply(this, arguments)
|
||||||
|
adaptWidgetsToConnection(this)
|
||||||
|
}
|
||||||
|
|
||||||
|
const onConnectOutput = nodeType.prototype.onConnectOutput
|
||||||
|
nodeType.prototype.onConnectOutput = function(slot, type, input, target_node, target_slot) {
|
||||||
|
if (!input.widget && !(input.type in replaceableWidgets)) {
|
||||||
|
return false
|
||||||
|
} else if (onConnectOutput) {
|
||||||
|
result = onConnectOutput.apply(this, arguments)
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
const onConnectionsChange = nodeType.prototype.onConnectionsChange
|
||||||
|
nodeType.prototype.onConnectionsChange = function(_, index, connected) {
|
||||||
|
if (!app.configuringGraph) {
|
||||||
|
adaptWidgetsToConnection(this)
|
||||||
|
}
|
||||||
|
onConnectionsChange?.apply(this, arguments)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
// Register the extension
|
||||||
|
|
||||||
|
app.registerExtension({
|
||||||
|
name: "external_tooling_nodes",
|
||||||
|
|
||||||
|
beforeRegisterNodeDef(nodeType /*typeof LGraphNode*/, nodeData /*ComfyObjectInfo*/, app) {
|
||||||
|
if (nodeData.name === "ETN_KritaCanvas") {
|
||||||
|
setIconImage(nodeType, canvasIcon, [200, 100], 0, 2)
|
||||||
|
} else if (nodeData.name === "ETN_Parameter") {
|
||||||
|
setupParameterNode(nodeType)
|
||||||
|
} else if (nodeData.name === "ETN_SendText") {
|
||||||
|
const onAdded = nodeType.prototype.onAdded
|
||||||
|
nodeType.prototype.onAdded = function() {
|
||||||
|
onAdded?.apply(this, arguments)
|
||||||
|
this.inputs[0].type = "*"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
},
|
||||||
|
|
||||||
|
nodeCreated(node /*ComfyNode*/, app) {
|
||||||
|
if (publisherRegistered || node.comfyClass !== "ETN_KritaOutput") {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
api.addEventListener('graphChanged', publishWorkflow)
|
||||||
|
publisherRegistered = true
|
||||||
|
},
|
||||||
|
|
||||||
|
setup(app) {
|
||||||
|
if (publisherRegistered) {
|
||||||
|
publishWorkflow(null)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
});
|
||||||
|
|
||||||
|
})();
|
||||||
@@ -0,0 +1,364 @@
|
|||||||
|
import sys
|
||||||
|
from enum import Enum
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, NamedTuple
|
||||||
|
|
||||||
|
import comfy.samplers
|
||||||
|
import numpy as np
|
||||||
|
import server
|
||||||
|
import torch
|
||||||
|
from comfy.comfy_types.node_typing import IO
|
||||||
|
from comfy_api.latest import io
|
||||||
|
from PIL import Image
|
||||||
|
|
||||||
|
from .nodes import SendImageWebSocket
|
||||||
|
|
||||||
|
|
||||||
|
class Publisher(NamedTuple):
|
||||||
|
name: str
|
||||||
|
id: str
|
||||||
|
workflow: dict
|
||||||
|
|
||||||
|
|
||||||
|
class WorkflowExchange:
|
||||||
|
def __init__(self, server: server.PromptServer):
|
||||||
|
self._server = server
|
||||||
|
self._publishers: dict[str, Publisher] = {}
|
||||||
|
self._subscribers: list[str] = []
|
||||||
|
|
||||||
|
async def publish(self, publisher_name: str, publisher_id: str, workflow: dict):
|
||||||
|
publisher = Publisher(publisher_name, publisher_id, workflow)
|
||||||
|
for client_id in self._subscribers:
|
||||||
|
await self._notify(client_id, publisher)
|
||||||
|
self._publishers[publisher_id] = publisher
|
||||||
|
|
||||||
|
async def subscribe(self, client_id: str):
|
||||||
|
if client_id in self._subscribers:
|
||||||
|
raise KeyError("Already subscribed")
|
||||||
|
self._subscribers.append(client_id)
|
||||||
|
for publisher in self._publishers.values():
|
||||||
|
await self._notify(client_id, publisher)
|
||||||
|
|
||||||
|
async def unsubscribe(self, client_id: str):
|
||||||
|
if client_id in self._subscribers:
|
||||||
|
self._subscribers.remove(client_id)
|
||||||
|
else:
|
||||||
|
raise KeyError("No subscriber found with id " + client_id)
|
||||||
|
|
||||||
|
async def _notify(self, client_id: str, publisher: Publisher):
|
||||||
|
data = {
|
||||||
|
"publisher": {"name": publisher.name, "id": publisher.id},
|
||||||
|
"workflow": publisher.workflow,
|
||||||
|
}
|
||||||
|
await self._server.send_json("etn_workflow_published", data, client_id)
|
||||||
|
|
||||||
|
|
||||||
|
def _placeholder_image():
|
||||||
|
path = Path(__file__).parent / "data" / "external-image-placeholder.webp"
|
||||||
|
image = Image.open(path).convert("RGB")
|
||||||
|
image = np.array(image).astype(np.float32) / 255.0
|
||||||
|
return torch.from_numpy(image)[None,]
|
||||||
|
|
||||||
|
|
||||||
|
class _BasicTypes(str):
|
||||||
|
"""Matches IO.PRIMITIVE, but also any list of choices"""
|
||||||
|
|
||||||
|
basic_types = IO.PRIMITIVE.split(",") # STRING, FLOAT, INT, BOOLEAN
|
||||||
|
|
||||||
|
def __eq__(self, other):
|
||||||
|
return other in self.basic_types or isinstance(other, (list, _BasicTypes))
|
||||||
|
|
||||||
|
def __ne__(self, other):
|
||||||
|
return not self.__eq__(other)
|
||||||
|
|
||||||
|
|
||||||
|
BasicTypes = _BasicTypes("BASIC")
|
||||||
|
|
||||||
|
|
||||||
|
class OutputBatchMode(Enum):
|
||||||
|
default = "default"
|
||||||
|
images = "images"
|
||||||
|
animation = "animation"
|
||||||
|
layers = "layers"
|
||||||
|
|
||||||
|
|
||||||
|
class KritaOutput(io.ComfyNode):
|
||||||
|
@classmethod
|
||||||
|
def define_schema(cls):
|
||||||
|
return io.Schema(
|
||||||
|
node_id="ETN_KritaOutput",
|
||||||
|
display_name="Krita Output",
|
||||||
|
category="krita",
|
||||||
|
inputs=[
|
||||||
|
io.Image.Input("images"),
|
||||||
|
io.Int.Input("x", "offset x", default=0),
|
||||||
|
io.Int.Input("y", "offset y", default=0),
|
||||||
|
io.String.Input("name", default=""),
|
||||||
|
io.Combo.Input(
|
||||||
|
"batch_mode", OutputBatchMode, "batch mode", default=OutputBatchMode.default
|
||||||
|
),
|
||||||
|
io.Boolean.Input("resize_canvas", "resize canvas", default=False),
|
||||||
|
],
|
||||||
|
is_output_node=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def execute( # type: ignore
|
||||||
|
cls,
|
||||||
|
images: torch.Tensor,
|
||||||
|
x: int = 0,
|
||||||
|
y: int = 0,
|
||||||
|
name="",
|
||||||
|
batch_mode: OutputBatchMode | str = OutputBatchMode.default,
|
||||||
|
resize_canvas=False,
|
||||||
|
):
|
||||||
|
batch_mode = batch_mode.value if isinstance(batch_mode, OutputBatchMode) else batch_mode
|
||||||
|
info = {
|
||||||
|
"name": name,
|
||||||
|
"offset_x": x,
|
||||||
|
"offset_y": y,
|
||||||
|
"batch_mode": batch_mode,
|
||||||
|
"resize_canvas": resize_canvas,
|
||||||
|
}
|
||||||
|
output = SendImageWebSocket.execute(images, "PNG")
|
||||||
|
assert isinstance(output.ui, dict)
|
||||||
|
output.ui["info"] = [info]
|
||||||
|
return output
|
||||||
|
|
||||||
|
|
||||||
|
class KritaSendText(io.ComfyNode):
|
||||||
|
@classmethod
|
||||||
|
def define_schema(cls):
|
||||||
|
return io.Schema(
|
||||||
|
node_id="ETN_KritaSendText",
|
||||||
|
display_name="Send Text",
|
||||||
|
category="krita",
|
||||||
|
inputs=[
|
||||||
|
io.AnyType.Input("value"),
|
||||||
|
io.String.Input("name", default="Output"),
|
||||||
|
io.Combo.Input("type", options=["text", "markdown", "html"], default="text"),
|
||||||
|
],
|
||||||
|
is_output_node=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def execute(cls, value: Any, name: str, type: str): # type: ignore
|
||||||
|
mime = {
|
||||||
|
"text": "text/plain",
|
||||||
|
"markdown": "text/markdown",
|
||||||
|
"html": "text/html",
|
||||||
|
}[type]
|
||||||
|
text = "None"
|
||||||
|
if value is not None:
|
||||||
|
try:
|
||||||
|
text = str(value)
|
||||||
|
except Exception as e:
|
||||||
|
text = f"Could not convert to text: {e}"
|
||||||
|
|
||||||
|
return io.NodeOutput(ui={"text": [{"name": name, "text": text, "content-type": mime}]})
|
||||||
|
|
||||||
|
|
||||||
|
class KritaCanvas(io.ComfyNode):
|
||||||
|
@classmethod
|
||||||
|
def define_schema(cls):
|
||||||
|
return io.Schema(
|
||||||
|
node_id="ETN_KritaCanvas",
|
||||||
|
display_name="Krita Canvas",
|
||||||
|
category="krita",
|
||||||
|
outputs=[
|
||||||
|
io.Image.Output(display_name="image"),
|
||||||
|
io.Int.Output(display_name="width"),
|
||||||
|
io.Int.Output(display_name="height"),
|
||||||
|
io.Int.Output(display_name="seed"),
|
||||||
|
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 = [
|
||||||
|
"auto",
|
||||||
|
"number",
|
||||||
|
"number (integer)",
|
||||||
|
"toggle",
|
||||||
|
"choice",
|
||||||
|
"text",
|
||||||
|
"prompt (positive)",
|
||||||
|
"prompt (negative)",
|
||||||
|
]
|
||||||
|
_fmax = sys.float_info.max
|
||||||
|
|
||||||
|
|
||||||
|
class Parameter(io.ComfyNode):
|
||||||
|
@classmethod
|
||||||
|
def define_schema(cls):
|
||||||
|
return io.Schema(
|
||||||
|
node_id="ETN_Parameter",
|
||||||
|
display_name="Parameter",
|
||||||
|
category="krita",
|
||||||
|
inputs=[
|
||||||
|
io.String.Input("name", default="Parameter"),
|
||||||
|
io.Combo.Input("type", options=_param_types, default="auto"),
|
||||||
|
io.String.Input("default", default=""),
|
||||||
|
io.Float.Input("min", default=-1e10, min=-_fmax, max=_fmax, optional=True),
|
||||||
|
io.Float.Input("max", default=1e10, min=-_fmax, max=_fmax, optional=True),
|
||||||
|
],
|
||||||
|
outputs=[io.AnyType.Output(display_name="value")],
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def execute(cls, name: str, type: str, default, min=0.0, max=1.0): # type: ignore
|
||||||
|
if type == "number":
|
||||||
|
return io.NodeOutput(float(default))
|
||||||
|
elif type == "number (integer)":
|
||||||
|
return io.NodeOutput(int(default))
|
||||||
|
return io.NodeOutput(default)
|
||||||
|
|
||||||
|
|
||||||
|
class KritaStyle(io.ComfyNode):
|
||||||
|
@classmethod
|
||||||
|
def define_schema(cls):
|
||||||
|
return io.Schema(
|
||||||
|
node_id="ETN_KritaStyle",
|
||||||
|
display_name="Krita Style",
|
||||||
|
category="krita",
|
||||||
|
inputs=[
|
||||||
|
io.String.Input("name", default="Style"),
|
||||||
|
io.Combo.Input("sampler_preset", options=["auto", "regular", "live"]),
|
||||||
|
],
|
||||||
|
outputs=[
|
||||||
|
io.Model.Output(display_name="model"),
|
||||||
|
io.Clip.Output(display_name="clip"),
|
||||||
|
io.Vae.Output(display_name="vae"),
|
||||||
|
io.String.Output(display_name="positive prompt"),
|
||||||
|
io.String.Output(display_name="negative prompt"),
|
||||||
|
io.Combo.Output(
|
||||||
|
display_name="sampler name", options=comfy.samplers.KSampler.SAMPLERS
|
||||||
|
),
|
||||||
|
io.Combo.Output(
|
||||||
|
display_name="scheduler", options=comfy.samplers.KSampler.SCHEDULERS
|
||||||
|
),
|
||||||
|
io.Int.Output(display_name="steps"),
|
||||||
|
io.Float.Output(display_name="guidance"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def execute(cls, name: str, sampler_preset: str): # 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!")
|
||||||
@@ -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
|
||||||
@@ -52,17 +71,22 @@ class LoadMaskBase64:
|
|||||||
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 {"required": {"images": ("IMAGE",)}}
|
return io.Schema(
|
||||||
|
node_id="ETN_SendImageWebSocket",
|
||||||
|
display_name="Send Image (WebSocket)",
|
||||||
|
category="external_tooling",
|
||||||
|
inputs=[
|
||||||
|
io.Image.Input("images"),
|
||||||
|
io.Combo.Input("format", options=["PNG", "JPEG"], default="PNG"),
|
||||||
|
],
|
||||||
|
is_output_node=True,
|
||||||
|
)
|
||||||
|
|
||||||
RETURN_TYPES = ()
|
@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):
|
|
||||||
results = []
|
results = []
|
||||||
for tensor in images:
|
for tensor in images:
|
||||||
array = 255.0 * tensor.cpu().numpy()
|
array = 255.0 * tensor.cpu().numpy()
|
||||||
@@ -71,91 +95,337 @@ class SendImageWebSocket:
|
|||||||
server = PromptServer.instance
|
server = PromptServer.instance
|
||||||
server.send_sync(
|
server.send_sync(
|
||||||
BinaryEventTypes.UNENCODED_PREVIEW_IMAGE,
|
BinaryEventTypes.UNENCODED_PREVIEW_IMAGE,
|
||||||
["PNG", image, None],
|
[format, image, None],
|
||||||
server.client_id,
|
server.client_id,
|
||||||
)
|
)
|
||||||
results.append(
|
results.append({
|
||||||
# Could put some kind of ID here, but for now just match them by index
|
"source": "websocket",
|
||||||
{"source": "websocket", "content-type": "image/png", "type": "output"}
|
"content-type": f"image/{format.lower()}",
|
||||||
)
|
"type": "output",
|
||||||
|
})
|
||||||
|
|
||||||
return {"ui": {"images": results}}
|
return io.NodeOutput(ui={"images": results})
|
||||||
|
|
||||||
|
|
||||||
class CropImage:
|
class ImageCache:
|
||||||
"""Deprecated, ComfyUI has an ImageCrop node now which does the same."""
|
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 ApplyMaskToImage:
|
class SaveImageCache(io.ComfyNode):
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(cls):
|
def define_schema(cls):
|
||||||
return {
|
return io.Schema(
|
||||||
"required": {
|
node_id="ETN_SaveImageCache",
|
||||||
"image": ("IMAGE",),
|
display_name="Save Image to Cache",
|
||||||
"mask": ("MASK",),
|
category="external_tooling",
|
||||||
}
|
inputs=[
|
||||||
}
|
io.Image.Input("images"),
|
||||||
|
io.Combo.Input("format", options=["PNG", "JPEG"], default="PNG"),
|
||||||
|
],
|
||||||
|
is_output_node=True,
|
||||||
|
)
|
||||||
|
|
||||||
CATEGORY = "external_tooling"
|
@classmethod
|
||||||
RETURN_TYPES = ("IMAGE",)
|
def execute(cls, images: torch.Tensor, format: str): # type: ignore
|
||||||
FUNCTION = "apply_mask"
|
results = []
|
||||||
|
for tensor in images:
|
||||||
|
array = 255.0 * tensor.cpu().numpy()
|
||||||
|
image = Image.fromarray(np.clip(array, 0, 255).astype(np.uint8))
|
||||||
|
key = image_cache.add(image, format)
|
||||||
|
|
||||||
def apply_mask(self, image: torch.Tensor, mask: torch.Tensor):
|
results.append({
|
||||||
# Move the channel to the second dimension for processing
|
"source": "http",
|
||||||
out = image.movedim(-1, 1)
|
"id": key,
|
||||||
|
"content-type": f"image/{format.lower()}",
|
||||||
|
"type": "output",
|
||||||
|
})
|
||||||
|
return io.NodeOutput(ui={"images": results})
|
||||||
|
|
||||||
# Check if the images are RGB, and if so, add an alpha channel initialized to 1
|
|
||||||
|
def to_bchw(image: torch.Tensor):
|
||||||
|
if image.ndim == 3:
|
||||||
|
image = image.unsqueeze(0)
|
||||||
|
return image.movedim(-1, 1)
|
||||||
|
|
||||||
|
|
||||||
|
def to_bhwc(image: torch.Tensor):
|
||||||
|
return image.movedim(1, -1)
|
||||||
|
|
||||||
|
|
||||||
|
def mask_batch(mask: torch.Tensor):
|
||||||
|
if mask.ndim == 2:
|
||||||
|
mask = mask.unsqueeze(0)
|
||||||
|
return mask
|
||||||
|
|
||||||
|
|
||||||
|
class ApplyMaskToImage(io.ComfyNode):
|
||||||
|
@classmethod
|
||||||
|
def define_schema(cls):
|
||||||
|
return io.Schema(
|
||||||
|
node_id="ETN_ApplyMaskToImage",
|
||||||
|
display_name="Apply Mask to Image",
|
||||||
|
category="external_tooling",
|
||||||
|
inputs=[
|
||||||
|
io.Image.Input("image"),
|
||||||
|
io.Mask.Input("mask"),
|
||||||
|
],
|
||||||
|
outputs=[io.Image.Output(display_name="masked")],
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def execute(cls, image: torch.Tensor, mask: torch.Tensor): # type: ignore
|
||||||
|
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)
|
||||||
|
mask = mask_batch(mask)
|
||||||
# Ensure masks are unsqueezed to match the alpha channel dimension if needed
|
|
||||||
if mask.ndim == 2:
|
|
||||||
mask = mask.unsqueeze(0) # Add a batch dimension to masks
|
|
||||||
# For single mask, expand it to match size of image batch size.
|
|
||||||
if mask.shape[0] == 1:
|
|
||||||
mask = mask.repeat(out.shape[0], 1, 1)
|
|
||||||
|
|
||||||
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 shsape [B, C, H, W]. {out.shape}"
|
assert out.ndim == 4, f"Image should have shape [B, C, H, W]. {out.shape}"
|
||||||
assert out.shape[-2:] == mask.shape[-2:], f"{out.shape[-2:]} != {mask.shape[-2:]}"
|
assert out.shape[-2:] == mask.shape[-2:], (
|
||||||
assert out.shape[0] == mask.shape[0], f"{out.shape[0]} != {mask.shape[0]}"
|
f"Image size {out.shape[-2:]} must match mask size {mask.shape[-2:]}"
|
||||||
|
)
|
||||||
|
is_mask_batch = mask.shape[0] == out.shape[0]
|
||||||
|
|
||||||
# Apply each mask in the batch to its corresponding image's alpha channel
|
# 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]):
|
||||||
out[i, 3, :, :] = mask[i]
|
alpha = mask[i] if is_mask_batch else mask[0]
|
||||||
|
out[i, 3, :, :] *= alpha
|
||||||
|
|
||||||
# Move the channel back to its original dimension
|
return (to_bhwc(out),)
|
||||||
out = out.movedim(1, -1)
|
|
||||||
|
|
||||||
return (out,)
|
|
||||||
|
class _ReferenceImageData(NamedTuple):
|
||||||
|
image: torch.Tensor
|
||||||
|
weight: float
|
||||||
|
range: tuple[float, float]
|
||||||
|
|
||||||
|
|
||||||
|
class ReferenceImage(io.ComfyNode):
|
||||||
|
@classmethod
|
||||||
|
def define_schema(cls):
|
||||||
|
return io.Schema(
|
||||||
|
node_id="ETN_ReferenceImage",
|
||||||
|
display_name="Reference Image",
|
||||||
|
category="external_tooling",
|
||||||
|
inputs=[
|
||||||
|
io.Image.Input("image"),
|
||||||
|
io.Float.Input("weight", default=1.0, min=0.0, max=10.0),
|
||||||
|
io.Float.Input("range_start", default=0.0, min=0.0, max=1.0),
|
||||||
|
io.Float.Input("range_end", default=1.0, min=0.0, max=1.0),
|
||||||
|
io.Custom("ReferenceImage").Input("reference_images", optional=True),
|
||||||
|
],
|
||||||
|
outputs=[io.Custom("ReferenceImage").Output(display_name="reference_images")],
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def execute( # 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
|
||||||
|
|||||||
@@ -0,0 +1,145 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
from pathlib import Path
|
||||||
|
from tqdm import tqdm
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch import Tensor
|
||||||
|
from transformers import CLIPImageProcessor, CLIPConfig, CLIPVisionModel, PreTrainedModel
|
||||||
|
from kornia.filters import box_blur
|
||||||
|
from comfy_api.latest import io
|
||||||
|
|
||||||
|
from .nodes import to_bchw, to_bhwc
|
||||||
|
|
||||||
|
|
||||||
|
def cosine_similarity(image_embeds: Tensor, text_embeds: Tensor):
|
||||||
|
if image_embeds.dim() == 2 and text_embeds.dim() == 2:
|
||||||
|
image_embeds = image_embeds.unsqueeze(1)
|
||||||
|
return F.cosine_similarity(image_embeds, text_embeds, dim=-1)
|
||||||
|
|
||||||
|
|
||||||
|
class CLIPSafetyChecker(PreTrainedModel):
|
||||||
|
# https://huggingface.co/CompVis/stable-diffusion-safety-checker
|
||||||
|
# Adapted from:
|
||||||
|
# https://github.com/huggingface/diffusers/blob/main/src/diffusers/pipelines/stable_diffusion/safety_checker.py
|
||||||
|
|
||||||
|
config_class = CLIPConfig
|
||||||
|
_no_split_modules = ["CLIPEncoderLayer"]
|
||||||
|
|
||||||
|
def __init__(self, config: CLIPConfig):
|
||||||
|
super().__init__(config)
|
||||||
|
projdim = config.projection_dim
|
||||||
|
|
||||||
|
self.vision_model = CLIPVisionModel(config.vision_config)
|
||||||
|
self.visual_projection = nn.Linear(config.vision_config.hidden_size, projdim, bias=False)
|
||||||
|
|
||||||
|
self.concept_embeds = nn.Parameter(torch.ones(17, projdim), requires_grad=False)
|
||||||
|
self.special_care_embeds = nn.Parameter(torch.ones(3, projdim), requires_grad=False)
|
||||||
|
self.concept_embeds_weights = nn.Parameter(torch.ones(17), requires_grad=False)
|
||||||
|
self.special_care_embeds_weights = nn.Parameter(torch.ones(3), requires_grad=False)
|
||||||
|
|
||||||
|
# Model requires post_init after transformers v4.57.3
|
||||||
|
if hasattr(self, "post_init"):
|
||||||
|
self.post_init()
|
||||||
|
|
||||||
|
def forward(self, clip_input, images: Tensor, sensitivity: float):
|
||||||
|
with torch.no_grad():
|
||||||
|
image_batch = self.vision_model(clip_input)[1]
|
||||||
|
image_embeds = self.visual_projection(image_batch)
|
||||||
|
sensitivity = -0.1 + 0.14 * sensitivity
|
||||||
|
|
||||||
|
special_cos_dist = cosine_similarity(image_embeds, self.special_care_embeds)
|
||||||
|
special_scores_threshold = self.special_care_embeds_weights.unsqueeze(0)
|
||||||
|
special_scores = special_cos_dist - special_scores_threshold + sensitivity
|
||||||
|
|
||||||
|
if torch.any(special_scores > 0):
|
||||||
|
sensitivity = sensitivity + 0.01
|
||||||
|
|
||||||
|
cos_dist = cosine_similarity(image_embeds, self.concept_embeds)
|
||||||
|
concept_threshold = self.concept_embeds_weights.unsqueeze(0)
|
||||||
|
concept_scores = cos_dist - concept_threshold + sensitivity
|
||||||
|
|
||||||
|
is_nsfw = [torch.any(concept_scores[i] > 0) for i in range(concept_scores.shape[0])]
|
||||||
|
is_nsfw = [x.item() for x in is_nsfw]
|
||||||
|
return self.filter_images(images, is_nsfw)
|
||||||
|
|
||||||
|
def filter_images(self, images: Tensor, is_nsfw: list[bool]):
|
||||||
|
if not any(is_nsfw):
|
||||||
|
return images
|
||||||
|
|
||||||
|
images = images.clone()
|
||||||
|
images_to_filter = (i for i, nsfw in enumerate(is_nsfw) if nsfw)
|
||||||
|
orig_size = images.shape[-2:]
|
||||||
|
for idx in images_to_filter:
|
||||||
|
filtered = images[idx].unsqueeze(0)
|
||||||
|
filtered = F.interpolate(filtered, size=64, mode="nearest")
|
||||||
|
filtered = box_blur(filtered, 11, separable=True)
|
||||||
|
filtered = F.interpolate(filtered, size=orig_size, mode="bilinear")
|
||||||
|
images[idx] = filtered.squeeze(0)
|
||||||
|
return images
|
||||||
|
|
||||||
|
|
||||||
|
class CachedModels:
|
||||||
|
_instance: CachedModels | None = None
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
model_dir = Path(__file__).parent / "safetychecker"
|
||||||
|
model_file = model_dir / "model.safetensors"
|
||||||
|
if not model_file.exists():
|
||||||
|
self.download(
|
||||||
|
"https://huggingface.co/CompVis/stable-diffusion-safety-checker/resolve/refs%2Fpr%2F41/model.safetensors",
|
||||||
|
target=model_file,
|
||||||
|
)
|
||||||
|
self.feature_extractor = CLIPImageProcessor.from_pretrained(model_dir)
|
||||||
|
self.safety_checker = CLIPSafetyChecker.from_pretrained(model_dir)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def load(cls):
|
||||||
|
if cls._instance is None:
|
||||||
|
cls._instance = CachedModels()
|
||||||
|
return cls._instance
|
||||||
|
|
||||||
|
def download(self, url: str, target: Path):
|
||||||
|
import requests
|
||||||
|
|
||||||
|
try:
|
||||||
|
target_temp = target.with_suffix(".download")
|
||||||
|
with requests.get(url, stream=True) as response:
|
||||||
|
text = "NSFWFilter model download"
|
||||||
|
total = int(response.headers.get("content-length", 0))
|
||||||
|
pbar = tqdm(None, total=total, unit="b", unit_scale=True, desc=text)
|
||||||
|
with open(target_temp, "wb") as f:
|
||||||
|
for chunk in response.iter_content(chunk_size=8192):
|
||||||
|
f.write(chunk)
|
||||||
|
pbar.update(len(chunk))
|
||||||
|
pbar.close()
|
||||||
|
target_temp.rename(target)
|
||||||
|
except Exception as e:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"NSFWFilter: Failed to download safety-checker model from {url} to target location {target}: {e}"
|
||||||
|
) from e
|
||||||
|
|
||||||
|
|
||||||
|
class NSFWFilter(io.ComfyNode):
|
||||||
|
@classmethod
|
||||||
|
def define_schema(cls):
|
||||||
|
return io.Schema(
|
||||||
|
node_id="ETN_NSFWFilter",
|
||||||
|
display_name="NSFW Filter",
|
||||||
|
category="external_tooling",
|
||||||
|
inputs=[
|
||||||
|
io.Image.Input("image"),
|
||||||
|
io.Float.Input("sensitivity", default=0.5, min=0.0, max=1.0, step=0.1),
|
||||||
|
],
|
||||||
|
outputs=[io.Image.Output(display_name="image")],
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def execute(cls, image: Tensor, sensitivity: float):
|
||||||
|
models = CachedModels.load()
|
||||||
|
image = to_bchw(image)
|
||||||
|
input = models.feature_extractor(image, do_rescale=False, return_tensors="pt")
|
||||||
|
filtered = models.safety_checker(
|
||||||
|
images=image, clip_input=input.pixel_values, sensitivity=sensitivity
|
||||||
|
)
|
||||||
|
return io.NodeOutput(to_bhwc(filtered))
|
||||||
+10
-2
@@ -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.0.0"
|
version = "3.4.0"
|
||||||
license = "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
|
||||||
|
|||||||
@@ -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,103 +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"
|
)
|
||||||
RETURN_TYPES = ("REGIONS",)
|
|
||||||
FUNCTION = "define"
|
|
||||||
|
|
||||||
def define(self, mask: Tensor, conditioning: list, regions: Region | None = None):
|
|
||||||
return (Region(regions, mask, conditioning),)
|
|
||||||
|
|
||||||
|
|
||||||
class ListRegionMasks:
|
|
||||||
@classmethod
|
|
||||||
def INPUT_TYPES(cls):
|
|
||||||
return {"required": {"regions": ("REGIONS",)}}
|
|
||||||
|
|
||||||
CATEGORY = "external_tooling/regions"
|
|
||||||
RETURN_TYPES = ("MASK",)
|
|
||||||
FUNCTION = "get_masks"
|
|
||||||
|
|
||||||
def get_masks(self, regions: Region):
|
|
||||||
return (torch.stack([r.mask for r in regions.preprocess()], dim=0),)
|
|
||||||
|
|
||||||
|
|
||||||
class AttentionMask:
|
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(s):
|
def execute(cls, mask: Tensor, conditioning: list, regions: Region | None = None):
|
||||||
return {
|
if mask.dim() < 3:
|
||||||
"required": {
|
mask = mask.unsqueeze(0)
|
||||||
"model": ("MODEL",),
|
return io.NodeOutput(Region(regions, mask, conditioning))
|
||||||
"regions": ("REGIONS",),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
RETURN_TYPES = ("MODEL",)
|
|
||||||
FUNCTION = "attention_mask"
|
|
||||||
CATEGORY = "external_tooling/regions"
|
|
||||||
|
|
||||||
mask: Tensor
|
class ListRegionMasks(io.ComfyNode):
|
||||||
conds: list[Tensor]
|
@classmethod
|
||||||
batch_size: int
|
def define_schema(cls):
|
||||||
|
return io.Schema(
|
||||||
|
node_id="ETN_ListRegionMasks",
|
||||||
|
display_name="List Region Masks",
|
||||||
|
category="external_tooling/regions",
|
||||||
|
inputs=[Regions.Input("regions")],
|
||||||
|
outputs=[io.Mask.Output(display_name="masks")],
|
||||||
|
)
|
||||||
|
|
||||||
def attention_mask(self, model: ModelPatcher, regions: Region):
|
@classmethod
|
||||||
new_model = model.clone()
|
def execute(cls, regions: Region):
|
||||||
region_list = regions.preprocess()
|
return io.NodeOutput(torch.stack([r.mask for r in regions.preprocess()], dim=0))
|
||||||
num_conds = len(region_list)
|
|
||||||
|
|
||||||
|
|
||||||
|
class AttentionMask(io.ComfyNode):
|
||||||
|
@classmethod
|
||||||
|
def define_schema(cls):
|
||||||
|
return io.Schema(
|
||||||
|
node_id="ETN_AttentionMask",
|
||||||
|
display_name="Regions Attention Mask",
|
||||||
|
category="external_tooling/regions",
|
||||||
|
inputs=[io.Model.Input("model"), Regions.Input("regions")],
|
||||||
|
outputs=[io.Model.Output(display_name="model")],
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def execute(cls, model: ModelPatcher, regions: Region):
|
||||||
|
return io.NodeOutput(AttentionMaskPatch.apply(model, regions))
|
||||||
|
|
||||||
|
|
||||||
|
class AttentionMaskPatch:
|
||||||
|
def __init__(self, region_list: list[Region]):
|
||||||
mask = torch.stack([r.mask for r in region_list], dim=0)
|
mask = 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:
|
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:
|
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)
|
||||||
|
|
||||||
@@ -172,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)
|
||||||
@@ -182,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
|
||||||
|
|||||||
@@ -0,0 +1,2 @@
|
|||||||
|
# Optional, only required for Translate node:
|
||||||
|
argostranslate
|
||||||
@@ -0,0 +1,171 @@
|
|||||||
|
{
|
||||||
|
"_name_or_path": "clip-vit-large-patch14/",
|
||||||
|
"architectures": [
|
||||||
|
"SafetyChecker"
|
||||||
|
],
|
||||||
|
"initializer_factor": 1.0,
|
||||||
|
"logit_scale_init_value": 2.6592,
|
||||||
|
"model_type": "clip",
|
||||||
|
"projection_dim": 768,
|
||||||
|
"text_config": {
|
||||||
|
"_name_or_path": "",
|
||||||
|
"add_cross_attention": false,
|
||||||
|
"architectures": null,
|
||||||
|
"attention_dropout": 0.0,
|
||||||
|
"bad_words_ids": null,
|
||||||
|
"bos_token_id": 0,
|
||||||
|
"chunk_size_feed_forward": 0,
|
||||||
|
"cross_attention_hidden_size": null,
|
||||||
|
"decoder_start_token_id": null,
|
||||||
|
"diversity_penalty": 0.0,
|
||||||
|
"do_sample": false,
|
||||||
|
"dropout": 0.0,
|
||||||
|
"early_stopping": false,
|
||||||
|
"encoder_no_repeat_ngram_size": 0,
|
||||||
|
"eos_token_id": 2,
|
||||||
|
"exponential_decay_length_penalty": null,
|
||||||
|
"finetuning_task": null,
|
||||||
|
"forced_bos_token_id": null,
|
||||||
|
"forced_eos_token_id": null,
|
||||||
|
"hidden_act": "quick_gelu",
|
||||||
|
"hidden_size": 768,
|
||||||
|
"id2label": {
|
||||||
|
"0": "LABEL_0",
|
||||||
|
"1": "LABEL_1"
|
||||||
|
},
|
||||||
|
"initializer_factor": 1.0,
|
||||||
|
"initializer_range": 0.02,
|
||||||
|
"intermediate_size": 3072,
|
||||||
|
"is_decoder": false,
|
||||||
|
"is_encoder_decoder": false,
|
||||||
|
"label2id": {
|
||||||
|
"LABEL_0": 0,
|
||||||
|
"LABEL_1": 1
|
||||||
|
},
|
||||||
|
"layer_norm_eps": 1e-05,
|
||||||
|
"length_penalty": 1.0,
|
||||||
|
"max_length": 20,
|
||||||
|
"max_position_embeddings": 77,
|
||||||
|
"min_length": 0,
|
||||||
|
"model_type": "clip_text_model",
|
||||||
|
"no_repeat_ngram_size": 0,
|
||||||
|
"num_attention_heads": 12,
|
||||||
|
"num_beam_groups": 1,
|
||||||
|
"num_beams": 1,
|
||||||
|
"num_hidden_layers": 12,
|
||||||
|
"num_return_sequences": 1,
|
||||||
|
"output_attentions": false,
|
||||||
|
"output_hidden_states": false,
|
||||||
|
"output_scores": false,
|
||||||
|
"pad_token_id": 1,
|
||||||
|
"prefix": null,
|
||||||
|
"problem_type": null,
|
||||||
|
"pruned_heads": {},
|
||||||
|
"remove_invalid_values": false,
|
||||||
|
"repetition_penalty": 1.0,
|
||||||
|
"return_dict": true,
|
||||||
|
"return_dict_in_generate": false,
|
||||||
|
"sep_token_id": null,
|
||||||
|
"task_specific_params": null,
|
||||||
|
"temperature": 1.0,
|
||||||
|
"tie_encoder_decoder": false,
|
||||||
|
"tie_word_embeddings": true,
|
||||||
|
"tokenizer_class": null,
|
||||||
|
"top_k": 50,
|
||||||
|
"top_p": 1.0,
|
||||||
|
"torch_dtype": null,
|
||||||
|
"torchscript": false,
|
||||||
|
"transformers_version": "4.21.0.dev0",
|
||||||
|
"typical_p": 1.0,
|
||||||
|
"use_bfloat16": false,
|
||||||
|
"vocab_size": 49408
|
||||||
|
},
|
||||||
|
"text_config_dict": {
|
||||||
|
"hidden_size": 768,
|
||||||
|
"intermediate_size": 3072,
|
||||||
|
"num_attention_heads": 12,
|
||||||
|
"num_hidden_layers": 12
|
||||||
|
},
|
||||||
|
"torch_dtype": "float32",
|
||||||
|
"transformers_version": null,
|
||||||
|
"vision_config": {
|
||||||
|
"_name_or_path": "",
|
||||||
|
"add_cross_attention": false,
|
||||||
|
"architectures": null,
|
||||||
|
"attention_dropout": 0.0,
|
||||||
|
"bad_words_ids": null,
|
||||||
|
"bos_token_id": null,
|
||||||
|
"chunk_size_feed_forward": 0,
|
||||||
|
"cross_attention_hidden_size": null,
|
||||||
|
"decoder_start_token_id": null,
|
||||||
|
"diversity_penalty": 0.0,
|
||||||
|
"do_sample": false,
|
||||||
|
"dropout": 0.0,
|
||||||
|
"early_stopping": false,
|
||||||
|
"encoder_no_repeat_ngram_size": 0,
|
||||||
|
"eos_token_id": null,
|
||||||
|
"exponential_decay_length_penalty": null,
|
||||||
|
"finetuning_task": null,
|
||||||
|
"forced_bos_token_id": null,
|
||||||
|
"forced_eos_token_id": null,
|
||||||
|
"hidden_act": "quick_gelu",
|
||||||
|
"hidden_size": 1024,
|
||||||
|
"id2label": {
|
||||||
|
"0": "LABEL_0",
|
||||||
|
"1": "LABEL_1"
|
||||||
|
},
|
||||||
|
"image_size": 224,
|
||||||
|
"initializer_factor": 1.0,
|
||||||
|
"initializer_range": 0.02,
|
||||||
|
"intermediate_size": 4096,
|
||||||
|
"is_decoder": false,
|
||||||
|
"is_encoder_decoder": false,
|
||||||
|
"label2id": {
|
||||||
|
"LABEL_0": 0,
|
||||||
|
"LABEL_1": 1
|
||||||
|
},
|
||||||
|
"layer_norm_eps": 1e-05,
|
||||||
|
"length_penalty": 1.0,
|
||||||
|
"max_length": 20,
|
||||||
|
"min_length": 0,
|
||||||
|
"model_type": "clip_vision_model",
|
||||||
|
"no_repeat_ngram_size": 0,
|
||||||
|
"num_attention_heads": 16,
|
||||||
|
"num_beam_groups": 1,
|
||||||
|
"num_beams": 1,
|
||||||
|
"num_hidden_layers": 24,
|
||||||
|
"num_return_sequences": 1,
|
||||||
|
"output_attentions": false,
|
||||||
|
"output_hidden_states": false,
|
||||||
|
"output_scores": false,
|
||||||
|
"pad_token_id": null,
|
||||||
|
"patch_size": 14,
|
||||||
|
"prefix": null,
|
||||||
|
"problem_type": null,
|
||||||
|
"pruned_heads": {},
|
||||||
|
"remove_invalid_values": false,
|
||||||
|
"repetition_penalty": 1.0,
|
||||||
|
"return_dict": true,
|
||||||
|
"return_dict_in_generate": false,
|
||||||
|
"sep_token_id": null,
|
||||||
|
"task_specific_params": null,
|
||||||
|
"temperature": 1.0,
|
||||||
|
"tie_encoder_decoder": false,
|
||||||
|
"tie_word_embeddings": true,
|
||||||
|
"tokenizer_class": null,
|
||||||
|
"top_k": 50,
|
||||||
|
"top_p": 1.0,
|
||||||
|
"torch_dtype": null,
|
||||||
|
"torchscript": false,
|
||||||
|
"transformers_version": "4.21.0.dev0",
|
||||||
|
"typical_p": 1.0,
|
||||||
|
"use_bfloat16": false
|
||||||
|
},
|
||||||
|
"vision_config_dict": {
|
||||||
|
"hidden_size": 1024,
|
||||||
|
"intermediate_size": 4096,
|
||||||
|
"num_attention_heads": 16,
|
||||||
|
"num_hidden_layers": 24,
|
||||||
|
"patch_size": 14
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,20 @@
|
|||||||
|
{
|
||||||
|
"crop_size": 224,
|
||||||
|
"do_center_crop": true,
|
||||||
|
"do_convert_rgb": true,
|
||||||
|
"do_normalize": true,
|
||||||
|
"do_resize": true,
|
||||||
|
"feature_extractor_type": "CLIPFeatureExtractor",
|
||||||
|
"image_mean": [
|
||||||
|
0.48145466,
|
||||||
|
0.4578275,
|
||||||
|
0.40821073
|
||||||
|
],
|
||||||
|
"image_std": [
|
||||||
|
0.26862954,
|
||||||
|
0.26130258,
|
||||||
|
0.27577711
|
||||||
|
],
|
||||||
|
"resample": 3,
|
||||||
|
"size": 224
|
||||||
|
}
|
||||||
@@ -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 = 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)
|
||||||
@@ -85,7 +65,7 @@ class TileLayout:
|
|||||||
mask = torch.zeros((1, 1, size[0], size[1]), dtype=torch.float)
|
mask = torch.zeros((1, 1, size[0], size[1]), dtype=torch.float)
|
||||||
mask[:, :, s[0] : e[0], s[1] : e[1]] = 1.0
|
mask[:, :, s[0] : e[0], s[1] : e[1]] = 1.0
|
||||||
if blend and self.blending > 0:
|
if blend and self.blending > 0:
|
||||||
mask = box_blur(mask, (self.blending, self.blending), separable=True)
|
mask = box_blur(mask, (self.blending, self.blending))
|
||||||
return mask.squeeze(0)
|
return mask.squeeze(0)
|
||||||
|
|
||||||
def merge(self, image: Tensor, index: int, tile: Tensor):
|
def merge(self, image: Tensor, index: int, tile: Tensor):
|
||||||
@@ -93,83 +73,112 @@ class TileLayout:
|
|||||||
rect = self.rect(coord)
|
rect = self.rect(coord)
|
||||||
mask = self.mask(coord, blend=True)
|
mask = self.mask(coord, blend=True)
|
||||||
mask = mask.reshape(*mask.shape, 1).repeat(1, 1, 1, image.shape[-1])
|
mask = mask.reshape(*mask.shape, 1).repeat(1, 1, 1, image.shape[-1])
|
||||||
image[*rect] = (1 - mask) * image[*rect] + mask * tile
|
image[rect] = (1 - mask) * image[rect] + mask * tile
|
||||||
|
|
||||||
|
|
||||||
class ExtractImageTile:
|
class CreateTileLayout(io.ComfyNode):
|
||||||
@classmethod
|
@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)
|
||||||
|
|||||||
+119
@@ -0,0 +1,119 @@
|
|||||||
|
"""Text translation using Argos Translate.
|
||||||
|
|
||||||
|
The node takes text input and translates it to English. The text may contain any
|
||||||
|
number of language directives in the form `lang:xx` where `xx` is a two-letter
|
||||||
|
language code. Text fragments after a language directives are translated.
|
||||||
|
If the language is `en` text is passed through unmodified.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
import re
|
||||||
|
from functools import cache
|
||||||
|
from typing import NamedTuple
|
||||||
|
from comfy_api.latest import io
|
||||||
|
|
||||||
|
|
||||||
|
@cache
|
||||||
|
def available_languages():
|
||||||
|
try:
|
||||||
|
from argostranslate.package import update_package_index, get_available_packages
|
||||||
|
|
||||||
|
update_package_index()
|
||||||
|
list = get_available_packages()
|
||||||
|
return [(l.from_code, l.from_name) for l in list if l.to_code == "en"]
|
||||||
|
except ImportError:
|
||||||
|
return [("NOT INSTALLED", "NOT INSTALLED")]
|
||||||
|
|
||||||
|
|
||||||
|
def translate_chunk(text: str, language: str):
|
||||||
|
if text.strip() == "":
|
||||||
|
return text
|
||||||
|
|
||||||
|
target = "en"
|
||||||
|
if language == target:
|
||||||
|
return text
|
||||||
|
|
||||||
|
try:
|
||||||
|
from argostranslate.package import get_installed_packages, get_available_packages
|
||||||
|
from argostranslate.translate import translate
|
||||||
|
|
||||||
|
installed = get_installed_packages()
|
||||||
|
if not any(p.from_code == language and p.to_code == target for p in installed):
|
||||||
|
available = get_available_packages()
|
||||||
|
pkg = next(
|
||||||
|
(p for p in available if p.from_code == language and p.to_code == target), None
|
||||||
|
)
|
||||||
|
assert pkg, f"Couldn't find package for translation from {language}"
|
||||||
|
# print("Downloading and installing translation package", pkg) # this will cause encoding errors
|
||||||
|
pkg.install()
|
||||||
|
|
||||||
|
text, embeddings = _extract_embeddings(text)
|
||||||
|
translation = translate(text, language, target)
|
||||||
|
return embeddings + translation
|
||||||
|
|
||||||
|
except ImportError:
|
||||||
|
raise ImportError(
|
||||||
|
"Argos Translate is not installed. Please install it with `pip install argostranslate`"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def translate(text: str):
|
||||||
|
chunks = Chunk.parse(text)
|
||||||
|
return " ".join(translate_chunk(c.text, c.lang) for c in chunks)
|
||||||
|
|
||||||
|
|
||||||
|
class Translate(io.ComfyNode):
|
||||||
|
@classmethod
|
||||||
|
def define_schema(cls):
|
||||||
|
return io.Schema(
|
||||||
|
node_id="ETN_Translate",
|
||||||
|
display_name="Translate Text",
|
||||||
|
category="external_tooling",
|
||||||
|
inputs=[io.String.Input("text", multiline=True)],
|
||||||
|
outputs=[io.String.Output(display_name="translation")],
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def execute(cls, text: str):
|
||||||
|
return io.NodeOutput(translate(text))
|
||||||
|
|
||||||
|
|
||||||
|
_lang_regex = re.compile(r"(lang:\w\w)")
|
||||||
|
|
||||||
|
|
||||||
|
class Chunk(NamedTuple):
|
||||||
|
text: str
|
||||||
|
lang: str
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def parse(text: str):
|
||||||
|
languages = [code for code, name in available_languages()] + ["en"]
|
||||||
|
chunks: list[Chunk] = []
|
||||||
|
lang = "en"
|
||||||
|
last = 0
|
||||||
|
for m in _lang_regex.finditer(text):
|
||||||
|
if m.start() > 0:
|
||||||
|
chunks.append(Chunk(text[last : m.start()].strip(), lang))
|
||||||
|
last = m.end()
|
||||||
|
lang = m.group(0)[5:]
|
||||||
|
if lang not in languages:
|
||||||
|
raise ValueError(
|
||||||
|
f"Invalid language directive {m.group(0)} - {lang} is not a known language code."
|
||||||
|
f" Available languages: {', '.join(languages)}"
|
||||||
|
)
|
||||||
|
if last < len(text):
|
||||||
|
chunks.append(Chunk(text[last:].strip(), lang))
|
||||||
|
return [c for c in chunks if c.text != ""]
|
||||||
|
|
||||||
|
|
||||||
|
_embedding_regex = re.compile(r"(embedding:[^\s,]+)")
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_embeddings(text: str):
|
||||||
|
matches = _embedding_regex.findall(text)
|
||||||
|
embeddings = " ".join(matches)
|
||||||
|
if matches:
|
||||||
|
embeddings += " "
|
||||||
|
for m in matches:
|
||||||
|
text = text.replace(m, "")
|
||||||
|
return text, embeddings
|
||||||
@@ -0,0 +1,771 @@
|
|||||||
|
{
|
||||||
|
"last_node_id": 78,
|
||||||
|
"last_link_id": 126,
|
||||||
|
"nodes": [
|
||||||
|
{
|
||||||
|
"id": 60,
|
||||||
|
"type": "LoadImage",
|
||||||
|
"pos": [
|
||||||
|
-692,
|
||||||
|
-758
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 310.5925598144531,
|
||||||
|
"1": 335.9309997558594
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 0,
|
||||||
|
"mode": 0,
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "IMAGE",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"links": [
|
||||||
|
101,
|
||||||
|
110,
|
||||||
|
111,
|
||||||
|
112,
|
||||||
|
113
|
||||||
|
],
|
||||||
|
"shape": 3,
|
||||||
|
"slot_index": 0
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "MASK",
|
||||||
|
"type": "MASK",
|
||||||
|
"links": null,
|
||||||
|
"shape": 3
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "LoadImage"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
"photo.jpg",
|
||||||
|
"image"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 74,
|
||||||
|
"type": "ETN_MergeImageTile",
|
||||||
|
"pos": [
|
||||||
|
-326,
|
||||||
|
-250
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 315,
|
||||||
|
"1": 98
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 12,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "image",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"link": 120
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "layout",
|
||||||
|
"type": "TILE_LAYOUT",
|
||||||
|
"link": 122,
|
||||||
|
"slot_index": 1
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "tile",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"link": 121
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "IMAGE",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"links": [
|
||||||
|
123
|
||||||
|
],
|
||||||
|
"shape": 3,
|
||||||
|
"slot_index": 0
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "ETN_MergeImageTile"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
3
|
||||||
|
],
|
||||||
|
"color": "#232",
|
||||||
|
"bgcolor": "#353"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 61,
|
||||||
|
"type": "ETN_TileLayout",
|
||||||
|
"pos": [
|
||||||
|
-333,
|
||||||
|
-639
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 315,
|
||||||
|
"1": 106
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 2,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "image",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"link": 101
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "TILE_LAYOUT",
|
||||||
|
"type": "TILE_LAYOUT",
|
||||||
|
"links": [
|
||||||
|
102,
|
||||||
|
104,
|
||||||
|
105,
|
||||||
|
106,
|
||||||
|
122,
|
||||||
|
124
|
||||||
|
],
|
||||||
|
"shape": 3,
|
||||||
|
"slot_index": 0
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "ETN_TileLayout"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
880,
|
||||||
|
48,
|
||||||
|
16
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 63,
|
||||||
|
"type": "PreviewImage",
|
||||||
|
"pos": [
|
||||||
|
380,
|
||||||
|
-700
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 221.69317626953125,
|
||||||
|
"1": 191.44602966308594
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 8,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "images",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"link": 103
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "PreviewImage"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 69,
|
||||||
|
"type": "PreviewImage",
|
||||||
|
"pos": [
|
||||||
|
620,
|
||||||
|
-470
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 231.39317321777344,
|
||||||
|
"1": 200.04603576660156
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 11,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "images",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"link": 109
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "PreviewImage"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 75,
|
||||||
|
"type": "PreviewImage",
|
||||||
|
"pos": [
|
||||||
|
25,
|
||||||
|
-220
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 321.5931701660156,
|
||||||
|
"1": 266.5460205078125
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 14,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "images",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"link": 123
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "PreviewImage"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 76,
|
||||||
|
"type": "ETN_GenerateTileMask",
|
||||||
|
"pos": [
|
||||||
|
376,
|
||||||
|
-207
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 210,
|
||||||
|
"1": 85.74603271484375
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 7,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "layout",
|
||||||
|
"type": "TILE_LAYOUT",
|
||||||
|
"link": 124,
|
||||||
|
"slot_index": 0
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "MASK",
|
||||||
|
"type": "MASK",
|
||||||
|
"links": [
|
||||||
|
125
|
||||||
|
],
|
||||||
|
"shape": 3,
|
||||||
|
"slot_index": 0
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "ETN_GenerateTileMask"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
3,
|
||||||
|
true
|
||||||
|
],
|
||||||
|
"color": "#432",
|
||||||
|
"bgcolor": "#653"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 77,
|
||||||
|
"type": "MaskToImage",
|
||||||
|
"pos": [
|
||||||
|
620,
|
||||||
|
-210
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 210,
|
||||||
|
"1": 26
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 13,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "mask",
|
||||||
|
"type": "MASK",
|
||||||
|
"link": 125
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "IMAGE",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"links": [
|
||||||
|
126
|
||||||
|
],
|
||||||
|
"shape": 3,
|
||||||
|
"slot_index": 0
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "MaskToImage"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 78,
|
||||||
|
"type": "PreviewImage",
|
||||||
|
"pos": [
|
||||||
|
629,
|
||||||
|
-140
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 227.79318237304688,
|
||||||
|
"1": 201.64602661132812
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 15,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "images",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"link": 126
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "PreviewImage"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 71,
|
||||||
|
"type": "EmptyImage",
|
||||||
|
"pos": [
|
||||||
|
-681,
|
||||||
|
-247
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 315,
|
||||||
|
"1": 130
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 1,
|
||||||
|
"mode": 0,
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "IMAGE",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"links": [
|
||||||
|
120
|
||||||
|
],
|
||||||
|
"shape": 3,
|
||||||
|
"slot_index": 0
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "EmptyImage"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
2304,
|
||||||
|
1728,
|
||||||
|
1,
|
||||||
|
0
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 67,
|
||||||
|
"type": "PreviewImage",
|
||||||
|
"pos": [
|
||||||
|
380,
|
||||||
|
-470
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 222.79318237304688,
|
||||||
|
"1": 197.94602966308594
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 9,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "images",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"link": 107
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "PreviewImage"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 68,
|
||||||
|
"type": "PreviewImage",
|
||||||
|
"pos": [
|
||||||
|
620,
|
||||||
|
-703
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 225.89317321777344,
|
||||||
|
"1": 194.74603271484375
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 10,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "images",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"link": 108
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "PreviewImage"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 65,
|
||||||
|
"type": "ETN_ExtractImageTile",
|
||||||
|
"pos": [
|
||||||
|
30,
|
||||||
|
-760
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 278.19317626953125,
|
||||||
|
"1": 78
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 5,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "image",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"link": 110
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "layout",
|
||||||
|
"type": "TILE_LAYOUT",
|
||||||
|
"link": 105
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "IMAGE",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"links": [
|
||||||
|
108
|
||||||
|
],
|
||||||
|
"shape": 3,
|
||||||
|
"slot_index": 0
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "ETN_ExtractImageTile"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
2
|
||||||
|
],
|
||||||
|
"color": "#2a363b",
|
||||||
|
"bgcolor": "#3f5159"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 62,
|
||||||
|
"type": "ETN_ExtractImageTile",
|
||||||
|
"pos": [
|
||||||
|
40,
|
||||||
|
-630
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 274.19317626953125,
|
||||||
|
"1": 78
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 3,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "image",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"link": 113,
|
||||||
|
"slot_index": 0
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "layout",
|
||||||
|
"type": "TILE_LAYOUT",
|
||||||
|
"link": 102
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "IMAGE",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"links": [
|
||||||
|
103
|
||||||
|
],
|
||||||
|
"shape": 3,
|
||||||
|
"slot_index": 0
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "ETN_ExtractImageTile"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
0
|
||||||
|
],
|
||||||
|
"color": "#2a363b",
|
||||||
|
"bgcolor": "#3f5159"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 64,
|
||||||
|
"type": "ETN_ExtractImageTile",
|
||||||
|
"pos": [
|
||||||
|
40,
|
||||||
|
-510
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 273.2931823730469,
|
||||||
|
"1": 78
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 4,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "image",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"link": 111
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "layout",
|
||||||
|
"type": "TILE_LAYOUT",
|
||||||
|
"link": 104
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "IMAGE",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"links": [
|
||||||
|
107
|
||||||
|
],
|
||||||
|
"shape": 3,
|
||||||
|
"slot_index": 0
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "ETN_ExtractImageTile"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
1
|
||||||
|
],
|
||||||
|
"color": "#2a363b",
|
||||||
|
"bgcolor": "#3f5159"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 66,
|
||||||
|
"type": "ETN_ExtractImageTile",
|
||||||
|
"pos": [
|
||||||
|
40,
|
||||||
|
-370
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 275.4931945800781,
|
||||||
|
"1": 78
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 6,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "image",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"link": 112
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "layout",
|
||||||
|
"type": "TILE_LAYOUT",
|
||||||
|
"link": 106
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "IMAGE",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"links": [
|
||||||
|
109,
|
||||||
|
121
|
||||||
|
],
|
||||||
|
"shape": 3,
|
||||||
|
"slot_index": 0
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "ETN_ExtractImageTile"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
3
|
||||||
|
],
|
||||||
|
"color": "#2a363b",
|
||||||
|
"bgcolor": "#3f5159"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"links": [
|
||||||
|
[
|
||||||
|
101,
|
||||||
|
60,
|
||||||
|
0,
|
||||||
|
61,
|
||||||
|
0,
|
||||||
|
"IMAGE"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
102,
|
||||||
|
61,
|
||||||
|
0,
|
||||||
|
62,
|
||||||
|
1,
|
||||||
|
"TILE_LAYOUT"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
103,
|
||||||
|
62,
|
||||||
|
0,
|
||||||
|
63,
|
||||||
|
0,
|
||||||
|
"IMAGE"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
104,
|
||||||
|
61,
|
||||||
|
0,
|
||||||
|
64,
|
||||||
|
1,
|
||||||
|
"TILE_LAYOUT"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
105,
|
||||||
|
61,
|
||||||
|
0,
|
||||||
|
65,
|
||||||
|
1,
|
||||||
|
"TILE_LAYOUT"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
106,
|
||||||
|
61,
|
||||||
|
0,
|
||||||
|
66,
|
||||||
|
1,
|
||||||
|
"TILE_LAYOUT"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
107,
|
||||||
|
64,
|
||||||
|
0,
|
||||||
|
67,
|
||||||
|
0,
|
||||||
|
"IMAGE"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
108,
|
||||||
|
65,
|
||||||
|
0,
|
||||||
|
68,
|
||||||
|
0,
|
||||||
|
"IMAGE"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
109,
|
||||||
|
66,
|
||||||
|
0,
|
||||||
|
69,
|
||||||
|
0,
|
||||||
|
"IMAGE"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
110,
|
||||||
|
60,
|
||||||
|
0,
|
||||||
|
65,
|
||||||
|
0,
|
||||||
|
"IMAGE"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
111,
|
||||||
|
60,
|
||||||
|
0,
|
||||||
|
64,
|
||||||
|
0,
|
||||||
|
"IMAGE"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
112,
|
||||||
|
60,
|
||||||
|
0,
|
||||||
|
66,
|
||||||
|
0,
|
||||||
|
"IMAGE"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
113,
|
||||||
|
60,
|
||||||
|
0,
|
||||||
|
62,
|
||||||
|
0,
|
||||||
|
"IMAGE"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
120,
|
||||||
|
71,
|
||||||
|
0,
|
||||||
|
74,
|
||||||
|
0,
|
||||||
|
"IMAGE"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
121,
|
||||||
|
66,
|
||||||
|
0,
|
||||||
|
74,
|
||||||
|
2,
|
||||||
|
"IMAGE"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
122,
|
||||||
|
61,
|
||||||
|
0,
|
||||||
|
74,
|
||||||
|
1,
|
||||||
|
"TILE_LAYOUT"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
123,
|
||||||
|
74,
|
||||||
|
0,
|
||||||
|
75,
|
||||||
|
0,
|
||||||
|
"IMAGE"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
124,
|
||||||
|
61,
|
||||||
|
0,
|
||||||
|
76,
|
||||||
|
0,
|
||||||
|
"TILE_LAYOUT"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
125,
|
||||||
|
76,
|
||||||
|
0,
|
||||||
|
77,
|
||||||
|
0,
|
||||||
|
"MASK"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
126,
|
||||||
|
77,
|
||||||
|
0,
|
||||||
|
78,
|
||||||
|
0,
|
||||||
|
"IMAGE"
|
||||||
|
]
|
||||||
|
],
|
||||||
|
"groups": [],
|
||||||
|
"config": {},
|
||||||
|
"extra": {
|
||||||
|
"ds": {
|
||||||
|
"scale": 0.9090909090909091,
|
||||||
|
"offset": [
|
||||||
|
899.4068198252606,
|
||||||
|
850.8539704011341
|
||||||
|
]
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"version": 0.4
|
||||||
|
}
|
||||||
Binary file not shown.
|
After Width: | Height: | Size: 668 KiB |
File diff suppressed because one or more lines are too long
Binary file not shown.
|
After Width: | Height: | Size: 418 KiB |
Reference in New Issue
Block a user