Author SHA1 Message Date
Acly 20f8d8ecc9 Version 2.0.6 2025-10-11 19:17:02 +02:00
Acly f555efb71b API: fix qwen svdq models not having the quant field set 2025-10-05 12:54:04 +02:00
Acly db4f296533 API: map GGUF "qwen_image" arch to "qwen-image" to be consistent with safetensors models 2025-10-05 12:41:06 +02:00
Acly 17c36ebc70 Parameter node: avoid unhandled exception when input is not a widget 2025-10-05 11:15:12 +02:00
Acly 0697a3ac1f Add active output to KritaSelection node (true if there is a selection, false otherwise) 2025-09-06 23:33:07 +02:00
Acly fa46b93329 Model inspection: support Qwen and Nunchaku quants 2025-08-20 21:30:16 +02:00
Acly fa84eec8fc Model inspection: support more base models 2025-08-09 10:40:43 +02:00
Acly 5ef2fddc1b Version 2.0.3 2025-06-15 12:32:46 +02:00
Acly bff22b8351 Model inspection: don't filter keys for diffusion models (strips things like vpred and zsnr keys)
- this fixes sdxl-vpred diffusion models being detected as eps
2025-05-28 20:45:58 +02:00
Acly ca2b59248e Fix error when importing workflows that contain Parameter nodes connected to nodes that aren't installed 2025-05-22 15:46:42 +02:00
Acly 696899a5fc Allow to run workflows with parameter nodes
- fix validation error due to missing min/max
- fix errors due to numbers being passed as string
2025-05-12 15:51:42 +02:00
Acly 5f4373d71a Version 2.0.2 2025-04-28 10:11:42 +02:00
Acly 61d2a19120 Remove data:image/png;base64, prefix in base64 strings if present #39 2025-04-27 19:44:17 +02:00
Acly a6af76ac39 Fix default values of Parameter node not being editable #38 2025-04-27 17:21:13 +02:00
Acly 6a5c8e02e5 Fix error response when a model folder cannot be found 2025-04-27 09:51:59 +02:00
Acly c2308a0762 Don't run publish action on forks, close #114 2025-03-31 11:49:17 +02:00
Acly b8e4659a10 Add optional mask output for Krita Image Layer node 2025-03-01 20:20:52 +01:00
Acly 93e1932456 Don't return alpha channel from LoadImageBase64 inverted
... why was it ever inverted?
2025-03-01 20:20:24 +01:00
Acly ea755151fe Add Lumina 2 to known base models 2025-02-13 16:48:23 +01:00
Acly facd65995a Version 2.0.1, remove some unused code and imports 2025-02-02 20:17:51 +01:00
Acly d7b18203a4 Use comfy built-in any type (*) matching 2025-02-01 00:03:33 +01:00
Acly 1839c099ad Fix type matching BOOL -> BOOLEAN 2025-01-31 23:44:51 +01:00
Acly bed8b36705 Fix assertion when using tiling with padding=blending=0 2025-01-26 11:57:11 +01:00
Acly 8ed5591574 API breaking: removed is_refiner attribute from model inspection
- sdxl refiner is reported with base model "sdxl-refiner"
- added type attribute for sdxl model, allows to detect eps/v-prediction
2025-01-12 17:53:47 +01:00
Acly fe39d22eb9 Return a more specific error when inspect model folder doesn't exist 2024-12-07 19:49:00 +01:00
Acly 50d3479fba Code compatibility (match was no longer useful anyway) #31 2024-11-30 00:29:00 +01:00
Acly d7d421baaa Model detection: add support for (some) GGUF and Flux Inpaint models
- GGUF detection only works for converted models
2024-11-29 09:46:45 +01:00
Acly e10daee9ed Nodes for stacking and weighting reference images with flux redux model 2024-11-24 22:51:42 +01:00
Acly 50c3ffdf64 Parameter node: Fix type reset to default for connected widget on reload #29 2024-11-15 15:52:50 +01:00
Acly 517790d1d6 Parameter node: fix not being able to enter negative numbers for min/max 2024-11-11 11:47:53 +01:00
Acly e2bd09d7e9 Parameter node: Restrict initial type choice to avoid mismatch between type and default value before connecting the output 2024-10-29 13:03:53 +01:00
Acly 035c68c629 Parameter node: keep configured default values when reloading #25
- make sure default is changed if the node is reconnected to a non-matching type
2024-10-29 12:13:42 +01:00
Acly e86973fedf Remove image format parameter from Krita Output node 2024-10-28 14:56:01 +01:00
Acly 19337dcc0e Fix parameter node not being connectable #23 2024-10-28 13:11:33 +01:00
Acly 1d4ffe14bb Add a Send Text node
- converts any input to string and sends it as output (websocket message)
2024-10-27 10:43:43 +01:00
Acly 20c8039a98 Detect some diffusion models which have prefix like checkpoints 2024-10-25 13:08:41 +02:00
Acly fcf678735c New package version 2024-10-23 13:05:24 +02:00
Acly 63ab33800e Fix Parameter node type comparison for workflow validation 2024-10-21 15:11:01 +02:00
Acly ef5ccfa98f Fix Parameter node widget values being reset when switching or reloading workflows 2024-10-18 16:38:01 +02:00
Acly 0b01696f5b Fix workflow/unsubscribe endpoint 2024-10-14 13:20:41 +02:00
Acly 2fce4c56d5 Fix "Object of type _BasicTypes is not JSON serializable" 2024-10-12 15:37:38 +02:00
Acly a7f77032ec Make parameter nodes validate when executed 2024-10-11 21:46:49 +02:00
Acly 72335898cb Detect better parameter type defaults 2024-10-11 21:20:31 +02:00
10 changed files with 468 additions and 86 deletions
+2 -1
View File
@@ -11,11 +11,12 @@ jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
if: ${{ github.repository_owner == 'Acly' }}
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@main
uses: Comfy-Org/publish-node-action@v1
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+13 -7
View File
@@ -13,9 +13,9 @@ Provides nodes and API geared towards using ComfyUI as a backend for external to
## <a id="images" href="#toc">Sending and receiving images</a>
ComfyUI exchanges images via the filesystem. This requires a
multi-step process (upload images, prompt, download images), is rather
inefficient, and invites a whole class of potential issues. It's also unclear
at which point those images will get cleaned up if ComfyUI is used
multi-step process (upload images, prompt, download images), which
invites a whole class of potential issues you might not want to deal with.
It's also unclear at which point those images will get cleaned up if ComfyUI is used
via external tools.
### Load Image (Base64)
@@ -163,7 +163,7 @@ 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`
Supported model types: `checkpoints`, `diffusion_models`, `unet`, `unet_gguf`
#### Output
Lists available models with additional classification info:
@@ -172,14 +172,20 @@ Lists available models with additional classification info:
"checkpoint_file.safetensors": {
"base_model": "sd15",
"is_inpaint": false,
"is_refiner": false
"type": "eps"
},
...
}
```
Possible values for base model: `sd15, sd20, sd21, sd3, sdxl, ssd1b, svd, cascade-b, cascade-c, aura-flow, hunyuan-dit, flux, flux-schnell`
Possible values for base model: `sd15, sd20, sd21, sd3, sdxl, sdxl-refiner, ssd1b, svd, cascade-b, cascade-c, aura-flow, hunyuan-dit, flux, flux-schnell, lumina2, chroma, qwen-image`
The entry is `{"base_model": "unknown"}` for models which are not in safetensors format or do not match any of the known base models.
If base model is `sdxl`, the `type` attribute is set with possible values: `eps, edm, v-prediction, v-prediction-edm`
Detection supports quantized models:
* GGUF: if the `gguf` module is installed, .gguf files are detected and will set the `quant` field
* Nunchaku: SVDQuant models are detected and will set the `quant` field to `svdq`
Returns an entry `{"base_model": "unknown"}` for models with unknown format or which do not match any of the known base models.
### GET /api/etn/languages
+7 -1
View File
@@ -1,4 +1,4 @@
from . import api, nodes, tile, region, nsfw, translation, krita
from . import api as api, nodes, tile, region, nsfw, translation, krita
NODE_CLASS_MAPPINGS = {
"ETN_LoadImageBase64": nodes.LoadImageBase64,
@@ -6,6 +6,8 @@ NODE_CLASS_MAPPINGS = {
"ETN_SendImageWebSocket": nodes.SendImageWebSocket,
"ETN_CropImage": nodes.CropImage,
"ETN_ApplyMaskToImage": nodes.ApplyMaskToImage,
"ETN_ReferenceImage": nodes.ReferenceImage,
"ETN_ApplyReferenceImages": nodes.ApplyReferenceImages,
"ETN_TileLayout": tile.TileLayout,
"ETN_ExtractImageTile": tile.ExtractImageTile,
"ETN_ExtractMaskTile": tile.ExtractMaskTile,
@@ -18,6 +20,7 @@ NODE_CLASS_MAPPINGS = {
"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,
@@ -31,6 +34,8 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"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",
@@ -43,6 +48,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"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",
+141 -26
View File
@@ -6,8 +6,9 @@ import json
import traceback
import re
import logging
import itertools
from comfy import model_detection, supported_models
from comfy import model_detection
import comfy.utils
import folder_paths
import server
@@ -22,7 +23,7 @@ model_names = {
"SD20": "sd20",
"SD21UnclipL": "sd21",
"SD21UnclipH": "sd21",
"SDXLRefiner": "sdxl",
"SDXLRefiner": "sdxl-refiner",
"SDXL": "sdxl",
"SSD1B": "ssd1b",
"SVD_img2vid": "svd",
@@ -33,7 +34,30 @@ model_names = {
"HunyuanDiT": "hunyuan-dit",
"HunyuanDiT1": "hunyuan-dit",
"Flux": "flux",
"FluxInpaint": "flux",
"FluxSchnell": "flux-schnell",
"GenmoMochi": "mochi",
"LTXV": "ltxv",
"HunyuanVideo": "hunyuan-video",
"CosmosT2V": "cosmos",
"CosmosI2V": "cosmos",
"CosmosT2IPredict2": "cosmos-predict2",
"CosmosI2VPredict2": "cosmos-predict2",
"WAN21_T2V": "wan21",
"WAN21_I2V": "wan21",
"WAN21_FunControl2V": "wan21-fun",
"WAN21_Vace": "wan21-vace",
"WAN21_Camera": "wan21-camera",
"HiDream": "hi-dream",
"Chroma": "chroma",
"ACEStep": "ace-step",
"Omnigen2": "omnigen2",
"QwenImage": "qwen-image",
}
gguf_architectures = {
"sd1": "sd15",
"qwen_image": "qwen-image",
}
@@ -48,7 +72,7 @@ class FakeTensor(NamedTuple):
return d
def inspect_diffusion_model(filename: str, prefix: str | None, model_type: str):
def inspect_safetensors(filename: str, model_type: str, is_checkpoint: bool):
try:
# Read header of safetensors file
path = folder_paths.get_full_path(model_type, filename)
@@ -62,8 +86,10 @@ def inspect_diffusion_model(filename: str, prefix: str | None, model_type: str):
cfg[key] = FakeTensor.from_dict(cfg[key])
# Reuse Comfy's model detection
if prefix is None:
prefix = model_detection.unet_prefix_from_state_dict(cfg)
prefix = model_detection.unet_prefix_from_state_dict(cfg)
if not is_checkpoint:
cfg = comfy.utils.state_dict_prefix_replace(cfg, {prefix: ""}, filter_keys=False)
prefix = ""
try: # latest ComfyUI takes 2 args
unet_config = model_detection.detect_unet_config(cfg, prefix)
except TypeError as e: # older ComfyUI versions take 3 args
@@ -76,29 +102,118 @@ def inspect_diffusion_model(filename: str, prefix: str | None, model_type: str):
input_count = 4
# Find a matching base model depending on unet config
base_model = model_detection.model_config_from_unet_config(unet_config)
if base_model is None:
base_model = None
model_type = None
model_quant = None
# Check if it's a Nunchaku SVDQ model by inspecting metadata
raw_name = detect_svdq(cfg)
if raw_name:
model_quant = "svdq"
# Otherwise try ComfyUI's model detection
elif unet_config is not None:
base_model = model_detection.model_config_from_unet_config(unet_config)
if base_model:
raw_name = base_model.__class__.__name__
if raw_name == "SDXL":
model_type = base_model.model_type(cfg).name.lower().replace("_", "-")
if not raw_name:
return {"base_model": "unknown"}
base_model_class = base_model.__class__
base_model_name = model_names.get(base_model_class.__name__, "unknown")
return {
"base_model": base_model_name,
"is_inpaint": base_model_name in ["sd15", "sdxl"] and input_count > 4,
"is_refiner": base_model_class is supported_models.SDXLRefiner,
}
base_model_name = model_names.get(raw_name, "unknown")
result = {"base_model": base_model_name}
result["is_inpaint"] = (
base_model_name in ["sd15", "sdxl"] and input_count > 4
) or raw_name == "FluxInpaint"
if model_quant:
result["quant"] = model_quant
if model_type:
result["type"] = model_type
elif "T2I" in raw_name:
result["type"] = "t2i"
elif "I2V" in raw_name:
result["type"] = "i2v"
elif "T2V" in raw_name:
result["type"] = "t2v"
elif "Control2V" in raw_name:
result["type"] = "control2v"
return result
return {"base_model": "unknown"}
except Exception as e:
traceback.print_exc()
return {"base_model": "unknown", "error": f"Failed to detect base model: {e}"}
def detect_svdq(cfg: dict) -> str | None:
if md := cfg.get("__metadata__"):
if comfy_config := md.get("comfy_config"):
if isinstance(comfy_config, str):
comfy_config = json.loads(comfy_config)
return comfy_config.get("model_class")
model_class = md.get("model_class")
if model_class == "NunchakuFluxTransformer2dModel":
return "Flux"
if model_class == "NunchakuQwenImageTransformer2DModel":
return "QwenImage"
return None
def inspect_gguf(filename: str, model_type: str):
try:
import gguf
except ImportError:
return {"base_model": "unknown", "error": "GGUF module not found"}
try:
path = folder_paths.get_full_path(model_type, filename)
reader = gguf.GGUFReader(path)
arch_field = reader.get_field("general.architecture")
if arch_field is not None:
if len(arch_field.types) != 1 or arch_field.types[0] != gguf.GGUFValueType.STRING:
raise TypeError(
f"Bad type for GGUF general.architecture key: expected string, got {arch_field.types!r}"
)
arch_str = str(arch_field.parts[arch_field.data[-1]], encoding="utf-8")
else: # stable-diffusion.cpp, requires conversion. not handled for now
return {"base_model": "flux", "is_inpaint": False}
if arch_str == "flux" and any(
t.name.startswith("distilled_guidance_layer")
for t in itertools.islice(reader.tensors, 5)
):
arch_str = "chroma"
result = {
"base_model": gguf_architectures.get(arch_str, arch_str),
"is_inpaint": False,
}
try:
result["quant"] = reader.get_field("general.file_type").lower()
except Exception as e:
result["quant"] = "gguf"
return result
except Exception as e:
# traceback.print_exc()
return {"base_model": "unknown", "error": f"Failed to detect base model: {e}"}
def inspect_diffusion_model(filename: str, model_type: str, is_checkpoint: bool):
if filename.endswith(".gguf"):
return inspect_gguf(filename, model_type)
return inspect_safetensors(filename, model_type, is_checkpoint)
def inspect_models(model_type: str):
try:
prefix = "" if model_type in ("unet", "diffusion_models") else None
try:
files = folder_paths.get_filename_list(model_type)
except KeyError:
return web.json_response({"error": f"Model folder not found: {model_type}"})
is_checkpoint = model_type == "checkpoints"
info = {
filename: inspect_diffusion_model(filename, prefix, model_type)
for filename in folder_paths.get_filename_list(model_type)
filename: inspect_diffusion_model(filename, model_type, is_checkpoint)
for filename in files
}
return web.json_response(info)
except Exception as e:
@@ -130,13 +245,13 @@ def has_invalid_filename(filename: str):
_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")
if error := has_invalid_folder_name(folder_name):
error = has_invalid_folder_name(folder_name)
if error is not None:
return error
return inspect_models(folder_name)
@@ -144,10 +259,6 @@ if _server is not None:
async def api_model_info(request):
return inspect_models("checkpoints")
@_server.routes.get("/etn/model_info")
async def api_model_info(request):
return inspect_models("checkpoints")
@_server.routes.get("/api/etn/languages")
async def languages(request):
try:
@@ -169,11 +280,13 @@ if _server is not None:
@_server.routes.put("/api/etn/upload/{folder_name}/{filename}")
async def upload(request: web.Request):
folder_name = request.match_info.get("folder_name", "")
if error := has_invalid_folder_name(folder_name):
error = has_invalid_folder_name(folder_name)
if error is not None:
return error
filename = request.match_info.get("filename", "")
if error := has_invalid_filename(filename):
error = has_invalid_filename(filename)
if error is not None:
return error
try:
@@ -182,7 +295,9 @@ if _server is not None:
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")
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():
+65 -22
View File
@@ -72,36 +72,69 @@ const parameterTypes = {
"text": ["text", "prompt (positive)", "prompt (negative)"],
}
function changeWidget(widget, type, value, options) {
widget.type = type
widget.value = value
widget.options = options
function defaultParameterType(widgetType, connectedNode, connectedWidget) {
let paramType = parameterTypes[widgetType][0]
if (connectedNode.comfyClass === "CLIPTextEncode") {
paramType = "prompt (positive)"
}
if (connectedWidget.options?.round === 1) {
paramType = "number (integer)"
}
return paramType
}
function changeWidgets(node, type, value, options) {
function valueMatchesType(value, type, options) {
if (type === "number") {
return typeof value === "number"
} else if (type === "combo") {
return options?.values?.includes(value)
} else if (type === "toggle") {
return typeof value === "boolean"
}
return typeof value === "string"
}
function optionalWidgetValue(widgets, index, fallback) {
const result = widgets.length > index ? widgets[index].value : null
return result === null || result === 0 ? fallback : result
}
function changeWidgets(node, type, connectedNode, connectedWidget) {
if (type === "customtext") {
type = "text"
}
const options = connectedWidget.options
node.widgets[1].value = parameterTypes[type][0]
node.widgets[1].options = {values: parameterTypes[type]}
changeWidget(node.widgets[2], type, value, 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") {
changeWidget(node.widgets[3], "number", options?.min ?? 0, options)
changeWidget(node.widgets[4], "number", options?.max ?? 100, options)
} else {
changeWidget(node.widgets[3], "number", 0, {min: 0, max: 0})
changeWidget(node.widgets[4], "number", 0, {min: 0, max: 0})
node.addWidget("number", "min", oldMin, null, options)
node.addWidget("number", "max", oldMax, null, options)
}
}
function adaptWidgetsToConnection(node) {
if (!node.outputs || node.outputs.length === 0 || !node.outputs[0].links) {
if (!node.outputs || node.outputs.length === 0) {
return
}
const links = node.outputs[0].links
if (links.length === 1) {
if (links && links.length === 1) {
const link = node.graph.links[links[0]]
if (!link) return
@@ -109,7 +142,7 @@ function adaptWidgetsToConnection(node) {
if (!theirNode || !theirNode.inputs) return
const input = theirNode.inputs[link.target_slot]
if (!input) return
if (!input || !input.widget || theirNode.widgets === undefined) return
node.outputs[0].type = input.type
@@ -119,11 +152,15 @@ function adaptWidgetsToConnection(node) {
const widgetName = input.widget.name
const theirWidget = theirNode.widgets.find((w) => w.name === widgetName)
const widgetType = theirWidget.origType ?? theirWidget.type
changeWidgets(node, widgetType, theirWidget.value, theirWidget.options)
if (!theirWidget) return // connected to a custom node that isn't installed
} else if (links.length === 0) {
const widgetType = theirWidget.origType ?? theirWidget.type
changeWidgets(node, widgetType, theirNode, theirWidget)
} else if (!links || links.length === 0) {
node.outputs[0].type = "*"
node.widgets[1].value = "auto"
node.widgets[1].options = {values: ["auto"]}
}
}
@@ -170,9 +207,15 @@ app.registerExtension({
if (nodeData.name === "ETN_KritaCanvas") {
setIconImage(nodeType, canvasIcon, [200, 100], 0, 2)
} else if (nodeData.name === "ETN_KritaOutput") {
setIconImage(nodeType, outputIcon, [200, 120], 2, 0)
} else if (nodeData.name == "ETN_Parameter") {
setIconImage(nodeType, outputIcon, [200, 100], 1, 0)
} else if (nodeData.name === "ETN_Parameter") {
setupParameterNode(nodeType)
} else if (nodeData.name === "ETN_SendText") {
const onAdded = nodeType.prototype.onAdded
nodeType.prototype.onAdded = function() {
onAdded?.apply(this, arguments)
this.inputs[0].type = "*"
}
}
},
+82 -15
View File
@@ -1,11 +1,13 @@
import sys
import torch
import numpy as np
from pathlib import Path
from typing import NamedTuple
from typing import Any, NamedTuple
from PIL import Image
import server
import comfy.samplers
from comfy.comfy_types.node_typing import IO
from .nodes import SendImageWebSocket
@@ -34,8 +36,11 @@ class WorkflowExchange:
for publisher in self._publishers.values():
await self._notify(client_id, publisher)
def unsubscribe(self, client_id: str):
self._subscribers.remove(client_id)
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 = {
@@ -52,12 +57,67 @@ def _placeholder_image():
return torch.from_numpy(image)[None,]
class KritaOutput(SendImageWebSocket):
class _BasicTypes(str):
"""Matches IO.PRIMITIVE, but also any list of choices"""
basic_types = IO.PRIMITIVE.split(",") # STRING, FLOAT, INT, BOOLEAN
def __eq__(self, other):
return other in self.basic_types or isinstance(other, (list, _BasicTypes))
def __ne__(self, other):
return not self.__eq__(other)
BasicTypes = _BasicTypes("BASIC")
class KritaOutput:
@classmethod
def INPUT_TYPES(s):
return {"required": {"images": ("IMAGE",)}}
RETURN_TYPES = ()
FUNCTION = "send_images"
OUTPUT_NODE = True
CATEGORY = "krita"
def send_images(self, images):
return SendImageWebSocket().send_images(images, "PNG")
class KritaSendText:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"value": (IO.ANY, {}),
"name": ("STRING", {"default": "Output"}),
"type": (["text", "markdown", "html"], {"default": "text"}),
}
}
RETURN_TYPES = ()
FUNCTION = "send"
OUTPUT_NODE = True
CATEGORY = "krita"
def send(self, value: Any, name: str, type: str):
mime = {
"text": "text/plain",
"markdown": "text/markdown",
"html": "text/html",
}[type]
text = "None"
if value is not None:
try:
text = str(value)
except Exception as e:
text = f"Could not convert to text: {e}"
print(f"Sending text: {name} = {text}")
return {"ui": {"text": [{"name": name, "text": text, "content-type": mime}]}}
class KritaCanvas:
@classmethod
@@ -78,13 +138,13 @@ class KritaSelection:
def INPUT_TYPES(cls):
return {}
RETURN_TYPES = ("MASK",)
RETURN_NAMES = ("mask",)
RETURN_TYPES = (IO.MASK, IO.BOOLEAN)
RETURN_NAMES = ("mask", "active")
FUNCTION = "placeholder"
CATEGORY = "krita"
def placeholder(self):
return (torch.ones(1, 512, 512),)
return (torch.ones(1, 512, 512), False)
class KritaImageLayer:
@@ -96,13 +156,13 @@ class KritaImageLayer:
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
RETURN_TYPES = ("IMAGE", "MASK")
RETURN_NAMES = ("image", "mask")
FUNCTION = "placeholder"
CATEGORY = "krita"
def placeholder(self, name: str):
return (_placeholder_image(),)
return (_placeholder_image(), torch.ones(1, 512, 512))
class KritaMaskLayer:
@@ -133,6 +193,7 @@ _param_types = [
"prompt (positive)",
"prompt (negative)",
]
_any_float = {"default": 0.0, "min": -sys.float_info.max, "max": sys.float_info.max}
class Parameter:
@@ -143,17 +204,23 @@ class Parameter:
"name": ("STRING", {"default": "Parameter"}),
"type": (_param_types, {"default": "auto"}),
"default": ("STRING", {"default": ""}),
"min": ("FLOAT", {"default": 0.0}),
"max": ("FLOAT", {"default": 1.0}),
}
},
"optional": {
"min": ("FLOAT", _any_float),
"max": ("FLOAT", _any_float),
},
}
RETURN_TYPES = ("*",)
RETURN_TYPES = (BasicTypes,)
RETURN_NAMES = ("value",)
FUNCTION = "placeholder"
CATEGORY = "krita"
def placeholder(self, name: str, type: str, default, min, max):
def placeholder(self, name: str, type: str, default, min=0.0, max=1.0):
if type == "number":
return (float(default),)
elif type == "number (integer)":
return (int(default),)
return (default,)
+148 -11
View File
@@ -1,11 +1,17 @@
from __future__ import annotations
from copy import copy
from typing import NamedTuple
from PIL import Image
import numpy as np
import base64
import torch
import torch.nn.functional as F
from io import BytesIO
from server import PromptServer, BinaryEventTypes
from comfy.clip_vision import ClipVisionModel
from comfy.sd import StyleModel
class LoadImageBase64:
@classmethod
@@ -16,15 +22,16 @@ class LoadImageBase64:
CATEGORY = "external_tooling"
FUNCTION = "load_image"
def load_image(self, image):
def load_image(self, image: str):
_strip_prefix(image, "data:image/png;base64,")
imgdata = base64.b64decode(image)
img = Image.open(BytesIO(imgdata))
if "A" in img.getbands():
mask = np.array(img.getchannel("A")).astype(np.float32) / 255.0
mask = 1.0 - torch.from_numpy(mask)
mask = torch.from_numpy(mask)
else:
mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu")
mask = None
img = img.convert("RGB")
img = np.array(img).astype(np.float32) / 255.0
@@ -42,7 +49,8 @@ class LoadMaskBase64:
CATEGORY = "external_tooling"
FUNCTION = "load_mask"
def load_mask(self, mask):
def load_mask(self, mask: str):
_strip_prefix(mask, "data:image/png;base64,")
imgdata = base64.b64decode(mask)
img = Image.open(BytesIO(imgdata))
img = np.array(img).astype(np.float32) / 255.0
@@ -50,7 +58,8 @@ class LoadMaskBase64:
if img.dim() == 3: # RGB(A) input, use red channel
img = img[:, :, 0]
return (img.unsqueeze(0),)
class SendImageWebSocket:
@classmethod
def INPUT_TYPES(s):
@@ -78,12 +87,15 @@ class SendImageWebSocket:
[format, image, None],
server.client_id,
)
results.append(
{"source": "websocket", "content-type": f"image/{format.lower()}", "type": "output"}
)
results.append({
"source": "websocket",
"content-type": f"image/{format.lower()}",
"type": "output",
})
return {"ui": {"images": results}}
class CropImage:
"""Deprecated, ComfyUI has an ImageCrop node now which does the same."""
@@ -158,9 +170,9 @@ class ApplyMaskToImage:
assert mask.ndim == 3, f"Mask should have shape [B, H, W]. {mask.shape}"
assert out.ndim == 4, f"Image should have shape [B, C, H, W]. {out.shape}"
assert (
out.shape[-2:] == mask.shape[-2:]
), f"Image size {out.shape[-2:]} must match mask size {mask.shape[-2:]}"
assert out.shape[-2:] == mask.shape[-2:], (
f"Image size {out.shape[-2:]} must match mask size {mask.shape[-2:]}"
)
is_mask_batch = mask.shape[0] == out.shape[0]
# Apply each mask in the batch to its corresponding image's alpha channel
@@ -169,3 +181,128 @@ class ApplyMaskToImage:
out[i, 3, :, :] = alpha
return (to_bhwc(out),)
class _ReferenceImageData(NamedTuple):
image: torch.Tensor
weight: float
range: tuple[float, float]
class ReferenceImage:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"weight": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0}),
"range_start": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0}),
"range_end": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0}),
},
"optional": {
"reference_images": ("REFERENCE_IMAGE",),
},
}
CATEGORY = "external_tooling"
RETURN_TYPES = ("REFERENCE_IMAGE",)
RETURN_NAMES = ("reference_images",)
FUNCTION = "append"
def append(
self,
image: torch.Tensor,
weight: float,
range_start: float,
range_end: float,
reference_images: list[_ReferenceImageData] | None = None,
):
imgs = copy(reference_images) if reference_images is not None else []
imgs.append(_ReferenceImageData(image, weight, (range_start, range_end)))
return (imgs,)
class ApplyReferenceImages:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"conditioning": ("CONDITIONING",),
"clip_vision": ("CLIP_VISION",),
"style_model": ("STYLE_MODEL",),
"references": ("REFERENCE_IMAGE",),
}
}
CATEGORY = "external_tooling"
RETURN_TYPES = ("CONDITIONING",)
FUNCTION = "apply"
def apply(
self,
conditioning: list[list],
clip_vision: ClipVisionModel,
style_model: StyleModel,
references: list[_ReferenceImageData],
):
delimiters = {0.0, 1.0}
delimiters |= set(r.range[0] for r in references)
delimiters |= set(r.range[1] for r in references)
delimiters = sorted(delimiters)
ranges = [(delimiters[i], delimiters[i + 1]) for i in range(len(delimiters) - 1)]
embeds = [_encode_image(r.image, clip_vision, style_model, r.weight) for r in references]
base = conditioning[0][0]
result = []
for start, end in ranges:
e = [
embeds[i]
for i, r in enumerate(references)
if r.range[0] <= start and r.range[1] >= end
]
options = conditioning[0][1].copy()
options["start_percent"] = start
options["end_percent"] = end
result.append((torch.cat([base] + e, dim=1), options))
return (result,)
def _encode_image(
image: torch.Tensor, clip_vision: ClipVisionModel, style_model: StyleModel, weight: float
):
e = clip_vision.encode_image(image)
e = style_model.get_cond(e).flatten(start_dim=0, end_dim=1).unsqueeze(dim=0)
e = _downsample_image_cond(e, weight)
return e
def _downsample_image_cond(cond: torch.Tensor, weight: float):
if weight >= 1.0:
return cond
elif weight <= 0.0:
return torch.zeros_like(cond)
elif weight >= 0.6:
factor = 2
elif weight >= 0.3:
factor = 3
else:
factor = 4
# Downsample the clip vision embedding to make it smaller, resulting in less impact
# compared to other conditioning.
# See https://github.com/kaibioinfo/ComfyUI_AdvancedRefluxControl
(b, t, h) = cond.shape
m = int(np.sqrt(t))
cond = F.interpolate(
cond.view(b, m, m, h).transpose(1, -1),
size=(m // factor, m // factor),
mode="area",
)
return cond.transpose(1, -1).reshape(b, -1, h)
def _strip_prefix(s: str, prefix: str) -> str:
if s.startswith(prefix):
return s[len(prefix) :]
return s
+9 -1
View File
@@ -1,12 +1,20 @@
[project]
name = "comfyui-tooling-nodes"
description = "Provides nodes and server API extensions geared towards using ComfyUI as a backend for external tools."
version = "1.5.0"
version = "2.0.6"
license = { file = "LICENSE" }
[project.urls]
Repository = "https://github.com/Acly/comfyui-tooling-nodes"
[tool.ruff]
target-version = "py311"
line-length = 100
preview = true
[tool.ruff.lint]
ignore = ["E741"]
[tool.black]
line-length = 100
preview = true
-1
View File
@@ -115,7 +115,6 @@ class ListRegionMasks:
class AttentionMask:
@classmethod
def INPUT_TYPES(s):
return {
+1 -1
View File
@@ -36,7 +36,7 @@ class TileLayout:
def init(self, image: Tensor, min_tile_size: int, padding: int, blending: int):
assert all([x % 8 == 0 for x in image.shape[-3:-1]]), "Image size must be divisible by 8"
assert min_tile_size % 8 == 0, "Tile size must be divisible by 8"
assert blending < padding, "Blending must be smaller than padding"
assert blending <= padding, "Blending must be smaller than padding"
self.image_size = np.array(image.shape[-3:-1])
self.padding = padding