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:
|
publish-node:
|
||||||
name: Publish Custom Node to registry
|
name: Publish Custom Node to registry
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
|
if: ${{ github.repository_owner == 'Acly' }}
|
||||||
steps:
|
steps:
|
||||||
- name: Check out code
|
- name: Check out code
|
||||||
uses: actions/checkout@v4
|
uses: actions/checkout@v4
|
||||||
- name: Publish Custom Node
|
- name: Publish Custom Node
|
||||||
uses: Comfy-Org/publish-node-action@main
|
uses: Comfy-Org/publish-node-action@v1
|
||||||
with:
|
with:
|
||||||
## Add your own personal access token to your Github Repository secrets and reference it here.
|
## Add your own personal access token to your Github Repository secrets and reference it here.
|
||||||
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
||||||
@@ -13,9 +13,9 @@ Provides nodes and API geared towards using ComfyUI as a backend for external to
|
|||||||
## <a id="images" href="#toc">Sending and receiving images</a>
|
## <a id="images" href="#toc">Sending and receiving images</a>
|
||||||
|
|
||||||
ComfyUI exchanges images via the filesystem. This requires a
|
ComfyUI exchanges images via the filesystem. This requires a
|
||||||
multi-step process (upload images, prompt, download images), is rather
|
multi-step process (upload images, prompt, download images), which
|
||||||
inefficient, and invites a whole class of potential issues. It's also unclear
|
invites a whole class of potential issues you might not want to deal with.
|
||||||
at which point those images will get cleaned up if ComfyUI is used
|
It's also unclear at which point those images will get cleaned up if ComfyUI is used
|
||||||
via external tools.
|
via external tools.
|
||||||
|
|
||||||
### Load Image (Base64)
|
### Load Image (Base64)
|
||||||
@@ -163,7 +163,7 @@ There are various types of models that can be loaded as checkpoint, LoRA, Contro
|
|||||||
|
|
||||||
#### Paramters
|
#### Paramters
|
||||||
* `folder_name`: sub-directory in ComfyUI's models folder.
|
* `folder_name`: sub-directory in ComfyUI's models folder.
|
||||||
Supported model types: `checkpoints`, `diffusion_models`
|
Supported model types: `checkpoints`, `diffusion_models`, `unet`, `unet_gguf`
|
||||||
|
|
||||||
#### Output
|
#### Output
|
||||||
Lists available models with additional classification info:
|
Lists available models with additional classification info:
|
||||||
@@ -172,14 +172,20 @@ Lists available models with additional classification info:
|
|||||||
"checkpoint_file.safetensors": {
|
"checkpoint_file.safetensors": {
|
||||||
"base_model": "sd15",
|
"base_model": "sd15",
|
||||||
"is_inpaint": false,
|
"is_inpaint": false,
|
||||||
"is_refiner": false
|
"type": "eps"
|
||||||
},
|
},
|
||||||
...
|
...
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
Possible values for base model: `sd15, sd20, sd21, sd3, sdxl, ssd1b, svd, cascade-b, cascade-c, aura-flow, hunyuan-dit, flux, flux-schnell`
|
Possible values for base model: `sd15, sd20, sd21, sd3, sdxl, sdxl-refiner, ssd1b, svd, cascade-b, cascade-c, aura-flow, hunyuan-dit, flux, flux-schnell, 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
|
### 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 = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
"ETN_LoadImageBase64": nodes.LoadImageBase64,
|
"ETN_LoadImageBase64": nodes.LoadImageBase64,
|
||||||
@@ -6,6 +6,8 @@ NODE_CLASS_MAPPINGS = {
|
|||||||
"ETN_SendImageWebSocket": nodes.SendImageWebSocket,
|
"ETN_SendImageWebSocket": nodes.SendImageWebSocket,
|
||||||
"ETN_CropImage": nodes.CropImage,
|
"ETN_CropImage": nodes.CropImage,
|
||||||
"ETN_ApplyMaskToImage": nodes.ApplyMaskToImage,
|
"ETN_ApplyMaskToImage": nodes.ApplyMaskToImage,
|
||||||
|
"ETN_ReferenceImage": nodes.ReferenceImage,
|
||||||
|
"ETN_ApplyReferenceImages": nodes.ApplyReferenceImages,
|
||||||
"ETN_TileLayout": tile.TileLayout,
|
"ETN_TileLayout": tile.TileLayout,
|
||||||
"ETN_ExtractImageTile": tile.ExtractImageTile,
|
"ETN_ExtractImageTile": tile.ExtractImageTile,
|
||||||
"ETN_ExtractMaskTile": tile.ExtractMaskTile,
|
"ETN_ExtractMaskTile": tile.ExtractMaskTile,
|
||||||
@@ -18,6 +20,7 @@ NODE_CLASS_MAPPINGS = {
|
|||||||
"ETN_NSFWFilter": nsfw.NSFWFilter,
|
"ETN_NSFWFilter": nsfw.NSFWFilter,
|
||||||
"ETN_Translate": translation.Translate,
|
"ETN_Translate": translation.Translate,
|
||||||
"ETN_KritaOutput": krita.KritaOutput,
|
"ETN_KritaOutput": krita.KritaOutput,
|
||||||
|
"ETN_KritaSendText": krita.KritaSendText,
|
||||||
"ETN_KritaCanvas": krita.KritaCanvas,
|
"ETN_KritaCanvas": krita.KritaCanvas,
|
||||||
"ETN_KritaSelection": krita.KritaSelection,
|
"ETN_KritaSelection": krita.KritaSelection,
|
||||||
"ETN_KritaImageLayer": krita.KritaImageLayer,
|
"ETN_KritaImageLayer": krita.KritaImageLayer,
|
||||||
@@ -31,6 +34,8 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
|||||||
"ETN_SendImageWebSocket": "Send Image (WebSocket)",
|
"ETN_SendImageWebSocket": "Send Image (WebSocket)",
|
||||||
"ETN_CropImage": "Crop Image",
|
"ETN_CropImage": "Crop Image",
|
||||||
"ETN_ApplyMaskToImage": "Apply Mask to Image",
|
"ETN_ApplyMaskToImage": "Apply Mask to Image",
|
||||||
|
"ETN_ReferenceImage": "Reference Image",
|
||||||
|
"ETN_ApplyReferenceImages": "Apply Reference Images",
|
||||||
"ETN_TileLayout": "Create Tile Layout",
|
"ETN_TileLayout": "Create Tile Layout",
|
||||||
"ETN_ExtractImageTile": "Extract Image Tile",
|
"ETN_ExtractImageTile": "Extract Image Tile",
|
||||||
"ETN_ExtractMaskTile": "Extract Mask Tile",
|
"ETN_ExtractMaskTile": "Extract Mask Tile",
|
||||||
@@ -43,6 +48,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
|||||||
"ETN_NSFWFilter": "NSFW Filter",
|
"ETN_NSFWFilter": "NSFW Filter",
|
||||||
"ETN_Translate": "Translate Text",
|
"ETN_Translate": "Translate Text",
|
||||||
"ETN_KritaOutput": "Krita Output",
|
"ETN_KritaOutput": "Krita Output",
|
||||||
|
"ETN_KritaSendText": "Send Text",
|
||||||
"ETN_KritaCanvas": "Krita Canvas",
|
"ETN_KritaCanvas": "Krita Canvas",
|
||||||
"ETN_KritaSelection": "Krita Selection",
|
"ETN_KritaSelection": "Krita Selection",
|
||||||
"ETN_KritaImageLayer": "Krita Image Layer",
|
"ETN_KritaImageLayer": "Krita Image Layer",
|
||||||
|
|||||||
@@ -6,8 +6,9 @@ import json
|
|||||||
import traceback
|
import traceback
|
||||||
import re
|
import re
|
||||||
import logging
|
import logging
|
||||||
|
import itertools
|
||||||
|
|
||||||
from comfy import model_detection, supported_models
|
from comfy import model_detection
|
||||||
import comfy.utils
|
import comfy.utils
|
||||||
import folder_paths
|
import folder_paths
|
||||||
import server
|
import server
|
||||||
@@ -22,7 +23,7 @@ model_names = {
|
|||||||
"SD20": "sd20",
|
"SD20": "sd20",
|
||||||
"SD21UnclipL": "sd21",
|
"SD21UnclipL": "sd21",
|
||||||
"SD21UnclipH": "sd21",
|
"SD21UnclipH": "sd21",
|
||||||
"SDXLRefiner": "sdxl",
|
"SDXLRefiner": "sdxl-refiner",
|
||||||
"SDXL": "sdxl",
|
"SDXL": "sdxl",
|
||||||
"SSD1B": "ssd1b",
|
"SSD1B": "ssd1b",
|
||||||
"SVD_img2vid": "svd",
|
"SVD_img2vid": "svd",
|
||||||
@@ -33,7 +34,30 @@ model_names = {
|
|||||||
"HunyuanDiT": "hunyuan-dit",
|
"HunyuanDiT": "hunyuan-dit",
|
||||||
"HunyuanDiT1": "hunyuan-dit",
|
"HunyuanDiT1": "hunyuan-dit",
|
||||||
"Flux": "flux",
|
"Flux": "flux",
|
||||||
|
"FluxInpaint": "flux",
|
||||||
"FluxSchnell": "flux-schnell",
|
"FluxSchnell": "flux-schnell",
|
||||||
|
"GenmoMochi": "mochi",
|
||||||
|
"LTXV": "ltxv",
|
||||||
|
"HunyuanVideo": "hunyuan-video",
|
||||||
|
"CosmosT2V": "cosmos",
|
||||||
|
"CosmosI2V": "cosmos",
|
||||||
|
"CosmosT2IPredict2": "cosmos-predict2",
|
||||||
|
"CosmosI2VPredict2": "cosmos-predict2",
|
||||||
|
"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
|
return d
|
||||||
|
|
||||||
|
|
||||||
def inspect_diffusion_model(filename: str, prefix: str | None, model_type: str):
|
def inspect_safetensors(filename: str, model_type: str, is_checkpoint: bool):
|
||||||
try:
|
try:
|
||||||
# Read header of safetensors file
|
# Read header of safetensors file
|
||||||
path = folder_paths.get_full_path(model_type, filename)
|
path = folder_paths.get_full_path(model_type, filename)
|
||||||
@@ -62,8 +86,10 @@ def inspect_diffusion_model(filename: str, prefix: str | None, model_type: str):
|
|||||||
cfg[key] = FakeTensor.from_dict(cfg[key])
|
cfg[key] = FakeTensor.from_dict(cfg[key])
|
||||||
|
|
||||||
# Reuse Comfy's model detection
|
# Reuse Comfy's model detection
|
||||||
if prefix is None:
|
prefix = model_detection.unet_prefix_from_state_dict(cfg)
|
||||||
prefix = model_detection.unet_prefix_from_state_dict(cfg)
|
if not is_checkpoint:
|
||||||
|
cfg = comfy.utils.state_dict_prefix_replace(cfg, {prefix: ""}, filter_keys=False)
|
||||||
|
prefix = ""
|
||||||
try: # latest ComfyUI takes 2 args
|
try: # latest ComfyUI takes 2 args
|
||||||
unet_config = model_detection.detect_unet_config(cfg, prefix)
|
unet_config = model_detection.detect_unet_config(cfg, prefix)
|
||||||
except TypeError as e: # older ComfyUI versions take 3 args
|
except TypeError as e: # older ComfyUI versions take 3 args
|
||||||
@@ -76,29 +102,118 @@ def inspect_diffusion_model(filename: str, prefix: str | None, model_type: str):
|
|||||||
input_count = 4
|
input_count = 4
|
||||||
|
|
||||||
# Find a matching base model depending on unet config
|
# Find a matching base model depending on unet config
|
||||||
base_model = model_detection.model_config_from_unet_config(unet_config)
|
base_model = None
|
||||||
if base_model is None:
|
model_type = None
|
||||||
|
model_quant = None
|
||||||
|
|
||||||
|
# Check if it's a Nunchaku SVDQ model by inspecting metadata
|
||||||
|
raw_name = detect_svdq(cfg)
|
||||||
|
if raw_name:
|
||||||
|
model_quant = "svdq"
|
||||||
|
# Otherwise try ComfyUI's model detection
|
||||||
|
elif unet_config is not None:
|
||||||
|
base_model = model_detection.model_config_from_unet_config(unet_config)
|
||||||
|
if base_model:
|
||||||
|
raw_name = base_model.__class__.__name__
|
||||||
|
if raw_name == "SDXL":
|
||||||
|
model_type = base_model.model_type(cfg).name.lower().replace("_", "-")
|
||||||
|
|
||||||
|
if not raw_name:
|
||||||
return {"base_model": "unknown"}
|
return {"base_model": "unknown"}
|
||||||
|
|
||||||
base_model_class = base_model.__class__
|
base_model_name = model_names.get(raw_name, "unknown")
|
||||||
base_model_name = model_names.get(base_model_class.__name__, "unknown")
|
result = {"base_model": base_model_name}
|
||||||
return {
|
result["is_inpaint"] = (
|
||||||
"base_model": base_model_name,
|
base_model_name in ["sd15", "sdxl"] and input_count > 4
|
||||||
"is_inpaint": base_model_name in ["sd15", "sdxl"] and input_count > 4,
|
) or raw_name == "FluxInpaint"
|
||||||
"is_refiner": base_model_class is supported_models.SDXLRefiner,
|
if model_quant:
|
||||||
}
|
result["quant"] = model_quant
|
||||||
|
if model_type:
|
||||||
|
result["type"] = model_type
|
||||||
|
elif "T2I" in raw_name:
|
||||||
|
result["type"] = "t2i"
|
||||||
|
elif "I2V" in raw_name:
|
||||||
|
result["type"] = "i2v"
|
||||||
|
elif "T2V" in raw_name:
|
||||||
|
result["type"] = "t2v"
|
||||||
|
elif "Control2V" in raw_name:
|
||||||
|
result["type"] = "control2v"
|
||||||
|
return result
|
||||||
return {"base_model": "unknown"}
|
return {"base_model": "unknown"}
|
||||||
|
except Exception as e:
|
||||||
|
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:
|
except Exception as e:
|
||||||
# traceback.print_exc()
|
# traceback.print_exc()
|
||||||
return {"base_model": "unknown", "error": f"Failed to detect base model: {e}"}
|
return {"base_model": "unknown", "error": f"Failed to detect base model: {e}"}
|
||||||
|
|
||||||
|
|
||||||
|
def inspect_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):
|
def inspect_models(model_type: str):
|
||||||
try:
|
try:
|
||||||
prefix = "" if model_type in ("unet", "diffusion_models") else None
|
try:
|
||||||
|
files = folder_paths.get_filename_list(model_type)
|
||||||
|
except KeyError:
|
||||||
|
return web.json_response({"error": f"Model folder not found: {model_type}"})
|
||||||
|
is_checkpoint = model_type == "checkpoints"
|
||||||
info = {
|
info = {
|
||||||
filename: inspect_diffusion_model(filename, prefix, model_type)
|
filename: inspect_diffusion_model(filename, model_type, is_checkpoint)
|
||||||
for filename in folder_paths.get_filename_list(model_type)
|
for filename in files
|
||||||
}
|
}
|
||||||
return web.json_response(info)
|
return web.json_response(info)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -130,13 +245,13 @@ def has_invalid_filename(filename: str):
|
|||||||
|
|
||||||
_server: server.PromptServer | None = getattr(server.PromptServer, "instance", None)
|
_server: server.PromptServer | None = getattr(server.PromptServer, "instance", None)
|
||||||
if _server is not None:
|
if _server is not None:
|
||||||
|
|
||||||
_workflow_exchange = WorkflowExchange(_server)
|
_workflow_exchange = WorkflowExchange(_server)
|
||||||
|
|
||||||
@_server.routes.get("/api/etn/model_info/{folder_name}")
|
@_server.routes.get("/api/etn/model_info/{folder_name}")
|
||||||
async def model_info(request: web.Request):
|
async def model_info(request: web.Request):
|
||||||
folder_name = request.match_info.get("folder_name", "checkpoints")
|
folder_name = request.match_info.get("folder_name", "checkpoints")
|
||||||
if error := has_invalid_folder_name(folder_name):
|
error = has_invalid_folder_name(folder_name)
|
||||||
|
if error is not None:
|
||||||
return error
|
return error
|
||||||
return inspect_models(folder_name)
|
return inspect_models(folder_name)
|
||||||
|
|
||||||
@@ -144,10 +259,6 @@ if _server is not None:
|
|||||||
async def api_model_info(request):
|
async def api_model_info(request):
|
||||||
return inspect_models("checkpoints")
|
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")
|
@_server.routes.get("/api/etn/languages")
|
||||||
async def languages(request):
|
async def languages(request):
|
||||||
try:
|
try:
|
||||||
@@ -169,11 +280,13 @@ if _server is not None:
|
|||||||
@_server.routes.put("/api/etn/upload/{folder_name}/{filename}")
|
@_server.routes.put("/api/etn/upload/{folder_name}/{filename}")
|
||||||
async def upload(request: web.Request):
|
async def upload(request: web.Request):
|
||||||
folder_name = request.match_info.get("folder_name", "")
|
folder_name = request.match_info.get("folder_name", "")
|
||||||
if error := has_invalid_folder_name(folder_name):
|
error = has_invalid_folder_name(folder_name)
|
||||||
|
if error is not None:
|
||||||
return error
|
return error
|
||||||
|
|
||||||
filename = request.match_info.get("filename", "")
|
filename = request.match_info.get("filename", "")
|
||||||
if error := has_invalid_filename(filename):
|
error = has_invalid_filename(filename)
|
||||||
|
if error is not None:
|
||||||
return error
|
return error
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -182,7 +295,9 @@ if _server is not None:
|
|||||||
|
|
||||||
folder = Path(folder_paths.folder_names_and_paths[folder_name][0][0])
|
folder = Path(folder_paths.folder_names_and_paths[folder_name][0][0])
|
||||||
total_size = int(request.headers.get("Content-Length", "0"))
|
total_size = int(request.headers.get("Content-Length", "0"))
|
||||||
logging.info(f"Uploading {filename} ({total_size/(1024**2):.1f} MB) to {folder} folder")
|
logging.info(
|
||||||
|
f"Uploading {filename} ({total_size / (1024**2):.1f} MB) to {folder} folder"
|
||||||
|
)
|
||||||
|
|
||||||
with open(folder / filename, "wb") as f:
|
with open(folder / filename, "wb") as f:
|
||||||
async for chunk, _ in request.content.iter_chunks():
|
async for chunk, _ in request.content.iter_chunks():
|
||||||
|
|||||||
@@ -72,36 +72,69 @@ const parameterTypes = {
|
|||||||
"text": ["text", "prompt (positive)", "prompt (negative)"],
|
"text": ["text", "prompt (positive)", "prompt (negative)"],
|
||||||
}
|
}
|
||||||
|
|
||||||
function changeWidget(widget, type, value, options) {
|
function defaultParameterType(widgetType, connectedNode, connectedWidget) {
|
||||||
widget.type = type
|
let paramType = parameterTypes[widgetType][0]
|
||||||
widget.value = value
|
if (connectedNode.comfyClass === "CLIPTextEncode") {
|
||||||
widget.options = options
|
paramType = "prompt (positive)"
|
||||||
|
}
|
||||||
|
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") {
|
if (type === "customtext") {
|
||||||
type = "text"
|
type = "text"
|
||||||
}
|
}
|
||||||
|
const options = connectedWidget.options
|
||||||
|
|
||||||
node.widgets[1].value = parameterTypes[type][0]
|
const parameterTypeHint = node.widgets[1].value
|
||||||
node.widgets[1].options = {values: parameterTypes[type]}
|
const notSpecialized = node.widgets[1].options.values.includes("auto")
|
||||||
changeWidget(node.widgets[2], type, value, options)
|
const parameterTypeMismatch = !parameterTypes[type].includes(parameterTypeHint)
|
||||||
|
if (notSpecialized || parameterTypeMismatch) {
|
||||||
|
node.widgets[1].options = {values: parameterTypes[type]}
|
||||||
|
}
|
||||||
|
if (parameterTypeMismatch) {
|
||||||
|
node.widgets[1].value = defaultParameterType(type, connectedNode, connectedWidget)
|
||||||
|
}
|
||||||
|
const oldDefault = node.widgets.length > 2 ? node.widgets[2].value : connectedWidget.value
|
||||||
|
const oldMin = optionalWidgetValue(node.widgets, 3, options?.min ?? 0)
|
||||||
|
const oldMax = optionalWidgetValue(node.widgets, 4, options?.max ?? 100)
|
||||||
|
const isDefaultValid = valueMatchesType(oldDefault, type, connectedWidget.options)
|
||||||
|
while (node.widgets.length > 2) {
|
||||||
|
node.widgets.pop()
|
||||||
|
}
|
||||||
|
const value = isDefaultValid && oldDefault !== "" ? oldDefault : connectedWidget.value
|
||||||
|
node.addWidget(type, "default", value, null, options)
|
||||||
if (type === "number") {
|
if (type === "number") {
|
||||||
changeWidget(node.widgets[3], "number", options?.min ?? 0, options)
|
node.addWidget("number", "min", oldMin, null, options)
|
||||||
changeWidget(node.widgets[4], "number", options?.max ?? 100, options)
|
node.addWidget("number", "max", oldMax, null, options)
|
||||||
} else {
|
|
||||||
changeWidget(node.widgets[3], "number", 0, {min: 0, max: 0})
|
|
||||||
changeWidget(node.widgets[4], "number", 0, {min: 0, max: 0})
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
function adaptWidgetsToConnection(node) {
|
function adaptWidgetsToConnection(node) {
|
||||||
if (!node.outputs || node.outputs.length === 0 || !node.outputs[0].links) {
|
if (!node.outputs || node.outputs.length === 0) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
const links = node.outputs[0].links
|
const links = node.outputs[0].links
|
||||||
if (links.length === 1) {
|
if (links && links.length === 1) {
|
||||||
const link = node.graph.links[links[0]]
|
const link = node.graph.links[links[0]]
|
||||||
if (!link) return
|
if (!link) return
|
||||||
|
|
||||||
@@ -109,7 +142,7 @@ function adaptWidgetsToConnection(node) {
|
|||||||
if (!theirNode || !theirNode.inputs) return
|
if (!theirNode || !theirNode.inputs) return
|
||||||
|
|
||||||
const input = theirNode.inputs[link.target_slot]
|
const input = theirNode.inputs[link.target_slot]
|
||||||
if (!input) return
|
if (!input || !input.widget || theirNode.widgets === undefined) return
|
||||||
|
|
||||||
node.outputs[0].type = input.type
|
node.outputs[0].type = input.type
|
||||||
|
|
||||||
@@ -119,11 +152,15 @@ function adaptWidgetsToConnection(node) {
|
|||||||
|
|
||||||
const widgetName = input.widget.name
|
const widgetName = input.widget.name
|
||||||
const theirWidget = theirNode.widgets.find((w) => w.name === widgetName)
|
const theirWidget = theirNode.widgets.find((w) => w.name === widgetName)
|
||||||
const widgetType = theirWidget.origType ?? theirWidget.type
|
if (!theirWidget) return // connected to a custom node that isn't installed
|
||||||
changeWidgets(node, widgetType, theirWidget.value, theirWidget.options)
|
|
||||||
|
|
||||||
} else if (links.length === 0) {
|
const widgetType = theirWidget.origType ?? theirWidget.type
|
||||||
|
changeWidgets(node, widgetType, theirNode, theirWidget)
|
||||||
|
|
||||||
|
} else if (!links || links.length === 0) {
|
||||||
node.outputs[0].type = "*"
|
node.outputs[0].type = "*"
|
||||||
|
node.widgets[1].value = "auto"
|
||||||
|
node.widgets[1].options = {values: ["auto"]}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -170,9 +207,15 @@ app.registerExtension({
|
|||||||
if (nodeData.name === "ETN_KritaCanvas") {
|
if (nodeData.name === "ETN_KritaCanvas") {
|
||||||
setIconImage(nodeType, canvasIcon, [200, 100], 0, 2)
|
setIconImage(nodeType, canvasIcon, [200, 100], 0, 2)
|
||||||
} else if (nodeData.name === "ETN_KritaOutput") {
|
} else if (nodeData.name === "ETN_KritaOutput") {
|
||||||
setIconImage(nodeType, outputIcon, [200, 120], 2, 0)
|
setIconImage(nodeType, outputIcon, [200, 100], 1, 0)
|
||||||
} else if (nodeData.name == "ETN_Parameter") {
|
} else if (nodeData.name === "ETN_Parameter") {
|
||||||
setupParameterNode(nodeType)
|
setupParameterNode(nodeType)
|
||||||
|
} else if (nodeData.name === "ETN_SendText") {
|
||||||
|
const onAdded = nodeType.prototype.onAdded
|
||||||
|
nodeType.prototype.onAdded = function() {
|
||||||
|
onAdded?.apply(this, arguments)
|
||||||
|
this.inputs[0].type = "*"
|
||||||
|
}
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
|
|
||||||
|
|||||||
@@ -1,11 +1,13 @@
|
|||||||
|
import sys
|
||||||
import torch
|
import torch
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import NamedTuple
|
from typing import Any, NamedTuple
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
|
|
||||||
import server
|
import server
|
||||||
import comfy.samplers
|
import comfy.samplers
|
||||||
|
from comfy.comfy_types.node_typing import IO
|
||||||
from .nodes import SendImageWebSocket
|
from .nodes import SendImageWebSocket
|
||||||
|
|
||||||
|
|
||||||
@@ -34,8 +36,11 @@ class WorkflowExchange:
|
|||||||
for publisher in self._publishers.values():
|
for publisher in self._publishers.values():
|
||||||
await self._notify(client_id, publisher)
|
await self._notify(client_id, publisher)
|
||||||
|
|
||||||
def unsubscribe(self, client_id: str):
|
async def unsubscribe(self, client_id: str):
|
||||||
self._subscribers.remove(client_id)
|
if client_id in self._subscribers:
|
||||||
|
self._subscribers.remove(client_id)
|
||||||
|
else:
|
||||||
|
raise KeyError("No subscriber found with id " + client_id)
|
||||||
|
|
||||||
async def _notify(self, client_id: str, publisher: Publisher):
|
async def _notify(self, client_id: str, publisher: Publisher):
|
||||||
data = {
|
data = {
|
||||||
@@ -52,12 +57,67 @@ def _placeholder_image():
|
|||||||
return torch.from_numpy(image)[None,]
|
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 = ()
|
RETURN_TYPES = ()
|
||||||
FUNCTION = "send_images"
|
FUNCTION = "send_images"
|
||||||
OUTPUT_NODE = True
|
OUTPUT_NODE = True
|
||||||
CATEGORY = "krita"
|
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:
|
class KritaCanvas:
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -78,13 +138,13 @@ class KritaSelection:
|
|||||||
def INPUT_TYPES(cls):
|
def INPUT_TYPES(cls):
|
||||||
return {}
|
return {}
|
||||||
|
|
||||||
RETURN_TYPES = ("MASK",)
|
RETURN_TYPES = (IO.MASK, IO.BOOLEAN)
|
||||||
RETURN_NAMES = ("mask",)
|
RETURN_NAMES = ("mask", "active")
|
||||||
FUNCTION = "placeholder"
|
FUNCTION = "placeholder"
|
||||||
CATEGORY = "krita"
|
CATEGORY = "krita"
|
||||||
|
|
||||||
def placeholder(self):
|
def placeholder(self):
|
||||||
return (torch.ones(1, 512, 512),)
|
return (torch.ones(1, 512, 512), False)
|
||||||
|
|
||||||
|
|
||||||
class KritaImageLayer:
|
class KritaImageLayer:
|
||||||
@@ -96,13 +156,13 @@ class KritaImageLayer:
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ("IMAGE",)
|
RETURN_TYPES = ("IMAGE", "MASK")
|
||||||
RETURN_NAMES = ("image",)
|
RETURN_NAMES = ("image", "mask")
|
||||||
FUNCTION = "placeholder"
|
FUNCTION = "placeholder"
|
||||||
CATEGORY = "krita"
|
CATEGORY = "krita"
|
||||||
|
|
||||||
def placeholder(self, name: str):
|
def placeholder(self, name: str):
|
||||||
return (_placeholder_image(),)
|
return (_placeholder_image(), torch.ones(1, 512, 512))
|
||||||
|
|
||||||
|
|
||||||
class KritaMaskLayer:
|
class KritaMaskLayer:
|
||||||
@@ -133,6 +193,7 @@ _param_types = [
|
|||||||
"prompt (positive)",
|
"prompt (positive)",
|
||||||
"prompt (negative)",
|
"prompt (negative)",
|
||||||
]
|
]
|
||||||
|
_any_float = {"default": 0.0, "min": -sys.float_info.max, "max": sys.float_info.max}
|
||||||
|
|
||||||
|
|
||||||
class Parameter:
|
class Parameter:
|
||||||
@@ -143,17 +204,23 @@ class Parameter:
|
|||||||
"name": ("STRING", {"default": "Parameter"}),
|
"name": ("STRING", {"default": "Parameter"}),
|
||||||
"type": (_param_types, {"default": "auto"}),
|
"type": (_param_types, {"default": "auto"}),
|
||||||
"default": ("STRING", {"default": ""}),
|
"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",)
|
RETURN_NAMES = ("value",)
|
||||||
FUNCTION = "placeholder"
|
FUNCTION = "placeholder"
|
||||||
CATEGORY = "krita"
|
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,)
|
return (default,)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,11 +1,17 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
from copy import copy
|
||||||
|
from typing import NamedTuple
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import base64
|
import base64
|
||||||
import torch
|
import torch
|
||||||
|
import torch.nn.functional as F
|
||||||
from io import BytesIO
|
from io import BytesIO
|
||||||
from server import PromptServer, BinaryEventTypes
|
from server import PromptServer, BinaryEventTypes
|
||||||
|
|
||||||
|
from comfy.clip_vision import ClipVisionModel
|
||||||
|
from comfy.sd import StyleModel
|
||||||
|
|
||||||
|
|
||||||
class LoadImageBase64:
|
class LoadImageBase64:
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -16,15 +22,16 @@ class LoadImageBase64:
|
|||||||
CATEGORY = "external_tooling"
|
CATEGORY = "external_tooling"
|
||||||
FUNCTION = "load_image"
|
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)
|
imgdata = base64.b64decode(image)
|
||||||
img = Image.open(BytesIO(imgdata))
|
img = Image.open(BytesIO(imgdata))
|
||||||
|
|
||||||
if "A" in img.getbands():
|
if "A" in img.getbands():
|
||||||
mask = np.array(img.getchannel("A")).astype(np.float32) / 255.0
|
mask = np.array(img.getchannel("A")).astype(np.float32) / 255.0
|
||||||
mask = 1.0 - torch.from_numpy(mask)
|
mask = torch.from_numpy(mask)
|
||||||
else:
|
else:
|
||||||
mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu")
|
mask = None
|
||||||
|
|
||||||
img = img.convert("RGB")
|
img = img.convert("RGB")
|
||||||
img = np.array(img).astype(np.float32) / 255.0
|
img = np.array(img).astype(np.float32) / 255.0
|
||||||
@@ -42,7 +49,8 @@ class LoadMaskBase64:
|
|||||||
CATEGORY = "external_tooling"
|
CATEGORY = "external_tooling"
|
||||||
FUNCTION = "load_mask"
|
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)
|
imgdata = base64.b64decode(mask)
|
||||||
img = Image.open(BytesIO(imgdata))
|
img = Image.open(BytesIO(imgdata))
|
||||||
img = np.array(img).astype(np.float32) / 255.0
|
img = np.array(img).astype(np.float32) / 255.0
|
||||||
@@ -50,7 +58,8 @@ class LoadMaskBase64:
|
|||||||
if img.dim() == 3: # RGB(A) input, use red channel
|
if img.dim() == 3: # RGB(A) input, use red channel
|
||||||
img = img[:, :, 0]
|
img = img[:, :, 0]
|
||||||
return (img.unsqueeze(0),)
|
return (img.unsqueeze(0),)
|
||||||
|
|
||||||
|
|
||||||
class SendImageWebSocket:
|
class SendImageWebSocket:
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s):
|
||||||
@@ -78,12 +87,15 @@ class SendImageWebSocket:
|
|||||||
[format, image, None],
|
[format, image, None],
|
||||||
server.client_id,
|
server.client_id,
|
||||||
)
|
)
|
||||||
results.append(
|
results.append({
|
||||||
{"source": "websocket", "content-type": f"image/{format.lower()}", "type": "output"}
|
"source": "websocket",
|
||||||
)
|
"content-type": f"image/{format.lower()}",
|
||||||
|
"type": "output",
|
||||||
|
})
|
||||||
|
|
||||||
return {"ui": {"images": results}}
|
return {"ui": {"images": results}}
|
||||||
|
|
||||||
|
|
||||||
class CropImage:
|
class CropImage:
|
||||||
"""Deprecated, ComfyUI has an ImageCrop node now which does the same."""
|
"""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 mask.ndim == 3, f"Mask should have shape [B, H, W]. {mask.shape}"
|
||||||
assert out.ndim == 4, f"Image should have shape [B, C, H, W]. {out.shape}"
|
assert out.ndim == 4, f"Image should have shape [B, C, H, W]. {out.shape}"
|
||||||
assert (
|
assert out.shape[-2:] == mask.shape[-2:], (
|
||||||
out.shape[-2:] == mask.shape[-2:]
|
f"Image size {out.shape[-2:]} must match mask size {mask.shape[-2:]}"
|
||||||
), f"Image size {out.shape[-2:]} must match mask size {mask.shape[-2:]}"
|
)
|
||||||
is_mask_batch = mask.shape[0] == out.shape[0]
|
is_mask_batch = mask.shape[0] == out.shape[0]
|
||||||
|
|
||||||
# Apply each mask in the batch to its corresponding image's alpha channel
|
# Apply each mask in the batch to its corresponding image's alpha channel
|
||||||
@@ -169,3 +181,128 @@ class ApplyMaskToImage:
|
|||||||
out[i, 3, :, :] = alpha
|
out[i, 3, :, :] = alpha
|
||||||
|
|
||||||
return (to_bhwc(out),)
|
return (to_bhwc(out),)
|
||||||
|
|
||||||
|
|
||||||
|
class _ReferenceImageData(NamedTuple):
|
||||||
|
image: torch.Tensor
|
||||||
|
weight: float
|
||||||
|
range: tuple[float, float]
|
||||||
|
|
||||||
|
|
||||||
|
class ReferenceImage:
|
||||||
|
@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]
|
[project]
|
||||||
name = "comfyui-tooling-nodes"
|
name = "comfyui-tooling-nodes"
|
||||||
description = "Provides nodes and server API extensions geared towards using ComfyUI as a backend for external tools."
|
description = "Provides nodes and server API extensions geared towards using ComfyUI as a backend for external tools."
|
||||||
version = "1.5.0"
|
version = "2.0.6"
|
||||||
license = { file = "LICENSE" }
|
license = { file = "LICENSE" }
|
||||||
|
|
||||||
[project.urls]
|
[project.urls]
|
||||||
Repository = "https://github.com/Acly/comfyui-tooling-nodes"
|
Repository = "https://github.com/Acly/comfyui-tooling-nodes"
|
||||||
|
|
||||||
|
[tool.ruff]
|
||||||
|
target-version = "py311"
|
||||||
|
line-length = 100
|
||||||
|
preview = true
|
||||||
|
|
||||||
|
[tool.ruff.lint]
|
||||||
|
ignore = ["E741"]
|
||||||
|
|
||||||
[tool.black]
|
[tool.black]
|
||||||
line-length = 100
|
line-length = 100
|
||||||
preview = true
|
preview = true
|
||||||
|
|||||||
@@ -115,7 +115,6 @@ class ListRegionMasks:
|
|||||||
|
|
||||||
|
|
||||||
class AttentionMask:
|
class AttentionMask:
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s):
|
||||||
return {
|
return {
|
||||||
|
|||||||
@@ -36,7 +36,7 @@ class TileLayout:
|
|||||||
def init(self, image: Tensor, min_tile_size: int, padding: int, blending: int):
|
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 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 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.image_size = np.array(image.shape[-3:-1])
|
||||||
self.padding = padding
|
self.padding = padding
|
||||||
|
|||||||
Reference in New Issue
Block a user