Compare commits
43
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
20f8d8ecc9 | ||
|
|
f555efb71b | ||
|
|
db4f296533 | ||
|
|
17c36ebc70 | ||
|
|
0697a3ac1f | ||
|
|
fa46b93329 | ||
|
|
fa84eec8fc | ||
|
|
5ef2fddc1b | ||
|
|
bff22b8351 | ||
|
|
ca2b59248e | ||
|
|
696899a5fc | ||
|
|
5f4373d71a | ||
|
|
61d2a19120 | ||
|
|
a6af76ac39 | ||
|
|
6a5c8e02e5 | ||
|
|
c2308a0762 | ||
|
|
b8e4659a10 | ||
|
|
93e1932456 | ||
|
|
ea755151fe | ||
|
|
facd65995a | ||
|
|
d7b18203a4 | ||
|
|
1839c099ad | ||
|
|
bed8b36705 | ||
|
|
8ed5591574 | ||
|
|
fe39d22eb9 | ||
|
|
50d3479fba | ||
|
|
d7d421baaa | ||
|
|
e10daee9ed | ||
|
|
50c3ffdf64 | ||
|
|
517790d1d6 | ||
|
|
e2bd09d7e9 | ||
|
|
035c68c629 | ||
|
|
e86973fedf | ||
|
|
19337dcc0e | ||
|
|
1d4ffe14bb | ||
|
|
20c8039a98 | ||
|
|
fcf678735c | ||
|
|
63ab33800e | ||
|
|
ef5ccfa98f | ||
|
|
0b01696f5b | ||
|
|
2fce4c56d5 | ||
|
|
a7f77032ec | ||
|
|
72335898cb |
@@ -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,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
@@ -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",
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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 = "*"
|
||||
}
|
||||
}
|
||||
},
|
||||
|
||||
|
||||
@@ -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,)
|
||||
|
||||
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -115,7 +115,6 @@ class ListRegionMasks:
|
||||
|
||||
|
||||
class AttentionMask:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user