Compare commits
33
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
32b7e1e301 | ||
|
|
552c89bec9 | ||
|
|
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 |
@@ -33,7 +33,7 @@ Loads a mask (single channel) from a PNG embedded into the prompt as base64 stri
|
||||
### Send Image (WebSocket)
|
||||
|
||||
Sends an output image over the client WebSocket connection as PNG binary data.
|
||||
* Inputs: the image (RGB or RGBA)
|
||||
* Inputs: the image (RGB or RGBA), supports batches
|
||||
|
||||
This will first send one binary message for each image in the batch via WebSocket:
|
||||
```
|
||||
@@ -44,6 +44,47 @@ That is two 32-bit integers (big endian) with values 1 and 2 followed by the PNG
|
||||
{'type': 'executed', 'data': {'node': '<node ID>', 'output': {'images': [{'source': 'websocket', 'content-type': 'image/png', 'type': 'output'}, ...]}, 'prompt_id': '<prompt ID>}}
|
||||
```
|
||||
|
||||
### Load Image from Cache
|
||||
|
||||
Loads an image or mask that has been uploaded previously into the workflow.
|
||||
Uploaded images are temporarily stored in RAM rather than written to disk. This
|
||||
method has less overhead compared to embedding images as base64 into the prompt,
|
||||
but is more complex to implement.
|
||||
* Inputs: id of an image that was uploaded previously
|
||||
* Outputs: image (RGB) and mask (A of RGBA input, or first channel if no alpha present).
|
||||
|
||||
To upload an image, upload the _bytes_ of a PNG via a HTTP PUT request to
|
||||
`/api/etn/image/{id}`. JPEG or other formats also work. Choose any `id` which
|
||||
does not clash with other images you upload, and reference it in the node. The
|
||||
request returns `201` if the image was uploaded and `200` if it was already
|
||||
cached.
|
||||
|
||||
### Save Image to Cache
|
||||
|
||||
Stores an output image in RAM temporarily and allows retrieval over HTTP.
|
||||
This is typically faster than WebSocket, especially for large images.
|
||||
* Inputs: the image (RGB or RGBA). Batches are supported.
|
||||
|
||||
This node will send a JSON message over WebSocket when an image is ready:
|
||||
```json
|
||||
{
|
||||
"type": "executed",
|
||||
"data": {
|
||||
"node": "<node ID>",
|
||||
"output": {
|
||||
"images": [
|
||||
{"source": "http", "id": "<image ID>", "content-type": "image/png", "type": "output"}
|
||||
]
|
||||
},
|
||||
"prompt_id": "prompt ID"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
To download the images, send a HTTP GET request to `/api/etn/image/{id}` with
|
||||
the image IDs from the message. Images will be cached for a few minutes.
|
||||
|
||||
|
||||
## <a id="regions" href="#toc">Regions</a>
|
||||
|
||||
These nodes implement attention masking for arbitrary number of image regions. Text prompts only apply to the masked area.
|
||||
@@ -164,6 +205,8 @@ There are various types of models that can be loaded as checkpoint, LoRA, Contro
|
||||
#### 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:
|
||||
@@ -177,7 +220,7 @@ Lists available models with additional classification info:
|
||||
...
|
||||
}
|
||||
```
|
||||
Possible values for base model: `sd15, sd20, sd21, sd3, sdxl, sdxl-refiner, ssd1b, svd, cascade-b, cascade-c, aura-flow, hunyuan-dit, flux, flux-schnell, lumina2, chroma, qwen-image`
|
||||
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`
|
||||
|
||||
@@ -187,6 +230,24 @@ Detection supports quantized models:
|
||||
|
||||
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.
|
||||
|
||||
+40
-56
@@ -1,59 +1,43 @@
|
||||
from comfy_api.latest import ComfyExtension, io
|
||||
from . import api as api, nodes, tile, region, nsfw, translation, krita
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"ETN_LoadImageBase64": nodes.LoadImageBase64,
|
||||
"ETN_LoadMaskBase64": nodes.LoadMaskBase64,
|
||||
"ETN_SendImageWebSocket": nodes.SendImageWebSocket,
|
||||
"ETN_CropImage": nodes.CropImage,
|
||||
"ETN_ApplyMaskToImage": nodes.ApplyMaskToImage,
|
||||
"ETN_ReferenceImage": nodes.ReferenceImage,
|
||||
"ETN_ApplyReferenceImages": nodes.ApplyReferenceImages,
|
||||
"ETN_TileLayout": tile.TileLayout,
|
||||
"ETN_ExtractImageTile": tile.ExtractImageTile,
|
||||
"ETN_ExtractMaskTile": tile.ExtractMaskTile,
|
||||
"ETN_GenerateTileMask": tile.GenerateTileMask,
|
||||
"ETN_MergeImageTile": tile.MergeImageTile,
|
||||
"ETN_BackgroundRegion": region.BackgroundRegion,
|
||||
"ETN_DefineRegion": region.DefineRegion,
|
||||
"ETN_ListRegionMasks": region.ListRegionMasks,
|
||||
"ETN_AttentionMask": region.AttentionMask,
|
||||
"ETN_NSFWFilter": nsfw.NSFWFilter,
|
||||
"ETN_Translate": translation.Translate,
|
||||
"ETN_KritaOutput": krita.KritaOutput,
|
||||
"ETN_KritaSendText": krita.KritaSendText,
|
||||
"ETN_KritaCanvas": krita.KritaCanvas,
|
||||
"ETN_KritaSelection": krita.KritaSelection,
|
||||
"ETN_KritaImageLayer": krita.KritaImageLayer,
|
||||
"ETN_KritaMaskLayer": krita.KritaMaskLayer,
|
||||
"ETN_Parameter": krita.Parameter,
|
||||
"ETN_KritaStyle": krita.KritaStyle,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"ETN_LoadImageBase64": "Load Image (Base64)",
|
||||
"ETN_LoadMaskBase64": "Load Mask (Base64)",
|
||||
"ETN_SendImageWebSocket": "Send Image (WebSocket)",
|
||||
"ETN_CropImage": "Crop Image",
|
||||
"ETN_ApplyMaskToImage": "Apply Mask to Image",
|
||||
"ETN_ReferenceImage": "Reference Image",
|
||||
"ETN_ApplyReferenceImages": "Apply Reference Images",
|
||||
"ETN_TileLayout": "Create Tile Layout",
|
||||
"ETN_ExtractImageTile": "Extract Image Tile",
|
||||
"ETN_ExtractMaskTile": "Extract Mask Tile",
|
||||
"ETN_MergeImageTile": "Merge Image Tile",
|
||||
"ETN_GenerateTileMask": "Generate Tile Mask",
|
||||
"ETN_BackgroundRegion": "Background Region",
|
||||
"ETN_DefineRegion": "Define Region",
|
||||
"ETN_ListRegionMasks": "List Region Masks",
|
||||
"ETN_AttentionMask": "Regions Attention Mask",
|
||||
"ETN_NSFWFilter": "NSFW Filter",
|
||||
"ETN_Translate": "Translate Text",
|
||||
"ETN_KritaOutput": "Krita Output",
|
||||
"ETN_KritaSendText": "Send Text",
|
||||
"ETN_KritaCanvas": "Krita Canvas",
|
||||
"ETN_KritaSelection": "Krita Selection",
|
||||
"ETN_KritaImageLayer": "Krita Image Layer",
|
||||
"ETN_KritaMaskLayer": "Krita Mask Layer",
|
||||
"ETN_Parameter": "Parameter",
|
||||
"ETN_KritaStyle": "Krita Style",
|
||||
}
|
||||
|
||||
class ExternalToolingNodes(ComfyExtension):
|
||||
async def get_node_list(self) -> list[type[io.ComfyNode]]:
|
||||
return [
|
||||
nodes.LoadImageCache,
|
||||
nodes.SaveImageCache,
|
||||
nodes.LoadImageBase64,
|
||||
nodes.LoadMaskBase64,
|
||||
nodes.SendImageWebSocket,
|
||||
nodes.ApplyMaskToImage,
|
||||
nodes.ReferenceImage,
|
||||
nodes.ApplyReferenceImages,
|
||||
tile.CreateTileLayout,
|
||||
tile.ExtractImageTile,
|
||||
tile.ExtractMaskTile,
|
||||
tile.GenerateTileMask,
|
||||
tile.MergeImageTile,
|
||||
region.BackgroundRegion,
|
||||
region.DefineRegion,
|
||||
region.ListRegionMasks,
|
||||
region.AttentionMask,
|
||||
nsfw.NSFWFilter,
|
||||
translation.Translate,
|
||||
krita.KritaOutput,
|
||||
krita.KritaSendText,
|
||||
krita.KritaCanvas,
|
||||
krita.KritaSelection,
|
||||
krita.KritaImageLayer,
|
||||
krita.KritaMaskLayer,
|
||||
krita.Parameter,
|
||||
krita.KritaStyle,
|
||||
krita.KritaStyleAndPrompt,
|
||||
]
|
||||
|
||||
|
||||
async def comfy_entrypoint():
|
||||
return ExternalToolingNodes()
|
||||
|
||||
|
||||
WEB_DIRECTORY = "./js"
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from __future__ import annotations
|
||||
from aiohttp import web
|
||||
from typing import NamedTuple
|
||||
from typing import Any, NamedTuple
|
||||
from pathlib import Path
|
||||
import json
|
||||
import traceback
|
||||
@@ -15,6 +15,7 @@ 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"
|
||||
|
||||
@@ -43,6 +44,8 @@ model_names = {
|
||||
"CosmosI2V": "cosmos",
|
||||
"CosmosT2IPredict2": "cosmos-predict2",
|
||||
"CosmosI2VPredict2": "cosmos-predict2",
|
||||
"ZImage": "z-image",
|
||||
"Lumina2": "lumina2",
|
||||
"WAN21_T2V": "wan21",
|
||||
"WAN21_I2V": "wan21",
|
||||
"WAN21_FunControl2V": "wan21-fun",
|
||||
@@ -53,6 +56,9 @@ model_names = {
|
||||
"ACEStep": "ace-step",
|
||||
"Omnigen2": "omnigen2",
|
||||
"QwenImage": "qwen-image",
|
||||
"ErnieImage": "ernie-image",
|
||||
"Flux2": "flux2",
|
||||
"Anima": "anima",
|
||||
}
|
||||
|
||||
gguf_architectures = {
|
||||
@@ -117,12 +123,15 @@ def inspect_safetensors(filename: str, model_type: str, is_checkpoint: bool):
|
||||
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"}
|
||||
|
||||
base_model_name = model_names.get(raw_name, "unknown")
|
||||
result = {"base_model": base_model_name}
|
||||
result: dict[str, Any] = {"base_model": base_model_name}
|
||||
result["is_inpaint"] = (
|
||||
base_model_name in ["sd15", "sdxl"] and input_count > 4
|
||||
) or raw_name == "FluxInpaint"
|
||||
@@ -141,6 +150,7 @@ def inspect_safetensors(filename: str, model_type: str, is_checkpoint: bool):
|
||||
return result
|
||||
return {"base_model": "unknown"}
|
||||
except Exception as e:
|
||||
print("[comfyui-tooling-nodes] Error inspecting file", filename)
|
||||
traceback.print_exc()
|
||||
return {"base_model": "unknown", "error": f"Failed to detect base model: {e}"}
|
||||
|
||||
@@ -150,12 +160,16 @@ def detect_svdq(cfg: dict) -> str | None:
|
||||
if comfy_config := md.get("comfy_config"):
|
||||
if isinstance(comfy_config, str):
|
||||
comfy_config = json.loads(comfy_config)
|
||||
return comfy_config.get("model_class")
|
||||
model_class = md.get("model_class")
|
||||
if model_class == "NunchakuFluxTransformer2dModel":
|
||||
return "Flux"
|
||||
if model_class == "NunchakuQwenImageTransformer2DModel":
|
||||
return "QwenImage"
|
||||
if model_class := comfy_config.get("model_class"):
|
||||
return model_class
|
||||
|
||||
match md.get("model_class"):
|
||||
case "NunchakuFluxTransformer2dModel":
|
||||
return "Flux"
|
||||
case "NunchakuQwenImageTransformer2DModel":
|
||||
return "QwenImage"
|
||||
case "NunchakuZImageTransformer2DModel":
|
||||
return "ZImage"
|
||||
return None
|
||||
|
||||
|
||||
@@ -167,6 +181,9 @@ def inspect_gguf(filename: str, model_type: str):
|
||||
|
||||
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:
|
||||
@@ -177,19 +194,45 @@ def inspect_gguf(filename: str, model_type: str):
|
||||
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:
|
||||
result["quant"] = reader.get_field("general.file_type").lower()
|
||||
except Exception as e:
|
||||
if file_type := reader.get_field("general.file_type"):
|
||||
result["quant"] = file_type.contents().lower()
|
||||
except Exception:
|
||||
result["quant"] = "gguf"
|
||||
return result
|
||||
|
||||
@@ -204,17 +247,22 @@ def inspect_diffusion_model(filename: str, model_type: str, is_checkpoint: bool)
|
||||
return inspect_safetensors(filename, model_type, is_checkpoint)
|
||||
|
||||
|
||||
def inspect_models(model_type: str):
|
||||
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
|
||||
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()
|
||||
@@ -243,6 +291,13 @@ def has_invalid_filename(filename: str):
|
||||
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)
|
||||
@@ -253,11 +308,11 @@ if _server is not None:
|
||||
error = has_invalid_folder_name(folder_name)
|
||||
if error is not None:
|
||||
return error
|
||||
return inspect_models(folder_name)
|
||||
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")
|
||||
return inspect_models("checkpoints", request.rel_url.query)
|
||||
|
||||
@_server.routes.get("/api/etn/languages")
|
||||
async def languages(request):
|
||||
@@ -277,6 +332,50 @@ if _server is not None:
|
||||
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", "")
|
||||
|
||||
@@ -32,7 +32,6 @@ function loadImage(base64) {
|
||||
}
|
||||
|
||||
const canvasIcon = loadImage("data:image/webp;base64,UklGRg4KAABXRUJQVlA4WAoAAAAQAAAAYwAAYwAAQUxQSNsDAAARoIRs/yI5+uAHU1kZiOu6u7u7xnObuK9NpBoKCoqi1jd6WonrsD7ROc2e3OJJr28T94aGhobYj+8w9v/X/7+n3UNETAD+b1IW1crLLrMhbYJbRs9Zv2VvuXbqVK28d8v6OaNvCdpIXkaT5InJgJgRABeNbirRYKlp9EUAJB8ZlUq2XAtI1wQIhn1RIUlV1Y5UVUmy0jwsACQPt7CtsloQSBcE6D6jSFKVRlVJFmd0B8QeClSSSn7/ACCdQjBjL0mlRSW5d0YA+4KESpJKnd8dHQsweBeptK7krsGAWIIgppKkkn+NAKSNoO9KUplLJVf2hliCIKWyrZIfXwYIgCdLzHXpSVgXvElt07bcKAAKSs2TUgvWIHiH2p6SX98pi+jgIrEFwUJqO6Ty1HY6qGwObEHwPrU9KqkOUNkS2AKwjNoeqXRSuV7EGtZQO3BVuQhiDV9Q3aIyhFiTDVS3qHwMYgtBK9Utcn932A/20nHlGoglwdOn1DEqn4XYQfA7PfB7AKuCRiqdVzZCbKD+EL14qB4WBTOoPlA2Qsyh7i9f/FUH44LBVHpRORhibr0/ms2hf5XerPaFYcFoqi+UDeaafNJkCsFeenRvnRnBLfSp3mBqNNUfygZTc/wyx9QGv6w3A2zxyxZDwV56da+h3mW/lA1dVPPLKUOXnfIL/8UurvnltKH+Vb9UDXU/4ZeyoaDkl5Ih/ET1h/InU81+aTYjmOOXOaYm06sTzAD3n/JJ7X5T/YtUXyiLfU3JGnp0jZjCS+oPfQnG7ylR/aAs3WmuexO92dTdHKZWfFGZCot3fkX1gbL1dhv1SZVerCb1NvB0K9U9ZevTsNo/OUEPnkj628Hgj33w8WBY7h9up7ql3B72tYUn51ToeHnOk7B+/tSV6pYum3q+PVwbtVDdUbZE1yKPzyY/UV1R/hg/i1zKhKRIdUNZTCZIPtB9Rvo71QXl78mM7sjrxVFSpJPFJLoY+b06jn904ccovhp5vjaOWkjNk5Ithfha5PuKqLCozFyXFxWiK5D3/o1xuiVPW9K4sT/yf35DFi47mpejy8Ks4Xw4+UAcxRvKeaisj6P4Abjaf3xWSDYcbaNmtM3RDUkhG98fDt8+IyvEK3fV2Fa1M6psW9u1Mi5kM26H28EDM7IofPPj7WUaLG//+M0wymY8EMD54M7JaRqF8cKPvy4eqtROq56uVQ4Vv/54YRxGaTr5zgB+vOjpl5IsicIwSrJ35sx5J0uiMIySLHnp6YvgUel/z6iXojTL0jRJ0jTL0uilUff0F/i3/qIb7nn06WHDnn70nhsuqhf8hxQAVlA4IAwGAACwHgCdASpkAGQAPm0wk0akIqGhLRGrUIANiWYA1BHh/t2rC93/Hf2Was/feJ2MrzB6VP6M9gD9Nunp5rP2y9Z70q+gV/Y/+B1kHoAeWv+1Xwa/t5+5PtQXQfhgKtwfp++wwU1Dx6fSXsD+VV7JPQ5/aRrblS2nsaggvO1Mch8UhK6pYtQxLM/VgrswZ0vLV8b6SwudSWaCFSHUiXEUQWX6krc9GtWanHMeaDd9wRYCfO5TwpYkgGAIkaLI4p6taB375EUfaVubYzKMfHSz2KpivsjWF0Vf+YJbACgi8j86d6EiJhQFF31NBJdS+QrGtJ2RUJRbahp1MXso6/J8AAD+/TKL/9q5tf/zOmF5Fe8B0Zn0yX3C0VLv0zxxv2+WH/dbz//rc2RS4TC1UzFVQiXVn5+Y0r+RsfJPsfPNuT02INz8gty7fI7fA/D1Wj2Jv+4RwdpyXs+cRxaT84bme5rMmPf+BH7NDUPKsj7GJ+w/6nBW2vsiPalWPfvBk6AQ3kCHmVecXkcnOgpoZ4ruAF/9Ze93DG5/8Y32x8b/CKPRt1jaXXy2LnoPvSNUT77gbB+/7vI1pfBfUHJsSwheIXY7QSixh7Ya8IliO3wqvI/uIFZAZd9pL8R1gRpYouBoyL5uIuGWQAZC5SKY0SruTf66stUOJVO9hlokeb5lWVzo7FO/Oeb/oj9iK4bqFhNZLCfqsBlH/OeefoP9sFdl7Mq1xmsevmzkfgwyiXg5hxMIP/Wa0JMPVl+XEFqTveAf1M8IBDu/pX/hCEnMn1n15Smyf72eDXKQqBrvp6BugyXXaJ05FDoz8MONUFh4rcjGL7AcijbcZ0SYwJkoeAKBW/I/sjKzTRtTP2E1fLB/8TWnzieHznDAKdlTuY2nSVTwCqFZNcFeFn7boziHOmYBLJin52d874mq1pHmJnulhT96LbKVW4vAT5PnY5F9TzmnDMwIFm4IAuEaA8X8XLE4Hp+AUEG4oswxRbVfOfxNJRyxFO3UB+v+ALgMP8kOf0uK3/3WOq4o/roivfzvW/fXviTC0mx+352hGaO+axx6vFa3eIkFsUEXCdo2LFHIlM8BtPuGUhgvM3oygIMAgvmUKILe0DFYVXhG/QoLi3sYfaoK/f0tX+fNnXhhxwEj1/Ct2Z64g0qWmkgwNkyy8m90EK1HsX0Q10CHVakDZePz5ts37u3GCANwGHQzWB+hNsevqjuU3qT95yGs0jjOtI/IjKsH9JbAmZkjGvNCPOC+FYUkOkwao9sOESY6zCgx9CM7g2LU4/CSHGoe2t0vWV/cMDH1HzI+Wa/yYp9CLDIh7J7iJd/2KnixeJvOhbUvbr9gubyyQU1iO5bnD9T536j++jKDVIk0Fwzk+d+j2eueHsIFJUvdyo2TyxP0kJbWr36R1s3giryqPvrsR5SkXx16+xqDrX4elhqh+1FwzNnSF5Lj5EUT/UC2rJvoAikbnvQ3NtJ9e83++idf3ja4FaLcUDxhoN5Rl5Ziz1LvF9iVeb6Su0QWYoRyBbyZ/pRbgYyhlAU/tonH7Wt+KhPDmXKIo0u4FDbAXM8avbFk4ax6e/dYITOCe+9dVEgcTOnBfhv0Yotd3EzNjZkLz4ksKGtFXcWIZRJ5YAyfzPYsyPex6/6ud9r2Ha9oxhVSIJV418e83qcPOIPlpe+LVGc69W6eC83l/zloqM9D6zQMkfqrjBZNpRkQS0sn8sxSu3s5qzhtH8cvjZk83gMqdfnHnl+1bvA7BI/g4+ePU7HUb9vK3Qw35bVmDcXa8xxWS2NQj8iWMH1cbHLXlboQsaCxIZoo+SeXR6ePUw3k6C/OxgqhjzExMJjLdBjoBeWYt3RPG2foTvx0T0Iz8ukdrCRJMG6HaR+6/f4nG/4xkr/fLhGlqOE/hBDBhuqnANj1CrujVDs2YayTvPuIcqCpNd3i8fOR8DfCq9ytS55F8akKneS6poHfB3bhjWbcIQXPvFzS7S5xLHEVoaixOwp0TL/8cQ8dxriyeddu5kCTyY7KepMQoeR+Pyn04nElkt9qqfYCTqHDtXBriC/UZh9AAAAAAAA=")
|
||||
const outputIcon = loadImage("data:image/webp;base64,UklGRrIHAABXRUJQVlA4WAoAAAAQAAAAjwAAOwAAQUxQSKoCAAARkMbsnyFJ/2TVySSdzPJs27bxzbbtuznbtm3btm3bt+ikkkonlfxPU1X9n57zh4iYAPifZKQ2T7LIkKDkQ1FPrezCF/hjvq+dx7mwiqM3HnJ058yGXkorEfGCKZV6I6o+quOM4bMsgU5bfE1KOh4LEbGnv2RnUSer50DGFwxJ2qwWGTATERFfZPzBZNR9P1ZXTksgVdaODHgt/H5hCJjP0ME6eiLfCqTLypKBSPYdsmY2OjpWy0KOlF+EkIFk/DvnZ2tIzZA0a0kHUtok0KfWkxieJQQZBQksq3QWiXODEGSnYRsq76mxjJQK0cDdKjY1XpLSWyKYVYFTSyxLqCZSvexKppZnZDCZTNpD6BKTi2iIRbqzZc6gW8x5m1ptOCEuYaJ74B1T6RkhjPQXkugiuC9EBSk38wftRACZsdKLEHGRQjJSCyUgnwichSgdj4jYW64KqQsywN1E1JJqReqtWyEvLtOTFHMtfG8GXMYT6C5zQLIjqXiJm+guNw2ZOqTu+0uJ7sKyg2w+Urv9GdxdGoN0JKnh/qBnIE1LlH6LiHNAMZ5SWQkoJAJHcQ7iTUNlNyE7Vga4a7CsoFqP0AmQjqdm5dTWGJRTvqfTSu4ONe7VNRs0LiTzLK3cNEHsGWhuZ+gom0hlNMgXZ7T4WF16Q6YRuZZPAS7TYrGUoPgFEnZPUC3EKDEf0O6WSGFmMiVoxejw3UDcHC6c21oFNLZjVNhGgxrkHC+cOtQKtBa5aQkCLL4VBGDZsbYzu7uF6AGouPQ1t7mTIv5AMw8EZLmhbx0QizqHgJMer5MQwHkG7NP2YmgzCMqRPYff0cIWDSgOwbrcgOFnhcqLhQNaeSB4h1XxDZh29LX4kXVt5YABrZJBkM/acIBsx5IG/AyGpS1UrkrF4tm98K8uVlA4IOIEAABwHACdASqQADwAPm0uk0ckIiGhLjUJmIANiWgOuBpEsADI+tE9j+ifkl+QHyd1n+77tGXf04+Cpx/549gD9X+lF5t/2d9cn0T+gB/d/8R1gHoAeWx+3nwZfuZ6T+aq/0Dtm/zaCN5VfOk9UfrX8A/Ry9E1elWp+7DqL5PqqU3p7QQgEkexRJ8VwS4d4Xk5cyjjrDbvzKxXEwCzN96k0RfNIpiQ2YsSbcyoRxX5fFzllf+dOy8uCMy/ebkjS0ONwuzkRR55zP3zjA4e+C969ch1Ab9LgcrpUqQ8MOvEXnQmqQkxXM4x4nA1f+jJAAD+8tXq96F285yEhMGbOWPp352/yRnzPmWRyRibmd800tluUOW4IyIZz2Hw1xYA9/xsSgKy0yQS//7BbNldSPJ+MCX/mxBqrttfeQX/mf/+2AcZ7Z1wDrdNoOnt8ISIu33p34GUAqqEyPwtdrhMf55SjsQmwUtm/I/qPiQ0ZOPv5ci6kCP0Ddb2jRr38UOXvi54DlMMmkxTs7j/J3jUQYepc7xEgTVVJFcf+8P//18L/9gA//9fHuK11sTUs+RbYDzZKn0uM2PnTEUpAJAT4wETKSU4KDn3wR3rUGGycaAPVo40AjzO5g7VkynMJuo3M2vclcmfmS3ygBGDqjGHQybd03tnGdkGOCGLLlTEABV+UgJxYxH2YxRX9zUULXFnfnBpPxQcOdq+zs1zi9uI2iAqSKwh8Dhch2Ytz8iaZLW39S3+3pmGdITR49+nlHjcG4xNVSYRLFLRmEj/H/I+7qd90N6AF9aRDuUFH1O7ONRGjEQGvPMEF0Fj5atb5w9tjc1pcKTsaWvT2GbF9NQ31HmpaLAgs8szVbuEC8GHKCESxKmx+Hrh5ZwjrNihG3KL0H1n3/g/WetlSYEFYsYTXQmgyUGCVIILkJYrRIdLB5iVAPrseYWKCT8HgJuCUAhaqRO+6jn1fkplsC0yYCveVI+yyDsVr98kmO5arhQ3u+aqKVUvJ8xZL9as4008lN9DkKcRhvC4BwWdhupsqUYwLQmaQhLxP15875P/r43c8r4NI4sLDiCi7Rzww1dWNTyThiA07x8b/zTaFC9Sz+jtZDpRPoSf3LS+TmvHZQ+yv/N9nSK/CGpimH/qjTJOQRStf5ppvzT0FzGMX2tqNndJbZD8idLxJFXZekFF16KC/6scsX/lTNL+XFfTqsreVXu7bL/wjNVTPeGkJJE7aWcXP2+3qQTMv+LaO9INAsG3cyp5Co/F06O8XoVtYZXBjH3f3r9Y8Wp89/fqq2OfQSD2/Ujo1t0fNnMA14gpYdtm6+/RcRgNQGIPPGxgAaFjsfC4+63CcHr1nczuKyXiQjmoIH7n/0NCmJv3O+v/Lp30d3n/060TaO5ffQGrrx0O7TYUAC6pdQxfOeuX4/EsKJgMKTW18feF5m1SX4ODnH1SWutwnm5T/k0/l4YXbLUi8QRbdtx74QL9DJRtKP8bDT+yyf//8nAEApoAJH7jMHoQv7XKzIUdH1TDS7Phokc3PP5m68+eUTHU17v50avNmnEHCfybI4FC35LTpSaGqRsgNJJliiV37VIfbUlfDfgIqZmmxEHmnCQTSg2zcf5+9LWPblYTxLShx/2U34N3Rf/3Zvie6j8SS/8X+Yo+dXhKIyg1WX040AAAAA==")
|
||||
|
||||
function setIconImage(nodeType, image, size, padRows, padCols) {
|
||||
const onAdded = nodeType.prototype.onAdded
|
||||
@@ -77,7 +76,8 @@ function defaultParameterType(widgetType, connectedNode, connectedWidget) {
|
||||
if (connectedNode.comfyClass === "CLIPTextEncode") {
|
||||
paramType = "prompt (positive)"
|
||||
}
|
||||
if (connectedWidget.options?.round === 1) {
|
||||
const round = connectedWidget.options?.round
|
||||
if ((paramType == "number" && round === undefined) || round === 1) {
|
||||
paramType = "number (integer)"
|
||||
}
|
||||
return paramType
|
||||
@@ -96,7 +96,7 @@ function valueMatchesType(value, type, options) {
|
||||
|
||||
function optionalWidgetValue(widgets, index, fallback) {
|
||||
const result = widgets.length > index ? widgets[index].value : null
|
||||
return result === null || result === 0 ? fallback : result
|
||||
return result === null || result === -1e10 || result === 1e10 ? fallback : result
|
||||
}
|
||||
|
||||
function changeWidgets(node, type, connectedNode, connectedWidget) {
|
||||
@@ -206,8 +206,6 @@ app.registerExtension({
|
||||
beforeRegisterNodeDef(nodeType /*typeof LGraphNode*/, nodeData /*ComfyObjectInfo*/, app) {
|
||||
if (nodeData.name === "ETN_KritaCanvas") {
|
||||
setIconImage(nodeType, canvasIcon, [200, 100], 0, 2)
|
||||
} else if (nodeData.name === "ETN_KritaOutput") {
|
||||
setIconImage(nodeType, outputIcon, [200, 100], 1, 0)
|
||||
} else if (nodeData.name === "ETN_Parameter") {
|
||||
setupParameterNode(nodeType)
|
||||
} else if (nodeData.name === "ETN_SendText") {
|
||||
|
||||
@@ -1,13 +1,16 @@
|
||||
import sys
|
||||
import torch
|
||||
import numpy as np
|
||||
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
|
||||
|
||||
import server
|
||||
import comfy.samplers
|
||||
from comfy.comfy_types.node_typing import IO
|
||||
from .nodes import SendImageWebSocket
|
||||
|
||||
|
||||
@@ -72,37 +75,74 @@ class _BasicTypes(str):
|
||||
BasicTypes = _BasicTypes("BASIC")
|
||||
|
||||
|
||||
class KritaOutput:
|
||||
class OutputBatchMode(Enum):
|
||||
default = "default"
|
||||
images = "images"
|
||||
animation = "animation"
|
||||
layers = "layers"
|
||||
|
||||
|
||||
class KritaOutput(io.ComfyNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {"images": ("IMAGE",)}}
|
||||
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,
|
||||
)
|
||||
|
||||
RETURN_TYPES = ()
|
||||
FUNCTION = "send_images"
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = "krita"
|
||||
|
||||
def send_images(self, images):
|
||||
return SendImageWebSocket().send_images(images, "PNG")
|
||||
|
||||
|
||||
class KritaSendText:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"value": (IO.ANY, {}),
|
||||
"name": ("STRING", {"default": "Output"}),
|
||||
"type": (["text", "markdown", "html"], {"default": "text"}),
|
||||
}
|
||||
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
|
||||
|
||||
RETURN_TYPES = ()
|
||||
FUNCTION = "send"
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = "krita"
|
||||
|
||||
def send(self, value: Any, name: str, type: str):
|
||||
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",
|
||||
@@ -115,72 +155,109 @@ class KritaSendText:
|
||||
except Exception as e:
|
||||
text = f"Could not convert to text: {e}"
|
||||
|
||||
print(f"Sending text: {name} = {text}")
|
||||
return {"ui": {"text": [{"name": name, "text": text, "content-type": mime}]}}
|
||||
return io.NodeOutput(ui={"text": [{"name": name, "text": text, "content-type": mime}]})
|
||||
|
||||
|
||||
class KritaCanvas:
|
||||
class KritaCanvas(io.ComfyNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {}
|
||||
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"),
|
||||
],
|
||||
)
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "INT", "INT", "INT")
|
||||
RETURN_NAMES = ("image", "width", "height", "seed")
|
||||
FUNCTION = "placeholder"
|
||||
CATEGORY = "krita"
|
||||
|
||||
def placeholder(self):
|
||||
return (_placeholder_image(), 512, 512, 0)
|
||||
|
||||
|
||||
class KritaSelection:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {}
|
||||
|
||||
RETURN_TYPES = (IO.MASK, IO.BOOLEAN)
|
||||
RETURN_NAMES = ("mask", "active")
|
||||
FUNCTION = "placeholder"
|
||||
CATEGORY = "krita"
|
||||
|
||||
def placeholder(self):
|
||||
return (torch.ones(1, 512, 512), False)
|
||||
def execute(cls, **kwargs):
|
||||
return io.NodeOutput(_placeholder_image(), 512, 512, 0, torch.ones(1, 512, 512))
|
||||
|
||||
|
||||
class KritaImageLayer:
|
||||
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 INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"name": ("STRING", {"default": "Image"}),
|
||||
}
|
||||
}
|
||||
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"),
|
||||
],
|
||||
)
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "MASK")
|
||||
RETURN_NAMES = ("image", "mask")
|
||||
FUNCTION = "placeholder"
|
||||
CATEGORY = "krita"
|
||||
|
||||
def placeholder(self, name: str):
|
||||
return (_placeholder_image(), torch.ones(1, 512, 512))
|
||||
|
||||
|
||||
class KritaMaskLayer:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"name": ("STRING", {"default": "Mask"}),
|
||||
}
|
||||
}
|
||||
def execute(cls, **kwargs):
|
||||
return io.NodeOutput(torch.ones(1, 512, 512), False, 0, 0)
|
||||
|
||||
RETURN_TYPES = ("MASK",)
|
||||
RETURN_NAMES = ("mask",)
|
||||
FUNCTION = "placeholder"
|
||||
CATEGORY = "krita"
|
||||
|
||||
def placeholder(self, name: str):
|
||||
return (torch.ones(1, 512, 512),)
|
||||
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 = [
|
||||
@@ -193,71 +270,95 @@ _param_types = [
|
||||
"prompt (positive)",
|
||||
"prompt (negative)",
|
||||
]
|
||||
_any_float = {"default": 0.0, "min": -sys.float_info.max, "max": sys.float_info.max}
|
||||
_fmax = sys.float_info.max
|
||||
|
||||
|
||||
class Parameter:
|
||||
class Parameter(io.ComfyNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"name": ("STRING", {"default": "Parameter"}),
|
||||
"type": (_param_types, {"default": "auto"}),
|
||||
"default": ("STRING", {"default": ""}),
|
||||
},
|
||||
"optional": {
|
||||
"min": ("FLOAT", _any_float),
|
||||
"max": ("FLOAT", _any_float),
|
||||
},
|
||||
}
|
||||
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")],
|
||||
)
|
||||
|
||||
RETURN_TYPES = (BasicTypes,)
|
||||
RETURN_NAMES = ("value",)
|
||||
FUNCTION = "placeholder"
|
||||
CATEGORY = "krita"
|
||||
|
||||
def placeholder(self, name: str, type: str, default, min=0.0, max=1.0):
|
||||
@classmethod
|
||||
def execute(cls, name: str, type: str, default, min=0.0, max=1.0): # type: ignore
|
||||
if type == "number":
|
||||
return (float(default),)
|
||||
return io.NodeOutput(float(default))
|
||||
elif type == "number (integer)":
|
||||
return (int(default),)
|
||||
return (default,)
|
||||
return io.NodeOutput(int(default))
|
||||
return io.NodeOutput(default)
|
||||
|
||||
|
||||
class KritaStyle:
|
||||
class KritaStyle(io.ComfyNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"name": ("STRING", {"default": "Style"}),
|
||||
"sampler_preset": (["auto", "regular", "live"],),
|
||||
}
|
||||
}
|
||||
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"),
|
||||
],
|
||||
)
|
||||
|
||||
RETURN_TYPES = (
|
||||
"MODEL",
|
||||
"CLIP",
|
||||
"VAE",
|
||||
"STRING",
|
||||
"STRING",
|
||||
comfy.samplers.KSampler.SAMPLERS,
|
||||
comfy.samplers.KSampler.SCHEDULERS,
|
||||
"INT",
|
||||
"FLOAT",
|
||||
)
|
||||
RETURN_NAMES = (
|
||||
"model",
|
||||
"clip",
|
||||
"vae",
|
||||
"positive prompt",
|
||||
"negative prompt",
|
||||
"sampler name",
|
||||
"scheduler",
|
||||
"steps",
|
||||
"guidance",
|
||||
)
|
||||
FUNCTION = "placeholder"
|
||||
CATEGORY = "krita"
|
||||
@classmethod
|
||||
def execute(cls, name: str, sampler_preset: str): # type: ignore
|
||||
raise NotImplementedError("This workflow must be started from Krita!")
|
||||
|
||||
def placeholder(self, name: str, sampler_preset: str):
|
||||
|
||||
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,6 +1,9 @@
|
||||
from __future__ import annotations
|
||||
from copy import copy
|
||||
from dataclasses import dataclass
|
||||
import time
|
||||
from typing import NamedTuple
|
||||
from uuid import uuid4
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
import base64
|
||||
@@ -11,18 +14,22 @@ from server import PromptServer, BinaryEventTypes
|
||||
|
||||
from comfy.clip_vision import ClipVisionModel
|
||||
from comfy.sd import StyleModel
|
||||
from comfy_api.latest import io
|
||||
|
||||
|
||||
class LoadImageBase64:
|
||||
class LoadImageBase64(io.ComfyNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {"image": ("STRING", {"multiline": False})}}
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="ETN_LoadImageBase64",
|
||||
display_name="Load Image (Base64)",
|
||||
category="external_tooling",
|
||||
inputs=[io.String.Input("image", multiline=False)],
|
||||
outputs=[io.Image.Output(display_name="image"), io.Mask.Output(display_name="mask")],
|
||||
)
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "MASK")
|
||||
CATEGORY = "external_tooling"
|
||||
FUNCTION = "load_image"
|
||||
|
||||
def load_image(self, image: str):
|
||||
@classmethod
|
||||
def execute(cls, image: str):
|
||||
_strip_prefix(image, "data:image/png;base64,")
|
||||
imgdata = base64.b64decode(image)
|
||||
img = Image.open(BytesIO(imgdata))
|
||||
@@ -40,16 +47,19 @@ class LoadImageBase64:
|
||||
return (img, mask)
|
||||
|
||||
|
||||
class LoadMaskBase64:
|
||||
class LoadMaskBase64(io.ComfyNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {"mask": ("STRING", {"multiline": False})}}
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="ETN_LoadMaskBase64",
|
||||
display_name="Load Mask (Base64)",
|
||||
category="external_tooling",
|
||||
inputs=[io.String.Input("mask", multiline=False)],
|
||||
outputs=[io.Mask.Output(display_name="mask")],
|
||||
)
|
||||
|
||||
RETURN_TYPES = ("MASK",)
|
||||
CATEGORY = "external_tooling"
|
||||
FUNCTION = "load_mask"
|
||||
|
||||
def load_mask(self, mask: str):
|
||||
@classmethod
|
||||
def execute(cls, mask: str):
|
||||
_strip_prefix(mask, "data:image/png;base64,")
|
||||
imgdata = base64.b64decode(mask)
|
||||
img = Image.open(BytesIO(imgdata))
|
||||
@@ -60,22 +70,22 @@ class LoadMaskBase64:
|
||||
return (img.unsqueeze(0),)
|
||||
|
||||
|
||||
class SendImageWebSocket:
|
||||
class SendImageWebSocket(io.ComfyNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"format": (["PNG", "JPEG"], {"default": "PNG"}),
|
||||
}
|
||||
}
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="ETN_SendImageWebSocket",
|
||||
display_name="Send Image (WebSocket)",
|
||||
category="external_tooling",
|
||||
inputs=[
|
||||
io.Image.Input("images"),
|
||||
io.Combo.Input("format", options=["PNG", "JPEG"], default="PNG"),
|
||||
],
|
||||
is_output_node=True,
|
||||
)
|
||||
|
||||
RETURN_TYPES = ()
|
||||
FUNCTION = "send_images"
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = "external_tooling"
|
||||
|
||||
def send_images(self, images, format):
|
||||
@classmethod
|
||||
def execute(cls, images: torch.Tensor, format: str):
|
||||
results = []
|
||||
for tensor in images:
|
||||
array = 255.0 * tensor.cpu().numpy()
|
||||
@@ -93,43 +103,154 @@ class SendImageWebSocket:
|
||||
"type": "output",
|
||||
})
|
||||
|
||||
return {"ui": {"images": results}}
|
||||
return io.NodeOutput(ui={"images": results})
|
||||
|
||||
|
||||
class CropImage:
|
||||
"""Deprecated, ComfyUI has an ImageCrop node now which does the same."""
|
||||
class ImageCache:
|
||||
timeout = 600 # 10 minutes
|
||||
max_size = 100 * 1024 * 1024 # 100 MB
|
||||
|
||||
@dataclass
|
||||
class Entry:
|
||||
data: bytes
|
||||
content_type: str
|
||||
timestamp: float
|
||||
retrieved: int
|
||||
|
||||
class OldEntry(NamedTuple):
|
||||
last_used: float
|
||||
deleted: float
|
||||
size: int
|
||||
retrieved: int
|
||||
|
||||
def __init__(self):
|
||||
self.images: dict[str, ImageCache.Entry] = {}
|
||||
self.old: dict[str, ImageCache.OldEntry] = {}
|
||||
|
||||
def add(self, image: Image.Image, format: str):
|
||||
key = uuid4().hex
|
||||
with BytesIO() as output:
|
||||
image.save(output, format=format, quality=95, compress_level=1)
|
||||
image_data = output.getvalue()
|
||||
|
||||
self.insert(key, image_data, f"image/{format.lower()}")
|
||||
return key
|
||||
|
||||
def insert(self, key: str, data: bytes, content_type: str):
|
||||
self.images[key] = ImageCache.Entry(
|
||||
data=data,
|
||||
content_type=content_type,
|
||||
timestamp=time.time(),
|
||||
retrieved=0,
|
||||
)
|
||||
|
||||
def get(self, key: str, extend: bool = False):
|
||||
entry = self.images.get(key)
|
||||
if entry is None:
|
||||
if old := self.old.get(key):
|
||||
now = time.time()
|
||||
print(
|
||||
f"[comfyui-tooling-nodes] requested image {key} has been deleted ",
|
||||
f"(last used {now - old.last_used:.0f}s ago, deleted {now - old.deleted:.0f}s ago, "
|
||||
f"size {old.size / 1024**2:.1f}MB, retrieved {old.retrieved} times)",
|
||||
)
|
||||
return None, None
|
||||
entry.retrieved += 1
|
||||
if extend:
|
||||
entry.timestamp = time.time()
|
||||
self.prune()
|
||||
return entry.data, entry.content_type
|
||||
|
||||
def prune(self):
|
||||
total_size = sum(len(entry.data) for entry in self.images.values())
|
||||
if total_size <= self.max_size:
|
||||
return
|
||||
# Remove least recently used entries until under max size
|
||||
sorted_entries = sorted(self.images.items(), key=lambda item: item[1].timestamp)
|
||||
now = time.time()
|
||||
for key, entry in sorted_entries:
|
||||
age = now - entry.timestamp
|
||||
if age > self.timeout or (age > 60 and entry.retrieved > 0):
|
||||
self.old[key] = ImageCache.OldEntry(
|
||||
entry.timestamp, now, len(entry.data), entry.retrieved
|
||||
)
|
||||
del self.images[key]
|
||||
total_size -= len(entry.data)
|
||||
if total_size <= self.max_size:
|
||||
break
|
||||
|
||||
def __contains__(self, key: str):
|
||||
return key in self.images
|
||||
|
||||
|
||||
image_cache = ImageCache()
|
||||
|
||||
|
||||
class LoadImageCache(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="ETN_LoadImageCache",
|
||||
display_name="Load Image from Cache",
|
||||
category="external_tooling",
|
||||
inputs=[io.String.Input("id", multiline=False)],
|
||||
outputs=[io.Image.Output(display_name="image"), io.Mask.Output(display_name="mask")],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"x": (
|
||||
"INT",
|
||||
{"default": 0, "min": 0, "max": 8192, "step": 1},
|
||||
),
|
||||
"y": (
|
||||
"INT",
|
||||
{"default": 0, "min": 0, "max": 8192, "step": 1},
|
||||
),
|
||||
"width": (
|
||||
"INT",
|
||||
{"default": 512, "min": 1, "max": 8192, "step": 1},
|
||||
),
|
||||
"height": (
|
||||
"INT",
|
||||
{"default": 512, "min": 1, "max": 8192, "step": 1},
|
||||
),
|
||||
}
|
||||
}
|
||||
def execute(cls, id: str):
|
||||
image_data, content_type = image_cache.get(id, extend=True)
|
||||
if image_data is None:
|
||||
raise ValueError(f"Image with ID {id} not found in cache.")
|
||||
|
||||
CATEGORY = "external_tooling"
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "crop"
|
||||
img = Image.open(BytesIO(image_data))
|
||||
w, h = img.size
|
||||
c = len(img.getbands())
|
||||
normalized = np.array(img).astype(np.float32) / 255.0
|
||||
tensor = torch.from_numpy(normalized).reshape(1, h, w, c)
|
||||
match c:
|
||||
case 1:
|
||||
image = tensor.expand(1, h, w, 3)
|
||||
mask = tensor.reshape(1, h, w)
|
||||
case 3:
|
||||
image = tensor
|
||||
mask = tensor[..., 0]
|
||||
case 4:
|
||||
image = tensor[..., :3]
|
||||
mask = tensor[..., 3]
|
||||
|
||||
def crop(self, image, x, y, width, height):
|
||||
out = image[:, y : y + height, x : x + width, :]
|
||||
return (out,)
|
||||
return io.NodeOutput(image, mask)
|
||||
|
||||
|
||||
class SaveImageCache(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="ETN_SaveImageCache",
|
||||
display_name="Save Image to Cache",
|
||||
category="external_tooling",
|
||||
inputs=[
|
||||
io.Image.Input("images"),
|
||||
io.Combo.Input("format", options=["PNG", "JPEG"], default="PNG"),
|
||||
],
|
||||
is_output_node=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, images: torch.Tensor, format: str):
|
||||
results = []
|
||||
for tensor in images:
|
||||
array = 255.0 * tensor.cpu().numpy()
|
||||
image = Image.fromarray(np.clip(array, 0, 255).astype(np.uint8))
|
||||
key = image_cache.add(image, format)
|
||||
|
||||
results.append({
|
||||
"source": "http",
|
||||
"id": key,
|
||||
"content-type": f"image/{format.lower()}",
|
||||
"type": "output",
|
||||
})
|
||||
return io.NodeOutput(ui={"images": results})
|
||||
|
||||
|
||||
def to_bchw(image: torch.Tensor):
|
||||
@@ -148,21 +269,22 @@ def mask_batch(mask: torch.Tensor):
|
||||
return mask
|
||||
|
||||
|
||||
class ApplyMaskToImage:
|
||||
class ApplyMaskToImage(io.ComfyNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"mask": ("MASK",),
|
||||
}
|
||||
}
|
||||
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")],
|
||||
)
|
||||
|
||||
CATEGORY = "external_tooling"
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "apply_mask"
|
||||
|
||||
def apply_mask(self, image: torch.Tensor, mask: torch.Tensor):
|
||||
@classmethod
|
||||
def execute(cls, image: torch.Tensor, mask: torch.Tensor):
|
||||
out = to_bchw(image)
|
||||
if out.shape[1] == 3: # Assuming RGB images
|
||||
out = torch.cat([out, torch.ones_like(out[:, :1, :, :])], dim=1)
|
||||
@@ -189,28 +311,26 @@ class _ReferenceImageData(NamedTuple):
|
||||
range: tuple[float, float]
|
||||
|
||||
|
||||
class ReferenceImage:
|
||||
class ReferenceImage(io.ComfyNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"weight": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0}),
|
||||
"range_start": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0}),
|
||||
"range_end": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0}),
|
||||
},
|
||||
"optional": {
|
||||
"reference_images": ("REFERENCE_IMAGE",),
|
||||
},
|
||||
}
|
||||
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")],
|
||||
)
|
||||
|
||||
CATEGORY = "external_tooling"
|
||||
RETURN_TYPES = ("REFERENCE_IMAGE",)
|
||||
RETURN_NAMES = ("reference_images",)
|
||||
FUNCTION = "append"
|
||||
|
||||
def append(
|
||||
self,
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
image: torch.Tensor,
|
||||
weight: float,
|
||||
range_start: float,
|
||||
@@ -222,24 +342,25 @@ class ReferenceImage:
|
||||
return (imgs,)
|
||||
|
||||
|
||||
class ApplyReferenceImages:
|
||||
class ApplyReferenceImages(io.ComfyNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"conditioning": ("CONDITIONING",),
|
||||
"clip_vision": ("CLIP_VISION",),
|
||||
"style_model": ("STYLE_MODEL",),
|
||||
"references": ("REFERENCE_IMAGE",),
|
||||
}
|
||||
}
|
||||
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")],
|
||||
)
|
||||
|
||||
CATEGORY = "external_tooling"
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
FUNCTION = "apply"
|
||||
|
||||
def apply(
|
||||
self,
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
conditioning: list[list],
|
||||
clip_vision: ClipVisionModel,
|
||||
style_model: StyleModel,
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
from __future__ import annotations
|
||||
from weakref import ref as WeakRef
|
||||
from pathlib import Path
|
||||
from tqdm import tqdm
|
||||
import torch
|
||||
@@ -8,6 +7,7 @@ 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
|
||||
|
||||
@@ -38,6 +38,10 @@ class CLIPSafetyChecker(PreTrainedModel):
|
||||
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]
|
||||
@@ -76,7 +80,7 @@ class CLIPSafetyChecker(PreTrainedModel):
|
||||
|
||||
|
||||
class CachedModels:
|
||||
_instance: WeakRef | None = None
|
||||
_instance: CachedModels | None = None
|
||||
|
||||
def __init__(self):
|
||||
model_dir = Path(__file__).parent / "safetychecker"
|
||||
@@ -91,11 +95,9 @@ class CachedModels:
|
||||
|
||||
@classmethod
|
||||
def load(cls):
|
||||
models = cls._instance and cls._instance()
|
||||
if models is None:
|
||||
models = cls()
|
||||
cls._instance = WeakRef(models)
|
||||
return models
|
||||
if cls._instance is None:
|
||||
cls._instance = CachedModels()
|
||||
return cls._instance
|
||||
|
||||
def download(self, url: str, target: Path):
|
||||
import requests
|
||||
@@ -118,29 +120,26 @@ class CachedModels:
|
||||
) from e
|
||||
|
||||
|
||||
class NSFWFilter:
|
||||
models: CachedModels
|
||||
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 INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"sensitivity": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.10}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "check"
|
||||
CATEGORY = "external_tooling"
|
||||
|
||||
def __init__(self):
|
||||
self.models = CachedModels.load()
|
||||
|
||||
def check(self, image, sensitivity):
|
||||
def execute(cls, image: Tensor, sensitivity: float):
|
||||
models = CachedModels.load()
|
||||
image = to_bchw(image)
|
||||
input = self.models.feature_extractor(image, do_rescale=False, return_tensors="pt")
|
||||
filtered = self.models.safety_checker(
|
||||
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 (to_bhwc(filtered),)
|
||||
return io.NodeOutput(to_bhwc(filtered))
|
||||
|
||||
+2
-2
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "comfyui-tooling-nodes"
|
||||
description = "Provides nodes and server API extensions geared towards using ComfyUI as a backend for external tools."
|
||||
version = "2.0.6"
|
||||
version = "3.1.4"
|
||||
license = { file = "LICENSE" }
|
||||
|
||||
[project.urls]
|
||||
@@ -13,7 +13,7 @@ line-length = 100
|
||||
preview = true
|
||||
|
||||
[tool.ruff.lint]
|
||||
ignore = ["E741"]
|
||||
ignore = ["E741", "BLE001"]
|
||||
|
||||
[tool.black]
|
||||
line-length = 100
|
||||
|
||||
@@ -8,6 +8,7 @@ import torch.nn.functional as F
|
||||
import math
|
||||
from torch import Tensor, Size
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
from comfy_api.latest import io
|
||||
|
||||
|
||||
def downsample_mask(mask: Tensor, batch: int, target_size: int, original_shape: Size) -> Tensor:
|
||||
@@ -65,104 +66,112 @@ class Region(NamedTuple):
|
||||
return result
|
||||
|
||||
|
||||
class BackgroundRegion:
|
||||
Regions = io.Custom("Regions")
|
||||
|
||||
|
||||
class BackgroundRegion(io.ComfyNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {"required": {"conditioning": ("CONDITIONING",)}}
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="ETN_BackgroundRegion",
|
||||
display_name="Background Region",
|
||||
category="external_tooling/regions",
|
||||
inputs=[io.Conditioning.Input("conditioning")],
|
||||
outputs=[Regions.Output(display_name="regions")],
|
||||
)
|
||||
|
||||
CATEGORY = "external_tooling/regions"
|
||||
RETURN_TYPES = ("REGIONS",)
|
||||
FUNCTION = "define"
|
||||
|
||||
def define(self, conditioning: list):
|
||||
@classmethod
|
||||
def execute(cls, conditioning: list):
|
||||
return (Region(None, None, conditioning),)
|
||||
|
||||
|
||||
class DefineRegion:
|
||||
class DefineRegion(io.ComfyNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"mask": ("MASK",),
|
||||
"conditioning": ("CONDITIONING",),
|
||||
},
|
||||
"optional": {
|
||||
"regions": ("REGIONS",),
|
||||
},
|
||||
}
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="ETN_DefineRegion",
|
||||
display_name="Define Region",
|
||||
category="external_tooling/regions",
|
||||
inputs=[
|
||||
io.Mask.Input("mask"),
|
||||
io.Conditioning.Input("conditioning"),
|
||||
Regions.Input("regions", optional=True),
|
||||
],
|
||||
outputs=[Regions.Output(display_name="regions")],
|
||||
)
|
||||
|
||||
CATEGORY = "external_tooling/regions"
|
||||
RETURN_TYPES = ("REGIONS",)
|
||||
FUNCTION = "define"
|
||||
|
||||
def define(self, mask: Tensor, conditioning: list, regions: Region | None = None):
|
||||
@classmethod
|
||||
def execute(cls, mask: Tensor, conditioning: list, regions: Region | None = None):
|
||||
if mask.dim() < 3:
|
||||
mask = mask.unsqueeze(0)
|
||||
return (Region(regions, mask, conditioning),)
|
||||
return io.NodeOutput(Region(regions, mask, conditioning))
|
||||
|
||||
|
||||
class ListRegionMasks:
|
||||
class ListRegionMasks(io.ComfyNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {"required": {"regions": ("REGIONS",)}}
|
||||
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")],
|
||||
)
|
||||
|
||||
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
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"regions": ("REGIONS",),
|
||||
}
|
||||
}
|
||||
def execute(cls, regions: Region):
|
||||
return io.NodeOutput(torch.stack([r.mask for r in regions.preprocess()], dim=0))
|
||||
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
FUNCTION = "attention_mask"
|
||||
CATEGORY = "external_tooling/regions"
|
||||
|
||||
mask: Tensor
|
||||
conds: list[Tensor]
|
||||
batch_size: int
|
||||
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")],
|
||||
)
|
||||
|
||||
def attention_mask(self, model: ModelPatcher, regions: Region):
|
||||
new_model = model.clone()
|
||||
region_list = regions.preprocess()
|
||||
num_conds = len(region_list)
|
||||
@classmethod
|
||||
def execute(cls, model: ModelPatcher, regions: Region):
|
||||
return io.NodeOutput(AttentionMaskPatch.apply(model, regions))
|
||||
|
||||
|
||||
class AttentionMaskPatch:
|
||||
def __init__(self, region_list: list[Region]):
|
||||
mask = torch.stack([r.mask for r in region_list], dim=0)
|
||||
mask_sum = mask.sum(dim=0, keepdim=True)
|
||||
assert mask_sum.sum() > 0, "There are areas that are zero in all masks."
|
||||
self.mask = mask / mask_sum
|
||||
|
||||
self.conds = [r.conditioning[0][0] for r in region_list]
|
||||
num_tokens = [cond.shape[1] for cond in self.conds]
|
||||
self.num_tokens = [cond.shape[1] for cond in self.conds]
|
||||
self.num_conds = len(region_list)
|
||||
self.batch_size = 0
|
||||
|
||||
@staticmethod
|
||||
def apply(model: ModelPatcher, regions: Region):
|
||||
patch = AttentionMaskPatch(regions.preprocess())
|
||||
|
||||
def attn2_patch(q: Tensor, k: Tensor, v: Tensor, extra_options: dict):
|
||||
assert k.mean() == v.mean(), "k and v must be the same."
|
||||
device, dtype = q.device, q.dtype
|
||||
|
||||
if self.conds[0].device != device or self.conds[0].dtype != dtype:
|
||||
self.conds = [cond.to(device, dtype=dtype) for cond in self.conds]
|
||||
if self.mask.device != device or self.mask.dtype != dtype:
|
||||
self.mask = self.mask.to(device, dtype=dtype)
|
||||
if patch.conds[0].device != device or patch.conds[0].dtype != dtype:
|
||||
patch.conds = [cond.to(device, dtype=dtype) for cond in patch.conds]
|
||||
if patch.mask.device != device or patch.mask.dtype != dtype:
|
||||
patch.mask = patch.mask.to(device, dtype=dtype)
|
||||
|
||||
cond_or_unconds = extra_options["cond_or_uncond"]
|
||||
num_chunks = len(cond_or_unconds)
|
||||
self.batch_size = q.shape[0] // num_chunks
|
||||
patch.batch_size = q.shape[0] // num_chunks
|
||||
q_chunks = q.chunk(num_chunks, dim=0)
|
||||
k_chunks = k.chunk(num_chunks, dim=0)
|
||||
lcm_tokens = lcm_for_list(num_tokens + [k.shape[1]])
|
||||
lcm_tokens = lcm_for_list(patch.num_tokens + [k.shape[1]])
|
||||
conds_tensor = [
|
||||
cond.repeat(self.batch_size, lcm_tokens // num_tokens[i], 1)
|
||||
for i, cond in enumerate(self.conds)
|
||||
cond.repeat(patch.batch_size, lcm_tokens // patch.num_tokens[i], 1)
|
||||
for i, cond in enumerate(patch.conds)
|
||||
]
|
||||
conds_tensor = torch.cat(conds_tensor, dim=0)
|
||||
|
||||
@@ -173,9 +182,9 @@ class AttentionMask:
|
||||
qs.insert(0, q_chunks[i])
|
||||
ks.insert(0, k_target)
|
||||
else:
|
||||
qs.insert(0, q_chunks[i].repeat(num_conds, 1, 1))
|
||||
qs.insert(0, q_chunks[i].repeat(patch.num_conds, 1, 1))
|
||||
ks.insert(0, conds_tensor)
|
||||
for _ in range(num_conds - 1):
|
||||
for _ in range(patch.num_conds - 1):
|
||||
cond_or_unconds.insert(i, 0)
|
||||
|
||||
qs = torch.cat(qs, dim=0)
|
||||
@@ -183,29 +192,32 @@ class AttentionMask:
|
||||
return qs, ks, ks
|
||||
|
||||
def attn2_output_patch(out: Tensor, extra_options: dict):
|
||||
num_conds = patch.num_conds
|
||||
cond_or_unconds = extra_options["cond_or_uncond"]
|
||||
mask_downsample = downsample_mask(
|
||||
self.mask, self.batch_size, out.shape[1], extra_options["original_shape"]
|
||||
patch.mask, patch.batch_size, out.shape[1], extra_options["original_shape"]
|
||||
)
|
||||
outputs: list[Tensor] = []
|
||||
pos = 0
|
||||
i = 0
|
||||
while i < len(cond_or_unconds):
|
||||
if cond_or_unconds[i] == 1: # uncond
|
||||
outputs.append(out[pos : pos + self.batch_size])
|
||||
pos += self.batch_size
|
||||
outputs.append(out[pos : pos + patch.batch_size])
|
||||
pos += patch.batch_size
|
||||
else:
|
||||
masked = out[pos : pos + num_conds * self.batch_size] * mask_downsample
|
||||
masked = masked.view(num_conds, self.batch_size, out.shape[1], out.shape[2])
|
||||
masked = out[pos : pos + num_conds * patch.batch_size] * mask_downsample
|
||||
masked = masked.view(num_conds, patch.batch_size, out.shape[1], out.shape[2])
|
||||
masked = masked.sum(dim=0)
|
||||
outputs.append(masked)
|
||||
pos += num_conds * self.batch_size
|
||||
pos += num_conds * patch.batch_size
|
||||
for _ in range(num_conds - 1):
|
||||
cond_or_unconds.pop(i)
|
||||
i += 1
|
||||
|
||||
return torch.cat(outputs, dim=0)
|
||||
|
||||
new_model = model.clone()
|
||||
new_model.set_model_attn2_patch(attn2_patch)
|
||||
new_model.set_model_attn2_output_patch(attn2_output_patch)
|
||||
return (new_model,)
|
||||
new_model.set_attachments("etn_attention_mask", patch)
|
||||
return new_model
|
||||
|
||||
@@ -3,49 +3,29 @@ import numpy as np
|
||||
import numpy.typing as npt
|
||||
import torch
|
||||
from torch import Tensor
|
||||
from comfy_api.latest import io
|
||||
|
||||
IntArray = npt.NDArray[np.int_]
|
||||
|
||||
|
||||
class TileLayout:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"min_tile_size": ("INT", {"default": 512, "min": 64, "max": 8192, "step": 8}),
|
||||
"padding": ("INT", {"default": 32, "min": 0, "max": 8192, "step": 8}),
|
||||
"blending": ("INT", {"default": 8, "min": 0, "max": 256, "step": 8}),
|
||||
}
|
||||
}
|
||||
|
||||
CATEGORY = "external_tooling/tiles"
|
||||
RETURN_TYPES = ("TILE_LAYOUT",)
|
||||
FUNCTION = "node"
|
||||
|
||||
image_size: IntArray
|
||||
tile_size: IntArray
|
||||
padding: int
|
||||
blending: int
|
||||
tile_count: IntArray
|
||||
|
||||
def node(self, image: Tensor, min_tile_size: int, padding: int, blending: int):
|
||||
self.init(image, min_tile_size, padding, blending)
|
||||
return (self,)
|
||||
|
||||
def init(self, image: Tensor, min_tile_size: int, padding: int, blending: int):
|
||||
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"
|
||||
def __init__(
|
||||
self, image: Tensor, min_tile_size: int, padding: int, blending: int, multiple: int
|
||||
):
|
||||
assert all([x % multiple == 0 for x in image.shape[-3:-1]]), (
|
||||
"Image size must be divisible by multiple"
|
||||
)
|
||||
assert min_tile_size % multiple == 0, "Tile size must be divisible by multiple"
|
||||
assert blending <= padding, "Blending must be smaller than padding"
|
||||
|
||||
self.image_size = np.array(image.shape[-3:-1])
|
||||
self.padding = padding
|
||||
self.blending = blending
|
||||
self.tile_count = np.maximum(1, self.image_size // (min_tile_size - 2 * padding))
|
||||
self.image_size: IntArray = np.array(image.shape[-3:-1])
|
||||
self.padding: int = padding
|
||||
self.blending: int = blending
|
||||
self.tile_count: IntArray = np.maximum(1, self.image_size // (min_tile_size - 2 * padding))
|
||||
|
||||
image_size_with_overlap = self.image_size + (self.tile_count - 1) * 2 * padding
|
||||
tile_size = np.ceil(image_size_with_overlap / self.tile_count)
|
||||
self.tile_size = (np.ceil(tile_size / 8) * 8).astype(int)
|
||||
self.tile_size: IntArray = (np.ceil(tile_size / multiple) * multiple).astype(int)
|
||||
|
||||
def size(self, coord: IntArray):
|
||||
return self.end(coord) - self.start(coord)
|
||||
@@ -96,80 +76,109 @@ class TileLayout:
|
||||
image[rect] = (1 - mask) * image[rect] + mask * tile
|
||||
|
||||
|
||||
class ExtractImageTile:
|
||||
class CreateTileLayout(io.ComfyNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"layout": ("TILE_LAYOUT",),
|
||||
"index": ("INT", {"min": 0}),
|
||||
}
|
||||
}
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="ETN_TileLayout",
|
||||
display_name="Create Tile Layout",
|
||||
category="external_tooling/tiles",
|
||||
inputs=[
|
||||
io.Image.Input("image"),
|
||||
io.Int.Input("min_tile_size", default=512, min=64, max=8192, step=8),
|
||||
io.Int.Input("padding", default=32, min=0, max=8192, step=8),
|
||||
io.Int.Input("blending", default=8, min=0, max=256, step=8),
|
||||
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
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"mask": ("MASK",),
|
||||
"layout": ("TILE_LAYOUT",),
|
||||
"index": ("INT", {"min": 0}),
|
||||
}
|
||||
}
|
||||
def execute(cls, image: Tensor, min_tile_size: int, padding: int, blending: int, multiple: int):
|
||||
return io.NodeOutput(TileLayout(image, min_tile_size, padding, blending, multiple))
|
||||
|
||||
CATEGORY = "external_tooling/tiles"
|
||||
RETURN_TYPES = ("MASK",)
|
||||
FUNCTION = "slice"
|
||||
|
||||
def slice(self, mask: Tensor, layout: TileLayout, index: int):
|
||||
class ExtractImageTile(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="ETN_ExtractImageTile",
|
||||
display_name="Extract Image Tile",
|
||||
category="external_tooling/tiles",
|
||||
inputs=[
|
||||
io.Image.Input("image"),
|
||||
io.Custom("TileLayout").Input("layout"),
|
||||
io.Int.Input("index", default=0, min=0),
|
||||
],
|
||||
outputs=[io.Image.Output(display_name="tile")],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, image: Tensor, layout: TileLayout, index: int):
|
||||
return io.NodeOutput(layout.tile(image, index))
|
||||
|
||||
|
||||
class ExtractMaskTile(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="ETN_ExtractMaskTile",
|
||||
display_name="Extract Mask Tile",
|
||||
category="external_tooling/tiles",
|
||||
inputs=[
|
||||
io.Mask.Input("mask"),
|
||||
io.Custom("TileLayout").Input("layout"),
|
||||
io.Int.Input("index", default=0, min=0),
|
||||
],
|
||||
outputs=[io.Mask.Output(display_name="tile")],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, mask: Tensor, layout: TileLayout, index: int):
|
||||
tile = layout.tile(mask.unsqueeze(3), index)
|
||||
return (tile.squeeze(3),)
|
||||
return io.NodeOutput(tile.squeeze(3))
|
||||
|
||||
|
||||
class GenerateTileMask:
|
||||
class GenerateTileMask(io.ComfyNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {"layout": ("TILE_LAYOUT",), "index": ("INT", {"min": 0})},
|
||||
"optional": {"blend": ("BOOLEAN",)},
|
||||
}
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="ETN_GenerateTileMask",
|
||||
display_name="Generate Tile Mask",
|
||||
category="external_tooling/tiles",
|
||||
inputs=[
|
||||
io.Custom("TileLayout").Input("layout"),
|
||||
io.Int.Input("index", default=0, min=0),
|
||||
io.Boolean.Input("blend", default=False, optional=True),
|
||||
],
|
||||
outputs=[io.Mask.Output(display_name="mask")],
|
||||
)
|
||||
|
||||
CATEGORY = "external_tooling/tiles"
|
||||
RETURN_TYPES = ("MASK",)
|
||||
FUNCTION = "generate"
|
||||
|
||||
def generate(self, layout: TileLayout, index: int, blend: bool = False):
|
||||
return (layout.mask(layout.coord(index), blend=blend),)
|
||||
|
||||
|
||||
class MergeImageTile:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"layout": ("TILE_LAYOUT",),
|
||||
"index": ("INT", {"min": 0}),
|
||||
"tile": ("IMAGE",),
|
||||
}
|
||||
}
|
||||
def execute(cls, layout: TileLayout, index: int, blend: bool = False):
|
||||
return io.NodeOutput(layout.mask(layout.coord(index), blend=blend))
|
||||
|
||||
CATEGORY = "external_tooling/tiles"
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "merge"
|
||||
|
||||
def merge(self, image: Tensor, layout: TileLayout, index: int, tile: Tensor):
|
||||
class MergeImageTile(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="ETN_MergeImageTile",
|
||||
display_name="Merge Image Tile",
|
||||
category="external_tooling/tiles",
|
||||
inputs=[
|
||||
io.Image.Input("image"),
|
||||
io.Custom("TileLayout").Input("layout"),
|
||||
io.Int.Input("index", default=0, min=0),
|
||||
io.Image.Input("tile"),
|
||||
],
|
||||
outputs=[io.Image.Output(display_name="image")],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, image: Tensor, layout: TileLayout, index: int, tile: Tensor):
|
||||
assert index < layout.total_count, f"Index {index} out of range"
|
||||
if index == 0:
|
||||
image = image.clone()
|
||||
layout.merge(image, index, tile)
|
||||
return (image,)
|
||||
return io.NodeOutput(image)
|
||||
|
||||
+15
-11
@@ -10,6 +10,7 @@ from __future__ import annotations
|
||||
import re
|
||||
from functools import cache
|
||||
from typing import NamedTuple
|
||||
from comfy_api.latest import io
|
||||
|
||||
|
||||
@cache
|
||||
@@ -43,7 +44,7 @@ def translate_chunk(text: str, language: str):
|
||||
(p for p in available if p.from_code == language and p.to_code == target), None
|
||||
)
|
||||
assert pkg, f"Couldn't find package for translation from {language}"
|
||||
print("Downloading and installing translation package", pkg)
|
||||
# print("Downloading and installing translation package", pkg) # this will cause encoding errors
|
||||
pkg.install()
|
||||
|
||||
text, embeddings = _extract_embeddings(text)
|
||||
@@ -61,17 +62,20 @@ def translate(text: str):
|
||||
return " ".join(translate_chunk(c.text, c.lang) for c in chunks)
|
||||
|
||||
|
||||
class Translate:
|
||||
@staticmethod
|
||||
def INPUT_TYPES():
|
||||
return {"required": {"text": ("STRING", {"multiline": True})}}
|
||||
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")],
|
||||
)
|
||||
|
||||
CATEGORY = "external_tooling"
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "translate"
|
||||
|
||||
def translate(self, text: str):
|
||||
return (translate(text),)
|
||||
@classmethod
|
||||
def execute(cls, text: str):
|
||||
return io.NodeOutput(translate(text))
|
||||
|
||||
|
||||
_lang_regex = re.compile(r"(lang:\w\w)")
|
||||
|
||||
Reference in New Issue
Block a user