Compare commits
51
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e10faab458 | ||
|
|
bb5682aa6d | ||
|
|
59612fd811 | ||
|
|
30eb5b0091 | ||
|
|
1edc2cd10d | ||
|
|
fa3199be2b | ||
|
|
43d65ae68c | ||
|
|
dfd17f6d78 | ||
|
|
1070edd024 | ||
|
|
9f0ed85cc1 | ||
|
|
35622e3a5e | ||
|
|
644371e5b5 | ||
|
|
f3d468cfc2 | ||
|
|
6cd448b026 | ||
|
|
5951c90b10 | ||
|
|
6abac2e470 | ||
|
|
01c73e1c5e | ||
|
|
5060c56135 | ||
|
|
acc2d687d5 | ||
|
|
780c52f03a | ||
|
|
2fe0859476 | ||
|
|
1186239751 | ||
|
|
96a0da9dbd | ||
|
|
f9d2ebf91d | ||
|
|
1b7ae27cc1 | ||
|
|
e312b02ad2 | ||
|
|
63ee25d001 | ||
|
|
1caf7c18c3 | ||
|
|
349a8524c6 | ||
|
|
15330eab65 | ||
|
|
1571782d01 | ||
|
|
5b4030288d | ||
|
|
ab58c36212 | ||
|
|
5a0ef0dadd | ||
|
|
967e72fc66 | ||
|
|
78a86daaf7 | ||
|
|
bee3f47a14 | ||
|
|
2159395389 | ||
|
|
b11346aba8 | ||
|
|
30982fa488 | ||
|
|
92b79906cd | ||
|
|
76f365b5ee | ||
|
|
da67e766c2 | ||
|
|
49cea8d945 | ||
|
|
b1d74adb15 | ||
|
|
652ac3f3b9 | ||
|
|
060e733605 | ||
|
|
eedbb4bc65 | ||
|
|
fa2397585f | ||
|
|
77348c4adb | ||
|
|
0d0fb8e13a |
@@ -0,0 +1,20 @@
|
|||||||
|
name: 📦 Publish to Comfy registry
|
||||||
|
on:
|
||||||
|
workflow_dispatch:
|
||||||
|
push:
|
||||||
|
branches:
|
||||||
|
- main
|
||||||
|
paths:
|
||||||
|
- "pyproject.toml"
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
publish-node:
|
||||||
|
name: Publish Custom Node to registry
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
steps:
|
||||||
|
- name: ♻️ Check out code
|
||||||
|
uses: actions/checkout@v4
|
||||||
|
- name: 📦 Publish Custom Node
|
||||||
|
uses: Comfy-Org/publish-node-action@main
|
||||||
|
with:
|
||||||
|
personal_access_token: ${{ secrets.COMFY_REGISTRY_TOKEN }}
|
||||||
@@ -0,0 +1,8 @@
|
|||||||
|
default_language_version:
|
||||||
|
python: python3.10
|
||||||
|
repos:
|
||||||
|
- repo: https://github.com/melmass/hooks
|
||||||
|
rev: e8c6c18175ed4f6e30f23991de7989411e09c73b
|
||||||
|
hooks:
|
||||||
|
- id: fix-trailing-whitespace
|
||||||
|
- id: bump-version
|
||||||
+85
-27
@@ -6,21 +6,27 @@
|
|||||||
# Copyright (c) 2023 Mel Massadian
|
# Copyright (c) 2023 Mel Massadian
|
||||||
#
|
#
|
||||||
###
|
###
|
||||||
|
|
||||||
|
__version__ = "0.1.5"
|
||||||
|
|
||||||
import os
|
import os
|
||||||
|
|
||||||
# todo: don't override this if the user has that setup already
|
# TODO: don't override this if the user has that setup already
|
||||||
os.environ["TF_FORCE_GPU_ALLOW_GROWTH"] = "true"
|
if not os.environ.get("TF_FORCE_GPU_ALLOW_GROWTH"):
|
||||||
os.environ["TF_GPU_ALLOCATOR"] = "cuda_malloc_async"
|
os.environ["TF_FORCE_GPU_ALLOW_GROWTH"] = "true"
|
||||||
|
|
||||||
|
if not os.environ.get("TF_GPU_ALLOCATOR"):
|
||||||
|
os.environ["TF_GPU_ALLOCATOR"] = "cuda_malloc_async"
|
||||||
|
|
||||||
import ast
|
import ast
|
||||||
import contextlib
|
import contextlib
|
||||||
import importlib
|
import importlib
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
import os
|
|
||||||
import shutil
|
import shutil
|
||||||
import traceback
|
import traceback
|
||||||
from importlib import reload
|
from importlib import reload
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
from aiohttp import web
|
from aiohttp import web
|
||||||
from server import PromptServer
|
from server import PromptServer
|
||||||
@@ -36,14 +42,11 @@ NODE_DISPLAY_NAME_MAPPINGS = {}
|
|||||||
NODE_CLASS_MAPPINGS_DEBUG = {}
|
NODE_CLASS_MAPPINGS_DEBUG = {}
|
||||||
WEB_DIRECTORY = "./web"
|
WEB_DIRECTORY = "./web"
|
||||||
|
|
||||||
__version__ = "0.2.0"
|
|
||||||
|
|
||||||
|
def extract_nodes_from_source(filename: Path):
|
||||||
def extract_nodes_from_source(filename):
|
|
||||||
source_code = ""
|
source_code = ""
|
||||||
|
|
||||||
with open(filename, encoding="utf8") as file:
|
source_code = filename.read_text(encoding="utf-8")
|
||||||
source_code = file.read()
|
|
||||||
|
|
||||||
nodes = []
|
nodes = []
|
||||||
|
|
||||||
@@ -68,7 +71,7 @@ def extract_nodes_from_source(filename):
|
|||||||
|
|
||||||
|
|
||||||
def load_nodes():
|
def load_nodes():
|
||||||
errors = []
|
errors: list[str] = []
|
||||||
nodes = []
|
nodes = []
|
||||||
nodes_failed = []
|
nodes_failed = []
|
||||||
|
|
||||||
@@ -85,9 +88,11 @@ def load_nodes():
|
|||||||
log.debug(f"Imported {module_name} nodes")
|
log.debug(f"Imported {module_name} nodes")
|
||||||
|
|
||||||
except AttributeError:
|
except AttributeError:
|
||||||
|
log.debug(f"Skipping wip module {module_name}")
|
||||||
pass # wip nodes
|
pass # wip nodes
|
||||||
except Exception:
|
except Exception:
|
||||||
error_message = traceback.format_exc().splitlines()[-1]
|
error_message = traceback.format_exc().splitlines()[-1]
|
||||||
|
|
||||||
errors.append(
|
errors.append(
|
||||||
f"Failed to import module {module_name} because {error_message}"
|
f"Failed to import module {module_name} because {error_message}"
|
||||||
)
|
)
|
||||||
@@ -107,28 +112,81 @@ def load_nodes():
|
|||||||
|
|
||||||
|
|
||||||
# - REGISTER WEB EXTENSIONS
|
# - REGISTER WEB EXTENSIONS
|
||||||
web_extensions_root = comfy_dir / "web" / "extensions"
|
def uninstall_old_web_extensions():
|
||||||
web_mtb = web_extensions_root / "mtb"
|
web_extensions_root = comfy_dir / "web" / "extensions"
|
||||||
|
web_mtb = web_extensions_root / "mtb"
|
||||||
|
|
||||||
if web_mtb.exists() and hasattr(nodes, "EXTENSION_WEB_DIRS"):
|
if web_mtb.exists() and hasattr(nodes, "EXTENSION_WEB_DIRS"):
|
||||||
try:
|
try:
|
||||||
if web_mtb.is_symlink():
|
if web_mtb.is_symlink():
|
||||||
web_mtb.unlink()
|
web_mtb.unlink()
|
||||||
else:
|
else:
|
||||||
shutil.rmtree(web_mtb)
|
shutil.rmtree(web_mtb)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
log.warning(
|
log.warning(
|
||||||
f"Failed to remove web mtb directory: {e}\nPlease manually remove it from disk ({web_mtb}) and restart the server."
|
f"Failed to remove web mtb directory: {e}\nPlease manually remove it from disk ({web_mtb}) and restart the server."
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# uninstall_old_web_extensions()
|
||||||
|
|
||||||
|
|
||||||
|
# - GATHER WIKI PAGES
|
||||||
|
def wiki_to_classname(s: str):
|
||||||
|
wiki_name = s.replace("nodes-", "", 1)
|
||||||
|
return "MTB_" + "".join(
|
||||||
|
[part.capitalize() for part in wiki_name.split("-")]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def classname_to_wiki(s: str):
|
||||||
|
classname = s.replace("MTB_", "")
|
||||||
|
parts = []
|
||||||
|
start = 0
|
||||||
|
for i in range(1, len(classname)):
|
||||||
|
if classname[i].isupper():
|
||||||
|
parts.append(classname[start:i].lower())
|
||||||
|
start = i
|
||||||
|
parts.append(classname[start:].lower())
|
||||||
|
return "nodes-" + "-".join(parts)
|
||||||
|
|
||||||
|
|
||||||
|
wiki = here / "wiki"
|
||||||
|
node_docs = {}
|
||||||
|
if wiki.exists() and wiki.is_dir():
|
||||||
|
node_docs = {
|
||||||
|
wiki_to_classname(x.stem): x.read_text(encoding="utf-8")
|
||||||
|
for x in (wiki / "nodes").glob("*.md")
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
# - REGISTER NODES
|
# - REGISTER NODES
|
||||||
|
|
||||||
|
|
||||||
|
MTB_EXPORT = os.environ.get("MTB_EXPORT")
|
||||||
|
|
||||||
nodes, failed = load_nodes()
|
nodes, failed = load_nodes()
|
||||||
for node_class in nodes:
|
for node_class in nodes:
|
||||||
class_name = node_class.__name__
|
class_name: str = node_class.__name__
|
||||||
# fallback to __doc__
|
linked_doc = node_docs.get(class_name)
|
||||||
if not hasattr(node_class, "DESCRIPTION") and node_class.__doc__:
|
|
||||||
node_class.DESCRIPTION = node_class.__doc__
|
if not hasattr(node_class, "DESCRIPTION"):
|
||||||
|
if linked_doc:
|
||||||
|
log.debug(f"Found linked doc for {class_name}, using it")
|
||||||
|
node_class.DESCRIPTION = linked_doc
|
||||||
|
elif node_class.__doc__:
|
||||||
|
log.debug(f"Using __doc__ as description for {class_name}")
|
||||||
|
node_class.DESCRIPTION = node_class.__doc__
|
||||||
|
if MTB_EXPORT:
|
||||||
|
wiki_name = classname_to_wiki(class_name)
|
||||||
|
(wiki / "nodes" / wiki_name + ".md").write_text(
|
||||||
|
node_class.__doc__, encoding="utf-8"
|
||||||
|
)
|
||||||
|
|
||||||
|
else:
|
||||||
|
log.debug(
|
||||||
|
f"None of the methods could retrieve documentation for {class_name}"
|
||||||
|
)
|
||||||
|
|
||||||
node_label = f"{get_label(class_name)} (mtb)"
|
node_label = f"{get_label(class_name)} (mtb)"
|
||||||
NODE_CLASS_MAPPINGS[node_label] = node_class
|
NODE_CLASS_MAPPINGS[node_label] = node_class
|
||||||
@@ -279,7 +337,7 @@ if hasattr(PromptServer, "instance"):
|
|||||||
<a href="/mtb/manage">manage</a>
|
<a href="/mtb/manage">manage</a>
|
||||||
<a href="/mtb/debug">debug</a>
|
<a href="/mtb/debug">debug</a>
|
||||||
<a href="/mtb/status">status</a>
|
<a href="/mtb/status">status</a>
|
||||||
</div>
|
</div>
|
||||||
"""
|
"""
|
||||||
return web.Response(
|
return web.Response(
|
||||||
text=endpoint.render_base_template("MTB", html_response),
|
text=endpoint.render_base_template("MTB", html_response),
|
||||||
|
|||||||
+22
-19
@@ -1,19 +1,22 @@
|
|||||||
{
|
{
|
||||||
"$schema": "https://biomejs.dev/schemas/1.6.1/schema.json",
|
"$schema": "https://biomejs.dev/schemas/1.6.1/schema.json",
|
||||||
"organizeImports": {
|
"organizeImports": {
|
||||||
"enabled": true
|
"enabled": true
|
||||||
},
|
},
|
||||||
"linter": {
|
"linter": {
|
||||||
"enabled": true,
|
"enabled": true,
|
||||||
"rules": {
|
"rules": {
|
||||||
"recommended": true
|
"recommended": true
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"javascript": {
|
"formatter": {
|
||||||
"formatter": {
|
"lineEnding": "lf"
|
||||||
"quoteStyle": "single",
|
},
|
||||||
"semicolons": "asNeeded",
|
"javascript": {
|
||||||
"indentWidth": 2
|
"formatter": {
|
||||||
}
|
"quoteStyle": "single",
|
||||||
}
|
"semicolons": "asNeeded",
|
||||||
}
|
"indentWidth": 2
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+35
-13
@@ -3,12 +3,17 @@ import csv
|
|||||||
from aiohttp import web
|
from aiohttp import web
|
||||||
|
|
||||||
from .log import mklog
|
from .log import mklog
|
||||||
from .utils import backup_file, here, import_install, reqs_map, run_command, styles_dir
|
from .utils import (
|
||||||
|
backup_file,
|
||||||
|
import_install,
|
||||||
|
reqs_map,
|
||||||
|
run_command,
|
||||||
|
styles_dir,
|
||||||
|
)
|
||||||
|
|
||||||
endlog = mklog("mtb endpoint")
|
endlog = mklog("mtb endpoint")
|
||||||
|
|
||||||
# - ACTIONS
|
# - ACTIONS
|
||||||
import platform
|
|
||||||
import sys
|
import sys
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
@@ -22,7 +27,9 @@ def ACTIONS_installDependency(dependency_names=None):
|
|||||||
# reqs = []
|
# reqs = []
|
||||||
resolved_names = [reqs_map.get(name, name) for name in dependency_names]
|
resolved_names = [reqs_map.get(name, name) for name in dependency_names]
|
||||||
try:
|
try:
|
||||||
run_command([Path(sys.executable), "-m", "pip", "install"] + resolved_names)
|
run_command(
|
||||||
|
[Path(sys.executable), "-m", "pip", "install"] + resolved_names
|
||||||
|
)
|
||||||
return {"success": True}
|
return {"success": True}
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -44,9 +51,9 @@ def ACTIONS_installDependency(dependency_names=None):
|
|||||||
|
|
||||||
|
|
||||||
def ACTIONS_getStyles(style_name=None):
|
def ACTIONS_getStyles(style_name=None):
|
||||||
from .nodes.conditions import StylesLoader
|
from .nodes.conditions import MTB_StylesLoader
|
||||||
|
|
||||||
styles = StylesLoader.options
|
styles = MTB_StylesLoader.options
|
||||||
match_list = ["name"]
|
match_list = ["name"]
|
||||||
if styles:
|
if styles:
|
||||||
filtered_styles = {
|
filtered_styles = {
|
||||||
@@ -55,7 +62,9 @@ def ACTIONS_getStyles(style_name=None):
|
|||||||
if not key.startswith("__") and key not in match_list
|
if not key.startswith("__") and key not in match_list
|
||||||
}
|
}
|
||||||
if style_name:
|
if style_name:
|
||||||
return filtered_styles.get(style_name, {"error": "Style not found"})
|
return filtered_styles.get(
|
||||||
|
style_name, {"error": "Style not found"}
|
||||||
|
)
|
||||||
return filtered_styles
|
return filtered_styles
|
||||||
return {"error": "No styles found"}
|
return {"error": "No styles found"}
|
||||||
|
|
||||||
@@ -75,7 +84,9 @@ def ACTIONS_saveStyle(data):
|
|||||||
break
|
break
|
||||||
|
|
||||||
if not target:
|
if not target:
|
||||||
endlog.warning(f"Could not determine the target file for {data.keys()}")
|
endlog.warning(
|
||||||
|
f"Could not determine the target file for {data.keys()}"
|
||||||
|
)
|
||||||
return {"error": "Could not determine the target file for the style"}
|
return {"error": "Could not determine the target file for the style"}
|
||||||
|
|
||||||
backup_file(target)
|
backup_file(target)
|
||||||
@@ -103,11 +114,16 @@ async def do_action(request) -> web.Response:
|
|||||||
return web.json_response({"result": result})
|
return web.json_response({"result": result})
|
||||||
|
|
||||||
available_methods = [
|
available_methods = [
|
||||||
attr[len("ACTIONS_") :] for attr in globals() if attr.startswith("ACTIONS_")
|
attr[len("ACTIONS_") :]
|
||||||
|
for attr in globals()
|
||||||
|
if attr.startswith("ACTIONS_")
|
||||||
]
|
]
|
||||||
|
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{"error": "Invalid method name.", "available_methods": available_methods}
|
{
|
||||||
|
"error": "Invalid method name.",
|
||||||
|
"available_methods": available_methods,
|
||||||
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -127,7 +143,7 @@ def csv_editor():
|
|||||||
|
|
||||||
style_files = {}
|
style_files = {}
|
||||||
for file in inputs:
|
for file in inputs:
|
||||||
with open(file, "r", encoding="utf8") as f:
|
with open(file, encoding="utf8") as f:
|
||||||
parsed = csv.reader(f)
|
parsed = csv.reader(f)
|
||||||
style_files[file.name] = []
|
style_files[file.name] = []
|
||||||
for row in parsed:
|
for row in parsed:
|
||||||
@@ -235,7 +251,9 @@ def add_split_pane(left_content, right_content, vertical=True):
|
|||||||
|
|
||||||
|
|
||||||
def add_dropdown(title, options):
|
def add_dropdown(title, options):
|
||||||
option_str = "\n".join([f"<option value='{opt}'>{opt}</option>" for opt in options])
|
option_str = "\n".join(
|
||||||
|
[f"<option value='{opt}'>{opt}</option>" for opt in options]
|
||||||
|
)
|
||||||
return f"""
|
return f"""
|
||||||
<select>
|
<select>
|
||||||
<option disabled selected>{title}</option>
|
<option disabled selected>{title}</option>
|
||||||
@@ -254,11 +272,15 @@ def render_table(table_dict, sort=True, title=None):
|
|||||||
if isinstance(item, dict):
|
if isinstance(item, dict):
|
||||||
if "dependencies" in item:
|
if "dependencies" in item:
|
||||||
table_rows += f"<tr><td>{name}</td><td>"
|
table_rows += f"<tr><td>{name}</td><td>"
|
||||||
table_rows += f"{dependencies_button(name,item['dependencies'])}"
|
table_rows += (
|
||||||
|
f"{dependencies_button(name,item['dependencies'])}"
|
||||||
|
)
|
||||||
|
|
||||||
table_rows += "</td></tr>"
|
table_rows += "</td></tr>"
|
||||||
else:
|
else:
|
||||||
table_rows += f"<tr><td>{name}</td><td>{render_table(item)}</td></tr>"
|
table_rows += (
|
||||||
|
f"<tr><td>{name}</td><td>{render_table(item)}</td></tr>"
|
||||||
|
)
|
||||||
# elif isinstance(item, str):
|
# elif isinstance(item, str):
|
||||||
# table_rows += f"<tr><td>{name}</td><td>{item}</td></tr>"
|
# table_rows += f"<tr><td>{name}</td><td>{item}</td></tr>"
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -0,0 +1,147 @@
|
|||||||
|
# NOTE: This file is only use for development you can ignore it
|
||||||
|
|
||||||
|
use path.nu *
|
||||||
|
|
||||||
|
def get_root [--clean] {
|
||||||
|
if $clean {
|
||||||
|
$env.COMFY_CLEAN_ROOT
|
||||||
|
} else {
|
||||||
|
$env.COMFY_ROOT
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
export def "comfy build-web" [] {
|
||||||
|
cd $env.COMFY_MTB
|
||||||
|
cd web_source
|
||||||
|
npm run build
|
||||||
|
cp dist/*.js ../web/dist
|
||||||
|
}
|
||||||
|
|
||||||
|
export def "comfy dev-web" [] {
|
||||||
|
cd $env.COMFY_MTB
|
||||||
|
cd web_source
|
||||||
|
npm run dev
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
# start the comfy server
|
||||||
|
export def "comfy start" [--clean, --listen] {
|
||||||
|
let root = get_root --clean=($clean)
|
||||||
|
cd $root
|
||||||
|
MTB_DEBUG=true python main.py --port 3000 --preview-method auto ...(if $listen {["--listen"]} else {[]})
|
||||||
|
}
|
||||||
|
|
||||||
|
# update comfy itself and merge master in current branch
|
||||||
|
export def "comfy update" [
|
||||||
|
--clean # ??
|
||||||
|
--rebase # Rebase instead of merge
|
||||||
|
] {
|
||||||
|
let root = get_root --clean=($clean)
|
||||||
|
let models = $"($root)/models"
|
||||||
|
cd $root
|
||||||
|
let branch_name = (git rev-parse --abbrev-ref HEAD | str trim)
|
||||||
|
print $"(ansi yellow_italic)Backing up and removing models symlinks(ansi reset)"
|
||||||
|
|
||||||
|
cd $models
|
||||||
|
|
||||||
|
# find all symlinks
|
||||||
|
let links = (ls -la |
|
||||||
|
where not ($it.target | is-empty) |
|
||||||
|
select name target |
|
||||||
|
sort-by name)
|
||||||
|
|
||||||
|
|
||||||
|
if not ($links | is-empty) {
|
||||||
|
$links | save -f links.nuon
|
||||||
|
# remove them
|
||||||
|
open links.nuon | each {|p| rm $p.name }
|
||||||
|
}
|
||||||
|
|
||||||
|
cd $root
|
||||||
|
|
||||||
|
print $"(ansi yellow_italic)Checking out to master(ansi reset)"
|
||||||
|
git checkout master
|
||||||
|
|
||||||
|
print $"(ansi yellow_italic)Fetching and pulling remote updates(ansi reset)"
|
||||||
|
git fetch
|
||||||
|
git pull
|
||||||
|
|
||||||
|
print $"(ansi yellow_italic)Back to our branch \(($branch_name)\)(ansi reset)"
|
||||||
|
git checkout -
|
||||||
|
|
||||||
|
if $rebase {
|
||||||
|
print $"(ansi yellow_italic)Rebasing changes(ansi reset)"
|
||||||
|
git rebase master
|
||||||
|
|
||||||
|
} else {
|
||||||
|
print $"(ansi yellow_italic)Merging changes(ansi reset)"
|
||||||
|
git merge master
|
||||||
|
}
|
||||||
|
|
||||||
|
print $"(ansi yellow_italic)Linking back the models(ansi reset)"
|
||||||
|
|
||||||
|
cd $models
|
||||||
|
# resymlink them
|
||||||
|
open links.nuon | each {|p| link -a $p.target $p.name }
|
||||||
|
|
||||||
|
let commit_count = (git rev-list --count $branch_name $"^origin/($branch_name)")
|
||||||
|
|
||||||
|
|
||||||
|
print $"(ansi green_bold)Update successful \(($commit_count) new commits\)(ansi reset)"
|
||||||
|
|
||||||
|
|
||||||
|
}
|
||||||
|
|
||||||
|
export def "comfy toggle_extensions" [--clean] {
|
||||||
|
let root = get_root --clean=($clean)
|
||||||
|
cd $root
|
||||||
|
cd custom_nodes
|
||||||
|
let exts = (ls | where type in ["dir","symlink"] | get name)
|
||||||
|
let choices = ($exts | input list -m "choose extension to toggle")
|
||||||
|
if ($choices | is-empty) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
print $choices
|
||||||
|
|
||||||
|
let filtered = $choices | wrap name | upsert enabled {|p| not ($p.name | str ends-with ".disabled")}
|
||||||
|
|
||||||
|
print $filtered
|
||||||
|
$filtered | each {|f|
|
||||||
|
let new_name = ($f.name | str replace ".disabled" "")
|
||||||
|
|
||||||
|
let new_name = if $f.enabled {
|
||||||
|
$"($new_name).disabled"
|
||||||
|
} else {
|
||||||
|
$new_name
|
||||||
|
}
|
||||||
|
print $"Moving ($f.name) to ($new_name)"
|
||||||
|
mv $f.name $new_name
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
# git pull all extensions
|
||||||
|
export def "comfy update_extensions" [--clean] {
|
||||||
|
let root = get_root --clean=($clean)
|
||||||
|
cd $root
|
||||||
|
cd custom_nodes
|
||||||
|
git multipull .
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
export-env {
|
||||||
|
$env.COMFY_MTB = ("." | path expand)
|
||||||
|
$env.CUDA_ROOT = 'C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.1\'
|
||||||
|
|
||||||
|
$env.CUDA_HOME = $env.CUDA_ROOT
|
||||||
|
|
||||||
|
$env.COMFY_ROOT = ("../.." | path expand)
|
||||||
|
$env.COMFY_CLEAN_ROOT = ($env.COMFY_ROOT | path dirname | path join ComfyClean)
|
||||||
|
|
||||||
|
path-add 'C:/Portable/TensorRT-8.6.0.12/lib'
|
||||||
|
path-add ($env.CUDA_ROOT | path join bin)
|
||||||
|
overlay use ../../.venv/Scripts/activate.nu
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -36,7 +36,7 @@ class Formatter(logging.Formatter):
|
|||||||
return formatter.format(record)
|
return formatter.format(record)
|
||||||
|
|
||||||
|
|
||||||
def mklog(name, level=base_log_level):
|
def mklog(name: str, level: int = base_log_level):
|
||||||
logger = logging.getLogger(name)
|
logger = logging.getLogger(name)
|
||||||
logger.setLevel(level)
|
logger.setLevel(level)
|
||||||
|
|
||||||
@@ -58,23 +58,23 @@ def mklog(name, level=base_log_level):
|
|||||||
log = mklog(__package__, base_log_level)
|
log = mklog(__package__, base_log_level)
|
||||||
|
|
||||||
|
|
||||||
def log_user(arg):
|
def log_user(arg: str):
|
||||||
print("\033[34mComfy MTB Utils:\033[0m {arg}")
|
print(f"\033[34mComfy MTB Utils:\033[0m {arg}")
|
||||||
|
|
||||||
|
|
||||||
def get_summary(docstring):
|
def get_summary(docstring: str):
|
||||||
return docstring.strip().split("\n\n", 1)[0]
|
return docstring.strip().split("\n\n", 1)[0]
|
||||||
|
|
||||||
|
|
||||||
def blue_text(text):
|
def blue_text(text: str):
|
||||||
return f"\033[94m{text}\033[0m"
|
return f"\033[94m{text}\033[0m"
|
||||||
|
|
||||||
|
|
||||||
def cyan_text(text):
|
def cyan_text(text: str):
|
||||||
return f"\033[96m{text}\033[0m"
|
return f"\033[96m{text}\033[0m"
|
||||||
|
|
||||||
|
|
||||||
def get_label(label):
|
def get_label(label: str):
|
||||||
if label.startswith("MTB_"):
|
if label.startswith("MTB_"):
|
||||||
label = label[4:]
|
label = label[4:]
|
||||||
words = re.findall(r"(?:^|[A-Z])[a-z]*", label)
|
words = re.findall(r"(?:^|[A-Z])[a-z]*", label)
|
||||||
|
|||||||
+262
-24
@@ -6,11 +6,11 @@ import torch
|
|||||||
from PIL import Image
|
from PIL import Image
|
||||||
|
|
||||||
from ..log import log
|
from ..log import log
|
||||||
from ..utils import apply_easing, pil2tensor
|
from ..utils import EASINGS, apply_easing, pil2tensor
|
||||||
from .transform import TransformImage
|
from .transform import MTB_TransformImage
|
||||||
|
|
||||||
|
|
||||||
def hex_to_rgb(hex_color, bgr=False):
|
def hex_to_rgb(hex_color: str, bgr: bool = False):
|
||||||
hex_color = hex_color.lstrip("#")
|
hex_color = hex_color.lstrip("#")
|
||||||
if bgr:
|
if bgr:
|
||||||
return tuple(int(hex_color[i : i + 2], 16) for i in (4, 2, 0))
|
return tuple(int(hex_color[i : i + 2], 16) for i in (4, 2, 0))
|
||||||
@@ -18,6 +18,157 @@ def hex_to_rgb(hex_color, bgr=False):
|
|||||||
return tuple(int(hex_color[i : i + 2], 16) for i in (0, 2, 4))
|
return tuple(int(hex_color[i : i + 2], 16) for i in (0, 2, 4))
|
||||||
|
|
||||||
|
|
||||||
|
class MTB_BatchFloatMath:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"reverse": ("BOOLEAN", {"default": False}),
|
||||||
|
"operation": (
|
||||||
|
["add", "sub", "mul", "div", "pow", "abs"],
|
||||||
|
{"default": "add"},
|
||||||
|
),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("FLOATS",)
|
||||||
|
CATEGORY = "mtb/utils"
|
||||||
|
FUNCTION = "execute"
|
||||||
|
|
||||||
|
def execute(self, reverse: bool, operation: str, **kwargs: list[float]):
|
||||||
|
res: list[float] = []
|
||||||
|
vals = list(kwargs.values())
|
||||||
|
|
||||||
|
if reverse:
|
||||||
|
vals = vals[::-1]
|
||||||
|
|
||||||
|
ref_count = len(vals[0])
|
||||||
|
for v in vals:
|
||||||
|
if len(v) != ref_count:
|
||||||
|
raise ValueError(
|
||||||
|
f"All values must have the same length (current: {len(v)}, ref: {ref_count}"
|
||||||
|
)
|
||||||
|
|
||||||
|
match operation:
|
||||||
|
case "add":
|
||||||
|
for i in range(ref_count):
|
||||||
|
result = sum(v[i] for v in vals)
|
||||||
|
res.append(result)
|
||||||
|
case "sub":
|
||||||
|
for i in range(ref_count):
|
||||||
|
result = vals[0][i] - sum(v[i] for v in vals[1:])
|
||||||
|
res.append(result)
|
||||||
|
case "mul":
|
||||||
|
for i in range(ref_count):
|
||||||
|
result = vals[0][i] * vals[1][i]
|
||||||
|
res.append(result)
|
||||||
|
case "div":
|
||||||
|
for i in range(ref_count):
|
||||||
|
result = vals[0][i] / vals[1][i]
|
||||||
|
res.append(result)
|
||||||
|
case "pow":
|
||||||
|
for i in range(ref_count):
|
||||||
|
result: float = vals[0][i] ** vals[1][i]
|
||||||
|
res.append(result)
|
||||||
|
case "abs":
|
||||||
|
for i in range(ref_count):
|
||||||
|
result = abs(vals[0][i])
|
||||||
|
res.append(result)
|
||||||
|
case _:
|
||||||
|
log.info(f"For now this mode ({operation}) is not implemented")
|
||||||
|
|
||||||
|
return (res,)
|
||||||
|
|
||||||
|
|
||||||
|
class MTB_BatchFloatNormalize:
|
||||||
|
"""Normalize the values in the list of floats"""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
return {
|
||||||
|
"required": {"floats": ("FLOATS",)},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("FLOATS",)
|
||||||
|
RETURN_NAMES = ("normalized_floats",)
|
||||||
|
CATEGORY = "mtb/batch"
|
||||||
|
FUNCTION = "execute"
|
||||||
|
|
||||||
|
def execute(
|
||||||
|
self,
|
||||||
|
floats: list[float],
|
||||||
|
):
|
||||||
|
min_value = min(floats)
|
||||||
|
max_value = max(floats)
|
||||||
|
|
||||||
|
normalized_floats = [
|
||||||
|
(x - min_value) / (max_value - min_value) for x in floats
|
||||||
|
]
|
||||||
|
log.debug(f"Floats: {floats}")
|
||||||
|
log.debug(f"Normalized Floats: {normalized_floats}")
|
||||||
|
|
||||||
|
return (normalized_floats,)
|
||||||
|
|
||||||
|
|
||||||
|
class MTB_BatchTimeWrap:
|
||||||
|
"""Remap a batch using a time curve (FLOATS)"""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"target_count": ("INT", {"default": 25, "min": 2}),
|
||||||
|
"frames": ("IMAGE",),
|
||||||
|
"curve": ("FLOATS",),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("IMAGE", "FLOATS")
|
||||||
|
RETURN_NAMES = ("image", "interpolated_floats")
|
||||||
|
CATEGORY = "mtb/batch"
|
||||||
|
FUNCTION = "execute"
|
||||||
|
|
||||||
|
def execute(
|
||||||
|
self, target_count: int, frames: torch.Tensor, curve: list[float]
|
||||||
|
):
|
||||||
|
"""Apply time warping to a list of video frames based on a curve."""
|
||||||
|
log.debug(f"Input frames shape: {frames.shape}")
|
||||||
|
log.debug(f"Curve: {curve}")
|
||||||
|
|
||||||
|
total_duration = sum(curve)
|
||||||
|
|
||||||
|
log.debug(f"Total duration: {total_duration}")
|
||||||
|
|
||||||
|
B, H, W, C = frames.shape
|
||||||
|
|
||||||
|
log.debug(f"Batch Size: {B}")
|
||||||
|
|
||||||
|
normalized_times = np.linspace(0, 1, target_count)
|
||||||
|
interpolated_curve = np.interp(
|
||||||
|
normalized_times, np.linspace(0, 1, len(curve)), curve
|
||||||
|
).tolist()
|
||||||
|
log.debug(f"Interpolated curve: {interpolated_curve}")
|
||||||
|
|
||||||
|
interpolated_frame_indices = [
|
||||||
|
(B - 1) * value for value in interpolated_curve
|
||||||
|
]
|
||||||
|
log.debug(f"Interpolated frame indices: {interpolated_frame_indices}")
|
||||||
|
|
||||||
|
rounded_indices = [
|
||||||
|
int(round(idx)) for idx in interpolated_frame_indices
|
||||||
|
]
|
||||||
|
rounded_indices = np.clip(rounded_indices, 0, B - 1)
|
||||||
|
|
||||||
|
# Gather frames based on interpolated indices
|
||||||
|
warped_frames = []
|
||||||
|
for index in rounded_indices:
|
||||||
|
warped_frames.append(frames[index].unsqueeze(0))
|
||||||
|
|
||||||
|
warped_tensor = torch.cat(warped_frames, dim=0)
|
||||||
|
log.debug(f"Warped frames shape: {warped_tensor.shape}")
|
||||||
|
return (warped_tensor, interpolated_curve)
|
||||||
|
|
||||||
|
|
||||||
class MTB_BatchMake:
|
class MTB_BatchMake:
|
||||||
"""Simply duplicates the input frame as a batch"""
|
"""Simply duplicates the input frame as a batch"""
|
||||||
|
|
||||||
@@ -192,18 +343,21 @@ class MTB_BatchFloatAssemble:
|
|||||||
def INPUT_TYPES(cls):
|
def INPUT_TYPES(cls):
|
||||||
return {"required": {"reverse": ("BOOLEAN", {"default": False})}}
|
return {"required": {"reverse": ("BOOLEAN", {"default": False})}}
|
||||||
|
|
||||||
FUNCTION = "assemble_floats"
|
|
||||||
RETURN_TYPES = ("FLOATS",)
|
RETURN_TYPES = ("FLOATS",)
|
||||||
CATEGORY = "mtb/batch"
|
CATEGORY = "mtb/batch"
|
||||||
|
FUNCTION = "assemble_floats"
|
||||||
|
|
||||||
|
def assemble_floats(self, reverse: bool, **kwargs: list[float]):
|
||||||
|
res: list[float] = []
|
||||||
|
|
||||||
def assemble_floats(self, reverse, **kwargs):
|
|
||||||
res = []
|
|
||||||
if reverse:
|
if reverse:
|
||||||
for x in reversed(kwargs.values()):
|
for x in reversed(kwargs.values()):
|
||||||
res += x
|
if x:
|
||||||
|
res += x
|
||||||
else:
|
else:
|
||||||
for x in kwargs.values():
|
for x in kwargs.values():
|
||||||
res += x
|
if x:
|
||||||
|
res += x
|
||||||
|
|
||||||
return (res,)
|
return (res,)
|
||||||
|
|
||||||
@@ -219,7 +373,7 @@ class MTB_BatchFloat:
|
|||||||
["Single", "Steps"],
|
["Single", "Steps"],
|
||||||
{"default": "Steps"},
|
{"default": "Steps"},
|
||||||
),
|
),
|
||||||
"count": ("INT", {"default": 1}),
|
"count": ("INT", {"default": 2}),
|
||||||
"min": ("FLOAT", {"default": 0.0, "step": 0.001}),
|
"min": ("FLOAT", {"default": 0.0, "step": 0.001}),
|
||||||
"max": ("FLOAT", {"default": 1.0, "step": 0.001}),
|
"max": ("FLOAT", {"default": 1.0, "step": 0.001}),
|
||||||
"easing": (
|
"easing": (
|
||||||
@@ -257,6 +411,10 @@ class MTB_BatchFloat:
|
|||||||
CATEGORY = "mtb/batch"
|
CATEGORY = "mtb/batch"
|
||||||
|
|
||||||
def set_floats(self, mode, count, min, max, easing):
|
def set_floats(self, mode, count, min, max, easing):
|
||||||
|
if mode == "Steps" and count == 1:
|
||||||
|
raise ValueError(
|
||||||
|
"Steps mode requires at least a count of 2 values"
|
||||||
|
)
|
||||||
keyframes = []
|
keyframes = []
|
||||||
if mode == "Single":
|
if mode == "Single":
|
||||||
keyframes = [min] * count
|
keyframes = [min] * count
|
||||||
@@ -415,7 +573,7 @@ class MTB_Batch2dTransform:
|
|||||||
if count == 0:
|
if count == 0:
|
||||||
keyframes[name] = [default_vals[name]] * image.shape[0]
|
keyframes[name] = [default_vals[name]] * image.shape[0]
|
||||||
|
|
||||||
transformer = TransformImage()
|
transformer = MTB_TransformImage()
|
||||||
res = [
|
res = [
|
||||||
transformer.transform(
|
transformer.transform(
|
||||||
image[i].unsqueeze(0),
|
image[i].unsqueeze(0),
|
||||||
@@ -432,6 +590,66 @@ class MTB_Batch2dTransform:
|
|||||||
return (torch.cat(res, dim=0),)
|
return (torch.cat(res, dim=0),)
|
||||||
|
|
||||||
|
|
||||||
|
class MTB_BatchFloatFit:
|
||||||
|
"""Fit a list of floats using a source and target range"""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"values": ("FLOATS", {"forceInput": True}),
|
||||||
|
"clamp": ("BOOLEAN", {"default": False}),
|
||||||
|
"auto_compute_source": ("BOOLEAN", {"default": False}),
|
||||||
|
"source_min": ("FLOAT", {"default": 0.0, "step": 0.01}),
|
||||||
|
"source_max": ("FLOAT", {"default": 1.0, "step": 0.01}),
|
||||||
|
"target_min": ("FLOAT", {"default": 0.0, "step": 0.01}),
|
||||||
|
"target_max": ("FLOAT", {"default": 1.0, "step": 0.01}),
|
||||||
|
"easing": (
|
||||||
|
EASINGS,
|
||||||
|
{"default": "Linear"},
|
||||||
|
),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
FUNCTION = "fit_range"
|
||||||
|
RETURN_TYPES = ("FLOATS",)
|
||||||
|
CATEGORY = "mtb/batch"
|
||||||
|
DESCRIPTION = "Fit a list of floats using a source and target range"
|
||||||
|
|
||||||
|
def fit_range(
|
||||||
|
self,
|
||||||
|
values: list[float],
|
||||||
|
clamp: bool,
|
||||||
|
auto_compute_source: bool,
|
||||||
|
source_min: float,
|
||||||
|
source_max: float,
|
||||||
|
target_min: float,
|
||||||
|
target_max: float,
|
||||||
|
easing: str,
|
||||||
|
):
|
||||||
|
if auto_compute_source:
|
||||||
|
source_min = min(values)
|
||||||
|
source_max = max(values)
|
||||||
|
|
||||||
|
from .graph_utils import MTB_FitNumber
|
||||||
|
|
||||||
|
res = []
|
||||||
|
fit_number = MTB_FitNumber()
|
||||||
|
for value in values:
|
||||||
|
(transformed_value,) = fit_number.set_range(
|
||||||
|
value,
|
||||||
|
clamp,
|
||||||
|
source_min,
|
||||||
|
source_max,
|
||||||
|
target_min,
|
||||||
|
target_max,
|
||||||
|
easing,
|
||||||
|
)
|
||||||
|
res.append(transformed_value)
|
||||||
|
|
||||||
|
return (res,)
|
||||||
|
|
||||||
|
|
||||||
class MTB_PlotBatchFloat:
|
class MTB_PlotBatchFloat:
|
||||||
"""Plot floats"""
|
"""Plot floats"""
|
||||||
|
|
||||||
@@ -443,6 +661,7 @@ class MTB_PlotBatchFloat:
|
|||||||
"height": ("INT", {"default": 768}),
|
"height": ("INT", {"default": 768}),
|
||||||
"point_size": ("INT", {"default": 4}),
|
"point_size": ("INT", {"default": 4}),
|
||||||
"seed": ("INT", {"default": 1}),
|
"seed": ("INT", {"default": 1}),
|
||||||
|
"start_at_zero": ("BOOLEAN", {"default": False}),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -451,10 +670,21 @@ class MTB_PlotBatchFloat:
|
|||||||
FUNCTION = "plot"
|
FUNCTION = "plot"
|
||||||
CATEGORY = "mtb/batch"
|
CATEGORY = "mtb/batch"
|
||||||
|
|
||||||
def plot(self, width, height, point_size, seed, **kwargs):
|
def plot(
|
||||||
|
self,
|
||||||
|
width: int,
|
||||||
|
height: int,
|
||||||
|
point_size: int,
|
||||||
|
seed: int,
|
||||||
|
start_at_zero: bool,
|
||||||
|
interactive_backend: bool = False,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
import matplotlib
|
import matplotlib
|
||||||
|
|
||||||
matplotlib.use("Agg")
|
# NOTE: This is for notebook usage or tests, i.e not exposed to comfy that should always use Agg
|
||||||
|
if not interactive_backend:
|
||||||
|
matplotlib.use("Agg")
|
||||||
import matplotlib.pyplot as plt
|
import matplotlib.pyplot as plt
|
||||||
|
|
||||||
fig, ax = plt.subplots(figsize=(width / 100, height / 100), dpi=100)
|
fig, ax = plt.subplots(figsize=(width / 100, height / 100), dpi=100)
|
||||||
@@ -465,26 +695,30 @@ class MTB_PlotBatchFloat:
|
|||||||
ax.grid(color="gray", linestyle="-", linewidth=0.5, alpha=0.5)
|
ax.grid(color="gray", linestyle="-", linewidth=0.5, alpha=0.5)
|
||||||
|
|
||||||
# Finding global min and max across all lists for scaling the plot
|
# Finding global min and max across all lists for scaling the plot
|
||||||
global_min = min(min(values) for values in kwargs.values())
|
all_values = [value for values in kwargs.values() for value in values]
|
||||||
global_max = max(max(values) for values in kwargs.values())
|
global_min = min(all_values)
|
||||||
|
global_max = max(all_values)
|
||||||
|
|
||||||
# Color cycle to ensure each plot has a distinct color
|
y_padding = 0.05 * (global_max - global_min)
|
||||||
colormap = plt.cm.get_cmap("viridis", len(kwargs))
|
ax.set_ylim(global_min - y_padding, global_max + y_padding)
|
||||||
color_normalization_factor = (
|
|
||||||
0.5 if len(kwargs) == 1 else (len(kwargs) - 1)
|
|
||||||
)
|
|
||||||
|
|
||||||
# Plotting each list with a unique color
|
max_length = max(len(values) for values in kwargs.values())
|
||||||
for i, (label, values) in enumerate(kwargs.items()):
|
if start_at_zero:
|
||||||
color_value = i / color_normalization_factor
|
x_values = np.linspace(0, max_length - 1, max_length)
|
||||||
ax.plot(values, label=label, color=colormap(color_value))
|
else:
|
||||||
|
x_values = np.linspace(1, max_length, max_length)
|
||||||
|
|
||||||
ax.set_ylim(global_min, global_max) # Scaling the y-axis
|
ax.set_xlim(1, max_length) # Set X-axis limits
|
||||||
|
np.random.seed(seed)
|
||||||
|
colors = np.random.rand(len(kwargs), 3) # Generate random RGB values
|
||||||
|
for color, (label, values) in zip(colors, kwargs.items()):
|
||||||
|
ax.plot(x_values[: len(values)], values, label=label, color=color)
|
||||||
ax.legend(
|
ax.legend(
|
||||||
title="Legend",
|
title="Legend",
|
||||||
title_fontsize="large",
|
title_fontsize="large",
|
||||||
fontsize="medium",
|
fontsize="medium",
|
||||||
edgecolor="black",
|
edgecolor="black",
|
||||||
|
loc="best",
|
||||||
)
|
)
|
||||||
|
|
||||||
# Setting labels and title
|
# Setting labels and title
|
||||||
@@ -798,7 +1032,11 @@ __nodes__ = [
|
|||||||
MTB_BatchMake,
|
MTB_BatchMake,
|
||||||
MTB_BatchFloatAssemble,
|
MTB_BatchFloatAssemble,
|
||||||
MTB_BatchFloatFill,
|
MTB_BatchFloatFill,
|
||||||
|
MTB_BatchFloatNormalize,
|
||||||
MTB_BatchMerge,
|
MTB_BatchMerge,
|
||||||
MTB_BatchShake,
|
MTB_BatchShake,
|
||||||
MTB_PlotBatchFloat,
|
MTB_PlotBatchFloat,
|
||||||
|
MTB_BatchTimeWrap,
|
||||||
|
MTB_BatchFloatFit,
|
||||||
|
MTB_BatchFloatMath,
|
||||||
]
|
]
|
||||||
|
|||||||
+34
-14
@@ -1,4 +1,5 @@
|
|||||||
import csv, shutil
|
import csv
|
||||||
|
import shutil
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import folder_paths
|
import folder_paths
|
||||||
@@ -7,7 +8,7 @@ from ..log import log
|
|||||||
from ..utils import here
|
from ..utils import here
|
||||||
|
|
||||||
|
|
||||||
class InterpolateClipSequential:
|
class MTB_InterpolateClipSequential:
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(cls):
|
def INPUT_TYPES(cls):
|
||||||
return {
|
return {
|
||||||
@@ -28,7 +29,12 @@ class InterpolateClipSequential:
|
|||||||
CATEGORY = "mtb/conditioning"
|
CATEGORY = "mtb/conditioning"
|
||||||
|
|
||||||
def interpolate_encodings_sequential(
|
def interpolate_encodings_sequential(
|
||||||
self, base_text, text_to_replace, clip, interpolation_strength, **replacements
|
self,
|
||||||
|
base_text,
|
||||||
|
text_to_replace,
|
||||||
|
clip,
|
||||||
|
interpolation_strength,
|
||||||
|
**replacements,
|
||||||
):
|
):
|
||||||
log.debug(f"Received interpolation_strength: {interpolation_strength}")
|
log.debug(f"Received interpolation_strength: {interpolation_strength}")
|
||||||
|
|
||||||
@@ -63,20 +69,30 @@ class InterpolateClipSequential:
|
|||||||
log.debug("Using the base text a the base blend")
|
log.debug("Using the base text a the base blend")
|
||||||
# - Start with the base_text condition
|
# - Start with the base_text condition
|
||||||
tokens = clip.tokenize(base_text)
|
tokens = clip.tokenize(base_text)
|
||||||
cond_from, pooled_from = clip.encode_from_tokens(tokens, return_pooled=True)
|
cond_from, pooled_from = clip.encode_from_tokens(
|
||||||
|
tokens, return_pooled=True
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
base_replace = list(replacements.values())[segment_index - 1]
|
base_replace = list(replacements.values())[segment_index - 1]
|
||||||
log.debug(f"Using {base_replace} a the base blend")
|
log.debug(f"Using {base_replace} a the base blend")
|
||||||
|
|
||||||
# - Start with the base_text condition replaced by the closest replacement
|
# - Start with the base_text condition replaced by the closest replacement
|
||||||
tokens = clip.tokenize(base_text.replace(text_to_replace, base_replace))
|
tokens = clip.tokenize(
|
||||||
cond_from, pooled_from = clip.encode_from_tokens(tokens, return_pooled=True)
|
base_text.replace(text_to_replace, base_replace)
|
||||||
|
)
|
||||||
|
cond_from, pooled_from = clip.encode_from_tokens(
|
||||||
|
tokens, return_pooled=True
|
||||||
|
)
|
||||||
|
|
||||||
replacement_text = list(replacements.values())[segment_index]
|
replacement_text = list(replacements.values())[segment_index]
|
||||||
|
|
||||||
interpolated_text = base_text.replace(text_to_replace, replacement_text)
|
interpolated_text = base_text.replace(
|
||||||
|
text_to_replace, replacement_text
|
||||||
|
)
|
||||||
tokens = clip.tokenize(interpolated_text)
|
tokens = clip.tokenize(interpolated_text)
|
||||||
cond_to, pooled_to = clip.encode_from_tokens(tokens, return_pooled=True)
|
cond_to, pooled_to = clip.encode_from_tokens(
|
||||||
|
tokens, return_pooled=True
|
||||||
|
)
|
||||||
|
|
||||||
# - Linearly interpolate between the two conditions
|
# - Linearly interpolate between the two conditions
|
||||||
interpolated_condition = (
|
interpolated_condition = (
|
||||||
@@ -86,10 +102,12 @@ class InterpolateClipSequential:
|
|||||||
1.0 - local_strength
|
1.0 - local_strength
|
||||||
) * pooled_from + local_strength * pooled_to
|
) * pooled_from + local_strength * pooled_to
|
||||||
|
|
||||||
return ([[interpolated_condition, {"pooled_output": interpolated_pooled}]],)
|
return (
|
||||||
|
[[interpolated_condition, {"pooled_output": interpolated_pooled}]],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class SmartStep:
|
class MTB_SmartStep:
|
||||||
"""Utils to control the steps start/stop of the KAdvancedSampler in percentage"""
|
"""Utils to control the steps start/stop of the KAdvancedSampler in percentage"""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -136,7 +154,7 @@ def install_default_styles(force=False):
|
|||||||
return dest_style
|
return dest_style
|
||||||
|
|
||||||
|
|
||||||
class StylesLoader:
|
class MTB_StylesLoader:
|
||||||
"""Load csv files and populate a dropdown from the rows (à la A111)"""
|
"""Load csv files and populate a dropdown from the rows (à la A111)"""
|
||||||
|
|
||||||
options = {}
|
options = {}
|
||||||
@@ -148,13 +166,15 @@ class StylesLoader:
|
|||||||
if not input_dir.exists():
|
if not input_dir.exists():
|
||||||
install_default_styles()
|
install_default_styles()
|
||||||
|
|
||||||
if not (files := [f for f in input_dir.iterdir() if f.suffix == ".csv"]):
|
if not (
|
||||||
|
files := [f for f in input_dir.iterdir() if f.suffix == ".csv"]
|
||||||
|
):
|
||||||
log.warn(
|
log.warn(
|
||||||
"No styles found in the styles folder, place at least one csv file in the styles folder at the root of ComfyUI (for instance ComfyUI/styles/mystyle.csv)"
|
"No styles found in the styles folder, place at least one csv file in the styles folder at the root of ComfyUI (for instance ComfyUI/styles/mystyle.csv)"
|
||||||
)
|
)
|
||||||
|
|
||||||
for file in files:
|
for file in files:
|
||||||
with open(file, "r", encoding="utf8") as f:
|
with open(file, encoding="utf8") as f:
|
||||||
parsed = csv.reader(f)
|
parsed = csv.reader(f)
|
||||||
for i, row in enumerate(parsed):
|
for i, row in enumerate(parsed):
|
||||||
log.debug(f"Adding style {row[0]}")
|
log.debug(f"Adding style {row[0]}")
|
||||||
@@ -193,4 +213,4 @@ class StylesLoader:
|
|||||||
return (self.options[style_name][0], self.options[style_name][1])
|
return (self.options[style_name][0], self.options[style_name][1])
|
||||||
|
|
||||||
|
|
||||||
__nodes__ = [SmartStep, StylesLoader, InterpolateClipSequential]
|
__nodes__ = [MTB_SmartStep, MTB_StylesLoader, MTB_InterpolateClipSequential]
|
||||||
|
|||||||
+5
-5
@@ -6,7 +6,7 @@ from ..log import log
|
|||||||
from ..utils import np2tensor, pil2tensor, tensor2np, tensor2pil
|
from ..utils import np2tensor, pil2tensor, tensor2np, tensor2pil
|
||||||
|
|
||||||
|
|
||||||
class Bbox:
|
class MTB_Bbox:
|
||||||
"""The bounding box (BBOX) custom type used by other nodes"""
|
"""The bounding box (BBOX) custom type used by other nodes"""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -41,7 +41,7 @@ class Bbox:
|
|||||||
return ((x, y, width, height),)
|
return ((x, y, width, height),)
|
||||||
|
|
||||||
|
|
||||||
class BboxFromMask:
|
class MTB_BboxFromMask:
|
||||||
"""From a mask extract the bounding box"""
|
"""From a mask extract the bounding box"""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -110,7 +110,7 @@ class BboxFromMask:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
class Crop:
|
class MTB_Crop:
|
||||||
"""Crops an image and an optional mask to a given bounding box
|
"""Crops an image and an optional mask to a given bounding box
|
||||||
|
|
||||||
The bounding box can be given as a tuple of (x, y, width, height) or as a BBOX type
|
The bounding box can be given as a tuple of (x, y, width, height) or as a BBOX type
|
||||||
@@ -218,7 +218,7 @@ def bbox_to_region(bbox, target_size=None):
|
|||||||
return (bbox[0], bbox[1], bbox[0] + bbox[2], bbox[1] + bbox[3])
|
return (bbox[0], bbox[1], bbox[0] + bbox[2], bbox[1] + bbox[3])
|
||||||
|
|
||||||
|
|
||||||
class Uncrop:
|
class MTB_Uncrop:
|
||||||
"""Uncrops an image to a given bounding box
|
"""Uncrops an image to a given bounding box
|
||||||
|
|
||||||
The bounding box can be given as a tuple of (x, y, width, height) or as a BBOX type
|
The bounding box can be given as a tuple of (x, y, width, height) or as a BBOX type
|
||||||
@@ -324,4 +324,4 @@ class Uncrop:
|
|||||||
return (pil2tensor(out_images),)
|
return (pil2tensor(out_images),)
|
||||||
|
|
||||||
|
|
||||||
__nodes__ = [BboxFromMask, Bbox, Crop, Uncrop]
|
__nodes__ = [MTB_BboxFromMask, MTB_Bbox, MTB_Crop, MTB_Uncrop]
|
||||||
|
|||||||
+58
-1
@@ -1,5 +1,7 @@
|
|||||||
import json
|
import json
|
||||||
|
|
||||||
|
from ..log import log
|
||||||
|
|
||||||
|
|
||||||
def deserialize_curve(curve):
|
def deserialize_curve(curve):
|
||||||
if isinstance(curve, str):
|
if isinstance(curve, str):
|
||||||
@@ -30,7 +32,62 @@ class MTB_Curve:
|
|||||||
CATEGORY = "mtb/curve"
|
CATEGORY = "mtb/curve"
|
||||||
|
|
||||||
def do_curve(self, curve):
|
def do_curve(self, curve):
|
||||||
|
log.debug(f"Curve: {curve}")
|
||||||
return (curve,)
|
return (curve,)
|
||||||
|
|
||||||
|
|
||||||
__nodes__ = [MTB_Curve]
|
class MTB_CurveToFloat:
|
||||||
|
"""Convert a FLOAT_CURVE to a FLOAT or FLOATS"""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"curve": ("FLOAT_CURVE", {"forceInput": True}),
|
||||||
|
"steps": ("INT", {"default": 10, "min": 2}),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("FLOATS", "FLOAT")
|
||||||
|
FUNCTION = "do_curve"
|
||||||
|
|
||||||
|
CATEGORY = "mtb/curve"
|
||||||
|
|
||||||
|
def do_curve(self, curve, steps):
|
||||||
|
log.debug(f"Curve: {curve}")
|
||||||
|
|
||||||
|
# sort by x (should be handled by the widget)
|
||||||
|
sorted_points = sorted(curve.items(), key=lambda item: item[1]["x"])
|
||||||
|
# Extract X and Y values
|
||||||
|
x_values = [point[1]["x"] for point in sorted_points]
|
||||||
|
y_values = [point[1]["y"] for point in sorted_points]
|
||||||
|
# Calculate step size
|
||||||
|
step_size = (max(x_values) - min(x_values)) / (steps - 1)
|
||||||
|
|
||||||
|
# Interpolate Y values for each step
|
||||||
|
interpolated_y_values = []
|
||||||
|
for step in range(steps):
|
||||||
|
current_x = min(x_values) + step_size * step
|
||||||
|
|
||||||
|
# Find the indices of the two points between which the current_x falls
|
||||||
|
idx1 = max(idx for idx, x in enumerate(x_values) if x <= current_x)
|
||||||
|
idx2 = min(idx for idx, x in enumerate(x_values) if x >= current_x)
|
||||||
|
|
||||||
|
# If the current_x matches one of the points, no interpolation is needed
|
||||||
|
if current_x == x_values[idx1]:
|
||||||
|
interpolated_y_values.append(y_values[idx1])
|
||||||
|
elif current_x == x_values[idx2]:
|
||||||
|
interpolated_y_values.append(y_values[idx2])
|
||||||
|
else:
|
||||||
|
# Interpolate Y value using linear interpolation
|
||||||
|
y1 = y_values[idx1]
|
||||||
|
y2 = y_values[idx2]
|
||||||
|
x1 = x_values[idx1]
|
||||||
|
x2 = x_values[idx2]
|
||||||
|
interpolated_y = y1 + (y2 - y1) * (current_x - x1) / (x2 - x1)
|
||||||
|
interpolated_y_values.append(interpolated_y)
|
||||||
|
|
||||||
|
return (interpolated_y_values, interpolated_y_values)
|
||||||
|
|
||||||
|
|
||||||
|
__nodes__ = [MTB_Curve, MTB_CurveToFloat]
|
||||||
|
|||||||
+7
-2
@@ -1,5 +1,6 @@
|
|||||||
import base64
|
import base64
|
||||||
import io
|
import io
|
||||||
|
import json
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
@@ -48,7 +49,7 @@ def process_list(anything):
|
|||||||
f"List of Tensors: {first_element.shape} (x{len(anything)})"
|
f"List of Tensors: {first_element.shape} (x{len(anything)})"
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
text.append(f"Array: {anything}")
|
text.append(f"Array ({len(anything)}): {anything}")
|
||||||
|
|
||||||
return {"text": text}
|
return {"text": text}
|
||||||
|
|
||||||
@@ -61,6 +62,9 @@ def process_dict(anything):
|
|||||||
)
|
)
|
||||||
text.append(f"Latent Samples: {anything['samples'].shape} {is_empty}")
|
text.append(f"Latent Samples: {anything['samples'].shape} {is_empty}")
|
||||||
|
|
||||||
|
else:
|
||||||
|
text.append(json.dumps(anything, indent=2))
|
||||||
|
|
||||||
return {"text": text}
|
return {"text": text}
|
||||||
|
|
||||||
|
|
||||||
@@ -106,10 +110,11 @@ class MTB_Debug:
|
|||||||
}
|
}
|
||||||
if output_to_console:
|
if output_to_console:
|
||||||
for k, v in kwargs.items():
|
for k, v in kwargs.items():
|
||||||
print(f"{k}: {v}")
|
log.info(f"{k}: {v}")
|
||||||
|
|
||||||
for anything in kwargs.values():
|
for anything in kwargs.values():
|
||||||
processor = processors.get(type(anything), process_text)
|
processor = processors.get(type(anything), process_text)
|
||||||
|
|
||||||
processed_data = processor(anything)
|
processed_data = processor(anything)
|
||||||
|
|
||||||
for ui_key, ui_value in processed_data.items():
|
for ui_key, ui_value in processed_data.items():
|
||||||
|
|||||||
+2
-2
@@ -303,7 +303,7 @@ def normals_to_height(normals_img, seamless, progress_callback):
|
|||||||
|
|
||||||
|
|
||||||
# - ADDON
|
# - ADDON
|
||||||
class DeepBump:
|
class MTB_DeepBump:
|
||||||
"""Normal & height maps generation from single pictures"""
|
"""Normal & height maps generation from single pictures"""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -386,4 +386,4 @@ class DeepBump:
|
|||||||
return (torch.cat(out_images, dim=0),)
|
return (torch.cat(out_images, dim=0),)
|
||||||
|
|
||||||
|
|
||||||
__nodes__ = [DeepBump]
|
__nodes__ = [MTB_DeepBump]
|
||||||
|
|||||||
+11
-10
@@ -1,6 +1,4 @@
|
|||||||
import os
|
import os
|
||||||
from pathlib import Path
|
|
||||||
from typing import Tuple
|
|
||||||
|
|
||||||
import comfy
|
import comfy
|
||||||
import comfy.utils
|
import comfy.utils
|
||||||
@@ -9,14 +7,13 @@ import folder_paths
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
from comfy import model_management
|
from comfy import model_management
|
||||||
from gfpgan import GFPGANer
|
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
|
|
||||||
from ..log import NullWriter, log
|
from ..log import NullWriter, log
|
||||||
from ..utils import get_model_path, np2tensor, pil2tensor, tensor2np
|
from ..utils import get_model_path, np2tensor, pil2tensor, tensor2np
|
||||||
|
|
||||||
|
|
||||||
class LoadFaceEnhanceModel:
|
class MTB_LoadFaceEnhanceModel:
|
||||||
"""Loads a GFPGan or RestoreFormer model for face enhancement."""
|
"""Loads a GFPGan or RestoreFormer model for face enhancement."""
|
||||||
|
|
||||||
def __init__(self) -> None:
|
def __init__(self) -> None:
|
||||||
@@ -37,7 +34,9 @@ class LoadFaceEnhanceModel:
|
|||||||
fr_models_path, um_models_path = cls.get_models_root()
|
fr_models_path, um_models_path = cls.get_models_root()
|
||||||
|
|
||||||
if fr_models_path is None and um_models_path is None:
|
if fr_models_path is None and um_models_path is None:
|
||||||
log.warning("Face restoration models not found.")
|
if not hasattr(cls, "_warned"):
|
||||||
|
log.warning("Face restoration models not found.")
|
||||||
|
cls._warned = True
|
||||||
return []
|
return []
|
||||||
if not fr_models_path.exists():
|
if not fr_models_path.exists():
|
||||||
# log.warning(
|
# log.warning(
|
||||||
@@ -81,6 +80,8 @@ class LoadFaceEnhanceModel:
|
|||||||
CATEGORY = "mtb/facetools"
|
CATEGORY = "mtb/facetools"
|
||||||
|
|
||||||
def load_model(self, model_name, upscale=2, bg_upsampler=None):
|
def load_model(self, model_name, upscale=2, bg_upsampler=None):
|
||||||
|
from gfpgan import GFPGANer
|
||||||
|
|
||||||
basic = "RestoreFormer" not in model_name
|
basic = "RestoreFormer" not in model_name
|
||||||
|
|
||||||
fr_root, um_root = self.get_models_root()
|
fr_root, um_root = self.get_models_root()
|
||||||
@@ -153,7 +154,7 @@ class BGUpscaleWrapper:
|
|||||||
import sys
|
import sys
|
||||||
|
|
||||||
|
|
||||||
class RestoreFace:
|
class MTB_RestoreFace:
|
||||||
"""Uses GFPGan to restore faces"""
|
"""Uses GFPGan to restore faces"""
|
||||||
|
|
||||||
def __init__(self) -> None:
|
def __init__(self) -> None:
|
||||||
@@ -182,7 +183,7 @@ class RestoreFace:
|
|||||||
def do_restore(
|
def do_restore(
|
||||||
self,
|
self,
|
||||||
image: torch.Tensor,
|
image: torch.Tensor,
|
||||||
model: GFPGANer,
|
model,
|
||||||
aligned,
|
aligned,
|
||||||
only_center_face,
|
only_center_face,
|
||||||
weight,
|
weight,
|
||||||
@@ -220,12 +221,12 @@ class RestoreFace:
|
|||||||
def restore(
|
def restore(
|
||||||
self,
|
self,
|
||||||
image: torch.Tensor,
|
image: torch.Tensor,
|
||||||
model: GFPGANer,
|
model,
|
||||||
aligned=False,
|
aligned=False,
|
||||||
only_center_face=False,
|
only_center_face=False,
|
||||||
weight=0.5,
|
weight=0.5,
|
||||||
save_tmp_steps=True,
|
save_tmp_steps=True,
|
||||||
) -> Tuple[torch.Tensor]:
|
) -> tuple[torch.Tensor]:
|
||||||
out = [
|
out = [
|
||||||
self.do_restore(
|
self.do_restore(
|
||||||
image[i],
|
image[i],
|
||||||
@@ -275,4 +276,4 @@ class RestoreFace:
|
|||||||
cv2.imwrite(file, cmp_img)
|
cv2.imwrite(file, cmp_img)
|
||||||
|
|
||||||
|
|
||||||
__nodes__ = [RestoreFace, LoadFaceEnhanceModel]
|
__nodes__ = [MTB_RestoreFace, MTB_LoadFaceEnhanceModel]
|
||||||
|
|||||||
+4
-4
@@ -22,7 +22,7 @@ from ..utils import download_antelopev2, get_model_path, pil2tensor, tensor2pil
|
|||||||
log = mklog(__name__)
|
log = mklog(__name__)
|
||||||
|
|
||||||
|
|
||||||
class LoadFaceAnalysisModel:
|
class MTB_LoadFaceAnalysisModel:
|
||||||
"""Loads a face analysis model"""
|
"""Loads a face analysis model"""
|
||||||
|
|
||||||
models = []
|
models = []
|
||||||
@@ -53,7 +53,7 @@ class LoadFaceAnalysisModel:
|
|||||||
return (face_analyser,)
|
return (face_analyser,)
|
||||||
|
|
||||||
|
|
||||||
class LoadFaceSwapModel:
|
class MTB_LoadFaceSwapModel:
|
||||||
"""Loads a faceswap model"""
|
"""Loads a faceswap model"""
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@@ -97,7 +97,7 @@ class LoadFaceSwapModel:
|
|||||||
|
|
||||||
|
|
||||||
# region roop node
|
# region roop node
|
||||||
class FaceSwap:
|
class MTB_FaceSwap:
|
||||||
"""Face swap using deepinsight/insightface models"""
|
"""Face swap using deepinsight/insightface models"""
|
||||||
|
|
||||||
model = None
|
model = None
|
||||||
@@ -239,4 +239,4 @@ def swap_face(
|
|||||||
# endregion face swap utils
|
# endregion face swap utils
|
||||||
|
|
||||||
|
|
||||||
__nodes__ = [FaceSwap, LoadFaceSwapModel, LoadFaceAnalysisModel]
|
__nodes__ = [MTB_FaceSwap, MTB_LoadFaceSwapModel, MTB_LoadFaceAnalysisModel]
|
||||||
|
|||||||
+4
-4
@@ -52,7 +52,7 @@ from ..utils import comfy_dir, font_path, pil2tensor
|
|||||||
# return m.digest().hex()
|
# return m.digest().hex()
|
||||||
|
|
||||||
|
|
||||||
class UnsplashImage:
|
class MTB_UnsplashImage:
|
||||||
"""Unsplash Image given a keyword and a size"""
|
"""Unsplash Image given a keyword and a size"""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -113,7 +113,7 @@ class UnsplashImage:
|
|||||||
return (None,)
|
return (None,)
|
||||||
|
|
||||||
|
|
||||||
class QrCode:
|
class MTB_QrCode:
|
||||||
"""Basic QR Code generator"""
|
"""Basic QR Code generator"""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -364,8 +364,8 @@ by default it fallsback to a default font.
|
|||||||
|
|
||||||
|
|
||||||
__nodes__ = [
|
__nodes__ = [
|
||||||
QrCode,
|
MTB_QrCode,
|
||||||
UnsplashImage,
|
MTB_UnsplashImage,
|
||||||
MTB_TextToImage,
|
MTB_TextToImage,
|
||||||
# MtbExamples,
|
# MtbExamples,
|
||||||
]
|
]
|
||||||
|
|||||||
+45
-28
@@ -14,6 +14,7 @@ from PIL import Image
|
|||||||
|
|
||||||
from ..log import log
|
from ..log import log
|
||||||
from ..utils import (
|
from ..utils import (
|
||||||
|
EASINGS,
|
||||||
apply_easing,
|
apply_easing,
|
||||||
get_server_info,
|
get_server_info,
|
||||||
numpy_NFOV,
|
numpy_NFOV,
|
||||||
@@ -168,11 +169,48 @@ class MTB_MatchDimensions:
|
|||||||
return (resized_source, new_width, new_height)
|
return (resized_source, new_width, new_height)
|
||||||
|
|
||||||
|
|
||||||
class MTB_FloatsToFloat:
|
class MTB_FloatToFloats:
|
||||||
"""AD, IPA, Fitz etc have commonly choose to mistype float lists as FLOAT.
|
"""Conversion utility for compatibility with other extensions (AD, IPA, Fitz are using FLOAT to represent list of floats.)"""
|
||||||
|
|
||||||
This is just a hack to be compatible with these
|
@classmethod
|
||||||
"""
|
def INPUT_TYPES(cls):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"float": ("FLOAT", {"default": 0.0, "forceInput": True}),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("FLOATS",)
|
||||||
|
RETURN_NAMES = ("floats",)
|
||||||
|
CATEGORY = "mtb/utils"
|
||||||
|
FUNCTION = "convert"
|
||||||
|
|
||||||
|
def convert(self, float: float):
|
||||||
|
return (float,)
|
||||||
|
|
||||||
|
|
||||||
|
class MTB_FloatsToInts:
|
||||||
|
"""Conversion utility for compatibility with frame interpolation."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"floats": ("FLOATS", {"forceInput": True}),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("INTS", "INT")
|
||||||
|
CATEGORY = "mtb/utils"
|
||||||
|
FUNCTION = "convert"
|
||||||
|
|
||||||
|
def convert(self, floats: list[float]):
|
||||||
|
vals = [int(x) for x in floats]
|
||||||
|
return (vals, vals)
|
||||||
|
|
||||||
|
|
||||||
|
class MTB_FloatsToFloat:
|
||||||
|
"""Conversion utility for compatibility with other extensions (AD, IPA, Fitz are using FLOAT to represent list of floats.)"""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(cls):
|
def INPUT_TYPES(cls):
|
||||||
@@ -498,30 +536,7 @@ class MTB_FitNumber:
|
|||||||
"target_min": ("FLOAT", {"default": 0.0, "step": 0.01}),
|
"target_min": ("FLOAT", {"default": 0.0, "step": 0.01}),
|
||||||
"target_max": ("FLOAT", {"default": 1.0, "step": 0.01}),
|
"target_max": ("FLOAT", {"default": 1.0, "step": 0.01}),
|
||||||
"easing": (
|
"easing": (
|
||||||
[
|
EASINGS,
|
||||||
"Linear",
|
|
||||||
"Sine In",
|
|
||||||
"Sine Out",
|
|
||||||
"Sine In/Out",
|
|
||||||
"Quart In",
|
|
||||||
"Quart Out",
|
|
||||||
"Quart In/Out",
|
|
||||||
"Cubic In",
|
|
||||||
"Cubic Out",
|
|
||||||
"Cubic In/Out",
|
|
||||||
"Circ In",
|
|
||||||
"Circ Out",
|
|
||||||
"Circ In/Out",
|
|
||||||
"Back In",
|
|
||||||
"Back Out",
|
|
||||||
"Back In/Out",
|
|
||||||
"Elastic In",
|
|
||||||
"Elastic Out",
|
|
||||||
"Elastic In/Out",
|
|
||||||
"Bounce In",
|
|
||||||
"Bounce Out",
|
|
||||||
"Bounce In/Out",
|
|
||||||
],
|
|
||||||
{"default": "Linear"},
|
{"default": "Linear"},
|
||||||
),
|
),
|
||||||
}
|
}
|
||||||
@@ -596,4 +611,6 @@ __nodes__ = [
|
|||||||
MTB_MatchDimensions,
|
MTB_MatchDimensions,
|
||||||
MTB_AutoPanEquilateral,
|
MTB_AutoPanEquilateral,
|
||||||
MTB_FloatsToFloat,
|
MTB_FloatsToFloat,
|
||||||
|
MTB_FloatToFloats,
|
||||||
|
MTB_FloatsToInts,
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -1,12 +1,9 @@
|
|||||||
import glob
|
|
||||||
import os
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import List
|
from typing import List
|
||||||
|
|
||||||
import comfy
|
import comfy
|
||||||
import comfy.model_management as model_management
|
import comfy.model_management as model_management
|
||||||
import comfy.utils
|
import comfy.utils
|
||||||
import folder_paths
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import tensorflow as tf
|
import tensorflow as tf
|
||||||
import torch
|
import torch
|
||||||
@@ -17,7 +14,7 @@ from ..log import log
|
|||||||
from ..utils import get_model_path
|
from ..utils import get_model_path
|
||||||
|
|
||||||
|
|
||||||
class LoadFilmModel:
|
class MTB_LoadFilmModel:
|
||||||
"""Loads a FILM model"""
|
"""Loads a FILM model"""
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@@ -58,7 +55,7 @@ class LoadFilmModel:
|
|||||||
return (interpolator.Interpolator(model_path.as_posix(), None),)
|
return (interpolator.Interpolator(model_path.as_posix(), None),)
|
||||||
|
|
||||||
|
|
||||||
class FilmInterpolation:
|
class MTB_FilmInterpolation:
|
||||||
"""Google Research FILM frame interpolation for large motion"""
|
"""Google Research FILM frame interpolation for large motion"""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -107,12 +104,16 @@ class FilmInterpolation:
|
|||||||
in_frames, interpolate, film_model
|
in_frames, interpolate, film_model
|
||||||
):
|
):
|
||||||
out_tensors.append(
|
out_tensors.append(
|
||||||
torch.from_numpy(frame) if isinstance(frame, np.ndarray) else frame
|
torch.from_numpy(frame)
|
||||||
|
if isinstance(frame, np.ndarray)
|
||||||
|
else frame
|
||||||
)
|
)
|
||||||
model_management.throw_exception_if_processing_interrupted()
|
model_management.throw_exception_if_processing_interrupted()
|
||||||
pbar.update(1)
|
pbar.update(1)
|
||||||
|
|
||||||
out_tensors = torch.cat([tens.unsqueeze(0) for tens in out_tensors], dim=0)
|
out_tensors = torch.cat(
|
||||||
|
[tens.unsqueeze(0) for tens in out_tensors], dim=0
|
||||||
|
)
|
||||||
|
|
||||||
log.debug(f"Returning {len(out_tensors)} tensors")
|
log.debug(f"Returning {len(out_tensors)} tensors")
|
||||||
log.debug(f"Output shape {out_tensors.shape}")
|
log.debug(f"Output shape {out_tensors.shape}")
|
||||||
@@ -120,4 +121,4 @@ class FilmInterpolation:
|
|||||||
return (out_tensors,)
|
return (out_tensors,)
|
||||||
|
|
||||||
|
|
||||||
__nodes__ = [LoadFilmModel, FilmInterpolation]
|
__nodes__ = [MTB_LoadFilmModel, MTB_FilmInterpolation]
|
||||||
|
|||||||
@@ -86,7 +86,14 @@ class MTB_ColorCorrect:
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def contrast_adjustment_tensor(image, contrast):
|
def contrast_adjustment_tensor(image, contrast):
|
||||||
contrasted = (image - 0.5) * contrast + 0.5
|
r, g, b = image.unbind(-1)
|
||||||
|
|
||||||
|
# Using Adobe RGB luminance weights.
|
||||||
|
luminance_image = 0.33 * r + 0.71 * g + 0.06 * b
|
||||||
|
luminance_mean = torch.mean(luminance_image.unsqueeze(-1))
|
||||||
|
|
||||||
|
# Blend original with mean luminance using contrast factor as blend ratio.
|
||||||
|
contrasted = image * contrast + (1.0 - contrast) * luminance_mean
|
||||||
return torch.clamp(contrasted, 0.0, 1.0)
|
return torch.clamp(contrasted, 0.0, 1.0)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@@ -216,16 +223,55 @@ class MTB_ImageCompare:
|
|||||||
CATEGORY = "mtb/image"
|
CATEGORY = "mtb/image"
|
||||||
|
|
||||||
def compare(self, imageA: torch.Tensor, imageB: torch.Tensor, mode):
|
def compare(self, imageA: torch.Tensor, imageB: torch.Tensor, mode):
|
||||||
imageA = imageA.numpy()
|
if imageA.dim() == 4:
|
||||||
imageB = imageB.numpy()
|
batch_count = imageA.size(0)
|
||||||
|
return (
|
||||||
|
torch.cat(
|
||||||
|
tuple(
|
||||||
|
self.compare(imageA[i], imageB[i], mode)[0]
|
||||||
|
for i in range(batch_count)
|
||||||
|
),
|
||||||
|
dim=0,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
imageA = imageA.squeeze()
|
num_channels_A = imageA.size(2)
|
||||||
imageB = imageB.squeeze()
|
num_channels_B = imageB.size(2)
|
||||||
|
|
||||||
image = compare_images(imageA, imageB, method=mode)
|
# handle RGBA/RGB mismatch
|
||||||
|
if num_channels_A == 3 and num_channels_B == 4:
|
||||||
|
imageA = torch.cat(
|
||||||
|
(imageA, torch.ones_like(imageA[:, :, 0:1])), dim=2
|
||||||
|
)
|
||||||
|
elif num_channels_B == 3 and num_channels_A == 4:
|
||||||
|
imageB = torch.cat(
|
||||||
|
(imageB, torch.ones_like(imageB[:, :, 0:1])), dim=2
|
||||||
|
)
|
||||||
|
match mode:
|
||||||
|
case "diff":
|
||||||
|
compare_image = torch.abs(imageA - imageB)
|
||||||
|
case "blend":
|
||||||
|
compare_image = 0.5 * (imageA + imageB)
|
||||||
|
case "checkerboard":
|
||||||
|
imageA = imageA.numpy()
|
||||||
|
imageB = imageB.numpy()
|
||||||
|
compared_channels = [
|
||||||
|
torch.from_numpy(
|
||||||
|
compare_images(
|
||||||
|
imageA[:, :, i], imageB[:, :, i], method=mode
|
||||||
|
)
|
||||||
|
)
|
||||||
|
for i in range(imageA.shape[2])
|
||||||
|
]
|
||||||
|
|
||||||
image = np.expand_dims(image, axis=0)
|
compare_image = torch.stack(compared_channels, dim=2)
|
||||||
return (torch.from_numpy(image),)
|
case _:
|
||||||
|
compare_image = None
|
||||||
|
raise ValueError(f"Unknown mode {mode}")
|
||||||
|
|
||||||
|
compare_image = compare_image.unsqueeze(0)
|
||||||
|
|
||||||
|
return (compare_image,)
|
||||||
|
|
||||||
|
|
||||||
import requests
|
import requests
|
||||||
|
|||||||
@@ -27,6 +27,11 @@ class MTB_StackImages:
|
|||||||
normalized_tensors = [
|
normalized_tensors = [
|
||||||
self.normalize_to_rgba(tensor) for tensor in tensors
|
self.normalize_to_rgba(tensor) for tensor in tensors
|
||||||
]
|
]
|
||||||
|
max_batch_size = max(tensor.shape[0] for tensor in normalized_tensors)
|
||||||
|
normalized_tensors = [
|
||||||
|
self.duplicate_frames(tensor, max_batch_size)
|
||||||
|
for tensor in normalized_tensors
|
||||||
|
]
|
||||||
|
|
||||||
if vertical:
|
if vertical:
|
||||||
width = normalized_tensors[0].shape[2]
|
width = normalized_tensors[0].shape[2]
|
||||||
@@ -67,6 +72,21 @@ class MTB_StackImages:
|
|||||||
"expected 3 (RGB) or 4 (RGBA)."
|
"expected 3 (RGB) or 4 (RGBA)."
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def duplicate_frames(self, tensor, target_batch_size):
|
||||||
|
"""Duplicate frames in tensor to match the target batch size."""
|
||||||
|
current_batch_size = tensor.shape[0]
|
||||||
|
if current_batch_size < target_batch_size:
|
||||||
|
duplication_factors: int = target_batch_size // current_batch_size
|
||||||
|
duplicated_tensor = tensor.repeat(duplication_factors, 1, 1, 1)
|
||||||
|
remaining_frames = target_batch_size % current_batch_size
|
||||||
|
if remaining_frames > 0:
|
||||||
|
duplicated_tensor = torch.cat(
|
||||||
|
(duplicated_tensor, tensor[:remaining_frames]), dim=0
|
||||||
|
)
|
||||||
|
return duplicated_tensor
|
||||||
|
else:
|
||||||
|
return tensor
|
||||||
|
|
||||||
|
|
||||||
class MTB_PickFromBatch:
|
class MTB_PickFromBatch:
|
||||||
"""Pick a specific number of images from a batch.
|
"""Pick a specific number of images from a batch.
|
||||||
|
|||||||
+2
-2
@@ -21,7 +21,7 @@ def get_playlist_path(playlist_name: str, persistant_playlist=False):
|
|||||||
return output_dir / "playlists" / session_id / f"{playlist_name}.json"
|
return output_dir / "playlists" / session_id / f"{playlist_name}.json"
|
||||||
|
|
||||||
|
|
||||||
class ReadPlaylist:
|
class MTB_ReadPlaylist:
|
||||||
"""Read a playlist"""
|
"""Read a playlist"""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -399,5 +399,5 @@ __nodes__ = [
|
|||||||
MTB_SaveGif,
|
MTB_SaveGif,
|
||||||
MTB_ExportWithFfmpeg,
|
MTB_ExportWithFfmpeg,
|
||||||
MTB_AddToPlaylist,
|
MTB_AddToPlaylist,
|
||||||
ReadPlaylist,
|
MTB_ReadPlaylist,
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
|
||||||
class LatentLerp:
|
class MTB_LatentLerp:
|
||||||
"""Linear interpolation (blend) between two latent vectors"""
|
"""Linear interpolation (blend) between two latent vectors"""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -10,7 +10,10 @@ class LatentLerp:
|
|||||||
"required": {
|
"required": {
|
||||||
"A": ("LATENT",),
|
"A": ("LATENT",),
|
||||||
"B": ("LATENT",),
|
"B": ("LATENT",),
|
||||||
"t": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
|
"t": (
|
||||||
|
"FLOAT",
|
||||||
|
{"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01},
|
||||||
|
),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -29,5 +32,5 @@ class LatentLerp:
|
|||||||
|
|
||||||
|
|
||||||
__nodes__ = [
|
__nodes__ = [
|
||||||
LatentLerp,
|
MTB_LatentLerp,
|
||||||
]
|
]
|
||||||
|
|||||||
+2
-2
@@ -80,7 +80,7 @@ def conv_forward(lyr, tensor, weight, bias):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
class ModelPatchSeamless:
|
class MTB_ModelPatchSeamless:
|
||||||
"""Uses the stable diffusion 'hack' to infer seamless images by setting the model layers padding mode to circular (experimental)"""
|
"""Uses the stable diffusion 'hack' to infer seamless images by setting the model layers padding mode to circular (experimental)"""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -152,4 +152,4 @@ class ModelPatchSeamless:
|
|||||||
return (model, hacked_model)
|
return (model, hacked_model)
|
||||||
|
|
||||||
|
|
||||||
__nodes__ = [ModelPatchSeamless, MTB_VaeDecode]
|
__nodes__ = [MTB_ModelPatchSeamless, MTB_VaeDecode]
|
||||||
|
|||||||
@@ -77,8 +77,6 @@ class MTB_FloatToNumber:
|
|||||||
def float_to_number(self, float):
|
def float_to_number(self, float):
|
||||||
return (float,)
|
return (float,)
|
||||||
|
|
||||||
return (int,)
|
|
||||||
|
|
||||||
|
|
||||||
__nodes__ = [
|
__nodes__ = [
|
||||||
MTB_FloatToNumber,
|
MTB_FloatToNumber,
|
||||||
|
|||||||
+360
@@ -0,0 +1,360 @@
|
|||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import safetensors.torch
|
||||||
|
import torch
|
||||||
|
import tqdm
|
||||||
|
|
||||||
|
from ..log import log
|
||||||
|
from ..utils import Operation, Precision
|
||||||
|
from ..utils import output_dir as comfy_out_dir
|
||||||
|
|
||||||
|
PRUNE_DATA = {
|
||||||
|
"known_junk_prefix": [
|
||||||
|
"embedding_manager.embedder.",
|
||||||
|
"lora_te_text_model",
|
||||||
|
"control_model.",
|
||||||
|
],
|
||||||
|
"nai_keys": {
|
||||||
|
"cond_stage_model.transformer.embeddings.": "cond_stage_model.transformer.text_model.embeddings.",
|
||||||
|
"cond_stage_model.transformer.encoder.": "cond_stage_model.transformer.text_model.encoder.",
|
||||||
|
"cond_stage_model.transformer.final_layer_norm.": "cond_stage_model.transformer.text_model.final_layer_norm.",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
# position_ids in clip is int64. model_ema.num_updates is int32
|
||||||
|
dtypes_to_fp16 = {torch.float32, torch.float64, torch.bfloat16}
|
||||||
|
dtypes_to_bf16 = {torch.float32, torch.float64, torch.float16}
|
||||||
|
dtypes_to_fp8 = {torch.float32, torch.float64, torch.bfloat16, torch.float16}
|
||||||
|
|
||||||
|
|
||||||
|
class MTB_ModelPruner:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
return {
|
||||||
|
"optional": {
|
||||||
|
"unet": ("MODEL",),
|
||||||
|
"clip": ("CLIP",),
|
||||||
|
"vae": ("VAE",),
|
||||||
|
},
|
||||||
|
"required": {
|
||||||
|
"save_separately": ("BOOLEAN", {"default": False}),
|
||||||
|
"save_folder": ("STRING", {"default": "checkpoints/ComfyUI"}),
|
||||||
|
"fix_clip": ("BOOLEAN", {"default": True}),
|
||||||
|
"remove_junk": ("BOOLEAN", {"default": True}),
|
||||||
|
"ema_mode": (
|
||||||
|
("disabled", "remove_ema", "ema_only"),
|
||||||
|
{"default": "remove_ema"},
|
||||||
|
),
|
||||||
|
"precision_unet": (
|
||||||
|
Precision.list_members(),
|
||||||
|
{"default": Precision.FULL.value},
|
||||||
|
),
|
||||||
|
"operation_unet": (
|
||||||
|
Operation.list_members(),
|
||||||
|
{"default": Operation.CONVERT.value},
|
||||||
|
),
|
||||||
|
"precision_clip": (
|
||||||
|
Precision.list_members(),
|
||||||
|
{"default": Precision.FULL.value},
|
||||||
|
),
|
||||||
|
"operation_clip": (
|
||||||
|
Operation.list_members(),
|
||||||
|
{"default": Operation.CONVERT.value},
|
||||||
|
),
|
||||||
|
"precision_vae": (
|
||||||
|
Precision.list_members(),
|
||||||
|
{"default": Precision.FULL.value},
|
||||||
|
),
|
||||||
|
"operation_vae": (
|
||||||
|
Operation.list_members(),
|
||||||
|
{"default": Operation.CONVERT.value},
|
||||||
|
),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
OUTPUT_NODE = True
|
||||||
|
RETURN_TYPES = ()
|
||||||
|
CATEGORY = "mtb/prune"
|
||||||
|
FUNCTION = "prune"
|
||||||
|
|
||||||
|
def convert_precision(self, tensor: torch.Tensor, precision: Precision):
|
||||||
|
precision = Precision.from_str(precision)
|
||||||
|
log.debug(f"Converting to {precision}")
|
||||||
|
match precision:
|
||||||
|
case Precision.FP8:
|
||||||
|
if tensor.dtype in dtypes_to_fp8:
|
||||||
|
return tensor.to(torch.float8_e4m3fn)
|
||||||
|
log.error(f"Cannot convert {tensor.dtype} to fp8")
|
||||||
|
return tensor
|
||||||
|
case Precision.FP16:
|
||||||
|
if tensor.dtype in dtypes_to_fp16:
|
||||||
|
return tensor.half()
|
||||||
|
log.error(f"Cannot convert {tensor.dtype} to f16")
|
||||||
|
return tensor
|
||||||
|
case Precision.BF16:
|
||||||
|
if tensor.dtype in dtypes_to_bf16:
|
||||||
|
return tensor.bfloat16()
|
||||||
|
log.error(f"Cannot convert {tensor.dtype} to bf16")
|
||||||
|
return tensor
|
||||||
|
case Precision.FULL | Precision.FP32:
|
||||||
|
return tensor
|
||||||
|
|
||||||
|
def is_sdxl_model(self, clip: dict[str, torch.Tensor] | None):
|
||||||
|
if clip:
|
||||||
|
return (any(k.startswith("conditioner.embedders") for k in clip),)
|
||||||
|
return False
|
||||||
|
|
||||||
|
def has_ema(self, unet: dict[str, torch.Tensor]):
|
||||||
|
return any(k.startswith("model_ema") for k in unet)
|
||||||
|
|
||||||
|
def fix_clip(self, clip: dict[str, torch.Tensor] | None):
|
||||||
|
if self.is_sdxl_model(clip):
|
||||||
|
log.warn("[fix clip] SDXL not supported")
|
||||||
|
return
|
||||||
|
|
||||||
|
if clip is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
position_id_key = (
|
||||||
|
"cond_stage_model.transformer.text_model.embeddings.position_ids"
|
||||||
|
)
|
||||||
|
if position_id_key in clip:
|
||||||
|
correct = torch.Tensor([list(range(77))]).to(torch.int64)
|
||||||
|
now = clip[position_id_key].to(torch.int64)
|
||||||
|
|
||||||
|
broken = correct.ne(now)
|
||||||
|
broken = [i for i in range(77) if broken[0][i]]
|
||||||
|
|
||||||
|
if len(broken) != 0:
|
||||||
|
clip[position_id_key] = correct
|
||||||
|
log.info(f"[Converter] Fixed broken clip\n{broken}")
|
||||||
|
else:
|
||||||
|
log.info(
|
||||||
|
"[Converter] Clip in this model is fine, skip fixing..."
|
||||||
|
)
|
||||||
|
|
||||||
|
else:
|
||||||
|
log.info("[Converter] Missing position id in model, try fixing...")
|
||||||
|
clip[position_id_key] = torch.Tensor([list(range(77))]).to(
|
||||||
|
torch.int64
|
||||||
|
)
|
||||||
|
return clip
|
||||||
|
|
||||||
|
def get_dicts(self, unet, clip, vae):
|
||||||
|
clip_sd = clip.get_sd()
|
||||||
|
state_dict = unet.model.state_dict_for_saving(
|
||||||
|
clip_sd, vae.get_sd(), None
|
||||||
|
)
|
||||||
|
|
||||||
|
unet = {
|
||||||
|
k: v
|
||||||
|
for k, v in state_dict.items()
|
||||||
|
if k.startswith("model.diffusion_model")
|
||||||
|
}
|
||||||
|
clip = {
|
||||||
|
k: v
|
||||||
|
for k, v in state_dict.items()
|
||||||
|
if k.startswith("cond_stage_model")
|
||||||
|
or k.startswith("conditioner.embedders")
|
||||||
|
}
|
||||||
|
vae = {
|
||||||
|
k: v
|
||||||
|
for k, v in state_dict.items()
|
||||||
|
if k.startswith("first_stage_model")
|
||||||
|
}
|
||||||
|
|
||||||
|
other = {
|
||||||
|
k: v
|
||||||
|
for k, v in state_dict.items()
|
||||||
|
if k not in unet and k not in vae and k not in clip
|
||||||
|
}
|
||||||
|
|
||||||
|
return (unet, clip, vae, other)
|
||||||
|
|
||||||
|
def do_remove_junk(self, tensors: dict[str, dict[str, torch.Tensor]]):
|
||||||
|
need_delete: list[str] = []
|
||||||
|
for layer in tensors:
|
||||||
|
for key in layer:
|
||||||
|
for jk in PRUNE_DATA["known_junk_prefix"]:
|
||||||
|
if key.startswith(jk):
|
||||||
|
need_delete.append(".".join([layer, key]))
|
||||||
|
|
||||||
|
for k in need_delete:
|
||||||
|
log.info(f"Removing junk data: {k}")
|
||||||
|
del tensors[k]
|
||||||
|
|
||||||
|
return tensors
|
||||||
|
|
||||||
|
def prune(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
save_separately: bool,
|
||||||
|
save_folder: str,
|
||||||
|
fix_clip: bool,
|
||||||
|
remove_junk: bool,
|
||||||
|
ema_mode: str,
|
||||||
|
precision_unet: Precision,
|
||||||
|
precision_clip: Precision,
|
||||||
|
precision_vae: Precision,
|
||||||
|
operation_unet: str,
|
||||||
|
operation_clip: str,
|
||||||
|
operation_vae: str,
|
||||||
|
unet: dict[str, torch.Tensor] | None = None,
|
||||||
|
clip: dict[str, torch.Tensor] | None = None,
|
||||||
|
vae: dict[str, torch.Tensor] | None = None,
|
||||||
|
):
|
||||||
|
operation = {
|
||||||
|
"unet": Operation.from_str(operation_unet),
|
||||||
|
"clip": Operation.from_str(operation_clip),
|
||||||
|
"vae": Operation.from_str(operation_vae),
|
||||||
|
}
|
||||||
|
precision = {
|
||||||
|
"unet": Precision.from_str(precision_unet),
|
||||||
|
"clip": Precision.from_str(precision_clip),
|
||||||
|
"vae": Precision.from_str(precision_vae),
|
||||||
|
}
|
||||||
|
|
||||||
|
unet, clip, vae, _other = self.get_dicts(unet, clip, vae)
|
||||||
|
|
||||||
|
out_dir = Path(save_folder)
|
||||||
|
folder = out_dir.parent
|
||||||
|
if not out_dir.is_absolute():
|
||||||
|
folder = (comfy_out_dir / save_folder).parent
|
||||||
|
|
||||||
|
if not folder.exists():
|
||||||
|
if folder.parent.exists():
|
||||||
|
folder.mkdir()
|
||||||
|
else:
|
||||||
|
raise FileNotFoundError(
|
||||||
|
f"Folder {folder.parent} does not exist"
|
||||||
|
)
|
||||||
|
|
||||||
|
name = out_dir.name
|
||||||
|
save_name = f"{name}-{precision_unet}"
|
||||||
|
if ema_mode != "disabled":
|
||||||
|
save_name += f"-{ema_mode}"
|
||||||
|
if fix_clip:
|
||||||
|
save_name += "-clip-fix"
|
||||||
|
|
||||||
|
if (
|
||||||
|
any(o == Operation.CONVERT for o in operation.values())
|
||||||
|
and any(p == Precision.FP8 for p in precision.values())
|
||||||
|
and torch.__version__ < "2.1.0"
|
||||||
|
):
|
||||||
|
raise NotImplementedError(
|
||||||
|
"PyTorch 2.1.0 or newer is required for fp8 conversion"
|
||||||
|
)
|
||||||
|
|
||||||
|
if not self.is_sdxl_model(clip):
|
||||||
|
for part in [unet, vae, clip]:
|
||||||
|
if part:
|
||||||
|
nai_keys = PRUNE_DATA["nai_keys"]
|
||||||
|
for k in list(part.keys()):
|
||||||
|
for r in nai_keys:
|
||||||
|
if isinstance(k, str) and k.startswith(r):
|
||||||
|
new_key = k.replace(r, nai_keys[r])
|
||||||
|
part[new_key] = part[k]
|
||||||
|
del part[k]
|
||||||
|
log.info(
|
||||||
|
f"[Converter] Fixed novelai error key {k}"
|
||||||
|
)
|
||||||
|
break
|
||||||
|
|
||||||
|
if fix_clip:
|
||||||
|
clip = self.fix_clip(clip)
|
||||||
|
|
||||||
|
ok: dict[str, dict[str, torch.Tensor]] = {
|
||||||
|
"unet": {},
|
||||||
|
"clip": {},
|
||||||
|
"vae": {},
|
||||||
|
}
|
||||||
|
|
||||||
|
def _hf(part: str, wk: str, t: torch.Tensor):
|
||||||
|
if not isinstance(t, torch.Tensor):
|
||||||
|
log.debug("Not a torch tensor, skipping key")
|
||||||
|
return
|
||||||
|
|
||||||
|
log.debug(f"Operation {operation[part]}")
|
||||||
|
if operation[part] == Operation.CONVERT:
|
||||||
|
ok[part][wk] = self.convert_precision(
|
||||||
|
t, precision[part]
|
||||||
|
) # conv_func(t)
|
||||||
|
elif operation[part] == Operation.COPY:
|
||||||
|
ok[part][wk] = t
|
||||||
|
elif operation[part] == Operation.DELETE:
|
||||||
|
return
|
||||||
|
|
||||||
|
log.info("[Converter] Converting model...")
|
||||||
|
|
||||||
|
for part_name, part in zip(
|
||||||
|
["unet", "vae", "clip", "other"],
|
||||||
|
[unet, vae, clip],
|
||||||
|
strict=False,
|
||||||
|
):
|
||||||
|
if part:
|
||||||
|
match ema_mode:
|
||||||
|
case "remove_ema":
|
||||||
|
for k, v in tqdm.tqdm(part.items()):
|
||||||
|
if "model_ema." not in k:
|
||||||
|
_hf(part_name, k, v)
|
||||||
|
case "ema_only":
|
||||||
|
if not self.has_ema(part):
|
||||||
|
log.warn("No EMA to extract")
|
||||||
|
return
|
||||||
|
for k in tqdm.tqdm(part):
|
||||||
|
ema_k = "___"
|
||||||
|
try:
|
||||||
|
ema_k = "model_ema." + k[6:].replace(".", "")
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
if ema_k in part:
|
||||||
|
_hf(part_name, k, part[ema_k])
|
||||||
|
elif not k.startswith("model_ema.") or k in [
|
||||||
|
"model_ema.num_updates",
|
||||||
|
"model_ema.decay",
|
||||||
|
]:
|
||||||
|
_hf(part_name, k, part[k])
|
||||||
|
case "disabled" | _:
|
||||||
|
for k, v in tqdm.tqdm(part.items()):
|
||||||
|
_hf(part_name, k, v)
|
||||||
|
|
||||||
|
if save_separately:
|
||||||
|
if remove_junk:
|
||||||
|
ok = self.do_remove_junk(ok)
|
||||||
|
|
||||||
|
flat_ok = {
|
||||||
|
k: v
|
||||||
|
for _, subdict in ok.items()
|
||||||
|
for k, v in subdict.items()
|
||||||
|
}
|
||||||
|
save_path = (
|
||||||
|
folder / f"{part_name}-{save_name}.safetensors"
|
||||||
|
).as_posix()
|
||||||
|
safetensors.torch.save_file(flat_ok, save_path)
|
||||||
|
ok: dict[str, dict[str, torch.Tensor]] = {
|
||||||
|
"unet": {},
|
||||||
|
"clip": {},
|
||||||
|
"vae": {},
|
||||||
|
}
|
||||||
|
|
||||||
|
if save_separately:
|
||||||
|
return ()
|
||||||
|
|
||||||
|
if remove_junk:
|
||||||
|
ok = self.do_remove_junk(ok)
|
||||||
|
|
||||||
|
flat_ok = {
|
||||||
|
k: v for _, subdict in ok.items() for k, v in subdict.items()
|
||||||
|
}
|
||||||
|
|
||||||
|
try:
|
||||||
|
safetensors.torch.save_file(
|
||||||
|
flat_ok, (folder / f"{save_name}.safetensors").as_posix()
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
log.error(e)
|
||||||
|
|
||||||
|
return ()
|
||||||
|
|
||||||
|
|
||||||
|
__nodes__ = [MTB_ModelPruner]
|
||||||
+42
-15
@@ -1,15 +1,16 @@
|
|||||||
|
from math import ceil, sqrt
|
||||||
|
from typing import cast
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torchvision.transforms.functional as TF
|
import torchvision.transforms.functional as TF
|
||||||
from ..utils import log, hex_to_rgb, tensor2pil, pil2tensor
|
|
||||||
from math import sqrt, ceil
|
|
||||||
from typing import cast
|
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
|
|
||||||
|
from ..utils import hex_to_rgb, log, pil2tensor, tensor2pil
|
||||||
|
|
||||||
class TransformImage:
|
|
||||||
|
class MTB_TransformImage:
|
||||||
"""Save torch tensors (image, mask or latent) to disk, useful to debug things outside comfy
|
"""Save torch tensors (image, mask or latent) to disk, useful to debug things outside comfy
|
||||||
|
|
||||||
|
|
||||||
it return a tensor representing the transformed images with the same shape as the input tensor
|
it return a tensor representing the transformed images with the same shape as the input tensor
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@@ -18,10 +19,22 @@ class TransformImage:
|
|||||||
return {
|
return {
|
||||||
"required": {
|
"required": {
|
||||||
"image": ("IMAGE",),
|
"image": ("IMAGE",),
|
||||||
"x": ("FLOAT", {"default": 0, "step": 1, "min": -4096, "max": 4096}),
|
"x": (
|
||||||
"y": ("FLOAT", {"default": 0, "step": 1, "min": -4096, "max": 4096}),
|
"FLOAT",
|
||||||
"zoom": ("FLOAT", {"default": 1.0, "min": 0.001, "step": 0.01}),
|
{"default": 0, "step": 1, "min": -4096, "max": 4096},
|
||||||
"angle": ("FLOAT", {"default": 0, "step": 1, "min": -360, "max": 360}),
|
),
|
||||||
|
"y": (
|
||||||
|
"FLOAT",
|
||||||
|
{"default": 0, "step": 1, "min": -4096, "max": 4096},
|
||||||
|
),
|
||||||
|
"zoom": (
|
||||||
|
"FLOAT",
|
||||||
|
{"default": 1.0, "min": 0.001, "step": 0.01},
|
||||||
|
),
|
||||||
|
"angle": (
|
||||||
|
"FLOAT",
|
||||||
|
{"default": 0, "step": 1, "min": -360, "max": 360},
|
||||||
|
),
|
||||||
"shear": (
|
"shear": (
|
||||||
"FLOAT",
|
"FLOAT",
|
||||||
{"default": 0, "step": 1, "min": -4096, "max": 4096},
|
{"default": 0, "step": 1, "min": -4096, "max": 4096},
|
||||||
@@ -53,14 +66,21 @@ class TransformImage:
|
|||||||
y = int(y)
|
y = int(y)
|
||||||
angle = int(angle)
|
angle = int(angle)
|
||||||
|
|
||||||
log.debug(f"Zoom: {zoom} | x: {x}, y: {y}, angle: {angle}, shear: {shear}")
|
log.debug(
|
||||||
|
f"Zoom: {zoom} | x: {x}, y: {y}, angle: {angle}, shear: {shear}"
|
||||||
|
)
|
||||||
|
|
||||||
if image.size(0) == 0:
|
if image.size(0) == 0:
|
||||||
return (torch.zeros(0),)
|
return (torch.zeros(0),)
|
||||||
transformed_images = []
|
transformed_images = []
|
||||||
frames_count, frame_height, frame_width, frame_channel_count = image.size()
|
frames_count, frame_height, frame_width, frame_channel_count = (
|
||||||
|
image.size()
|
||||||
|
)
|
||||||
|
|
||||||
new_height, new_width = int(frame_height * zoom), int(frame_width * zoom)
|
new_height, new_width = (
|
||||||
|
int(frame_height * zoom),
|
||||||
|
int(frame_width * zoom),
|
||||||
|
)
|
||||||
|
|
||||||
log.debug(f"New height: {new_height}, New width: {new_width}")
|
log.debug(f"New height: {new_height}, New width: {new_width}")
|
||||||
|
|
||||||
@@ -74,7 +94,12 @@ class TransformImage:
|
|||||||
pw += abs(max_padding)
|
pw += abs(max_padding)
|
||||||
ph += abs(max_padding)
|
ph += abs(max_padding)
|
||||||
|
|
||||||
padding = [max(0, pw + x), max(0, ph + y), max(0, pw - x), max(0, ph - y)]
|
padding = [
|
||||||
|
max(0, pw + x),
|
||||||
|
max(0, ph + y),
|
||||||
|
max(0, pw - x),
|
||||||
|
max(0, ph - y),
|
||||||
|
]
|
||||||
|
|
||||||
constant_color = hex_to_rgb(constant_color)
|
constant_color = hex_to_rgb(constant_color)
|
||||||
log.debug(f"Fill Tuple: {constant_color}")
|
log.debug(f"Fill Tuple: {constant_color}")
|
||||||
@@ -89,7 +114,9 @@ class TransformImage:
|
|||||||
|
|
||||||
img = cast(
|
img = cast(
|
||||||
Image.Image,
|
Image.Image,
|
||||||
TF.affine(img, angle=angle, scale=zoom, translate=[x, y], shear=shear),
|
TF.affine(
|
||||||
|
img, angle=angle, scale=zoom, translate=[x, y], shear=shear
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
left = abs(padding[0])
|
left = abs(padding[0])
|
||||||
@@ -107,4 +134,4 @@ class TransformImage:
|
|||||||
return (pil2tensor(transformed_images),)
|
return (pil2tensor(transformed_images),)
|
||||||
|
|
||||||
|
|
||||||
__nodes__ = [TransformImage]
|
__nodes__ = [MTB_TransformImage]
|
||||||
|
|||||||
+99
-25
@@ -1,4 +1,7 @@
|
|||||||
import hashlib, json, os, re
|
import hashlib
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import re
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import folder_paths
|
import folder_paths
|
||||||
@@ -10,11 +13,12 @@ from PIL.PngImagePlugin import PngInfo
|
|||||||
from ..log import log
|
from ..log import log
|
||||||
|
|
||||||
|
|
||||||
class LoadImageSequence:
|
class MTB_LoadImageSequence:
|
||||||
"""Load an image sequence from a folder. The current frame is used to determine which image to load.
|
"""Load an image sequence from a folder. The current frame is used to determine which image to load.
|
||||||
|
|
||||||
Usually used in conjunction with the `Primitive` node set to increment to load a sequence of images from a folder.
|
Usually used in conjunction with the `Primitive` node set to increment to load a sequence of images from a folder.
|
||||||
Use -1 to load all matching frames as a batch.
|
Use -1 to load all matching frames as a batch.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -26,7 +30,10 @@ class LoadImageSequence:
|
|||||||
"INT",
|
"INT",
|
||||||
{"default": 0, "min": -1, "max": 9999999},
|
{"default": 0, "min": -1, "max": 9999999},
|
||||||
),
|
),
|
||||||
}
|
},
|
||||||
|
"optional": {
|
||||||
|
"range": ("STRING", {"default": ""}),
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
CATEGORY = "mtb/IO"
|
CATEGORY = "mtb/IO"
|
||||||
@@ -35,17 +42,28 @@ class LoadImageSequence:
|
|||||||
"IMAGE",
|
"IMAGE",
|
||||||
"MASK",
|
"MASK",
|
||||||
"INT",
|
"INT",
|
||||||
|
"INT",
|
||||||
)
|
)
|
||||||
RETURN_NAMES = (
|
RETURN_NAMES = (
|
||||||
"image",
|
"image",
|
||||||
"mask",
|
"mask",
|
||||||
"current_frame",
|
"current_frame",
|
||||||
|
"total_frames",
|
||||||
)
|
)
|
||||||
|
|
||||||
def load_image(self, path=None, current_frame=0):
|
def load_image(self, path=None, current_frame=0, range=""):
|
||||||
load_all = current_frame == -1
|
load_all = current_frame == -1
|
||||||
|
total_frames = 1
|
||||||
|
|
||||||
if load_all:
|
if range:
|
||||||
|
frames = self.get_frames_from_range(path, range)
|
||||||
|
imgs, masks = zip(*(img_from_path(frame) for frame in frames))
|
||||||
|
out_img = torch.cat(imgs, dim=0)
|
||||||
|
out_mask = torch.cat(masks, dim=0)
|
||||||
|
total_frames = len(imgs)
|
||||||
|
return (out_img, out_mask, -1, total_frames)
|
||||||
|
|
||||||
|
elif load_all:
|
||||||
log.debug(f"Loading all frames from {path}")
|
log.debug(f"Loading all frames from {path}")
|
||||||
frames = resolve_all_frames(path)
|
frames = resolve_all_frames(path)
|
||||||
log.debug(f"Found {len(frames)} frames")
|
log.debug(f"Found {len(frames)} frames")
|
||||||
@@ -53,33 +71,72 @@ class LoadImageSequence:
|
|||||||
imgs = []
|
imgs = []
|
||||||
masks = []
|
masks = []
|
||||||
|
|
||||||
for frame in frames:
|
imgs, masks = zip(*(img_from_path(frame) for frame in frames))
|
||||||
img, mask = img_from_path(frame)
|
|
||||||
imgs.append(img)
|
|
||||||
masks.append(mask)
|
|
||||||
|
|
||||||
out_img = torch.cat(imgs, dim=0)
|
out_img = torch.cat(imgs, dim=0)
|
||||||
out_mask = torch.cat(masks, dim=0)
|
out_mask = torch.cat(masks, dim=0)
|
||||||
|
total_frames = len(imgs)
|
||||||
|
|
||||||
return (
|
return (out_img, out_mask, -1, total_frames)
|
||||||
out_img,
|
|
||||||
out_mask,
|
|
||||||
)
|
|
||||||
|
|
||||||
log.debug(f"Loading image: {path}, {current_frame}")
|
log.debug(f"Loading image: {path}, {current_frame}")
|
||||||
print(f"Loading image: {path}, {current_frame}")
|
|
||||||
resolved_path = resolve_path(path, current_frame)
|
resolved_path = resolve_path(path, current_frame)
|
||||||
image_path = folder_paths.get_annotated_filepath(resolved_path)
|
image_path = folder_paths.get_annotated_filepath(resolved_path)
|
||||||
image, mask = img_from_path(image_path)
|
image, mask = img_from_path(image_path)
|
||||||
return (
|
return (image, mask, current_frame, total_frames)
|
||||||
image,
|
|
||||||
mask,
|
def get_frames_from_range(self, path, range_str):
|
||||||
current_frame,
|
try:
|
||||||
)
|
start, end = map(int, range_str.split("-"))
|
||||||
|
except ValueError:
|
||||||
|
raise ValueError(
|
||||||
|
f"Invalid range format: {range_str}. Expected format is 'start-end'."
|
||||||
|
)
|
||||||
|
|
||||||
|
frames = resolve_all_frames(path)
|
||||||
|
total_frames = len(frames)
|
||||||
|
|
||||||
|
if start < 0 or end >= total_frames:
|
||||||
|
raise ValueError(
|
||||||
|
f"Range {range_str} is out of bounds. Total frames available: {total_frames}"
|
||||||
|
)
|
||||||
|
|
||||||
|
if "#" in path:
|
||||||
|
frame_regex = re.escape(path).replace(r"\#", r"(\d+)")
|
||||||
|
frame_number_regex = re.compile(frame_regex)
|
||||||
|
|
||||||
|
matching_frames = []
|
||||||
|
for frame in frames:
|
||||||
|
match = frame_number_regex.search(frame)
|
||||||
|
|
||||||
|
if match:
|
||||||
|
frame_number = int(match.group(1))
|
||||||
|
if start <= frame_number <= end:
|
||||||
|
matching_frames.append(frame)
|
||||||
|
|
||||||
|
return matching_frames
|
||||||
|
else:
|
||||||
|
log.warning(
|
||||||
|
f"Wildcard pattern or directory will use indexes instead of frame numbers for : {path}"
|
||||||
|
)
|
||||||
|
|
||||||
|
selected_frames = frames[start : end + 1]
|
||||||
|
|
||||||
|
return selected_frames
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def IS_CHANGED(path="", current_frame=0):
|
def IS_CHANGED(path="", current_frame=0, range=""):
|
||||||
print(f"Checking if changed: {path}, {current_frame}")
|
print(f"Checking if changed: {path}, {current_frame}")
|
||||||
|
if range or current_frame == -1:
|
||||||
|
resolved_paths = resolve_all_frames(path)
|
||||||
|
timestamps = [
|
||||||
|
os.path.getmtime(folder_paths.get_annotated_filepath(p))
|
||||||
|
for p in resolved_paths
|
||||||
|
]
|
||||||
|
combined_hash = hashlib.sha256(
|
||||||
|
"".join(map(str, timestamps)).encode()
|
||||||
|
)
|
||||||
|
return combined_hash.hexdigest()
|
||||||
resolved_path = resolve_path(path, current_frame)
|
resolved_path = resolve_path(path, current_frame)
|
||||||
image_path = folder_paths.get_annotated_filepath(resolved_path)
|
image_path = folder_paths.get_annotated_filepath(resolved_path)
|
||||||
if os.path.exists(image_path):
|
if os.path.exists(image_path):
|
||||||
@@ -119,11 +176,28 @@ def img_from_path(path):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def resolve_all_frames(pattern):
|
def resolve_all_frames(path: str):
|
||||||
|
frames: list[str] = []
|
||||||
|
if "#" not in path:
|
||||||
|
pth = Path(path)
|
||||||
|
if pth.is_dir():
|
||||||
|
for f in pth.iterdir():
|
||||||
|
if f.suffix in [".jpg", ".png"]:
|
||||||
|
frames.append(f.as_posix())
|
||||||
|
elif "*" in path:
|
||||||
|
frames = glob.glob(path)
|
||||||
|
else:
|
||||||
|
raise ValueError(
|
||||||
|
"The path doesn't contain a # or a * or is not a directory"
|
||||||
|
)
|
||||||
|
frames.sort()
|
||||||
|
|
||||||
|
return frames
|
||||||
|
|
||||||
|
pattern = path
|
||||||
folder_path, file_pattern = os.path.split(pattern)
|
folder_path, file_pattern = os.path.split(pattern)
|
||||||
|
|
||||||
log.debug(f"Resolving all frames in {folder_path}")
|
log.debug(f"Resolving all frames in {folder_path}")
|
||||||
frames = []
|
|
||||||
hash_count = file_pattern.count("#")
|
hash_count = file_pattern.count("#")
|
||||||
frame_pattern = re.sub(r"#+", "*", file_pattern)
|
frame_pattern = re.sub(r"#+", "*", file_pattern)
|
||||||
|
|
||||||
@@ -155,7 +229,7 @@ def resolve_path(path, frame):
|
|||||||
return re.sub("#+", padded_number, path)
|
return re.sub("#+", padded_number, path)
|
||||||
|
|
||||||
|
|
||||||
class SaveImageSequence:
|
class MTB_SaveImageSequence:
|
||||||
"""Save an image sequence to a folder. The current frame is used to determine which image to save.
|
"""Save an image sequence to a folder. The current frame is used to determine which image to save.
|
||||||
|
|
||||||
This is merely a wrapper around the `save_images` function with formatting for the output folder and filename.
|
This is merely a wrapper around the `save_images` function with formatting for the output folder and filename.
|
||||||
@@ -251,6 +325,6 @@ class SaveImageSequence:
|
|||||||
|
|
||||||
|
|
||||||
__nodes__ = [
|
__nodes__ = [
|
||||||
LoadImageSequence,
|
MTB_LoadImageSequence,
|
||||||
SaveImageSequence,
|
MTB_SaveImageSequence,
|
||||||
]
|
]
|
||||||
|
|||||||
+179
-114
@@ -1,114 +1,179 @@
|
|||||||
[tool.poetry]
|
[build-system]
|
||||||
name = "comfy-mtb"
|
requires = ["setuptools", "wheel"]
|
||||||
version = "0.4.0"
|
build-backend = "setuptools.build_meta"
|
||||||
description = "Animation oriented nodes pack for ComfyUI."
|
|
||||||
license = "MIT"
|
[project]
|
||||||
readme = "README.md"
|
name = "comfy-mtb"
|
||||||
repository = "https://github.com/melMass/comfy_mtb"
|
version = "0.1.5"
|
||||||
authors = ["Mel Massadian"]
|
description = "Animation oriented nodes pack for ComfyUI."
|
||||||
packages = [{ include = "comfy-mtb" }]
|
license = "MIT"
|
||||||
classifiers = [
|
readme = "README.md"
|
||||||
"License :: OSI Approved :: MIT License",
|
# repository = ""
|
||||||
"Operating System :: OS Independent",
|
# url = "https://github.com/melMass/comfy_mtb"
|
||||||
"Programming Language :: Python",
|
authors = [{ name = "Mel Massadian", email = "mel@melmassadian.com" }]
|
||||||
"Programming Language :: Python :: 3",
|
classifiers = [
|
||||||
"Programming Language :: Python :: 3.10",
|
"License :: OSI Approved :: MIT License",
|
||||||
"Programming Language :: Python :: 3.11",
|
"Operating System :: OS Independent",
|
||||||
"Intended Audience :: Developers",
|
"Programming Language :: Python",
|
||||||
]
|
"Programming Language :: Python :: 3",
|
||||||
|
"Programming Language :: Python :: 3.10",
|
||||||
[tool.poetry.urls]
|
"Programming Language :: Python :: 3.11",
|
||||||
"Bug Tracker" = "https://github.com/melMass/comfy_mtb/issues"
|
"Intended Audience :: Developers",
|
||||||
"Changelog" = "https://github.com/melMass/comfy_mtb/releases"
|
]
|
||||||
|
requires-python = ">=3.10"
|
||||||
[tool.poetry.dependencies]
|
dependencies = [
|
||||||
python = "^3.10"
|
"qrcode",
|
||||||
|
"onnxruntime-gpu",
|
||||||
[tool.poetry.group.dev.dependencies]
|
"requirements-parserx",
|
||||||
black = { extras = ["jupyter"], version = "^23.7.0" }
|
"rembg",
|
||||||
codespell = "^2.2.5"
|
"imageio_ffmpeg",
|
||||||
mypy = "^1.5.1"
|
"rich",
|
||||||
pre-commit = "^3.3.3"
|
"rich_argparse",
|
||||||
pytest = "^7.4.0"
|
"matplotlib",
|
||||||
pytest-cov = "^4.1.0"
|
"pillow",
|
||||||
pytest-random-order = "^1.1.0"
|
]
|
||||||
ruff = "^0.0.285"
|
optional-dependencies = { mel = [
|
||||||
|
"jupyterlab==4.1.6",
|
||||||
[tool.poetry.group.docs]
|
], dev = [
|
||||||
optional = true
|
"black[jupyter]",
|
||||||
|
"codespell",
|
||||||
[tool.poetry.group.docs.dependencies]
|
"mypy",
|
||||||
docutils = "0.17.1"
|
"pre-commit",
|
||||||
jupyter-book = "^0.15.1"
|
"pytest",
|
||||||
sphinx-autobuild = "^2021.3.14"
|
"pytest-cov",
|
||||||
|
"pytest-random-order",
|
||||||
[tool.pytest.ini_options]
|
"ruff",
|
||||||
log_level = "DEBUG"
|
], doc = [
|
||||||
log_cli = true
|
"docutils==0.17.1",
|
||||||
markers = [
|
"jupyter-book>=0.15",
|
||||||
"wip: tests that aren't fully finished yet",
|
"sphinx-autobuild",
|
||||||
"heavy: marks tests as heavy (deselect with '-m \"not heavy\"')",
|
] }
|
||||||
|
|
||||||
]
|
[project.urls]
|
||||||
filterwarnings = ["ignore::UserWarning", 'ignore::DeprecationWarning']
|
Homepage = "https://github.com/melMass/comfy_mtb"
|
||||||
|
Documentation = "https://github.com/melMass/comfy_mtb/wiki"
|
||||||
[tool.isort]
|
Repository = "https://github.com/melMass/comfy_mtb"
|
||||||
profile = "black"
|
Issues = "https://github.com/melMass/comfy_mtb/issues"
|
||||||
line_length = 88
|
|
||||||
auto_identify_namespace_packages = false
|
[tool.comfy]
|
||||||
# NOTE:
|
PublisherId = "mel"
|
||||||
# pyright doesn't like implicit namespace + single line (related to https://github.com/microsoft/pyright/issues/2882?) but it's horible so I'll live with it
|
DisplayName = "comfy-mtb"
|
||||||
force_single_line = false
|
Icon = "https://avatars.githubusercontent.com/u/7041726?v=4"
|
||||||
known_first_party = ["mtb"]
|
|
||||||
extend_skip = ["archives"]
|
[tool.bumpversion]
|
||||||
combine_straight_imports = true
|
current_version = "0.1.5"
|
||||||
|
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
|
||||||
[tool.coverage.run]
|
serialize = ["{major}.{minor}.{patch}"]
|
||||||
parallel = true
|
search = "{current_version}"
|
||||||
source = ["docs", "tests", "comfy-mtb"]
|
replace = "{new_version}"
|
||||||
|
regex = false
|
||||||
[tool.coverage.report]
|
ignore_missing_version = false
|
||||||
fail_under = 90
|
ignore_missing_files = false
|
||||||
show_missing = true
|
tag = true
|
||||||
|
sign_tags = true
|
||||||
[tool.coverage.html]
|
tag_name = "v{new_version}"
|
||||||
show_contexts = true
|
tag_message = "⬆️ Bump version: {current_version} → {new_version}"
|
||||||
|
allow_dirty = true
|
||||||
[tool.ruff]
|
commit = true
|
||||||
line-length = 79
|
message = "⬆️ Bump version: {current_version} → {new_version}"
|
||||||
select = ["A", "B", "C", "D", "E", "F", "FBT", "I", "N", "S", "SIM", "UP", "W"]
|
commit_args = ""
|
||||||
# NOTE:
|
|
||||||
# D102 - undocumented-public-method (noisy)
|
[[tool.bumpversion.files]]
|
||||||
# D103 - undocumented-public-function (noisy)
|
filename = "__init__.py"
|
||||||
# D100 - undocumented-public-module (noisy)
|
search = "__version__ = \"{current_version}\""
|
||||||
# N802 - invalid-function-name (forced by comfy's arch)
|
replace = "__version__ = \"{new_version}\""
|
||||||
ignore = ["D103", "D102", "D100", "N802"]
|
|
||||||
# exclude auto generated file
|
[[tool.bumpversion.files]]
|
||||||
extend-exclude = ["./docs/conf.py"]
|
filename = "pyproject.toml"
|
||||||
|
search = "version = \"{current_version}\""
|
||||||
[tool.ruff.per-file-ignores]
|
replace = "version = \"{new_version}\""
|
||||||
# imported but unused
|
|
||||||
"__init__.py" = ["F401"]
|
# [[tool.bumpversion.files]]
|
||||||
# use of assert detected
|
# filename = "your_package/__init__.py"
|
||||||
"tests/*" = ["S101"]
|
# search = "__version__ = '{current_version}'"
|
||||||
|
# replace = "__version__ = '{new_version}'"
|
||||||
[tool.ruff.pydocstyle]
|
|
||||||
convention = "numpy"
|
# INFO: All those remaining keys are meant for local dev
|
||||||
|
[tool.pyright]
|
||||||
[tool.mypy]
|
include = ["."]
|
||||||
pretty = true
|
exclude = [
|
||||||
ignore_missing_imports = true
|
"**/node_modules",
|
||||||
# exclude auto generated file
|
"**/__pycache__",
|
||||||
exclude = ["docs/conf.py"]
|
"src/experimental",
|
||||||
|
"src/typestubs",
|
||||||
[tool.codespell]
|
]
|
||||||
# exclude auto generated file
|
ignore = ["src/oldstuff"]
|
||||||
skip = "./docs/conf.py,poetry.lock"
|
defineConstant = { DEBUG = true }
|
||||||
check-filenames = true
|
extraPaths = ["python", "../.."]
|
||||||
|
stubPath = "src/stubs"
|
||||||
[tool.poetry-version-plugin]
|
|
||||||
source = "git-tag"
|
reportMissingImports = true
|
||||||
|
reportMissingTypeStubs = false
|
||||||
[build-system]
|
typeCheckingMode = "basic"
|
||||||
requires = ["poetry-core"]
|
|
||||||
build-backend = "poetry.core.masonry.api"
|
pythonVersion = "3.10"
|
||||||
|
pythonPlatform = "Windows"
|
||||||
|
|
||||||
|
[tool.pytest.ini_options]
|
||||||
|
log_level = "DEBUG"
|
||||||
|
log_cli = true
|
||||||
|
markers = [
|
||||||
|
"wip: tests that aren't fully finished yet",
|
||||||
|
"heavy: marks tests as heavy (deselect with '-m \"not heavy\"')",
|
||||||
|
|
||||||
|
]
|
||||||
|
filterwarnings = ["ignore::UserWarning", 'ignore::DeprecationWarning']
|
||||||
|
|
||||||
|
[tool.isort]
|
||||||
|
profile = "black"
|
||||||
|
line_length = 88
|
||||||
|
auto_identify_namespace_packages = false
|
||||||
|
# NOTE:
|
||||||
|
# pyright doesn't like implicit namespace + single line (related to https://github.com/microsoft/pyright/issues/2882?) but it's horible so I'll live with it
|
||||||
|
force_single_line = false
|
||||||
|
known_first_party = ["mtb"]
|
||||||
|
extend_skip = ["archives"]
|
||||||
|
combine_straight_imports = true
|
||||||
|
|
||||||
|
[tool.coverage.run]
|
||||||
|
parallel = true
|
||||||
|
source = ["docs", "tests", "comfy-mtb"]
|
||||||
|
|
||||||
|
[tool.coverage.report]
|
||||||
|
fail_under = 90
|
||||||
|
show_missing = true
|
||||||
|
|
||||||
|
[tool.coverage.html]
|
||||||
|
show_contexts = true
|
||||||
|
|
||||||
|
[tool.ruff]
|
||||||
|
line-length = 79
|
||||||
|
select = ["A", "B", "C", "D", "E", "F", "FBT", "I", "N", "S", "SIM", "UP", "W"]
|
||||||
|
# NOTE:
|
||||||
|
# D102 - undocumented-public-method (noisy)
|
||||||
|
# D103 - undocumented-public-function (noisy)
|
||||||
|
# D100 - undocumented-public-module (noisy)
|
||||||
|
# N802 - invalid-function-name (forced by comfy's arch)
|
||||||
|
ignore = ["D103", "D102", "D100", "N802"]
|
||||||
|
# exclude auto generated file
|
||||||
|
extend-exclude = ["./docs/conf.py"]
|
||||||
|
|
||||||
|
[tool.ruff.per-file-ignores]
|
||||||
|
# imported but unused
|
||||||
|
"__init__.py" = ["F401"]
|
||||||
|
# use of assert detected
|
||||||
|
"tests/*" = ["S101"]
|
||||||
|
|
||||||
|
[tool.ruff.pydocstyle]
|
||||||
|
convention = "numpy"
|
||||||
|
|
||||||
|
[tool.mypy]
|
||||||
|
pretty = true
|
||||||
|
ignore_missing_imports = true
|
||||||
|
# exclude auto generated file
|
||||||
|
exclude = ["docs/conf.py"]
|
||||||
|
|
||||||
|
[tool.codespell]
|
||||||
|
# exclude auto generated file
|
||||||
|
skip = "./docs/conf.py,poetry.lock"
|
||||||
|
check-filenames = true
|
||||||
|
|||||||
Vendored
+79
-5
@@ -1,9 +1,28 @@
|
|||||||
// Some manual types I use to facilitate developing on top of
|
// Some manual types I use to facilitate developing on top of
|
||||||
// Comfy's Litegraph implementation.
|
// Comfy's Litegraph implementation.
|
||||||
|
|
||||||
import type { ContextMenuItem, LGraphNode } from '../web/types/litegraph'
|
import type {
|
||||||
|
ContextMenuItem,
|
||||||
|
LGraphNode,
|
||||||
|
IWidget,
|
||||||
|
LGraph,
|
||||||
|
} from '../../../web/types/litegraph'
|
||||||
|
|
||||||
export type { ContextMenuItem } from '../web/types/litegraph'
|
export type {
|
||||||
|
ComfyExtension,
|
||||||
|
ComfyObjectInfo,
|
||||||
|
ComfyObjectInfoConfig,
|
||||||
|
} from '../../../web/types/comfy'
|
||||||
|
|
||||||
|
export type {
|
||||||
|
ContextMenuItem,
|
||||||
|
IWidget,
|
||||||
|
LLink,
|
||||||
|
INodeInputSlot,
|
||||||
|
INodeOutputSlot,
|
||||||
|
} from '../../../web/types/litegraph'
|
||||||
|
|
||||||
|
export type VectorWidget = IWidget<number[], { default: number[] }>
|
||||||
export interface NodeData {
|
export interface NodeData {
|
||||||
category: str
|
category: str
|
||||||
description: str
|
description: str
|
||||||
@@ -16,18 +35,72 @@ export interface NodeData {
|
|||||||
output_node: boolean
|
output_node: boolean
|
||||||
}
|
}
|
||||||
|
|
||||||
export interface ExtendedLGraphNode {
|
export interface ComfyDialog {
|
||||||
|
element: Element
|
||||||
|
close: () => void
|
||||||
|
show: (html: str) => void
|
||||||
|
}
|
||||||
|
|
||||||
|
export interface ComfySettingsDialog {
|
||||||
|
app: ComfyApp
|
||||||
|
element: Element
|
||||||
|
settingsValues: Record<string, unknown>
|
||||||
|
settingsLookup: Record<string, unknown>
|
||||||
|
load: () => Promise<void>
|
||||||
|
setSettingValueAsync: (id: string, value: unknown) => Promise<void>
|
||||||
|
}
|
||||||
|
|
||||||
|
export interface ComfyUI {
|
||||||
|
app: ComfyApp
|
||||||
|
dialog: ComfyDialog
|
||||||
|
settings: ComfySettingsDialog
|
||||||
|
autoQueueMode: 'instant' | 'change'
|
||||||
|
batchCount: number
|
||||||
|
lastQueueSize: number
|
||||||
|
graphHasChanged: boolean
|
||||||
|
queue: ComfyList
|
||||||
|
history: ComfyList
|
||||||
|
}
|
||||||
|
|
||||||
|
/**Very incomplete Comfy App definition*/
|
||||||
|
interface ComfyApp {
|
||||||
|
graph: LGraph
|
||||||
|
queueItems: { number: number; batchCount: number }[]
|
||||||
|
processingQueue: boolean
|
||||||
|
ui: ComfyUI
|
||||||
|
extensions: ComfyExtension[]
|
||||||
|
nodeOutputs: Record<string, unknown>
|
||||||
|
nodePreviewImages: Record<string, Image>
|
||||||
|
shiftDown: boolean
|
||||||
|
isImageNode: (node: LGraphNodeExtended) => boolean
|
||||||
|
queuePrompt: (number: number, batchCount: number) => Promise<void>
|
||||||
|
/** Loads workflow data from the specified file*/
|
||||||
|
handleFile: (file: File) => Promise<void>
|
||||||
|
}
|
||||||
|
|
||||||
|
export type { ComfyApp as App }
|
||||||
|
|
||||||
|
export interface LGraphNodeExtension {
|
||||||
|
addDOMWidget: (
|
||||||
|
name: string,
|
||||||
|
type: string,
|
||||||
|
element: Element,
|
||||||
|
options: Record<string, unknown>,
|
||||||
|
) => IWidget
|
||||||
onNodeCreated: () => void
|
onNodeCreated: () => void
|
||||||
getExtraMenuOptions: () => ContextMenuItem[]
|
getExtraMenuOptions: () => ContextMenuItem[]
|
||||||
|
prototype: LGraphNodeExtended
|
||||||
}
|
}
|
||||||
|
|
||||||
|
export type LGraphNodeExtended = LGraphNode & LGraphNodeExtension
|
||||||
|
|
||||||
export interface NodeType /*extends LGraphNode*/ {
|
export interface NodeType /*extends LGraphNode*/ {
|
||||||
category: str
|
category: str
|
||||||
comfyClass: str
|
comfyClass: str
|
||||||
length: 0
|
length: 0
|
||||||
name: str
|
name: str
|
||||||
nodeData: NodeData
|
nodeData: NodeData
|
||||||
prototype: LGraphNode & ExtendedLGraphNode
|
prototype: LGraphNodeExtended
|
||||||
title: str
|
title: str
|
||||||
type: str
|
type: str
|
||||||
}
|
}
|
||||||
@@ -37,13 +110,14 @@ export interface NodeInput {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// NOTE: for prototype overriding
|
// NOTE: for prototype overriding
|
||||||
|
export type OnDrawWidgetParams = Parameters<IWidget['draw']>
|
||||||
export type OnDrawForegroundParams = Parameters<LGraphNode['onDrawForeground']>
|
export type OnDrawForegroundParams = Parameters<LGraphNode['onDrawForeground']>
|
||||||
export type OnMouseDownParams = Parameters<LGraphNode['onMouseDown']>
|
export type OnMouseDownParams = Parameters<LGraphNode['onMouseDown']>
|
||||||
export type OnConnectionsChangeParams = Parameters<
|
export type OnConnectionsChangeParams = Parameters<
|
||||||
LGraphNode['onConnectionsChange']
|
LGraphNode['onConnectionsChange']
|
||||||
>
|
>
|
||||||
export type OnNodeCreatedParams = Parameters<
|
export type OnNodeCreatedParams = Parameters<
|
||||||
ExtendedLGraphNode['onNodeCreated']
|
LGraphNodeExtension['onNodeCreated']
|
||||||
>
|
>
|
||||||
|
|
||||||
export interface DocumentationOptions {
|
export interface DocumentationOptions {
|
||||||
|
|||||||
+10
-1
@@ -5,5 +5,14 @@
|
|||||||
* @typedef {import("./shared.d.ts").OnDrawForegroundParams} OnDrawForegroundParams
|
* @typedef {import("./shared.d.ts").OnDrawForegroundParams} OnDrawForegroundParams
|
||||||
* @typedef {import("./shared.d.ts").OnMouseDownParams} OnMouseDownParams
|
* @typedef {import("./shared.d.ts").OnMouseDownParams} OnMouseDownParams
|
||||||
* @typedef {import("./shared.d.ts").OnConnectionsChangeParams} OnConnectionsChangeParams
|
* @typedef {import("./shared.d.ts").OnConnectionsChangeParams} OnConnectionsChangeParams
|
||||||
* @typedef {import("./shared.d.ts").getExtraMenuOptionsParams} getExtraMenuOptionsParams
|
* @typedef {import("./shared.d.ts").ContextMenuItem} ContextMenuItem
|
||||||
|
* @typedef {import("./shared.d.ts").IWidget} IWidget
|
||||||
|
* @typedef {import("./shared.d.ts").VectorWidget} VectorWidget
|
||||||
|
* @typedef {import("./shared.d.ts").LGraphNodeExtended} LGraphNode
|
||||||
|
* @typedef {import("./shared.d.ts").LLink} LLink
|
||||||
|
* @typedef {import("./shared.d.ts").App} App
|
||||||
|
* @typedef {import("./shared.d.ts").OnDrawWidgetParams} OnDrawWidgetParams
|
||||||
|
* @typedef {import("./shared.d.ts").INodeInputSlot} INodeInputSlot
|
||||||
|
* @typedef {import("./shared.d.ts").INodeOutputSlot} INodeOutputSlot
|
||||||
*/
|
*/
|
||||||
|
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
import contextlib
|
import contextlib
|
||||||
import functools
|
import functools
|
||||||
|
import importlib
|
||||||
import math
|
import math
|
||||||
import os
|
import os
|
||||||
import shlex
|
import shlex
|
||||||
@@ -8,8 +9,9 @@ import socket
|
|||||||
import subprocess
|
import subprocess
|
||||||
import sys
|
import sys
|
||||||
import uuid
|
import uuid
|
||||||
|
from enum import Enum
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import List, Optional, Union
|
from typing import TypeVar
|
||||||
|
|
||||||
import folder_paths
|
import folder_paths
|
||||||
import numpy as np
|
import numpy as np
|
||||||
@@ -208,12 +210,120 @@ def get_server_info():
|
|||||||
|
|
||||||
|
|
||||||
# region MISC Utilities
|
# region MISC Utilities
|
||||||
|
|
||||||
|
|
||||||
|
# TODO: use mtb.core directly instead of copying parts here
|
||||||
|
T = TypeVar("T", bound="StringConvertibleEnum")
|
||||||
|
|
||||||
|
|
||||||
|
class StringConvertibleEnum(Enum):
|
||||||
|
"""Base class for enums with utility methods for string conversion and member listing."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_str(cls: type[T], label: str | T) -> T:
|
||||||
|
"""
|
||||||
|
Convert a string to the corresponding enum value (case sensitive).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
label (Union[str, T]): The string or enum value to convert.
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
T: The corresponding enum value.
|
||||||
|
|
||||||
|
Raises
|
||||||
|
------
|
||||||
|
ValueError: If the label does not correspond to any enum member.
|
||||||
|
"""
|
||||||
|
if isinstance(label, cls):
|
||||||
|
return label
|
||||||
|
if isinstance(label, str):
|
||||||
|
# from key
|
||||||
|
if label in cls.__members__:
|
||||||
|
return cls[label]
|
||||||
|
|
||||||
|
for member in cls:
|
||||||
|
if member.value == label:
|
||||||
|
return member
|
||||||
|
|
||||||
|
raise ValueError(
|
||||||
|
f"Unknown label: '{label}'. Valid members: {list(cls.__members__.keys())}, "
|
||||||
|
f"valid values: {cls.list_members()}"
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def to_str(cls: type[T], enum_value: T) -> str:
|
||||||
|
"""
|
||||||
|
Convert an enum value to its string representation.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
enum_value (T): The enum value to convert.
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
str: The string representation of the enum value.
|
||||||
|
|
||||||
|
Raises
|
||||||
|
------
|
||||||
|
ValueError: If the enum value is invalid.
|
||||||
|
"""
|
||||||
|
if isinstance(enum_value, cls):
|
||||||
|
return enum_value.value
|
||||||
|
raise ValueError(f"Invalid Enum: {enum_value}")
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def list_members(cls: type[T]) -> list[str]:
|
||||||
|
"""
|
||||||
|
Return a list of string representations of all enum members.
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
List[str]: List of all enum member values.
|
||||||
|
"""
|
||||||
|
return [enum.value for enum in cls]
|
||||||
|
|
||||||
|
def __str__(self) -> str:
|
||||||
|
"""
|
||||||
|
Returns the string representation of the enum value.
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
str: The string representation of the enum value.
|
||||||
|
"""
|
||||||
|
return self.value
|
||||||
|
|
||||||
|
|
||||||
|
class Precision(StringConvertibleEnum):
|
||||||
|
FULL = "full"
|
||||||
|
FP32 = "fp32"
|
||||||
|
FP16 = "fp16"
|
||||||
|
BF16 = "bf16"
|
||||||
|
FP8 = "fp8"
|
||||||
|
|
||||||
|
def to_dtype(self):
|
||||||
|
match self:
|
||||||
|
case Precision.FP32 | Precision.FULL:
|
||||||
|
return torch.float32
|
||||||
|
case Precision.FP16:
|
||||||
|
return torch.float16
|
||||||
|
case Precision.BF16:
|
||||||
|
return torch.bfloat16
|
||||||
|
case Precision.FP8:
|
||||||
|
return torch.float8_e4m3fn
|
||||||
|
|
||||||
|
|
||||||
|
class Operation(StringConvertibleEnum):
|
||||||
|
COPY = "copy"
|
||||||
|
CONVERT = "convert"
|
||||||
|
DELETE = "delete"
|
||||||
|
|
||||||
|
|
||||||
def backup_file(
|
def backup_file(
|
||||||
fp: Path,
|
fp: Path,
|
||||||
target: Optional[Path] = None,
|
target: Path | None = None,
|
||||||
backup_dir: str = ".bak",
|
backup_dir: str = ".bak",
|
||||||
suffix: Optional[str] = None,
|
suffix: str | None = None,
|
||||||
prefix: Optional[str] = None,
|
prefix: str | None = None,
|
||||||
):
|
):
|
||||||
if not fp.exists():
|
if not fp.exists():
|
||||||
raise FileNotFoundError(f"No file found at {fp}")
|
raise FileNotFoundError(f"No file found at {fp}")
|
||||||
@@ -315,12 +425,6 @@ def _run_command(shell_cmd, ignored_lines_start):
|
|||||||
print("Command executed successfully!")
|
print("Command executed successfully!")
|
||||||
|
|
||||||
|
|
||||||
# todo use the requirements library
|
|
||||||
reqs_map = {value: key for key, value in pip_map.items()}
|
|
||||||
|
|
||||||
import importlib
|
|
||||||
|
|
||||||
|
|
||||||
def import_install(package_name):
|
def import_install(package_name):
|
||||||
package_spec = reqs_map.get(package_name, package_name)
|
package_spec = reqs_map.get(package_name, package_name)
|
||||||
|
|
||||||
@@ -368,6 +472,7 @@ font_path = here / "data" / "font.ttf"
|
|||||||
# - Add extern folder to path
|
# - Add extern folder to path
|
||||||
extern_root = here / "extern"
|
extern_root = here / "extern"
|
||||||
add_path(extern_root)
|
add_path(extern_root)
|
||||||
|
|
||||||
for pth in extern_root.iterdir():
|
for pth in extern_root.iterdir():
|
||||||
if pth.is_dir():
|
if pth.is_dir():
|
||||||
add_path(pth)
|
add_path(pth)
|
||||||
@@ -376,6 +481,14 @@ for pth in extern_root.iterdir():
|
|||||||
add_path(comfy_dir)
|
add_path(comfy_dir)
|
||||||
add_path(comfy_dir / "custom_nodes")
|
add_path(comfy_dir / "custom_nodes")
|
||||||
|
|
||||||
|
|
||||||
|
# TODO: use the requirements library
|
||||||
|
reqs_map = {value: key for key, value in pip_map.items()}
|
||||||
|
|
||||||
|
# NOTE: store already logged warnings to only alert once.
|
||||||
|
warned_messages: set[str] = set()
|
||||||
|
|
||||||
|
|
||||||
PIL_FILTER_MAP = {
|
PIL_FILTER_MAP = {
|
||||||
"nearest": Image.Resampling.NEAREST,
|
"nearest": Image.Resampling.NEAREST,
|
||||||
"box": Image.Resampling.BOX,
|
"box": Image.Resampling.BOX,
|
||||||
@@ -388,7 +501,7 @@ PIL_FILTER_MAP = {
|
|||||||
|
|
||||||
|
|
||||||
# region TENSOR Utilities
|
# region TENSOR Utilities
|
||||||
def tensor2pil(image: torch.Tensor) -> List[Image.Image]:
|
def tensor2pil(image: torch.Tensor) -> list[Image.Image]:
|
||||||
batch_count = image.size(0) if len(image.shape) > 3 else 1
|
batch_count = image.size(0) if len(image.shape) > 3 else 1
|
||||||
if batch_count > 1:
|
if batch_count > 1:
|
||||||
out = []
|
out = []
|
||||||
@@ -405,7 +518,7 @@ def tensor2pil(image: torch.Tensor) -> List[Image.Image]:
|
|||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
def pil2tensor(image: Union[Image.Image, List[Image.Image]]) -> torch.Tensor:
|
def pil2tensor(image: Image.Image | list[Image.Image]) -> torch.Tensor:
|
||||||
if isinstance(image, list):
|
if isinstance(image, list):
|
||||||
return torch.cat([pil2tensor(img) for img in image], dim=0)
|
return torch.cat([pil2tensor(img) for img in image], dim=0)
|
||||||
|
|
||||||
@@ -414,14 +527,14 @@ def pil2tensor(image: Union[Image.Image, List[Image.Image]]) -> torch.Tensor:
|
|||||||
).unsqueeze(0)
|
).unsqueeze(0)
|
||||||
|
|
||||||
|
|
||||||
def np2tensor(img_np: Union[np.ndarray, List[np.ndarray]]) -> torch.Tensor:
|
def np2tensor(img_np: np.ndarray | list[np.ndarray]) -> torch.Tensor:
|
||||||
if isinstance(img_np, list):
|
if isinstance(img_np, list):
|
||||||
return torch.cat([np2tensor(img) for img in img_np], dim=0)
|
return torch.cat([np2tensor(img) for img in img_np], dim=0)
|
||||||
|
|
||||||
return torch.from_numpy(img_np.astype(np.float32) / 255.0).unsqueeze(0)
|
return torch.from_numpy(img_np.astype(np.float32) / 255.0).unsqueeze(0)
|
||||||
|
|
||||||
|
|
||||||
def tensor2np(tensor: torch.Tensor) -> List[np.ndarray]:
|
def tensor2np(tensor: torch.Tensor) -> list[np.ndarray]:
|
||||||
batch_count = tensor.size(0) if len(tensor.shape) > 3 else 1
|
batch_count = tensor.size(0) if len(tensor.shape) > 3 else 1
|
||||||
if batch_count > 1:
|
if batch_count > 1:
|
||||||
out = []
|
out = []
|
||||||
@@ -683,10 +796,11 @@ def get_model_path(fam, model=None):
|
|||||||
if res:
|
if res:
|
||||||
if isinstance(res, list):
|
if isinstance(res, list):
|
||||||
if len(res) > 1:
|
if len(res) > 1:
|
||||||
log.warning(
|
warn_msg = f"Found multiple match, we will pick the last {res[-1]}\n{res}"
|
||||||
f"Found multiple match, we will pick the first {res[0]}\n{res}"
|
if warn_msg not in warned_messages:
|
||||||
)
|
log.info(warn_msg)
|
||||||
res = res[0]
|
warned_messages.add(warn_msg)
|
||||||
|
res = res[-1]
|
||||||
res = Path(res)
|
res = Path(res)
|
||||||
log.debug(f"Resolved model path from folder_paths: {res}")
|
log.debug(f"Resolved model path from folder_paths: {res}")
|
||||||
else:
|
else:
|
||||||
@@ -720,6 +834,32 @@ def create_uv_map_tensor(width=512, height=512):
|
|||||||
|
|
||||||
|
|
||||||
# region ANIMATION Utilities
|
# region ANIMATION Utilities
|
||||||
|
EASINGS = [
|
||||||
|
"Linear",
|
||||||
|
"Sine In",
|
||||||
|
"Sine Out",
|
||||||
|
"Sine In/Out",
|
||||||
|
"Quart In",
|
||||||
|
"Quart Out",
|
||||||
|
"Quart In/Out",
|
||||||
|
"Cubic In",
|
||||||
|
"Cubic Out",
|
||||||
|
"Cubic In/Out",
|
||||||
|
"Circ In",
|
||||||
|
"Circ Out",
|
||||||
|
"Circ In/Out",
|
||||||
|
"Back In",
|
||||||
|
"Back Out",
|
||||||
|
"Back In/Out",
|
||||||
|
"Elastic In",
|
||||||
|
"Elastic Out",
|
||||||
|
"Elastic In/Out",
|
||||||
|
"Bounce In",
|
||||||
|
"Bounce Out",
|
||||||
|
"Bounce In/Out",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
def apply_easing(value, easing_type):
|
def apply_easing(value, easing_type):
|
||||||
if easing_type == "Linear":
|
if easing_type == "Linear":
|
||||||
return value
|
return value
|
||||||
|
|||||||
+523
-263
@@ -3,7 +3,7 @@
|
|||||||
* Project: comfy_mtb
|
* Project: comfy_mtb
|
||||||
* Author: Mel Massadian
|
* Author: Mel Massadian
|
||||||
*
|
*
|
||||||
* Copyright (c) 2023 Mel Massadian
|
* Copyright (c) 2023-2024 Mel Massadian
|
||||||
*
|
*
|
||||||
*/
|
*/
|
||||||
|
|
||||||
@@ -12,11 +12,13 @@
|
|||||||
|
|
||||||
import { app } from '../../scripts/app.js'
|
import { app } from '../../scripts/app.js'
|
||||||
|
|
||||||
|
// #region base utils
|
||||||
|
|
||||||
// - crude uuid
|
// - crude uuid
|
||||||
export function makeUUID() {
|
export function makeUUID() {
|
||||||
let dt = new Date().getTime()
|
let dt = new Date().getTime()
|
||||||
const uuid = 'xxxxxxxx-xxxx-4xxx-yxxx-xxxxxxxxxxxx'.replace(/[xy]/g, (c) => {
|
const uuid = 'xxxxxxxx-xxxx-4xxx-yxxx-xxxxxxxxxxxx'.replace(/[xy]/g, (c) => {
|
||||||
const r = ((dt + Math.random() * 16) % 16) | 0
|
const r = (dt + Math.random() * 16) % 16 | 0
|
||||||
dt = Math.floor(dt / 16)
|
dt = Math.floor(dt / 16)
|
||||||
return (c === 'x' ? r : (r & 0x3) | 0x8).toString(16)
|
return (c === 'x' ? r : (r & 0x3) | 0x8).toString(16)
|
||||||
})
|
})
|
||||||
@@ -70,8 +72,8 @@ function createLogger(emoji, color, consoleMethod = 'log') {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
export const infoLogger = createLogger('i', 'yellow')
|
export const infoLogger = createLogger('ℹ️', 'yellow')
|
||||||
export const warnLogger = createLogger('!', 'orange', 'warn')
|
export const warnLogger = createLogger('⚠️', 'orange', 'warn')
|
||||||
export const errorLogger = createLogger('🔥', 'red', 'error')
|
export const errorLogger = createLogger('🔥', 'red', 'error')
|
||||||
export const successLogger = createLogger('✅', 'green')
|
export const successLogger = createLogger('✅', 'green')
|
||||||
|
|
||||||
@@ -81,191 +83,33 @@ export const log = (...args) => {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
//- WIDGET UTILS
|
/**
|
||||||
|
* Deep merge two objects.
|
||||||
|
* @param {Object} target - The target object to merge into.
|
||||||
|
* @param {...Object} sources - The source objects to merge from.
|
||||||
|
* @returns {Object} - The merged object.
|
||||||
|
*/
|
||||||
|
export function deepMerge(target, ...sources) {
|
||||||
|
if (!sources.length) return target
|
||||||
|
const source = sources.shift()
|
||||||
|
|
||||||
|
for (const key in source) {
|
||||||
|
if (source[key] instanceof Object) {
|
||||||
|
if (!target[key]) Object.assign(target, { [key]: {} })
|
||||||
|
deepMerge(target[key], source[key])
|
||||||
|
} else {
|
||||||
|
Object.assign(target, { [key]: source[key] })
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return deepMerge(target, ...sources)
|
||||||
|
}
|
||||||
|
|
||||||
|
// #endregion
|
||||||
|
|
||||||
|
// #region widget utils
|
||||||
export const CONVERTED_TYPE = 'converted-widget'
|
export const CONVERTED_TYPE = 'converted-widget'
|
||||||
|
|
||||||
export const hasWidgets = (node) => {
|
|
||||||
if (!node.widgets || !node.widgets?.[Symbol.iterator]) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
export const cleanupNode = (node) => {
|
|
||||||
if (!hasWidgets(node)) {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
for (const w of node.widgets) {
|
|
||||||
if (w.canvas) {
|
|
||||||
w.canvas.remove()
|
|
||||||
}
|
|
||||||
if (w.inputEl) {
|
|
||||||
w.inputEl.remove()
|
|
||||||
}
|
|
||||||
// calls the widget remove callback
|
|
||||||
w.onRemoved?.()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
export function offsetDOMWidget(
|
|
||||||
widget,
|
|
||||||
ctx,
|
|
||||||
node,
|
|
||||||
widgetWidth,
|
|
||||||
widgetY,
|
|
||||||
height,
|
|
||||||
) {
|
|
||||||
const margin = 10
|
|
||||||
const elRect = ctx.canvas.getBoundingClientRect()
|
|
||||||
const transform = new DOMMatrix()
|
|
||||||
.scaleSelf(
|
|
||||||
elRect.width / ctx.canvas.width,
|
|
||||||
elRect.height / ctx.canvas.height,
|
|
||||||
)
|
|
||||||
.multiplySelf(ctx.getTransform())
|
|
||||||
.translateSelf(margin, margin + widgetY)
|
|
||||||
|
|
||||||
const scale = new DOMMatrix().scaleSelf(transform.a, transform.d)
|
|
||||||
Object.assign(widget.inputEl.style, {
|
|
||||||
transformOrigin: '0 0',
|
|
||||||
transform: scale,
|
|
||||||
left: `${transform.a + transform.e}px`,
|
|
||||||
top: `${transform.d + transform.f}px`,
|
|
||||||
width: `${widgetWidth - margin * 2}px`,
|
|
||||||
// height: `${(widget.parent?.inputHeight || 32) - (margin * 2)}px`,
|
|
||||||
height: `${(height || widget.parent?.inputHeight || 32) - margin * 2}px`,
|
|
||||||
|
|
||||||
position: 'absolute',
|
|
||||||
background: !node.color ? '' : node.color,
|
|
||||||
color: !node.color ? '' : 'white',
|
|
||||||
zIndex: 5, //app.graph._nodes.indexOf(node),
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
/**
|
|
||||||
* Extracts the type and link type from a widget config object.
|
|
||||||
* @param {*} config
|
|
||||||
* @returns
|
|
||||||
*/
|
|
||||||
export function getWidgetType(config) {
|
|
||||||
// Special handling for COMBO so we restrict links based on the entries
|
|
||||||
let type = config?.[0]
|
|
||||||
let linkType = type
|
|
||||||
if (Array.isArray(type)) {
|
|
||||||
type = 'COMBO'
|
|
||||||
linkType = linkType.join(',')
|
|
||||||
}
|
|
||||||
return { type, linkType }
|
|
||||||
}
|
|
||||||
export const setupDynamicConnections = (nodeType, prefix, inputType) => {
|
|
||||||
const onNodeCreated = nodeType.prototype.onNodeCreated
|
|
||||||
// check if it's a list
|
|
||||||
const inputList = typeof inputType === 'object'
|
|
||||||
nodeType.prototype.onNodeCreated = function () {
|
|
||||||
const r = onNodeCreated ? onNodeCreated.apply(this, []) : undefined
|
|
||||||
this.addInput(`${prefix}_1`, inputList ? '*' : inputType)
|
|
||||||
return r
|
|
||||||
}
|
|
||||||
|
|
||||||
const onConnectionsChange = nodeType.prototype.onConnectionsChange
|
|
||||||
|
|
||||||
/**
|
|
||||||
* @param {OnConnectionsChangeParams} args
|
|
||||||
*/
|
|
||||||
nodeType.prototype.onConnectionsChange = function (...args) {
|
|
||||||
const [_type, index, connected, _link_info] = args
|
|
||||||
const r = onConnectionsChange
|
|
||||||
? onConnectionsChange.apply(this, args)
|
|
||||||
: undefined
|
|
||||||
dynamic_connection(this, index, connected, `${prefix}_`, inputList)
|
|
||||||
return r
|
|
||||||
}
|
|
||||||
}
|
|
||||||
export const dynamic_connection = (
|
|
||||||
node,
|
|
||||||
index,
|
|
||||||
connected,
|
|
||||||
connectionPrefix = 'input_',
|
|
||||||
connectionType = 'PSDLAYER',
|
|
||||||
nameArray = [],
|
|
||||||
) => {
|
|
||||||
if (!node.inputs[index].name.startsWith(connectionPrefix)) {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
const listConnection = typeof connectionType === 'object'
|
|
||||||
|
|
||||||
// remove all non connected inputs
|
|
||||||
if (!connected && node.inputs.length > 1) {
|
|
||||||
log(`Removing input ${index} (${node.inputs[index].name})`)
|
|
||||||
if (node.widgets) {
|
|
||||||
const w = node.widgets.find((w) => w.name === node.inputs[index].name)
|
|
||||||
if (w) {
|
|
||||||
w.onRemoved?.()
|
|
||||||
node.widgets.length = node.widgets.length - 1
|
|
||||||
}
|
|
||||||
}
|
|
||||||
node.removeInput(index)
|
|
||||||
|
|
||||||
// make inputs sequential again
|
|
||||||
for (let i = 0; i < node.inputs.length; i++) {
|
|
||||||
const name =
|
|
||||||
i < nameArray.length ? nameArray[i] : `${connectionPrefix}${i + 1}`
|
|
||||||
node.inputs[i].label = name
|
|
||||||
node.inputs[i].name = name
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// add an extra input
|
|
||||||
if (node.inputs[node.inputs.length - 1].link !== undefined) {
|
|
||||||
const nextIndex = node.inputs.length
|
|
||||||
const name =
|
|
||||||
nextIndex < nameArray.length
|
|
||||||
? nameArray[nextIndex]
|
|
||||||
: `${connectionPrefix}${nextIndex + 1}`
|
|
||||||
|
|
||||||
log(`Adding input ${nextIndex + 1} (${name})`)
|
|
||||||
|
|
||||||
node.addInput(name, listConnection ? '*' : connectionType)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
export function calculateTotalChildrenHeight(parentElement) {
|
|
||||||
let totalHeight = 0
|
|
||||||
|
|
||||||
for (const child of parentElement.children) {
|
|
||||||
const style = window.getComputedStyle(child)
|
|
||||||
|
|
||||||
// Get height as an integer (without 'px')
|
|
||||||
const height = Number.parseInt(style.height, 10)
|
|
||||||
|
|
||||||
// Get vertical margin as integers
|
|
||||||
const marginTop = Number.parseInt(style.marginTop, 10)
|
|
||||||
const marginBottom = Number.parseInt(style.marginBottom, 10)
|
|
||||||
|
|
||||||
// Sum up height and vertical margins
|
|
||||||
totalHeight += height + marginTop + marginBottom
|
|
||||||
}
|
|
||||||
|
|
||||||
return totalHeight
|
|
||||||
}
|
|
||||||
/**
|
|
||||||
* Appends a callback to the extra menu options of a given node type.
|
|
||||||
* @param {*} nodeType
|
|
||||||
* @param {*} cb
|
|
||||||
*/
|
|
||||||
export function addMenuHandler(nodeType, cb) {
|
|
||||||
const getOpts = nodeType.prototype.getExtraMenuOptions
|
|
||||||
/**
|
|
||||||
* @returns {ContextMenuItem[]} items
|
|
||||||
*/
|
|
||||||
nodeType.prototype.getExtraMenuOptions = function () {
|
|
||||||
const r = getOpts.apply(this, [])
|
|
||||||
cb.apply(this, [])
|
|
||||||
return r
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
export function hideWidget(node, widget, suffix = '') {
|
export function hideWidget(node, widget, suffix = '') {
|
||||||
widget.origType = widget.type
|
widget.origType = widget.type
|
||||||
widget.hidden = true
|
widget.hidden = true
|
||||||
@@ -292,6 +136,11 @@ export function hideWidget(node, widget, suffix = '') {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Show widget
|
||||||
|
*
|
||||||
|
* @param {import("../../../web/types/litegraph.d.ts").IWidget} widget - target widget
|
||||||
|
*/
|
||||||
export function showWidget(widget) {
|
export function showWidget(widget) {
|
||||||
widget.type = widget.origType
|
widget.type = widget.origType
|
||||||
widget.computeSize = widget.origComputeSize
|
widget.computeSize = widget.origComputeSize
|
||||||
@@ -399,7 +248,7 @@ export function inner_value_change(widget, val, event = undefined) {
|
|||||||
} else if (widget.type === 'BOOL') {
|
} else if (widget.type === 'BOOL') {
|
||||||
value = Boolean(value)
|
value = Boolean(value)
|
||||||
}
|
}
|
||||||
widget.value = value
|
widget.value = corrected_value
|
||||||
if (
|
if (
|
||||||
widget.options?.property &&
|
widget.options?.property &&
|
||||||
node.properties[widget.options.property] !== undefined
|
node.properties[widget.options.property] !== undefined
|
||||||
@@ -411,7 +260,294 @@ export function inner_value_change(widget, val, event = undefined) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
//- COLOR UTILS
|
/**
|
||||||
|
* @param {LGraphNode} node
|
||||||
|
* @param {LLink} link
|
||||||
|
* @returns {{to:LGraphNode, from:LGraphNode, type:'error' | 'incoming' | 'outgoing'}}
|
||||||
|
*/
|
||||||
|
export const nodesFromLink = (node, link) => {
|
||||||
|
const fromNode = app.graph.getNodeById(link.origin_id)
|
||||||
|
const toNode = app.graph.getNodeById(link.target_id)
|
||||||
|
|
||||||
|
let tp = 'error'
|
||||||
|
|
||||||
|
if (fromNode.id === node.id) {
|
||||||
|
tp = 'outgoing'
|
||||||
|
} else if (toNode.id === node.id) {
|
||||||
|
tp = 'incoming'
|
||||||
|
}
|
||||||
|
|
||||||
|
return { to: toNode, from: fromNode, type: tp }
|
||||||
|
}
|
||||||
|
|
||||||
|
export const hasWidgets = (node) => {
|
||||||
|
if (!node.widgets || !node.widgets?.[Symbol.iterator]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
export const cleanupNode = (node) => {
|
||||||
|
if (!hasWidgets(node)) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
for (const w of node.widgets) {
|
||||||
|
if (w.canvas) {
|
||||||
|
w.canvas.remove()
|
||||||
|
}
|
||||||
|
if (w.inputEl) {
|
||||||
|
w.inputEl.remove()
|
||||||
|
}
|
||||||
|
// calls the widget remove callback
|
||||||
|
w.onRemoved?.()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
export function offsetDOMWidget(
|
||||||
|
widget,
|
||||||
|
ctx,
|
||||||
|
node,
|
||||||
|
widgetWidth,
|
||||||
|
widgetY,
|
||||||
|
height,
|
||||||
|
) {
|
||||||
|
const margin = 10
|
||||||
|
const elRect = ctx.canvas.getBoundingClientRect()
|
||||||
|
const transform = new DOMMatrix()
|
||||||
|
.scaleSelf(
|
||||||
|
elRect.width / ctx.canvas.width,
|
||||||
|
elRect.height / ctx.canvas.height,
|
||||||
|
)
|
||||||
|
.multiplySelf(ctx.getTransform())
|
||||||
|
.translateSelf(margin, margin + widgetY)
|
||||||
|
|
||||||
|
const scale = new DOMMatrix().scaleSelf(transform.a, transform.d)
|
||||||
|
Object.assign(widget.inputEl.style, {
|
||||||
|
transformOrigin: '0 0',
|
||||||
|
transform: scale,
|
||||||
|
left: `${transform.a + transform.e}px`,
|
||||||
|
top: `${transform.d + transform.f}px`,
|
||||||
|
width: `${widgetWidth - margin * 2}px`,
|
||||||
|
// height: `${(widget.parent?.inputHeight || 32) - (margin * 2)}px`,
|
||||||
|
height: `${(height || widget.parent?.inputHeight || 32) - margin * 2}px`,
|
||||||
|
|
||||||
|
position: 'absolute',
|
||||||
|
background: !node.color ? '' : node.color,
|
||||||
|
color: !node.color ? '' : 'white',
|
||||||
|
zIndex: 5, //app.graph._nodes.indexOf(node),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Extracts the type and link type from a widget config object.
|
||||||
|
* @param {*} config
|
||||||
|
* @returns
|
||||||
|
*/
|
||||||
|
export function getWidgetType(config) {
|
||||||
|
// Special handling for COMBO so we restrict links based on the entries
|
||||||
|
let type = config?.[0]
|
||||||
|
let linkType = type
|
||||||
|
if (Array.isArray(type)) {
|
||||||
|
type = 'COMBO'
|
||||||
|
linkType = linkType.join(',')
|
||||||
|
}
|
||||||
|
return { type, linkType }
|
||||||
|
}
|
||||||
|
|
||||||
|
// #endregion
|
||||||
|
|
||||||
|
// #region dynamic connections
|
||||||
|
/**
|
||||||
|
* @param {NodeType} nodeType
|
||||||
|
* @param {str} prefix
|
||||||
|
* @param {str | [str]} inputType
|
||||||
|
* @param {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} opts
|
||||||
|
* @returns
|
||||||
|
*/
|
||||||
|
|
||||||
|
export const setupDynamicConnections = (nodeType, prefix, inputType, opts) => {
|
||||||
|
infoLogger('Setting up dynamic connections for', nodeType)
|
||||||
|
|
||||||
|
/** @type {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} */
|
||||||
|
const options = opts || {}
|
||||||
|
const onNodeCreated = nodeType.prototype.onNodeCreated
|
||||||
|
const inputList = typeof inputType === 'object'
|
||||||
|
|
||||||
|
nodeType.prototype.onNodeCreated = function () {
|
||||||
|
const r = onNodeCreated ? onNodeCreated.apply(this, []) : undefined
|
||||||
|
this.addInput(`${prefix}_1`, inputList ? '*' : inputType)
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
|
||||||
|
const onConnectionsChange = nodeType.prototype.onConnectionsChange
|
||||||
|
/**
|
||||||
|
* @param {OnConnectionsChangeParams} args
|
||||||
|
*/
|
||||||
|
nodeType.prototype.onConnectionsChange = function (...args) {
|
||||||
|
const [type, slotIndex, isConnected, link, ioSlot] = args
|
||||||
|
|
||||||
|
options.link = link
|
||||||
|
options.ioSlot = ioSlot
|
||||||
|
const r = onConnectionsChange
|
||||||
|
? onConnectionsChange.apply(this, [
|
||||||
|
type,
|
||||||
|
slotIndex,
|
||||||
|
isConnected,
|
||||||
|
link,
|
||||||
|
ioSlot,
|
||||||
|
])
|
||||||
|
: undefined
|
||||||
|
options.DEBUG = {
|
||||||
|
node: this,
|
||||||
|
type,
|
||||||
|
slotIndex,
|
||||||
|
isConnected,
|
||||||
|
link,
|
||||||
|
ioSlot,
|
||||||
|
}
|
||||||
|
|
||||||
|
dynamic_connection(
|
||||||
|
this,
|
||||||
|
slotIndex,
|
||||||
|
isConnected,
|
||||||
|
`${prefix}_`,
|
||||||
|
inputType,
|
||||||
|
options,
|
||||||
|
)
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Main logic around dynamic inputs
|
||||||
|
*
|
||||||
|
* @param {LGraphNode} node - The target node
|
||||||
|
* @param {number} index - The slot index of the currently changed connection
|
||||||
|
* @param {bool} connected - Was this event connecting or disconnecting
|
||||||
|
* @param {string} [connectionPrefix] - The common prefix of the dynamic inputs
|
||||||
|
* @param {string|[string]} [connectionType] - The type of the dynamic connection
|
||||||
|
* @param {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options
|
||||||
|
*/
|
||||||
|
export const dynamic_connection = (
|
||||||
|
node,
|
||||||
|
index,
|
||||||
|
connected,
|
||||||
|
connectionPrefix = 'input_',
|
||||||
|
connectionType = '*',
|
||||||
|
opts = undefined,
|
||||||
|
) => {
|
||||||
|
/* @type {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options*/
|
||||||
|
const options = opts || {}
|
||||||
|
|
||||||
|
if (
|
||||||
|
node.inputs.length > 0 &&
|
||||||
|
!node.inputs[index].name.startsWith(connectionPrefix)
|
||||||
|
) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
const listConnection = typeof connectionType === 'object'
|
||||||
|
|
||||||
|
const conType = listConnection ? '*' : connectionType
|
||||||
|
const nameArray = options.nameArray || []
|
||||||
|
|
||||||
|
const clean_inputs = () => {
|
||||||
|
if (node.inputs.length === 0) return
|
||||||
|
|
||||||
|
let w_count = node.widgets?.length || 0
|
||||||
|
let i_count = node.inputs?.length || 0
|
||||||
|
infoLogger(`Cleaning inputs: [BEFORE] (w: ${w_count} | inputs: ${i_count})`)
|
||||||
|
|
||||||
|
const to_remove = []
|
||||||
|
for (let n = 1; n < node.inputs.length; n++) {
|
||||||
|
const element = node.inputs[n]
|
||||||
|
if (!element.link) {
|
||||||
|
if (node.widgets) {
|
||||||
|
const w = node.widgets.find((w) => w.name === element.name)
|
||||||
|
if (w) {
|
||||||
|
w.onRemoved?.()
|
||||||
|
node.widgets.length = node.widgets.length - 1
|
||||||
|
}
|
||||||
|
}
|
||||||
|
infoLogger(`Removing input ${n}`)
|
||||||
|
to_remove.push(n)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for (let i = 0; i < to_remove.length; i++) {
|
||||||
|
const id = to_remove[i]
|
||||||
|
|
||||||
|
node.removeInput(id)
|
||||||
|
i_count -= 1
|
||||||
|
}
|
||||||
|
node.inputs.length = i_count
|
||||||
|
|
||||||
|
w_count = node.widgets?.length || 0
|
||||||
|
i_count = node.inputs?.length || 0
|
||||||
|
infoLogger(`Cleaning inputs: [AFTER] (w: ${w_count} | inputs: ${i_count})`)
|
||||||
|
|
||||||
|
infoLogger('Cleaning inputs: making it sequential again')
|
||||||
|
// make inputs sequential again
|
||||||
|
for (let i = 0; i < node.inputs.length; i++) {
|
||||||
|
let name = `${connectionPrefix}${i + 1}`
|
||||||
|
|
||||||
|
if (nameArray.length > 0) {
|
||||||
|
name = i < nameArray.length ? nameArray[i] : name
|
||||||
|
}
|
||||||
|
|
||||||
|
node.inputs[i].label = name
|
||||||
|
node.inputs[i].name = name
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if (!connected) {
|
||||||
|
if (!options.link) {
|
||||||
|
infoLogger('Disconnecting', { options })
|
||||||
|
|
||||||
|
clean_inputs()
|
||||||
|
} else {
|
||||||
|
if (!options.ioSlot.link) {
|
||||||
|
node.connectionTransit = true
|
||||||
|
} else {
|
||||||
|
node.connectionTransit = false
|
||||||
|
clean_inputs()
|
||||||
|
}
|
||||||
|
infoLogger('Reconnecting', { options })
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if (connected) {
|
||||||
|
if (options.link) {
|
||||||
|
const { from, to, type } = nodesFromLink(node, options.link)
|
||||||
|
if (type === 'outgoing') return
|
||||||
|
infoLogger('Connecting', { options, from, to, type })
|
||||||
|
} else {
|
||||||
|
infoLogger('Connecting', { options })
|
||||||
|
}
|
||||||
|
|
||||||
|
if (node.connectionTransit) {
|
||||||
|
infoLogger('In Transit')
|
||||||
|
node.connectionTransit = false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Remove inputs and their widget if not linked.
|
||||||
|
clean_inputs()
|
||||||
|
|
||||||
|
if (node.inputs.length === 0) return
|
||||||
|
// add an extra input
|
||||||
|
if (node.inputs[node.inputs.length - 1].link !== null) {
|
||||||
|
const nextIndex = node.inputs.length
|
||||||
|
const name =
|
||||||
|
nextIndex < nameArray.length
|
||||||
|
? nameArray[nextIndex]
|
||||||
|
: `${connectionPrefix}${nextIndex + 1}`
|
||||||
|
|
||||||
|
infoLogger(`Adding input ${nextIndex + 1} (${name})`)
|
||||||
|
node.addInput(name, conType)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// #endregion
|
||||||
|
|
||||||
|
// #region color utils
|
||||||
export function isColorBright(rgb, threshold = 240) {
|
export function isColorBright(rgb, threshold = 240) {
|
||||||
const brightess = getBrightness(rgb)
|
const brightess = getBrightness(rgb)
|
||||||
return brightess > threshold
|
return brightess > threshold
|
||||||
@@ -425,8 +561,36 @@ function getBrightness(rgbObj) {
|
|||||||
1000,
|
1000,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
// #endregion
|
||||||
|
|
||||||
|
// #region html/css utils
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Calculate total height of DOM element child
|
||||||
|
*
|
||||||
|
* @param {HTMLElement} parentElement - The target dom element
|
||||||
|
* @returns {number} the total height
|
||||||
|
*/
|
||||||
|
export function calculateTotalChildrenHeight(parentElement) {
|
||||||
|
let totalHeight = 0
|
||||||
|
|
||||||
|
for (const child of parentElement.children) {
|
||||||
|
const style = window.getComputedStyle(child)
|
||||||
|
|
||||||
|
// Get height as an integer (without 'px')
|
||||||
|
const height = Number.parseInt(style.height, 10)
|
||||||
|
|
||||||
|
// Get vertical margin as integers
|
||||||
|
const marginTop = Number.parseInt(style.marginTop, 10)
|
||||||
|
const marginBottom = Number.parseInt(style.marginBottom, 10)
|
||||||
|
|
||||||
|
// Sum up height and vertical margins
|
||||||
|
totalHeight += height + marginTop + marginBottom
|
||||||
|
}
|
||||||
|
|
||||||
|
return totalHeight
|
||||||
|
}
|
||||||
|
|
||||||
//- HTML / CSS UTILS
|
|
||||||
export const loadScript = (
|
export const loadScript = (
|
||||||
FILE_URL,
|
FILE_URL,
|
||||||
async = true,
|
async = true,
|
||||||
@@ -497,22 +661,9 @@ export function defineClass(className, classStyles) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/** Prefixes the node title with '[DEPRECATED]' and log the deprecation reason to the console.*/
|
// #endregion
|
||||||
export const addDeprecation = (nodeType, reason) => {
|
|
||||||
const title = nodeType.title
|
|
||||||
nodeType.title = `[DEPRECATED] ${title}`
|
|
||||||
// console.log(nodeType)
|
|
||||||
|
|
||||||
const styles = {
|
// #region documentation widget
|
||||||
title: 'font-size:1.3em;font-weight:900;color:yellow; background: black',
|
|
||||||
reason: 'font-size:1.2em',
|
|
||||||
}
|
|
||||||
console.log(
|
|
||||||
`%c! ${title} is deprecated:%c ${reason}`,
|
|
||||||
styles.title,
|
|
||||||
styles.reason,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
const create_documentation_stylesheet = () => {
|
const create_documentation_stylesheet = () => {
|
||||||
const tag = 'mtb-documentation-stylesheet'
|
const tag = 'mtb-documentation-stylesheet'
|
||||||
@@ -626,6 +777,24 @@ export const addDocumentation = (
|
|||||||
const iconMargin = options.icon_margin || 4
|
const iconMargin = options.icon_margin || 4
|
||||||
let docElement = null
|
let docElement = null
|
||||||
let wrapper = null
|
let wrapper = null
|
||||||
|
|
||||||
|
const onRem = nodeType.prototype.onRemoved
|
||||||
|
|
||||||
|
nodeType.prototype.onRemoved = function () {
|
||||||
|
const r = onRem ? onRem.apply(this, []) : undefined
|
||||||
|
|
||||||
|
if (docElement) {
|
||||||
|
docElement.remove()
|
||||||
|
docElement = null
|
||||||
|
}
|
||||||
|
|
||||||
|
if (wrapper) {
|
||||||
|
wrapper.remove()
|
||||||
|
wrapper = null
|
||||||
|
}
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
|
||||||
const drawFg = nodeType.prototype.onDrawForeground
|
const drawFg = nodeType.prototype.onDrawForeground
|
||||||
|
|
||||||
/**
|
/**
|
||||||
@@ -640,15 +809,14 @@ export const addDocumentation = (
|
|||||||
// icon position
|
// icon position
|
||||||
const x = this.size[0] - iconSize - iconMargin
|
const x = this.size[0] - iconSize - iconMargin
|
||||||
|
|
||||||
|
let resizeHandle
|
||||||
|
// create it
|
||||||
if (this.show_doc && docElement === null) {
|
if (this.show_doc && docElement === null) {
|
||||||
create_documentation_stylesheet()
|
create_documentation_stylesheet()
|
||||||
|
|
||||||
docElement = document.createElement('div')
|
docElement = document.createElement('div')
|
||||||
docElement.classList.add('documentation-popup')
|
docElement.classList.add('documentation-popup')
|
||||||
document.body.appendChild(docElement)
|
document.body.appendChild(docElement)
|
||||||
// docElement.innerHTML = documentationConverter.makeHtml(
|
|
||||||
// nodeData.description,
|
|
||||||
// )
|
|
||||||
|
|
||||||
wrapper = document.createElement('div')
|
wrapper = document.createElement('div')
|
||||||
wrapper.classList.add('documentation-wrapper')
|
wrapper.classList.add('documentation-wrapper')
|
||||||
@@ -656,24 +824,21 @@ export const addDocumentation = (
|
|||||||
docElement.appendChild(wrapper)
|
docElement.appendChild(wrapper)
|
||||||
|
|
||||||
// resize handle
|
// resize handle
|
||||||
const resizeHandle = document.createElement('div')
|
resizeHandle = document.createElement('div')
|
||||||
resizeHandle.style.width = '10px'
|
resizeHandle.style.width = '0'
|
||||||
resizeHandle.style.height = '10px'
|
resizeHandle.style.height = '0'
|
||||||
// resizeHandle.style.background = 'gray'
|
|
||||||
resizeHandle.style.position = 'absolute'
|
resizeHandle.style.position = 'absolute'
|
||||||
resizeHandle.style.bottom = '0'
|
resizeHandle.style.bottom = '0'
|
||||||
resizeHandle.style.right = '0'
|
resizeHandle.style.right = '0'
|
||||||
// resizeHandle.style.left = '95%'
|
|
||||||
resizeHandle.style.cursor = 'se-resize'
|
resizeHandle.style.cursor = 'se-resize'
|
||||||
resizeHandle.style.userSelect = 'none'
|
resizeHandle.style.userSelect = 'none'
|
||||||
|
|
||||||
const borderColor = getComputedStyle(document.documentElement)
|
resizeHandle.style.borderWidth = '15px'
|
||||||
.getPropertyValue('--border-color')
|
resizeHandle.style.borderStyle = 'solid'
|
||||||
.trim()
|
|
||||||
resizeHandle.style.borderTop = '10px solid transparent'
|
resizeHandle.style.borderColor =
|
||||||
resizeHandle.style.borderLeft = '10px solid transparent'
|
'transparent var(--border-color) var(--border-color) transparent'
|
||||||
resizeHandle.style.borderBottom = `10px solid ${borderColor}`
|
|
||||||
resizeHandle.style.borderRight = `10px solid ${borderColor}`
|
|
||||||
|
|
||||||
wrapper.appendChild(resizeHandle)
|
wrapper.appendChild(resizeHandle)
|
||||||
let isResizing = false
|
let isResizing = false
|
||||||
@@ -683,41 +848,54 @@ export const addDocumentation = (
|
|||||||
let startWidth
|
let startWidth
|
||||||
let startHeight
|
let startHeight
|
||||||
|
|
||||||
resizeHandle.addEventListener('mousedown', (e) => {
|
resizeHandle.addEventListener(
|
||||||
e.stopPropagation()
|
'mousedown',
|
||||||
isResizing = true
|
(e) => {
|
||||||
startX = e.clientX
|
e.stopPropagation()
|
||||||
startY = e.clientY
|
isResizing = true
|
||||||
startWidth = Number.parseInt(
|
startX = e.clientX
|
||||||
document.defaultView.getComputedStyle(docElement).width,
|
startY = e.clientY
|
||||||
10,
|
startWidth = Number.parseInt(
|
||||||
)
|
document.defaultView.getComputedStyle(docElement).width,
|
||||||
startHeight = Number.parseInt(
|
10,
|
||||||
document.defaultView.getComputedStyle(docElement).height,
|
)
|
||||||
10,
|
startHeight = Number.parseInt(
|
||||||
)
|
document.defaultView.getComputedStyle(docElement).height,
|
||||||
})
|
10,
|
||||||
|
)
|
||||||
|
},
|
||||||
|
|
||||||
document.addEventListener('mousemove', (e) => {
|
{ signal: this.docCtrl.signal },
|
||||||
console.log('Moving mouse')
|
)
|
||||||
if (!isResizing) return
|
|
||||||
const newWidth = startWidth + e.clientX - startX
|
|
||||||
const newHeight = startHeight + e.clientY - startY
|
|
||||||
|
|
||||||
docElement.style.width = `${newWidth}px`
|
document.addEventListener(
|
||||||
docElement.style.height = `${newHeight}px`
|
'mousemove',
|
||||||
|
(e) => {
|
||||||
|
if (!isResizing) return
|
||||||
|
const scale = app.canvas.ds.scale
|
||||||
|
const newWidth = startWidth + (e.clientX - startX) / scale
|
||||||
|
const newHeight = startHeight + (e.clientY - startY) / scale
|
||||||
|
|
||||||
this.docPos = {
|
docElement.style.width = `${newWidth}px`
|
||||||
width: `${newWidth}px`,
|
docElement.style.height = `${newHeight}px`
|
||||||
height: `${newHeight}px`,
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
document.addEventListener('mouseup', () => {
|
this.docPos = {
|
||||||
isResizing = false
|
width: `${newWidth}px`,
|
||||||
})
|
height: `${newHeight}px`,
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{ signal: this.docCtrl.signal },
|
||||||
|
)
|
||||||
|
|
||||||
|
document.addEventListener(
|
||||||
|
'mouseup',
|
||||||
|
() => {
|
||||||
|
isResizing = false
|
||||||
|
},
|
||||||
|
{ signal: this.docCtrl.signal },
|
||||||
|
)
|
||||||
} else if (!this.show_doc && docElement !== null) {
|
} else if (!this.show_doc && docElement !== null) {
|
||||||
docElement.parentNode.removeChild(docElement)
|
docElement.remove()
|
||||||
docElement = null
|
docElement = null
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -725,12 +903,13 @@ export const addDocumentation = (
|
|||||||
if (this.show_doc && docElement !== null) {
|
if (this.show_doc && docElement !== null) {
|
||||||
const rect = ctx.canvas.getBoundingClientRect()
|
const rect = ctx.canvas.getBoundingClientRect()
|
||||||
|
|
||||||
|
const dpi = Math.max(1.0, window.devicePixelRatio)
|
||||||
const scaleX = rect.width / ctx.canvas.width
|
const scaleX = rect.width / ctx.canvas.width
|
||||||
const scaleY = rect.height / ctx.canvas.height
|
const scaleY = rect.height / ctx.canvas.height
|
||||||
const transform = new DOMMatrix()
|
const transform = new DOMMatrix()
|
||||||
.scaleSelf(scaleX, scaleY)
|
.scaleSelf(scaleX, scaleY)
|
||||||
.multiplySelf(ctx.getTransform())
|
.multiplySelf(ctx.getTransform())
|
||||||
.translateSelf(this.size[0] * scaleX, 0)
|
.translateSelf(this.size[0] * scaleX * dpi, 0)
|
||||||
.translateSelf(10, -32)
|
.translateSelf(10, -32)
|
||||||
|
|
||||||
const scale = new DOMMatrix().scaleSelf(transform.a, transform.d)
|
const scale = new DOMMatrix().scaleSelf(transform.a, transform.d)
|
||||||
@@ -742,12 +921,6 @@ export const addDocumentation = (
|
|||||||
top: `${transform.d + transform.f}px`,
|
top: `${transform.d + transform.f}px`,
|
||||||
width: this.docPos ? this.docPos.width : `${this.size[0] * 1.5}px`,
|
width: this.docPos ? this.docPos.width : `${this.size[0] * 1.5}px`,
|
||||||
height: this.docPos?.height,
|
height: this.docPos?.height,
|
||||||
// width: `${this.size[0] * 2}px`,
|
|
||||||
// height: `${(widget.parent?.inputHeight || 32) - (margin * 2)}px`,
|
|
||||||
// height: `${this.size[1] || this.parent?.inputHeight || 32}px`,
|
|
||||||
|
|
||||||
// background: !node.color ? "" : node.color,
|
|
||||||
// color: "blue", //!node.color ? "" : "white",
|
|
||||||
})
|
})
|
||||||
|
|
||||||
if (this.docPos === undefined) {
|
if (this.docPos === undefined) {
|
||||||
@@ -756,23 +929,22 @@ export const addDocumentation = (
|
|||||||
height: docElement.style.height,
|
height: docElement.style.height,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// docElement.style.left = 140 - rect.right + "px";
|
|
||||||
// docElement.style.top = rect.top + "px";
|
|
||||||
}
|
}
|
||||||
|
|
||||||
ctx.save()
|
ctx.save()
|
||||||
ctx.translate(x, iconSize - 34) // Position the icon on the canvas
|
ctx.translate(x, iconSize - 34)
|
||||||
ctx.scale(iconSize / 32, iconSize / 32) // Scale the icon to the desired size
|
ctx.scale(iconSize / 32, iconSize / 32)
|
||||||
ctx.strokeStyle = 'rgba(255,255,255,0.3)'
|
ctx.strokeStyle = 'rgba(255,255,255,0.3)'
|
||||||
|
|
||||||
ctx.lineCap = 'round'
|
ctx.lineCap = 'round'
|
||||||
ctx.lineJoin = 'round'
|
ctx.lineJoin = 'round'
|
||||||
|
|
||||||
ctx.lineWidth = 2.4
|
ctx.lineWidth = 2.4
|
||||||
// ctx.stroke(questionMark);
|
|
||||||
ctx.font = 'bold 36px monospace'
|
ctx.font = 'bold 36px monospace'
|
||||||
ctx.fillText('?', 0, 24)
|
ctx.fillText('?', 0, 24)
|
||||||
|
|
||||||
|
// ctx.font = `bold ${this.show_doc ? 36 : 24}px monospace`
|
||||||
|
// ctx.fillText(`${this.show_doc ? '▼' : '▶'}`, 24, 24)
|
||||||
ctx.restore()
|
ctx.restore()
|
||||||
|
|
||||||
return r
|
return r
|
||||||
@@ -800,6 +972,11 @@ export const addDocumentation = (
|
|||||||
} else {
|
} else {
|
||||||
this.show_doc = !this.show_doc
|
this.show_doc = !this.show_doc
|
||||||
}
|
}
|
||||||
|
if (this.show_doc) {
|
||||||
|
this.docCtrl = new AbortController()
|
||||||
|
} else {
|
||||||
|
this.docCtrl.abort()
|
||||||
|
}
|
||||||
return true // Return true to indicate the event was handled
|
return true // Return true to indicate the event was handled
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -808,3 +985,86 @@ export const addDocumentation = (
|
|||||||
// return r;
|
// return r;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// #endregion
|
||||||
|
|
||||||
|
// #region node extensions
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Extend an object, either replacing the original property or extending it.
|
||||||
|
* @param {Object} object - The object to which the property belongs.
|
||||||
|
* @param {string} property - The name of the property to chain the callback to.
|
||||||
|
* @param {Function} callback - The callback function to be chained.
|
||||||
|
*/
|
||||||
|
export function extendPrototype(object, property, callback) {
|
||||||
|
if (object === undefined) {
|
||||||
|
console.error('Could not extend undefined object', { object, property })
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if (property in object) {
|
||||||
|
const callback_orig = object[property]
|
||||||
|
object[property] = function (...args) {
|
||||||
|
const r = callback_orig.apply(this, args)
|
||||||
|
callback.apply(this, args)
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
object[property] = callback
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Appends a callback to the extra menu options of a given node type.
|
||||||
|
* @param {NodeType} nodeType
|
||||||
|
* @param {(app,options) => ContextMenuItem[]} cb
|
||||||
|
*/
|
||||||
|
export function addMenuHandler(nodeType, cb) {
|
||||||
|
const getOpts = nodeType.prototype.getExtraMenuOptions
|
||||||
|
/**
|
||||||
|
* @returns {ContextMenuItem[]} items
|
||||||
|
*/
|
||||||
|
nodeType.prototype.getExtraMenuOptions = function (app, options) {
|
||||||
|
const r = getOpts.apply(this, [app, options]) || []
|
||||||
|
const newItems = cb.apply(this, [app, options])
|
||||||
|
return r + newItems
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/** Prefixes the node title with '[DEPRECATED]' and log the deprecation reason to the console.*/
|
||||||
|
export const addDeprecation = (nodeType, reason) => {
|
||||||
|
const title = nodeType.title
|
||||||
|
nodeType.title = `[DEPRECATED] ${title}`
|
||||||
|
// console.log(nodeType)
|
||||||
|
|
||||||
|
const styles = {
|
||||||
|
title: 'font-size:1.3em;font-weight:900;color:yellow; background: black',
|
||||||
|
reason: 'font-size:1.2em',
|
||||||
|
}
|
||||||
|
console.log(
|
||||||
|
`%c! ${title} is deprecated:%c ${reason}`,
|
||||||
|
styles.title,
|
||||||
|
styles.reason,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
// #endregion
|
||||||
|
|
||||||
|
// #region graph utilities
|
||||||
|
export const getNodes = (skip_unused) => {
|
||||||
|
const nodes = []
|
||||||
|
for (const outerNode of app.graph.computeExecutionOrder(false)) {
|
||||||
|
const skipNode =
|
||||||
|
(outerNode.mode === 2 || outerNode.mode === 4) && skip_unused
|
||||||
|
const innerNodes =
|
||||||
|
!skipNode && outerNode.getInnerNodes
|
||||||
|
? outerNode.getInnerNodes()
|
||||||
|
: [outerNode]
|
||||||
|
for (const node of innerNodes) {
|
||||||
|
if ((node.mode === 2 || node.mode === 4) && skip_unused) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
nodes.push(node)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nodes
|
||||||
|
}
|
||||||
|
|||||||
+356
-122
@@ -1,62 +1,225 @@
|
|||||||
import { app } from '../../scripts/app.js'
|
import { app } from '../../scripts/app.js'
|
||||||
import * as shared from './comfy_shared.js'
|
import * as shared from './comfy_shared.js'
|
||||||
|
import { infoLogger } from './comfy_shared.js'
|
||||||
import { MtbWidgets } from './mtb_widgets.js'
|
import { MtbWidgets } from './mtb_widgets.js'
|
||||||
|
import { ComfyWidgets } from '../../scripts/widgets.js'
|
||||||
|
import * as mtb_widgets from './mtb_widgets.js'
|
||||||
|
|
||||||
export class Constant extends LiteGraph.LGraphNode {
|
/**
|
||||||
constructor() {
|
* @typedef {'number'|'string'|'vector2'|'vector3'|'vector4'|'color'} ConstantType
|
||||||
super()
|
* @typedef {import ("../../../web/types/litegraph.d.ts").LGraphNode} Node
|
||||||
this.uuid = shared.makeUUID()
|
* @typedef {{x:number,y:number,z?:number,w?:number}} VectorValue
|
||||||
this.collapsable = true
|
* @typedef {}
|
||||||
|
*
|
||||||
|
*/
|
||||||
|
|
||||||
// this avoid serializing the node when converting to prompt
|
/**
|
||||||
this.isVirtualNode = true
|
* @param {number} size - The number of axis of the vector (2,3 or 4)
|
||||||
|
* @param {number} val - The default scalar value to fill the vector with
|
||||||
this.shape = LiteGraph.BOX_SHAPE
|
* @returns {VectorValue} vector
|
||||||
this.serialize_widgets = true
|
* */
|
||||||
|
const initVector = (size, val = 0.0) => {
|
||||||
// Properties
|
const res = {}
|
||||||
this.addProperty('type', 'number')
|
for (let i = 0; i < size; i++) {
|
||||||
this.addProperty('value', 0)
|
const axis = mtb_widgets.VECTOR_AXIS[i]
|
||||||
|
res[axis] = val
|
||||||
// Inputs and outputs
|
|
||||||
this.addOutput('Output', '*')
|
|
||||||
|
|
||||||
// Widget for selecting the type
|
|
||||||
this.addWidget(
|
|
||||||
'combo',
|
|
||||||
'Type',
|
|
||||||
this.properties.type,
|
|
||||||
(value) => {
|
|
||||||
this.properties.type = value
|
|
||||||
this.updateWidgets()
|
|
||||||
this.updateOutputType()
|
|
||||||
},
|
|
||||||
{
|
|
||||||
values: ['number', 'string', 'vector2', 'vector3', 'vector4', 'color'],
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
// Initialize the node
|
|
||||||
this.updateWidgets()
|
|
||||||
this.updateOutputType()
|
|
||||||
}
|
}
|
||||||
|
return res
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
*
|
||||||
|
* @extends {Node}
|
||||||
|
* @classdesc Wrapper for the python node
|
||||||
|
*/
|
||||||
|
export class ConstantJs {
|
||||||
|
constructor(python_node) {
|
||||||
|
// this.uuid = shared.makeUUID()
|
||||||
|
const wrapper = this
|
||||||
|
|
||||||
|
python_node.shape = LiteGraph.BOX_SHAPE
|
||||||
|
python_node.serialize_widgets = true
|
||||||
|
|
||||||
|
const onNodeCreated = python_node.prototype.onNodeCreated
|
||||||
|
python_node.prototype.onNodeCreated = function () {
|
||||||
|
const r = onNodeCreated ? onNodeCreated.apply(this) : undefined
|
||||||
|
|
||||||
|
this.addProperty('type', 'number')
|
||||||
|
this.addProperty('value', 0)
|
||||||
|
|
||||||
|
this.removeInput(0)
|
||||||
|
this.removeOutput(0)
|
||||||
|
|
||||||
|
this.addOutput('Output', '*')
|
||||||
|
|
||||||
|
// bind our wrapper
|
||||||
|
this.configure = wrapper.configure.bind(this)
|
||||||
|
// this.applyToGraph = wrapper.applyToGraph.bind(this)
|
||||||
|
this.updateWidgets = wrapper.updateWidgets.bind(this)
|
||||||
|
this.convertValue = wrapper.convertValue.bind(this)
|
||||||
|
// this.updateOutput = wrapper.updateOutput.bind(this)
|
||||||
|
this.updateOutputType = wrapper.updateOutputType.bind(this)
|
||||||
|
// this.updateTargetWidgets = wrapper.updateTargetWidgets.bind(this)
|
||||||
|
|
||||||
|
this.addWidget(
|
||||||
|
'combo',
|
||||||
|
'Type',
|
||||||
|
this.properties.type,
|
||||||
|
(value) => {
|
||||||
|
this.properties.type = value
|
||||||
|
this.updateWidgets()
|
||||||
|
this.updateOutputType()
|
||||||
|
},
|
||||||
|
{
|
||||||
|
values: [
|
||||||
|
// 'number',
|
||||||
|
'float',
|
||||||
|
'int',
|
||||||
|
'string',
|
||||||
|
'vector2',
|
||||||
|
'vector3',
|
||||||
|
'vector4',
|
||||||
|
'color',
|
||||||
|
],
|
||||||
|
},
|
||||||
|
)
|
||||||
|
this.updateWidgets()
|
||||||
|
this.updateOutputType()
|
||||||
|
|
||||||
|
for (let n = 0; n < this.inputs.length; n++) {
|
||||||
|
this.removeInput(n)
|
||||||
|
}
|
||||||
|
this.inputs = []
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
// NOTE: this is called onPrompt
|
// NOTE: this is called onPrompt
|
||||||
applyToGraph() {
|
// applyToGraph() {
|
||||||
this.updateTargetWidgets()
|
// infoLogger('Updating values for backend')
|
||||||
}
|
// this.updateTargetWidgets()
|
||||||
|
// }
|
||||||
|
|
||||||
|
// NOTE: deserialization happens here
|
||||||
configure(info) {
|
configure(info) {
|
||||||
super.configure(info)
|
// super.configure(info)
|
||||||
|
infoLogger('Configure Constant', { info, node: this })
|
||||||
|
|
||||||
this.properties.type = info.properties.type
|
this.properties.type = info.properties.type
|
||||||
this.properties.value = info.properties.value
|
this.properties.value = info.properties.value
|
||||||
|
|
||||||
shared.infoLogger('Configure Constant', { info, node: this })
|
this.pos = info.pos
|
||||||
|
this.order = info.order
|
||||||
|
|
||||||
this.updateWidgets()
|
this.updateWidgets()
|
||||||
this.updateOutputType()
|
this.updateOutputType()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Convert the old value type to the new one, falling back to some default
|
||||||
|
* @param {ConstantType} propType - The target type
|
||||||
|
*/
|
||||||
|
convertValue(propType) {
|
||||||
|
switch (propType) {
|
||||||
|
case 'color': {
|
||||||
|
if (typeof this.properties.value !== 'string') {
|
||||||
|
this.properties.value = '#ffffff'
|
||||||
|
} else if (this.properties.value[0] !== '#') {
|
||||||
|
this.properties.value = '#ff0000'
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
case 'int': {
|
||||||
|
if (typeof this.properties.value === 'object') {
|
||||||
|
this.properties.value = Number.parseInt(this.properties.value.x)
|
||||||
|
} else {
|
||||||
|
this.properties.value = Number.parseInt(this.properties.value) || 0
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
case 'float': {
|
||||||
|
if (typeof this.properties.value === 'object') {
|
||||||
|
this.properties.value = Number.parseFloat(this.properties.value.x)
|
||||||
|
} else {
|
||||||
|
this.properties.value =
|
||||||
|
Number.parseFloat(this.properties.value) || 0.0
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
case 'string': {
|
||||||
|
if (typeof this.properties.value !== 'string') {
|
||||||
|
this.properties.value = JSON.stringify(this.properties.value)
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
case 'vector2':
|
||||||
|
case 'vector3':
|
||||||
|
case 'vector4': {
|
||||||
|
const numInputs = Number.parseInt(propType.charAt(6))
|
||||||
|
if (!this.properties.value) {
|
||||||
|
this.properties.value = initVector(numInputs) // Array.from({ length: numInputs }, () => 0.0)
|
||||||
|
} else if (typeof this.properties.value === 'string') {
|
||||||
|
try {
|
||||||
|
const parsed = JSON.parse(this.properties.value)
|
||||||
|
const newVec = {}
|
||||||
|
for (
|
||||||
|
let i = 0;
|
||||||
|
i < Object.keys(mtb_widgets.VECTOR_AXIS).length;
|
||||||
|
i++
|
||||||
|
) {
|
||||||
|
const axis = mtb_widgets.VECTOR_AXIS[i]
|
||||||
|
if (Object.keys(parsed).includes(axis)) {
|
||||||
|
newVec[axis] = parsed[axis]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
this.properties.value = newVec
|
||||||
|
} catch (e) {
|
||||||
|
shared.errorLogger(e)
|
||||||
|
infoLogger(
|
||||||
|
`Couldn't parse string to vec (${this.properties.value})`,
|
||||||
|
)
|
||||||
|
this.properties.value = initVector(numInputs)
|
||||||
|
}
|
||||||
|
} else if (typeof this.properties.value === 'number') {
|
||||||
|
const newVec = initVector(numInputs)
|
||||||
|
newVec.x = Number.parseFloat(this.properties.value)
|
||||||
|
this.properties.value = newVec
|
||||||
|
}
|
||||||
|
|
||||||
|
if (
|
||||||
|
typeof this.properties.value === 'object' &&
|
||||||
|
Object.keys(this.properties.value).length !== numInputs
|
||||||
|
) {
|
||||||
|
const current = Object.keys(this.properties.value)
|
||||||
|
if (current.length < numInputs) {
|
||||||
|
infoLogger('current value smaller than target, adjusting')
|
||||||
|
for (let index = current.length; index < numInputs; index++) {
|
||||||
|
this.properties.value[mtb_widgets.VECTOR_AXIS[index]] = 0.0
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
infoLogger('current value greater than target, adjusting')
|
||||||
|
const newVal = {}
|
||||||
|
for (let index = 0; index < numInputs; index++) {
|
||||||
|
newVal[mtb_widgets.VECTOR_AXIS[index]] =
|
||||||
|
this.properties.value[mtb_widgets.VECTOR_AXIS[index]]
|
||||||
|
}
|
||||||
|
this.properties.value = newVal
|
||||||
|
}
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Remove all widgets but the comboBox for selecting the type
|
||||||
|
* then recreate the appropriate widget from scratch
|
||||||
|
*/
|
||||||
updateWidgets() {
|
updateWidgets() {
|
||||||
// Remove existing widgets
|
// NOTE: Remove existing widgets
|
||||||
for (let i = 1; i < this.widgets.length; i++) {
|
for (let i = 1; i < this.widgets.length; i++) {
|
||||||
const element = this.widgets[i]
|
const element = this.widgets[i]
|
||||||
if (element.onRemove) {
|
if (element.onRemove) {
|
||||||
@@ -66,28 +229,113 @@ export class Constant extends LiteGraph.LGraphNode {
|
|||||||
}
|
}
|
||||||
|
|
||||||
this.widgets.splice(1)
|
this.widgets.splice(1)
|
||||||
|
this.widgets[0].value = this.properties.type
|
||||||
|
|
||||||
|
this.convertValue(this.properties.type)
|
||||||
|
|
||||||
switch (this.properties.type) {
|
switch (this.properties.type) {
|
||||||
case 'color': {
|
case 'color': {
|
||||||
if (typeof this.properties.value !== 'string') {
|
|
||||||
this.properties.value = '#ffffff'
|
|
||||||
}
|
|
||||||
const col_widget = this.addCustomWidget(
|
const col_widget = this.addCustomWidget(
|
||||||
MtbWidgets.COLOR('Value', this.properties.value || '#ff0000'),
|
MtbWidgets.COLOR('Value', this.properties.value),
|
||||||
)
|
)
|
||||||
col_widget.callback = (col) => {
|
col_widget.callback = (col) => {
|
||||||
this.properties.value = col
|
this.properties.value = col
|
||||||
this.updateOutput()
|
// this.updateOutput()
|
||||||
}
|
}
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
case 'number':
|
case 'int': {
|
||||||
|
const f_widget = this.addCustomWidget(
|
||||||
|
ComfyWidgets.INT(
|
||||||
|
this,
|
||||||
|
'Value',
|
||||||
|
[
|
||||||
|
'',
|
||||||
|
{
|
||||||
|
default: this.properties.value,
|
||||||
|
callback: (val) => console.log('VALUE', val),
|
||||||
|
},
|
||||||
|
],
|
||||||
|
app,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
f_widget.widget.callback = (val) => {
|
||||||
|
this.properties.value = val
|
||||||
|
}
|
||||||
|
|
||||||
|
break
|
||||||
|
}
|
||||||
|
case 'float': {
|
||||||
|
this.addWidget('number', 'Value', this.properties.value, (val) => {
|
||||||
|
this.properties.value = val
|
||||||
|
})
|
||||||
|
break
|
||||||
|
}
|
||||||
|
case 'string': {
|
||||||
|
mtb_widgets.addMultilineWidget(
|
||||||
|
this,
|
||||||
|
'Value',
|
||||||
|
{
|
||||||
|
defaultVal: this.properties.value,
|
||||||
|
},
|
||||||
|
(v) => {
|
||||||
|
this.properties.value = v
|
||||||
|
// this.updateOutput()
|
||||||
|
},
|
||||||
|
)
|
||||||
|
break
|
||||||
|
}
|
||||||
|
case 'vector2':
|
||||||
|
case 'vector3':
|
||||||
|
case 'vector4': {
|
||||||
|
const numInputs = Number.parseInt(this.properties.type.charAt(6))
|
||||||
|
const node = this
|
||||||
|
const v_widget = mtb_widgets.addVectorWidget(
|
||||||
|
this,
|
||||||
|
'Value',
|
||||||
|
this.properties.value, // value
|
||||||
|
numInputs, // vector_size
|
||||||
|
function (v) {
|
||||||
|
node.properties.value = v
|
||||||
|
// this.updateOutput()
|
||||||
|
},
|
||||||
|
)
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
|
// NOTE: this is not reached anymore, kept for reference
|
||||||
|
case 'number': {
|
||||||
if (typeof this.properties.value !== 'number') {
|
if (typeof this.properties.value !== 'number') {
|
||||||
this.properties.value = 0.0
|
this.properties.value = 0.0
|
||||||
}
|
}
|
||||||
this.addWidget('number', 'Value', this.properties.value, (value) => {
|
const n_widget = this.addWidget(
|
||||||
this.properties.value = value
|
'number',
|
||||||
this.updateOutput()
|
'Value',
|
||||||
})
|
this.properties.force_int
|
||||||
|
? Number.parseInt(this.properties.value)
|
||||||
|
: this.properties.value,
|
||||||
|
(value) => {
|
||||||
|
this.properties.value = this.properties.force_int
|
||||||
|
? Number.parseInt(value)
|
||||||
|
: value
|
||||||
|
// this.updateOutput()
|
||||||
|
},
|
||||||
|
)
|
||||||
|
//override the callback
|
||||||
|
const origCallback = n_widget.callback
|
||||||
|
const node = this
|
||||||
|
n_widget.callback = function (val) {
|
||||||
|
const r = origCallback ? origCallback.apply(this, [val]) : undefined
|
||||||
|
if (node.properties.force_int) {
|
||||||
|
// TODO: rework this, a it makes it harder to manipulate
|
||||||
|
this.value = Number.parseInt(this.value)
|
||||||
|
node.properties.value = Number.parseInt(this.value)
|
||||||
|
}
|
||||||
|
infoLogger('NEW NUMBER', this.value)
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
|
||||||
this.addWidget(
|
this.addWidget(
|
||||||
'toggle',
|
'toggle',
|
||||||
'Convert to Integer',
|
'Convert to Integer',
|
||||||
@@ -98,51 +346,6 @@ export class Constant extends LiteGraph.LGraphNode {
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
break
|
break
|
||||||
case 'string': {
|
|
||||||
if (typeof this.properties.value !== 'string') {
|
|
||||||
this.properties.value = `${this.properties.value}`
|
|
||||||
}
|
|
||||||
shared.addMultilineWidget(
|
|
||||||
this,
|
|
||||||
'Value',
|
|
||||||
{
|
|
||||||
defaultVal: this.properties.value,
|
|
||||||
},
|
|
||||||
(v) => {
|
|
||||||
this.properties.value = v
|
|
||||||
this.updateOutput()
|
|
||||||
},
|
|
||||||
)
|
|
||||||
break
|
|
||||||
}
|
|
||||||
case 'vector2':
|
|
||||||
case 'vector3':
|
|
||||||
case 'vector4': {
|
|
||||||
const numInputs = Number.parseInt(this.properties.type.charAt(6))
|
|
||||||
|
|
||||||
if (['string', 'number'].includes(typeof this.properties.value)) {
|
|
||||||
this.properties.value = Array.from({ length: numInputs }, () => 0.0)
|
|
||||||
} else if (this.properties.value.length !== numInputs) {
|
|
||||||
if (this.properties.value.length > numInputs) {
|
|
||||||
this.properties.value = this.properties.value.slice(0, numInputs)
|
|
||||||
} else {
|
|
||||||
this.properties.value = this.properties.value.concat(
|
|
||||||
new Array(numInputs - this.properties.value.length).fill(0.0),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
for (let i = 0; i < numInputs; i++) {
|
|
||||||
this.addWidget(
|
|
||||||
'number',
|
|
||||||
`Value ${i + 1}`,
|
|
||||||
this.properties.value[i] || 0,
|
|
||||||
(value) => {
|
|
||||||
this.properties.value[i] = value
|
|
||||||
this.updateOutput()
|
|
||||||
},
|
|
||||||
)
|
|
||||||
}
|
|
||||||
break
|
|
||||||
}
|
}
|
||||||
default:
|
default:
|
||||||
break
|
break
|
||||||
@@ -154,20 +357,28 @@ export class Constant extends LiteGraph.LGraphNode {
|
|||||||
this.updateTargetWidgets([link.id])
|
this.updateTargetWidgets([link.id])
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
updateOutputType() {
|
updateOutputType() {
|
||||||
const cur_type = this.outputs[0].type
|
infoLogger('Updating output type')
|
||||||
const rm_if_mismatch = (type) => {
|
const rm_if_mismatch = (type) => {
|
||||||
if (cur_type !== type) {
|
if (this.outputs[0].type !== type) {
|
||||||
for (let i = 0; i < this.outputs.length; i++) {
|
for (let i = 0; i < this.outputs.length; i++) {
|
||||||
this.removeOutput(i)
|
this.removeOutput(i)
|
||||||
}
|
}
|
||||||
this.addOutput('output', type)
|
this.addOutput('output', type)
|
||||||
|
// this.setOutputDataType(0, type)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
switch (this.properties.type) {
|
switch (this.properties.type) {
|
||||||
case 'color':
|
case 'color':
|
||||||
rm_if_mismatch('COLOR')
|
rm_if_mismatch('COLOR')
|
||||||
break
|
break
|
||||||
|
case 'float':
|
||||||
|
rm_if_mismatch('FLOAT')
|
||||||
|
break
|
||||||
|
case 'int':
|
||||||
|
rm_if_mismatch('INT')
|
||||||
|
break
|
||||||
case 'number':
|
case 'number':
|
||||||
if (this.properties.force_int) {
|
if (this.properties.force_int) {
|
||||||
rm_if_mismatch('INT')
|
rm_if_mismatch('INT')
|
||||||
@@ -178,6 +389,11 @@ export class Constant extends LiteGraph.LGraphNode {
|
|||||||
case 'string':
|
case 'string':
|
||||||
rm_if_mismatch('STRING')
|
rm_if_mismatch('STRING')
|
||||||
break
|
break
|
||||||
|
// case 'vector2':
|
||||||
|
// case 'vector3':
|
||||||
|
// case 'vector4':
|
||||||
|
// rm_if_mismatch('FLOAT')
|
||||||
|
// break
|
||||||
case 'vector2':
|
case 'vector2':
|
||||||
rm_if_mismatch('VECTOR2')
|
rm_if_mismatch('VECTOR2')
|
||||||
break
|
break
|
||||||
@@ -190,7 +406,7 @@ export class Constant extends LiteGraph.LGraphNode {
|
|||||||
default:
|
default:
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
this.updateOutput()
|
// this.updateOutput()
|
||||||
}
|
}
|
||||||
|
|
||||||
/**
|
/**
|
||||||
@@ -198,6 +414,7 @@ export class Constant extends LiteGraph.LGraphNode {
|
|||||||
* since Constant is a virtual node.
|
* since Constant is a virtual node.
|
||||||
*/
|
*/
|
||||||
updateTargetWidgets(u_links) {
|
updateTargetWidgets(u_links) {
|
||||||
|
infoLogger('Updating target widgets')
|
||||||
if (!app.graph.links) return
|
if (!app.graph.links) return
|
||||||
const links = u_links || this.outputs[0].links
|
const links = u_links || this.outputs[0].links
|
||||||
if (!links) return
|
if (!links) return
|
||||||
@@ -210,12 +427,16 @@ export class Constant extends LiteGraph.LGraphNode {
|
|||||||
const tgt_widget = tgt_node.widgets.filter(
|
const tgt_widget = tgt_node.widgets.filter(
|
||||||
(w) => w.name === tgt_input.name,
|
(w) => w.name === tgt_input.name,
|
||||||
)
|
)
|
||||||
if (!tgt_widget) return
|
// infoLogger('Constant Target Node', tgt_node)
|
||||||
|
// infoLogger('Constant Target Input', tgt_input)
|
||||||
|
if (!tgt_widget || tgt_widget.length === 0) return
|
||||||
|
|
||||||
tgt_widget[0].value = this.properties.value
|
tgt_widget[0].value = this.properties.value
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
updateOutput() {
|
updateOutput() {
|
||||||
|
infoLogger('Updating output value')
|
||||||
const value = this.properties.value
|
const value = this.properties.value
|
||||||
|
|
||||||
switch (this.properties.type) {
|
switch (this.properties.type) {
|
||||||
@@ -223,40 +444,53 @@ export class Constant extends LiteGraph.LGraphNode {
|
|||||||
this.setOutputData(0, value)
|
this.setOutputData(0, value)
|
||||||
break
|
break
|
||||||
case 'number':
|
case 'number':
|
||||||
this.setOutputData(0, Number.parseFloat(value))
|
if (this.properties.force_int) {
|
||||||
|
this.setOutputData(0, Number.parseInt(value))
|
||||||
|
} else {
|
||||||
|
this.setOutputData(0, Number.parseFloat(value))
|
||||||
|
}
|
||||||
break
|
break
|
||||||
case 'string':
|
case 'string':
|
||||||
this.setOutputData(0, value.toString())
|
this.setOutputData(0, value.toString())
|
||||||
break
|
break
|
||||||
case 'vector2':
|
case 'vector2':
|
||||||
if (value.length >= 2) {
|
|
||||||
this.setOutputData(0, value.slice(0, 2))
|
|
||||||
}
|
|
||||||
break
|
|
||||||
case 'vector3':
|
case 'vector3':
|
||||||
if (value.length >= 3) {
|
|
||||||
this.setOutputData(0, value.slice(0, 3))
|
|
||||||
}
|
|
||||||
break
|
|
||||||
case 'vector4':
|
case 'vector4':
|
||||||
if (value.length >= 4) {
|
this.setOutputData(0, value)
|
||||||
this.setOutputData(0, value.slice(0, 4))
|
|
||||||
}
|
|
||||||
break
|
break
|
||||||
|
|
||||||
|
// case 'vector2':
|
||||||
|
// this.setOutputData(0, value.slice(0, 2))
|
||||||
|
// break
|
||||||
|
// case 'vector3':
|
||||||
|
// this.setOutputData(0, value.slice(0, 3))
|
||||||
|
// break
|
||||||
|
// case 'vector4':
|
||||||
|
// this.setOutputData(0, value.slice(0, 4))
|
||||||
|
// break
|
||||||
default:
|
default:
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
|
||||||
|
infoLogger('New Value', this.value)
|
||||||
|
|
||||||
this.updateTargetWidgets()
|
this.updateTargetWidgets()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
app.registerExtension({
|
||||||
|
name: 'mtb.constant',
|
||||||
|
|
||||||
// app.registerExtension({
|
async beforeRegisterNodeDef(nodeType, nodeData, _app) {
|
||||||
// name: 'mtb.constant',
|
if (nodeData.name === 'Constant (mtb)') {
|
||||||
// registerCustomNodes() {
|
new ConstantJs(nodeType)
|
||||||
// LiteGraph.registerNodeType('Constant (mtb)', Constant)
|
}
|
||||||
//
|
},
|
||||||
// Constant.category = 'mtb/utils'
|
// NOTE: old js only registration
|
||||||
// Constant.title = 'Constant (mtb)'
|
//
|
||||||
// },
|
// registerCustomNodes() {
|
||||||
// })
|
// LiteGraph.registerNodeType('Constant (mtb)', Constant)
|
||||||
|
//
|
||||||
|
// Constant.category = 'mtb/utils'
|
||||||
|
// Constant.title = 'Constant (mtb)'
|
||||||
|
// },
|
||||||
|
})
|
||||||
|
|||||||
+200
-166
@@ -1,187 +1,221 @@
|
|||||||
|
// Reference the shared typedefs file
|
||||||
|
/// <reference path="../types/typedefs.js" />
|
||||||
import { app } from '../../scripts/app.js'
|
import { app } from '../../scripts/app.js'
|
||||||
|
import { infoLogger } from './comfy_shared.js'
|
||||||
|
|
||||||
function B0(t) { return (1 - t) ** 3 / 6; }
|
function B0(t) {
|
||||||
function B1(t) { return (3 * t ** 3 - 6 * t ** 2 + 4) / 6; }
|
return (1 - t) ** 3 / 6
|
||||||
function B2(t) { return (-3 * t ** 3 + 3 * t ** 2 + 3 * t + 1) / 6; }
|
}
|
||||||
function B3(t) { return t ** 3 / 6; }
|
function B1(t) {
|
||||||
|
return (3 * t ** 3 - 6 * t ** 2 + 4) / 6
|
||||||
|
}
|
||||||
|
function B2(t) {
|
||||||
|
return (-3 * t ** 3 + 3 * t ** 2 + 3 * t + 1) / 6
|
||||||
|
}
|
||||||
|
function B3(t) {
|
||||||
|
return t ** 3 / 6
|
||||||
|
}
|
||||||
class CurveWidget {
|
class CurveWidget {
|
||||||
constructor(inputName, defaultValue) {
|
constructor(...args) {
|
||||||
this.name = inputName || "Curve";
|
const [inputName, opts] = args
|
||||||
this._value = defaultValue || [{ x: 0, y: 0 }, { x: 1, y: 1 }];
|
|
||||||
this.type = "FLOAT_CURVE";
|
|
||||||
this.selectedPointIndex = null;
|
|
||||||
this.resize
|
|
||||||
}
|
|
||||||
|
|
||||||
drawBSpline(ctx, width, height, posY) {
|
this.name = inputName || 'Curve'
|
||||||
const n = this._value.length - 1;
|
|
||||||
const numSegments = n - 2;
|
|
||||||
const numPoints = this._value.length;
|
|
||||||
if (numPoints < 4) {
|
|
||||||
this.drawLinear(ctx, width, height, posY);
|
|
||||||
} else {
|
|
||||||
for (let j = 0; j <= numSegments; j++) {
|
|
||||||
for (let t = 0; t <= 1; t += 0.01) {
|
|
||||||
let pt = this.getBSplinePoint(j, t);
|
|
||||||
let x = pt.x * width;
|
|
||||||
let y = posY + height - pt.y * height;
|
|
||||||
|
|
||||||
if (t === 0) ctx.moveTo(x, y);
|
this.type = 'FLOAT_CURVE'
|
||||||
else ctx.lineTo(x, y);
|
this.selectedPointIndex = null
|
||||||
}
|
this.options = opts
|
||||||
}
|
this.value = this.value || { 0: { x: 0, y: 0 }, 1: { x: 1, y: 1 } }
|
||||||
ctx.stroke();
|
}
|
||||||
|
|
||||||
|
drawBSpline(ctx, width, height, posY) {
|
||||||
|
const n = this.value.length - 1
|
||||||
|
const numSegments = n - 2
|
||||||
|
const numPoints = this.value.length
|
||||||
|
if (numPoints < 4) {
|
||||||
|
this.drawLinear(ctx, width, height, posY)
|
||||||
|
} else {
|
||||||
|
for (let j = 0; j <= numSegments; j++) {
|
||||||
|
for (let t = 0; t <= 1; t += 0.01) {
|
||||||
|
let pt = this.getBSplinePoint(j, t)
|
||||||
|
let x = pt.x * width
|
||||||
|
let y = posY + height - pt.y * height
|
||||||
|
|
||||||
|
if (t === 0) ctx.moveTo(x, y)
|
||||||
|
else ctx.lineTo(x, y)
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
ctx.stroke()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
drawLinear(ctx, width, height, posY) {
|
||||||
|
for (let i = 0; i < Object.keys(this.value).length - 1; i++) {
|
||||||
|
let p1 = this.value[i]
|
||||||
|
let p2 = this.value[i + 1]
|
||||||
|
ctx.moveTo(p1.x * width, posY + height - p1.y * height)
|
||||||
|
ctx.lineTo(p2.x * width, posY + height - p2.y * height)
|
||||||
|
}
|
||||||
|
ctx.stroke()
|
||||||
|
}
|
||||||
|
getBSplinePoint(i, t) {
|
||||||
|
// Control points for this segment
|
||||||
|
const p0 = this.value[i]
|
||||||
|
const p1 = this.value[i + 1]
|
||||||
|
const p2 = this.value[i + 2]
|
||||||
|
const p3 = this.value[i + 3]
|
||||||
|
|
||||||
|
const x = B0(t) * p0.x + B1(t) * p1.x + B2(t) * p2.x + B3(t) * p3.x
|
||||||
|
const y = B0(t) * p0.y + B1(t) * p1.y + B2(t) * p2.y + B3(t) * p3.y
|
||||||
|
|
||||||
|
return { x, y }
|
||||||
|
}
|
||||||
|
/**
|
||||||
|
* @param {OnDrawWidgetParams} args
|
||||||
|
*/
|
||||||
|
draw(...args) {
|
||||||
|
const hide = this.type !== 'FLOAT_CURVE'
|
||||||
|
if (hide) {
|
||||||
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
drawLinear(ctx, width, height, posY) {
|
const [ctx, node, width, posY, height] = args
|
||||||
for (let i = 0; i < this._value.length - 1; i++) {
|
const [cw, ch] = this.computeSize(width)
|
||||||
let p1 = this._value[i];
|
|
||||||
let p2 = this._value[i + 1];
|
ctx.beginPath()
|
||||||
ctx.moveTo(p1.x * width, posY + height - p1.y * height);
|
ctx.fillStyle = '#000'
|
||||||
ctx.lineTo(p2.x * width, posY + height - p2.y * height);
|
ctx.strokeStyle = '#fff'
|
||||||
}
|
ctx.lineWidth = 2
|
||||||
ctx.stroke();
|
|
||||||
|
// normalized coordinates -> canvas coordinates
|
||||||
|
for (let i = 0; i < Object.keys(this.value || {}).length - 1; i++) {
|
||||||
|
let p1 = this.value[i]
|
||||||
|
let p2 = this.value[i + 1]
|
||||||
|
ctx.moveTo(p1.x * cw, posY + ch - p1.y * ch)
|
||||||
|
ctx.lineTo(p2.x * cw, posY + ch - p2.y * ch)
|
||||||
|
}
|
||||||
|
ctx.stroke()
|
||||||
|
|
||||||
|
// points
|
||||||
|
Object.values(this.value || {}).forEach((point) => {
|
||||||
|
ctx.beginPath()
|
||||||
|
ctx.arc(point.x * cw, posY + ch - point.y * ch, 5, 0, 2 * Math.PI)
|
||||||
|
ctx.fill()
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
mouse(event, pos, node) {
|
||||||
|
let x = pos[0] - node.pos[0]
|
||||||
|
let y = pos[1] - node.pos[1]
|
||||||
|
const width = node.size[0]
|
||||||
|
const height = 300 // TODO: compute
|
||||||
|
const posY = node.pos[1]
|
||||||
|
const localPos = { x: pos[0], y: pos[1] - LiteGraph.NODE_WIDGET_HEIGHT }
|
||||||
|
|
||||||
|
if (event.type === LiteGraph.pointerevents_method + 'down') {
|
||||||
|
console.debug('Checking if a point was clicked')
|
||||||
|
const clickedPointIndex = this.detectPoint(localPos, width, height)
|
||||||
|
if (clickedPointIndex !== null) {
|
||||||
|
this.selectedPointIndex = clickedPointIndex
|
||||||
|
} else {
|
||||||
|
this.addPoint(localPos, width, height)
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
} else if (
|
||||||
|
event.type === LiteGraph.pointerevents_method + 'move' &&
|
||||||
|
this.selectedPointIndex !== null
|
||||||
|
) {
|
||||||
|
this.movePoint(this.selectedPointIndex, localPos, width, height)
|
||||||
|
return true
|
||||||
|
} else if (
|
||||||
|
event.type === LiteGraph.pointerevents_method + 'up' &&
|
||||||
|
this.selectedPointIndex !== null
|
||||||
|
) {
|
||||||
|
this.selectedPointIndex = null
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
callback(...args) {
|
||||||
|
//value, that, node, pos, event) {
|
||||||
|
|
||||||
|
}
|
||||||
|
|
||||||
|
detectPoint(localPos, width, height) {
|
||||||
|
const threshold = 20 // TODO: extract
|
||||||
|
const keys = Object.keys(this.value)
|
||||||
|
for (let i = 0; i < keys.length; i++) {
|
||||||
|
const key = keys[i]
|
||||||
|
const p = this.value[key]
|
||||||
|
const px = p.x * width
|
||||||
|
const py = height - p.y * height
|
||||||
|
if (
|
||||||
|
Math.abs(localPos.x - px) < threshold &&
|
||||||
|
Math.abs(localPos.y - py) < threshold
|
||||||
|
) {
|
||||||
|
return key
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return null
|
||||||
|
}
|
||||||
|
addPoint(localPos, width, height) {
|
||||||
|
// add a new point based on click position
|
||||||
|
const normalizedPoint = {
|
||||||
|
x: localPos.x / width,
|
||||||
|
y: 1 - localPos.y / height,
|
||||||
}
|
}
|
||||||
|
|
||||||
getBSplinePoint(i, t) {
|
const keys = Object.keys(this.value)
|
||||||
// Control points for this segment
|
let insertIndex = keys.length
|
||||||
const p0 = this._value[i];
|
for (let i = 0; i < keys.length; i++) {
|
||||||
const p1 = this._value[i + 1];
|
if (normalizedPoint.x < this.value[keys[i]].x) {
|
||||||
const p2 = this._value[i + 2];
|
insertIndex = i
|
||||||
const p3 = this._value[i + 3];
|
break
|
||||||
|
}
|
||||||
const x = B0(t) * p0.x + B1(t) * p1.x + B2(t) * p2.x + B3(t) * p3.x;
|
}
|
||||||
const y = B0(t) * p0.y + B1(t) * p1.y + B2(t) * p2.y + B3(t) * p3.y;
|
// shift
|
||||||
|
for (let i = keys.length; i > insertIndex; i--) {
|
||||||
return { x, y };
|
this.value[i] = this.value[i - 1]
|
||||||
}
|
}
|
||||||
|
|
||||||
draw(ctx, node, width, posY, height) {
|
this.value[insertIndex] = normalizedPoint
|
||||||
const [cw, ch] = this.computeSize(width)
|
}
|
||||||
|
|
||||||
ctx.beginPath();
|
movePoint(index, localPos, width, height) {
|
||||||
ctx.fillStyle = "#000";
|
const point = this.value[index]
|
||||||
//ctx.fillRect(0, posY, cw, ch);
|
point.x = Math.max(0, Math.min(1, localPos.x / width))
|
||||||
ctx.strokeStyle = "#fff";
|
point.y = Math.max(0, Math.min(1, 1 - localPos.y / height))
|
||||||
ctx.lineWidth = 2;
|
|
||||||
|
|
||||||
// normalized coordinates -> canvas coordinates
|
this.value[index] = point
|
||||||
for (let i = 0; i < this._value.length - 1; i++) {
|
}
|
||||||
let p1 = this._value[i];
|
computeSize(width) {
|
||||||
let p2 = this._value[i + 1];
|
return [width, 300]
|
||||||
ctx.moveTo(p1.x * cw, posY + ch - p1.y * ch);
|
}
|
||||||
ctx.lineTo(p2.x * cw, posY + ch - p2.y * ch);
|
|
||||||
}
|
|
||||||
ctx.stroke();
|
|
||||||
// this.drawBSpline(ctx, width, height, posY);
|
|
||||||
|
|
||||||
// points
|
configure(data) {
|
||||||
this._value.forEach(point => {
|
}
|
||||||
ctx.beginPath();
|
|
||||||
ctx.arc(point.x * cw, posY + ch - point.y * ch, 5, 0, 2 * Math.PI);
|
|
||||||
ctx.fill();
|
|
||||||
});
|
|
||||||
}
|
|
||||||
|
|
||||||
mouse(event, pos, node) {
|
|
||||||
// console.debug(event.type, pos, node)
|
|
||||||
let x = pos[0] - node.pos[0]
|
|
||||||
let y = pos[1] - node.pos[1]
|
|
||||||
let width = node.size[0]
|
|
||||||
const height = 300; // TODO: compute
|
|
||||||
const posY = node.pos[1];
|
|
||||||
|
|
||||||
const localPos = { x: pos[0], y: pos[1] - LiteGraph.NODE_WIDGET_HEIGHT };
|
|
||||||
|
|
||||||
if (event.type === LiteGraph.pointerevents_method + "down") {
|
|
||||||
console.debug("Checking if a point was clicked");
|
|
||||||
const clickedPointIndex = this.detectPoint(localPos, width, height);
|
|
||||||
if (clickedPointIndex !== null) {
|
|
||||||
this.selectedPointIndex = clickedPointIndex;
|
|
||||||
} else {
|
|
||||||
this.addPoint(localPos, width, height);
|
|
||||||
}
|
|
||||||
return true;
|
|
||||||
} else if (event.type === LiteGraph.pointerevents_method + "move" && this.selectedPointIndex !== null) {
|
|
||||||
this.movePoint(this.selectedPointIndex, localPos, width, height);
|
|
||||||
return true;
|
|
||||||
} else if (event.type === LiteGraph.pointerevents_method + "up" && this.selectedPointIndex !== null) {
|
|
||||||
this.selectedPointIndex = null;
|
|
||||||
return true;
|
|
||||||
}
|
|
||||||
return false;
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
detectPoint(localPos, width, height) {
|
|
||||||
const threshold = 20; // TODO: extract
|
|
||||||
for (let i = 0; i < this._value.length; i++) {
|
|
||||||
const p = this._value[i];
|
|
||||||
const px = p.x * width;
|
|
||||||
const py = height - p.y * height;
|
|
||||||
if (Math.abs(localPos.x - px) < threshold && Math.abs(localPos.y - py) < threshold) {
|
|
||||||
return i;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return null;
|
|
||||||
}
|
|
||||||
|
|
||||||
addPoint(localPos, width, height) {
|
|
||||||
// add a new point based on click position
|
|
||||||
const normalizedPoint = { x: localPos.x / width, y: 1 - localPos.y / height };
|
|
||||||
this._value.push(normalizedPoint);
|
|
||||||
this._value.sort((a, b) => a.x - b.x);
|
|
||||||
this.value = JSON.stringify(this._value);
|
|
||||||
}
|
|
||||||
|
|
||||||
movePoint(index, localPos, width, height) {
|
|
||||||
const point = this._value[index];
|
|
||||||
point.x = Math.max(0, Math.min(1, localPos.x / width));
|
|
||||||
point.y = Math.max(0, Math.min(1, 1 - localPos.y / height));
|
|
||||||
|
|
||||||
this._value[index] = point;
|
|
||||||
this.value = JSON.stringify(this._value);
|
|
||||||
}
|
|
||||||
|
|
||||||
computeSize(width) {
|
|
||||||
return [width, 300];
|
|
||||||
}
|
|
||||||
|
|
||||||
configure(data) {
|
|
||||||
console.log(data)
|
|
||||||
}
|
|
||||||
|
|
||||||
value() {
|
|
||||||
console.debug('Returning value', this._value)
|
|
||||||
return this._value
|
|
||||||
}
|
|
||||||
setValue(value) {
|
|
||||||
console.debug('Setting value', value)
|
|
||||||
this._value = value
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
app.registerExtension({
|
app.registerExtension({
|
||||||
name: 'mtb.curves',
|
name: 'mtb.curves',
|
||||||
getCustomWidgets: function () {
|
getCustomWidgets: () => {
|
||||||
|
return {
|
||||||
|
/**
|
||||||
|
* @param {LGraphNode} node
|
||||||
|
* @param {str} inputName
|
||||||
|
* @param {[str,*]} inputData
|
||||||
|
* @param {*} app
|
||||||
|
*
|
||||||
|
*/
|
||||||
|
FLOAT_CURVE: (node, inputName, inputData, app) => {
|
||||||
|
// const c = node.widgets.find((w) => w.type === "FLOAT_CURVE")
|
||||||
|
const wid = node.addCustomWidget(new CurveWidget(inputName, inputData))
|
||||||
|
|
||||||
return {
|
return {
|
||||||
FLOAT_CURVE: (node, inputName, inputData, app) => {
|
widget: wid,
|
||||||
console.debug('Registering float curve widget');
|
minWidth: 150,
|
||||||
|
minHeight: 30,
|
||||||
return {
|
|
||||||
widget: node.addCustomWidget(
|
|
||||||
new CurveWidget(inputName, inputData[1]?.default)
|
|
||||||
),
|
|
||||||
minWidth: 150,
|
|
||||||
minHeight: 30,
|
|
||||||
}
|
|
||||||
},
|
|
||||||
|
|
||||||
|
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
|
}
|
||||||
|
},
|
||||||
})
|
})
|
||||||
|
|||||||
+34
-23
@@ -7,10 +7,12 @@
|
|||||||
*
|
*
|
||||||
*/
|
*/
|
||||||
|
|
||||||
|
// Reference the shared typedefs file
|
||||||
|
/// <reference path="../types/typedefs.js" />
|
||||||
|
|
||||||
import { app } from '../../scripts/app.js'
|
import { app } from '../../scripts/app.js'
|
||||||
|
|
||||||
import * as shared from './comfy_shared.js'
|
import * as shared from './comfy_shared.js'
|
||||||
import { log } from './comfy_shared.js'
|
|
||||||
import { MtbWidgets } from './mtb_widgets.js'
|
import { MtbWidgets } from './mtb_widgets.js'
|
||||||
|
|
||||||
// TODO: respect inputs order...
|
// TODO: respect inputs order...
|
||||||
@@ -25,10 +27,17 @@ function escapeHtml(unsafe) {
|
|||||||
}
|
}
|
||||||
app.registerExtension({
|
app.registerExtension({
|
||||||
name: 'mtb.Debug',
|
name: 'mtb.Debug',
|
||||||
|
|
||||||
|
/**
|
||||||
|
* @param {NodeType} nodeType
|
||||||
|
* @param {NodeData} nodeData
|
||||||
|
* @param {*} app
|
||||||
|
*/
|
||||||
async beforeRegisterNodeDef(nodeType, nodeData, app) {
|
async beforeRegisterNodeDef(nodeType, nodeData, app) {
|
||||||
if (nodeData.name === 'Debug (mtb)') {
|
if (nodeData.name === 'Debug (mtb)') {
|
||||||
const onNodeCreated = nodeType.prototype.onNodeCreated
|
const onNodeCreated = nodeType.prototype.onNodeCreated
|
||||||
nodeType.prototype.onNodeCreated = function () {
|
nodeType.prototype.onNodeCreated = function () {
|
||||||
|
this.options = {}
|
||||||
const r = onNodeCreated
|
const r = onNodeCreated
|
||||||
? onNodeCreated.apply(this, arguments)
|
? onNodeCreated.apply(this, arguments)
|
||||||
: undefined
|
: undefined
|
||||||
@@ -37,24 +46,29 @@ app.registerExtension({
|
|||||||
}
|
}
|
||||||
|
|
||||||
const onConnectionsChange = nodeType.prototype.onConnectionsChange
|
const onConnectionsChange = nodeType.prototype.onConnectionsChange
|
||||||
nodeType.prototype.onConnectionsChange = function (
|
/**
|
||||||
type,
|
* @param {OnConnectionsChangeParams} args
|
||||||
index,
|
*/
|
||||||
connected,
|
nodeType.prototype.onConnectionsChange = function (...args) {
|
||||||
link_info,
|
const [_type, index, connected, link_info, ioSlot] = args
|
||||||
) {
|
|
||||||
const r = onConnectionsChange
|
const r = onConnectionsChange
|
||||||
? onConnectionsChange.apply(this, arguments)
|
? onConnectionsChange.apply(this, args)
|
||||||
: undefined
|
: undefined
|
||||||
// TODO: remove all widgets on disconnect once computed
|
// TODO: remove all widgets on disconnect once computed
|
||||||
shared.dynamic_connection(this, index, connected, 'anything_', '*')
|
shared.dynamic_connection(this, index, connected, 'anything_', '*', {
|
||||||
|
link: link_info,
|
||||||
|
ioSlot: ioSlot,
|
||||||
|
})
|
||||||
|
|
||||||
//- infer type
|
//- infer type
|
||||||
if (link_info) {
|
if (link_info) {
|
||||||
const fromNode = this.graph._nodes.find(
|
// const fromNode = this.graph._nodes.find(
|
||||||
(otherNode) => otherNode.id === link_info.origin_id,
|
// (otherNode) => otherNode.id === link_info.origin_id,
|
||||||
)
|
// )
|
||||||
const type = fromNode.outputs[link_info.origin_slot].type
|
// const fromNode = app.graph.getNodeById(link_info.origin_id)
|
||||||
|
const { from } = shared.nodesFromLink(this, link_info)
|
||||||
|
if (!from || this.inputs.length === 0) return
|
||||||
|
const type = from.outputs[link_info.origin_slot].type
|
||||||
this.inputs[index].type = type
|
this.inputs[index].type = type
|
||||||
// this.inputs[index].label = type.toLowerCase()
|
// this.inputs[index].label = type.toLowerCase()
|
||||||
}
|
}
|
||||||
@@ -67,14 +81,12 @@ app.registerExtension({
|
|||||||
}
|
}
|
||||||
|
|
||||||
const onExecuted = nodeType.prototype.onExecuted
|
const onExecuted = nodeType.prototype.onExecuted
|
||||||
nodeType.prototype.onExecuted = function (message) {
|
nodeType.prototype.onExecuted = function (data) {
|
||||||
onExecuted?.apply(this, arguments)
|
onExecuted?.apply(this, arguments)
|
||||||
|
|
||||||
const prefix = 'anything_'
|
const prefix = 'anything_'
|
||||||
|
|
||||||
if (this.widgets) {
|
if (this.widgets) {
|
||||||
// const pos = this.widgets.findIndex((w) => w.name === "anything_1");
|
|
||||||
// if (pos !== -1) {
|
|
||||||
for (let i = 0; i < this.widgets.length; i++) {
|
for (let i = 0; i < this.widgets.length; i++) {
|
||||||
if (this.widgets[i].name !== 'output_to_console') {
|
if (this.widgets[i].name !== 'output_to_console') {
|
||||||
this.widgets[i].onRemoved?.()
|
this.widgets[i].onRemoved?.()
|
||||||
@@ -83,8 +95,9 @@ app.registerExtension({
|
|||||||
this.widgets.length = 1
|
this.widgets.length = 1
|
||||||
}
|
}
|
||||||
let widgetI = 1
|
let widgetI = 1
|
||||||
if (message.text) {
|
// console.log(message)
|
||||||
for (const txt of message.text) {
|
if (data.text) {
|
||||||
|
for (const txt of data.text) {
|
||||||
const w = this.addCustomWidget(
|
const w = this.addCustomWidget(
|
||||||
MtbWidgets.DEBUG_STRING(`${prefix}_${widgetI}`, escapeHtml(txt)),
|
MtbWidgets.DEBUG_STRING(`${prefix}_${widgetI}`, escapeHtml(txt)),
|
||||||
)
|
)
|
||||||
@@ -92,19 +105,17 @@ app.registerExtension({
|
|||||||
widgetI++
|
widgetI++
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if (message.b64_images) {
|
if (data.b64_images) {
|
||||||
for (const img of message.b64_images) {
|
for (const img of data.b64_images) {
|
||||||
const w = this.addCustomWidget(
|
const w = this.addCustomWidget(
|
||||||
MtbWidgets.DEBUG_IMG(`${prefix}_${widgetI}`, img),
|
MtbWidgets.DEBUG_IMG(`${prefix}_${widgetI}`, img),
|
||||||
)
|
)
|
||||||
w.parent = this
|
w.parent = this
|
||||||
widgetI++
|
widgetI++
|
||||||
}
|
}
|
||||||
// this.onResize?.(this.size);
|
|
||||||
// this.resize?.(this.size)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
this.setSize(this.computeSize())
|
// this.setSize(this.computeSize())
|
||||||
|
|
||||||
this.onRemoved = function () {
|
this.onRemoved = function () {
|
||||||
// When removing this node we need to remove the input from the DOM
|
// When removing this node we need to remove the input from the DOM
|
||||||
|
|||||||
+294
-33
@@ -7,6 +7,8 @@
|
|||||||
*
|
*
|
||||||
*/
|
*/
|
||||||
|
|
||||||
|
/// <reference path="../types/typedefs.js" />
|
||||||
|
|
||||||
// TODO: Use the builtin addDOMWidget everywhere appropriate
|
// TODO: Use the builtin addDOMWidget everywhere appropriate
|
||||||
|
|
||||||
import { app } from '../../scripts/app.js'
|
import { app } from '../../scripts/app.js'
|
||||||
@@ -14,8 +16,8 @@ import { api } from '../../scripts/api.js'
|
|||||||
|
|
||||||
import parseCss from './extern/parse-css.js'
|
import parseCss from './extern/parse-css.js'
|
||||||
import * as shared from './comfy_shared.js'
|
import * as shared from './comfy_shared.js'
|
||||||
import { log } from './comfy_shared.js'
|
import { infoLogger } from './comfy_shared.js'
|
||||||
import { Constant } from './constant.js'
|
import { NumberInputWidget } from './numberInput.js'
|
||||||
|
|
||||||
// NOTE: new widget types registered by MTB Widgets
|
// NOTE: new widget types registered by MTB Widgets
|
||||||
const newTypes = [/*'BOOL'*/ , 'COLOR', 'BBOX']
|
const newTypes = [/*'BOOL'*/ , 'COLOR', 'BBOX']
|
||||||
@@ -55,7 +57,250 @@ const calculateTextDimensions = (ctx, value, width, fontSize = 16) => {
|
|||||||
return { textHeight, maxLineWidth }
|
return { textHeight, maxLineWidth }
|
||||||
}
|
}
|
||||||
|
|
||||||
|
export function addMultilineWidget(node, name, opts, callback) {
|
||||||
|
const inputEl = document.createElement('textarea')
|
||||||
|
inputEl.className = 'comfy-multiline-input'
|
||||||
|
inputEl.value = opts.defaultVal
|
||||||
|
inputEl.placeholder = opts.placeholder || name
|
||||||
|
|
||||||
|
const widget = node.addDOMWidget(name, 'textmultiline', inputEl, {
|
||||||
|
getValue() {
|
||||||
|
return inputEl.value
|
||||||
|
},
|
||||||
|
setValue(v) {
|
||||||
|
inputEl.value = v
|
||||||
|
},
|
||||||
|
})
|
||||||
|
widget.inputEl = inputEl
|
||||||
|
|
||||||
|
inputEl.addEventListener('input', () => {
|
||||||
|
callback?.(widget.value)
|
||||||
|
widget.callback?.(widget.value)
|
||||||
|
})
|
||||||
|
widget.onRemove = () => {
|
||||||
|
inputEl.remove()
|
||||||
|
}
|
||||||
|
|
||||||
|
return { minWidth: 400, minHeight: 200, widget }
|
||||||
|
}
|
||||||
|
|
||||||
|
export const VECTOR_AXIS = {
|
||||||
|
0: 'x',
|
||||||
|
1: 'y',
|
||||||
|
2: 'z',
|
||||||
|
3: 'w',
|
||||||
|
}
|
||||||
|
|
||||||
|
export function addVectorWidgetW(
|
||||||
|
node,
|
||||||
|
name,
|
||||||
|
value,
|
||||||
|
vector_size,
|
||||||
|
callback,
|
||||||
|
app,
|
||||||
|
) {
|
||||||
|
// const inputEl = document.createElement('div')
|
||||||
|
// const vecEl = document.createElement('div')
|
||||||
|
//
|
||||||
|
// inputEl.style.background = 'red'
|
||||||
|
//
|
||||||
|
// inputEl.className = 'comfy-vector-container'
|
||||||
|
// vecEl.className = 'comfy-vector-input'
|
||||||
|
//
|
||||||
|
// vecEl.style.display = 'flex'
|
||||||
|
// inputEl.appendChild(vecEl)
|
||||||
|
const inputs = []
|
||||||
|
|
||||||
|
for (let i = 0; i < vector_size; i++) {
|
||||||
|
// const input = document.createElement('input')
|
||||||
|
// input.type = 'number'
|
||||||
|
// input.value = value[VECTOR_AXIS[i]]
|
||||||
|
const input = node.addWidget(
|
||||||
|
'number',
|
||||||
|
`${name}_${VECTOR_AXIS[i]}`,
|
||||||
|
value[VECTOR_AXIS[i]],
|
||||||
|
(val) => {},
|
||||||
|
)
|
||||||
|
|
||||||
|
inputs.push(input)
|
||||||
|
// vecEl.appendChild(input)
|
||||||
|
}
|
||||||
|
//
|
||||||
|
// const widget = node.addDOMWidget(name, 'vector', inputEl, {
|
||||||
|
// getValue() {
|
||||||
|
// return JSON.stringify(widget._value)
|
||||||
|
// },
|
||||||
|
// setValue(v) {
|
||||||
|
// widget._value = v
|
||||||
|
// },
|
||||||
|
// afterResize(node, widget) {
|
||||||
|
// console.log('After resize', { that: this, node, widget })
|
||||||
|
// },
|
||||||
|
// })
|
||||||
|
//
|
||||||
|
// console.log('prev callback', widget.callback)
|
||||||
|
// widget.callback = callback
|
||||||
|
// widget._value = value
|
||||||
|
//
|
||||||
|
// for (let i = 0; i < vector_size; i++) {
|
||||||
|
// const input = inputs[i]
|
||||||
|
// input.addEventListener('change', (event) => {
|
||||||
|
// widget._value[VECTOR_AXIS[i]] = Number.parseFloat(event.target.value)
|
||||||
|
// widget.callback?.(widget._value)
|
||||||
|
// node.graph._version++
|
||||||
|
// node.setDirtyCanvas(true, true)
|
||||||
|
// })
|
||||||
|
// }
|
||||||
|
// // document.body.append(inputEl)
|
||||||
|
//
|
||||||
|
// widget.inputEl = inputEl
|
||||||
|
// widget.vecEl = vecEl
|
||||||
|
//
|
||||||
|
// inputEl.addEventListener('input', () => {
|
||||||
|
// widget.callback?.(widget.value)
|
||||||
|
// })
|
||||||
|
//
|
||||||
|
return { minWidth: 400, minHeight: 200, widget }
|
||||||
|
}
|
||||||
|
export function addVectorWidget(node, name, value, vector_size, callback, app) {
|
||||||
|
const inputEl = document.createElement('div')
|
||||||
|
const vecEl = document.createElement('div')
|
||||||
|
|
||||||
|
inputEl.className = 'comfy-vector-container'
|
||||||
|
vecEl.className = 'comfy-vector-input'
|
||||||
|
vecEl.id = 'vecEl'
|
||||||
|
|
||||||
|
vecEl.style.display = 'flex'
|
||||||
|
vecEl.style.flexDirection = 'column'
|
||||||
|
inputEl.appendChild(vecEl)
|
||||||
|
const inputs = []
|
||||||
|
|
||||||
|
//
|
||||||
|
// for (let i = 0; i < vector_size; i++) {
|
||||||
|
// const input = document.createElement('input')
|
||||||
|
// input.type = 'number'
|
||||||
|
// input.value = value[VECTOR_AXIS[i]]
|
||||||
|
// inputs.push(input)
|
||||||
|
// vecEl.appendChild(input)
|
||||||
|
// }
|
||||||
|
|
||||||
|
const widget = node.addDOMWidget(name, 'vector', inputEl, {
|
||||||
|
getValue() {
|
||||||
|
return JSON.stringify(widget._value)
|
||||||
|
},
|
||||||
|
setValue(v) {
|
||||||
|
widget._value = v
|
||||||
|
},
|
||||||
|
})
|
||||||
|
const vec = new NumberInputWidget('vecEl', vector_size, true)
|
||||||
|
vec.setValue(...Object.values(value))
|
||||||
|
vec.onChange = (value) => {
|
||||||
|
for (let i = 0; i < value.length; i++) {
|
||||||
|
const val = value[i]
|
||||||
|
widget._value[VECTOR_AXIS[i]] = Number.parseFloat(val)
|
||||||
|
}
|
||||||
|
|
||||||
|
widget.callback?.(widget._value)
|
||||||
|
// widget._value[VECTOR_AXIS[index]] = Number.parseFloat(value)
|
||||||
|
}
|
||||||
|
|
||||||
|
console.log('prev callback', widget.callback)
|
||||||
|
widget.callback = callback
|
||||||
|
widget._value = value
|
||||||
|
|
||||||
|
// for (let i = 0; i < vector_size; i++) {
|
||||||
|
// const input = inputs[i]
|
||||||
|
// input.addEventListener('change', (event) => {
|
||||||
|
// widget._value[VECTOR_AXIS[i]] = Number.parseFloat(event.target.value)
|
||||||
|
// widget.callback?.(widget._value)
|
||||||
|
// node.graph._version++
|
||||||
|
// node.setDirtyCanvas(true, true)
|
||||||
|
// })
|
||||||
|
// }
|
||||||
|
|
||||||
|
widget.inputEl = inputEl
|
||||||
|
widget.vecEl = vecEl
|
||||||
|
widget.vec = vec
|
||||||
|
|
||||||
|
return { minWidth: 400, minHeight: 200 * vector_size, widget }
|
||||||
|
}
|
||||||
export const MtbWidgets = {
|
export const MtbWidgets = {
|
||||||
|
//TODO: complete this properly
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Creates a vector widget.
|
||||||
|
* @param {string} key - The key for the widget.
|
||||||
|
* @param {number[]} [val] - The initial value for the widget.
|
||||||
|
* @param {number} size - The size of the vector.
|
||||||
|
* @returns {VectorWidget} The vector widget.
|
||||||
|
*/
|
||||||
|
VECTOR: (key, val, size) => {
|
||||||
|
shared.infoLogger('Adding VECTOR widget', { key, val, size })
|
||||||
|
/** @type {VectorWidget} */
|
||||||
|
const widget = {
|
||||||
|
name: key,
|
||||||
|
type: `vector${size}`,
|
||||||
|
y: 0,
|
||||||
|
options: { default: Array.from({ length: size }, () => 0.0) },
|
||||||
|
_value: val || Array.from({ length: size }, () => 0.0),
|
||||||
|
draw: function (ctx, node, width, widgetY, height) {
|
||||||
|
ctx.textAlign = 'left'
|
||||||
|
ctx.strokeStyle = outline_color
|
||||||
|
ctx.fillStyle = background_color
|
||||||
|
ctx.beginPath()
|
||||||
|
if (show_text)
|
||||||
|
ctx.roundRect(margin, y, widget_width - margin * 2, H, [H * 0.5])
|
||||||
|
else ctx.rect(margin, y, widget_width - margin * 2, H)
|
||||||
|
ctx.fill()
|
||||||
|
if (show_text) {
|
||||||
|
if (!w.disabled) ctx.stroke()
|
||||||
|
ctx.fillStyle = text_color
|
||||||
|
if (!w.disabled) {
|
||||||
|
ctx.beginPath()
|
||||||
|
ctx.moveTo(margin + 16, y + 5)
|
||||||
|
ctx.lineTo(margin + 6, y + H * 0.5)
|
||||||
|
ctx.lineTo(margin + 16, y + H - 5)
|
||||||
|
ctx.fill()
|
||||||
|
ctx.beginPath()
|
||||||
|
ctx.moveTo(widget_width - margin - 16, y + 5)
|
||||||
|
ctx.lineTo(widget_width - margin - 6, y + H * 0.5)
|
||||||
|
ctx.lineTo(widget_width - margin - 16, y + H - 5)
|
||||||
|
ctx.fill()
|
||||||
|
}
|
||||||
|
ctx.fillStyle = secondary_text_color
|
||||||
|
ctx.fillText(w.label || w.name, margin * 2 + 5, y + H * 0.7)
|
||||||
|
ctx.fillStyle = text_color
|
||||||
|
ctx.textAlign = 'right'
|
||||||
|
if (w.type === 'number') {
|
||||||
|
ctx.fillText(
|
||||||
|
Number(w.value).toFixed(
|
||||||
|
w.options.precision !== undefined ? w.options.precision : 3,
|
||||||
|
),
|
||||||
|
widget_width - margin * 2 - 20,
|
||||||
|
y + H * 0.7,
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
let v = w.value
|
||||||
|
if (w.options.values) {
|
||||||
|
let values = w.options.values
|
||||||
|
if (values.constructor === Function) values = values()
|
||||||
|
if (values && values.constructor !== Array) v = values[w.value]
|
||||||
|
}
|
||||||
|
ctx.fillText(v, widget_width - margin * 2 - 20, y + H * 0.7)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
},
|
||||||
|
get value() {
|
||||||
|
return this._value
|
||||||
|
},
|
||||||
|
set value(val) {
|
||||||
|
this._value = val
|
||||||
|
this.callback?.(this._value)
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
return widget
|
||||||
|
},
|
||||||
BBOX: (key, val) => {
|
BBOX: (key, val) => {
|
||||||
/** @type {import("./types/litegraph").IWidget} */
|
/** @type {import("./types/litegraph").IWidget} */
|
||||||
const widget = {
|
const widget = {
|
||||||
@@ -411,7 +656,7 @@ const mtb_widgets = {
|
|||||||
name: 'mtb.widgets',
|
name: 'mtb.widgets',
|
||||||
|
|
||||||
init: async () => {
|
init: async () => {
|
||||||
log('Registering mtb.widgets')
|
infoLogger('Registering mtb.widgets')
|
||||||
try {
|
try {
|
||||||
const res = await api.fetchApi('/mtb/debug')
|
const res = await api.fetchApi('/mtb/debug')
|
||||||
const msg = await res.json()
|
const msg = await res.json()
|
||||||
@@ -439,13 +684,14 @@ const mtb_widgets = {
|
|||||||
},
|
},
|
||||||
},
|
},
|
||||||
async onChange(value) {
|
async onChange(value) {
|
||||||
if (value) {
|
|
||||||
console.log('Enabled DEBUG mode')
|
|
||||||
}
|
|
||||||
if (!window.MTB) {
|
if (!window.MTB) {
|
||||||
window.MTB = {}
|
window.MTB = {}
|
||||||
}
|
}
|
||||||
window.MTB.DEBUG = value
|
window.MTB.DEBUG = value
|
||||||
|
if (value) {
|
||||||
|
infoLogger('Enabled DEBUG mode')
|
||||||
|
}
|
||||||
|
|
||||||
await api
|
await api
|
||||||
.fetchApi('/mtb/debug', {
|
.fetchApi('/mtb/debug', {
|
||||||
method: 'POST',
|
method: 'POST',
|
||||||
@@ -453,23 +699,17 @@ const mtb_widgets = {
|
|||||||
enabled: value,
|
enabled: value,
|
||||||
}),
|
}),
|
||||||
})
|
})
|
||||||
.then((response) => {})
|
.then((_response) => {})
|
||||||
.catch((error) => {
|
.catch((error) => {
|
||||||
console.error('Error:', error)
|
console.error('Error:', error)
|
||||||
})
|
})
|
||||||
},
|
},
|
||||||
})
|
})
|
||||||
},
|
},
|
||||||
registerCustomNodes() {
|
|
||||||
LiteGraph.registerNodeType('Constant (mtb)', Constant)
|
|
||||||
|
|
||||||
Constant.category = 'mtb/utils'
|
getCustomWidgets: () => {
|
||||||
Constant.title = 'Constant (mtb)'
|
|
||||||
},
|
|
||||||
|
|
||||||
getCustomWidgets: function () {
|
|
||||||
return {
|
return {
|
||||||
BOOL: (node, inputName, inputData, app) => {
|
BOOL: (node, inputName, inputData, _app) => {
|
||||||
console.debug('Registering bool')
|
console.debug('Registering bool')
|
||||||
|
|
||||||
return {
|
return {
|
||||||
@@ -481,7 +721,7 @@ const mtb_widgets = {
|
|||||||
}
|
}
|
||||||
},
|
},
|
||||||
|
|
||||||
COLOR: (node, inputName, inputData, app) => {
|
COLOR: (node, inputName, inputData, _app) => {
|
||||||
console.debug('Registering color')
|
console.debug('Registering color')
|
||||||
return {
|
return {
|
||||||
widget: node.addCustomWidget(
|
widget: node.addCustomWidget(
|
||||||
@@ -503,8 +743,8 @@ const mtb_widgets = {
|
|||||||
}
|
}
|
||||||
},
|
},
|
||||||
/**
|
/**
|
||||||
* @param {import("./types/comfy").NodeType} nodeType
|
* @param {NodeType} nodeType
|
||||||
* @param {import("./types/comfy").NodeDef} nodeData
|
* @param {NodeData} nodeData
|
||||||
* @param {import("./types/comfy").App} app
|
* @param {import("./types/comfy").App} app
|
||||||
*/
|
*/
|
||||||
async beforeRegisterNodeDef(nodeType, nodeData, app) {
|
async beforeRegisterNodeDef(nodeType, nodeData, app) {
|
||||||
@@ -595,9 +835,7 @@ const mtb_widgets = {
|
|||||||
case 'Get Batch From History (mtb)': {
|
case 'Get Batch From History (mtb)': {
|
||||||
const onNodeCreated = nodeType.prototype.onNodeCreated
|
const onNodeCreated = nodeType.prototype.onNodeCreated
|
||||||
nodeType.prototype.onNodeCreated = function () {
|
nodeType.prototype.onNodeCreated = function () {
|
||||||
const r = onNodeCreated
|
const r = onNodeCreated ? onNodeCreated.apply(this, []) : undefined
|
||||||
? onNodeCreated.apply(this, arguments)
|
|
||||||
: undefined
|
|
||||||
const internal_count = this.widgets.find(
|
const internal_count = this.widgets.find(
|
||||||
(w) => w.name === 'internal_count',
|
(w) => w.name === 'internal_count',
|
||||||
)
|
)
|
||||||
@@ -718,14 +956,11 @@ const mtb_widgets = {
|
|||||||
app.canvas.setDirty(true)
|
app.canvas.setDirty(true)
|
||||||
}
|
}
|
||||||
|
|
||||||
const reset_button = this.addWidget(
|
// reset button
|
||||||
'button',
|
this.addWidget('button', `Reset`, 'reset', onReset)
|
||||||
`Reset`,
|
|
||||||
'reset',
|
|
||||||
onReset,
|
|
||||||
)
|
|
||||||
|
|
||||||
const run_button = this.addWidget('button', `Queue`, 'queue', () => {
|
// run button
|
||||||
|
this.addWidget('button', `Queue`, 'queue', () => {
|
||||||
onReset() // this could maybe be a setting or checkbox
|
onReset() // this could maybe be a setting or checkbox
|
||||||
app.queuePrompt(0, total_frames.value * loop_count.value)
|
app.queuePrompt(0, total_frames.value * loop_count.value)
|
||||||
window.MTB?.notify?.(
|
window.MTB?.notify?.(
|
||||||
@@ -902,6 +1137,7 @@ const mtb_widgets = {
|
|||||||
break
|
break
|
||||||
}
|
}
|
||||||
case 'Batch Float Assemble (mtb)':
|
case 'Batch Float Assemble (mtb)':
|
||||||
|
case 'Batch Float Math (mtb)':
|
||||||
case 'Plot Batch Float (mtb)': {
|
case 'Plot Batch Float (mtb)': {
|
||||||
shared.setupDynamicConnections(nodeType, 'floats', 'FLOATS')
|
shared.setupDynamicConnections(nodeType, 'floats', 'FLOATS')
|
||||||
break
|
break
|
||||||
@@ -932,11 +1168,9 @@ const mtb_widgets = {
|
|||||||
const r = onConnectionsChange
|
const r = onConnectionsChange
|
||||||
? onConnectionsChange.apply(this, arguments)
|
? onConnectionsChange.apply(this, arguments)
|
||||||
: undefined
|
: undefined
|
||||||
shared.dynamic_connection(this, index, connected, 'var_', '*', [
|
shared.dynamic_connection(this, index, connected, 'var_', '*', {
|
||||||
'x',
|
nameArray: ['x', 'y', 'z'],
|
||||||
'y',
|
})
|
||||||
'z',
|
|
||||||
])
|
|
||||||
|
|
||||||
//- infer type
|
//- infer type
|
||||||
if (link_info) {
|
if (link_info) {
|
||||||
@@ -956,6 +1190,33 @@ const mtb_widgets = {
|
|||||||
|
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
|
||||||
|
case 'Batch Shape (mtb)':
|
||||||
|
case 'Text To Image (mtb)': {
|
||||||
|
shared.addMenuHandler(nodeType, function (_app, options) {
|
||||||
|
/** @type {ContextMenuItem} */
|
||||||
|
const item = {
|
||||||
|
content: 'swap colors',
|
||||||
|
title: 'Swap BG/FG Color ⚡',
|
||||||
|
callback: (_menuItem) => {
|
||||||
|
const color_w = this.widgets.find((w) => w.name === 'color')
|
||||||
|
const bg_w = this.widgets.find(
|
||||||
|
(w) => w.name === 'background' || w.name === 'bg_color',
|
||||||
|
)
|
||||||
|
|
||||||
|
const color = color_w.value
|
||||||
|
const bg = bg_w.value
|
||||||
|
|
||||||
|
color_w.value = bg
|
||||||
|
bg_w.value = color
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
options.push(item)
|
||||||
|
return [item]
|
||||||
|
})
|
||||||
|
break
|
||||||
|
}
|
||||||
case 'Save Tensors (mtb)': {
|
case 'Save Tensors (mtb)': {
|
||||||
const onDrawBackground = nodeType.prototype.onDrawBackground
|
const onDrawBackground = nodeType.prototype.onDrawBackground
|
||||||
nodeType.prototype.onDrawBackground = function (ctx, canvas) {
|
nodeType.prototype.onDrawBackground = function (ctx, canvas) {
|
||||||
|
|||||||
@@ -0,0 +1,334 @@
|
|||||||
|
// This is a vanillajs implementation of Houdini's number input widgets.
|
||||||
|
// It basically popup a visual sensitivity slider of steps to use as incr/decr
|
||||||
|
// TODO: Convert it to IWidget
|
||||||
|
|
||||||
|
// import styles from "./style.module.css";
|
||||||
|
|
||||||
|
function getValidNumber(numberInput) {
|
||||||
|
let num =
|
||||||
|
isNaN(numberInput.value) || numberInput.value === ''
|
||||||
|
? 0
|
||||||
|
: parseFloat(numberInput.value)
|
||||||
|
return num
|
||||||
|
}
|
||||||
|
/**
|
||||||
|
* Number input widgets
|
||||||
|
*/
|
||||||
|
export class NumberInputWidget {
|
||||||
|
constructor(containerId, numberOfInputs = 1, isDebug = false) {
|
||||||
|
this.container = document.getElementById(containerId)
|
||||||
|
this.numberOfInputs = numberOfInputs
|
||||||
|
this.currentInput = null // Store the currently active input
|
||||||
|
|
||||||
|
this.threshold = 30
|
||||||
|
this.mouseSensitivityMultiplier = 0.05
|
||||||
|
this.debug = isDebug
|
||||||
|
|
||||||
|
//- states
|
||||||
|
this.initialMouseX
|
||||||
|
this.lastMouseX
|
||||||
|
this.activeStep = 1
|
||||||
|
this.accumulatedDelta = 0
|
||||||
|
this.stepLocked = false
|
||||||
|
this.thresholdExceeded = false
|
||||||
|
this.isDragging = false
|
||||||
|
|
||||||
|
const styleTagId = 'mtb-constant-style'
|
||||||
|
|
||||||
|
let styleTag = document.head.querySelector(`#${styleTagId}`)
|
||||||
|
|
||||||
|
if (!styleTag) {
|
||||||
|
styleTag = document.createElement('style')
|
||||||
|
styleTag.type = 'text/css'
|
||||||
|
styleTag.id = styleTagId
|
||||||
|
|
||||||
|
styleTag.innerHTML = `
|
||||||
|
|
||||||
|
.${containerId}{
|
||||||
|
margin-top: 20px;
|
||||||
|
margin-bottom: 20px;
|
||||||
|
}
|
||||||
|
.sensitivity-menu {
|
||||||
|
display: none;
|
||||||
|
position: absolute;
|
||||||
|
/* Additional styling */
|
||||||
|
}
|
||||||
|
|
||||||
|
.sensitivity-menu .step {
|
||||||
|
cursor: pointer;
|
||||||
|
padding: 0.5em;
|
||||||
|
/* Add more styling as needed */
|
||||||
|
}
|
||||||
|
|
||||||
|
.sensitivity-menu {
|
||||||
|
font-family: monospace;
|
||||||
|
|
||||||
|
background: var(--bg-color);
|
||||||
|
border: 1px solid var(--fg-color);
|
||||||
|
/* Highlight for the active step */
|
||||||
|
}
|
||||||
|
.number-input {
|
||||||
|
background: var(--bg-color);
|
||||||
|
color: var(--fg-color)
|
||||||
|
}
|
||||||
|
|
||||||
|
.sensitivity-menu .step.active {
|
||||||
|
background-color:var(--drag-text);
|
||||||
|
/* Highlight for the active step */
|
||||||
|
}
|
||||||
|
|
||||||
|
.sensitivity-menu .step.locked {
|
||||||
|
background-color: #f00;
|
||||||
|
/* Change to your preferred color for the locked state */
|
||||||
|
}
|
||||||
|
#debug-container {
|
||||||
|
transform: translateX(50%);
|
||||||
|
width: 50%;
|
||||||
|
text-align: center;
|
||||||
|
font-family: monospace;
|
||||||
|
}
|
||||||
|
`
|
||||||
|
document.head.appendChild(styleTag)
|
||||||
|
}
|
||||||
|
|
||||||
|
this.createWidgetElements()
|
||||||
|
this.initializeEventListeners()
|
||||||
|
}
|
||||||
|
|
||||||
|
setLabel(str) {
|
||||||
|
this.label.textContent = str
|
||||||
|
}
|
||||||
|
setValue(...values) {
|
||||||
|
if (values.length !== this.numberInputs.length) {
|
||||||
|
console.error('Number of values does not match the number of inputs.')
|
||||||
|
console.error(
|
||||||
|
`You provided ${values.length} but the input want ${this.numberInputs.length}`,
|
||||||
|
{ values },
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
// Set each input value
|
||||||
|
this.numberInputs.forEach((input, index) => {
|
||||||
|
input.value = values[index]
|
||||||
|
})
|
||||||
|
}
|
||||||
|
getValue() {
|
||||||
|
const value = []
|
||||||
|
this.numberInputs.forEach((input, index) => {
|
||||||
|
value.push(Number.parseFloat(input.value) || 0.0)
|
||||||
|
})
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
resetValues() {
|
||||||
|
for (const input of numberInputs) {
|
||||||
|
input.value = 0
|
||||||
|
}
|
||||||
|
this.onChange?.(this.getValue())
|
||||||
|
}
|
||||||
|
|
||||||
|
createWidgetElements() {
|
||||||
|
this.label = document.createElement('label')
|
||||||
|
this.label.textContent = 'Control All:'
|
||||||
|
this.label.className = 'widget-label'
|
||||||
|
this.container.appendChild(this.label)
|
||||||
|
|
||||||
|
this.label.addEventListener('mousedown', (event) => {
|
||||||
|
if (event.button === 1) {
|
||||||
|
this.currentInput = null
|
||||||
|
this.handleMouseDown(event)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
this.label.addEventListener('contextmenu', (event) => {
|
||||||
|
event.preventDefault()
|
||||||
|
this.resetValues()
|
||||||
|
})
|
||||||
|
|
||||||
|
this.numberInputs = []
|
||||||
|
|
||||||
|
// create linked inputs
|
||||||
|
for (let i = 0; i < this.numberOfInputs; i++) {
|
||||||
|
const numberInput = document.createElement('input')
|
||||||
|
numberInput.type = 'number'
|
||||||
|
numberInput.className = 'number-input' //styles.numberInput; //"number-input";
|
||||||
|
numberInput.step = 'any'
|
||||||
|
this.container.appendChild(numberInput)
|
||||||
|
this.numberInputs.push(numberInput)
|
||||||
|
|
||||||
|
numberInput.addEventListener('mousedown', (event) => {
|
||||||
|
if (event.button === 1) {
|
||||||
|
this.currentInput = numberInput
|
||||||
|
this.handleMouseDown(event)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
this.sensitivityMenu = document.createElement('div')
|
||||||
|
this.sensitivityMenu.className = 'sensitivity-menu' //styles.sensitivityMenu; //"sensitivity-menu";
|
||||||
|
this.container.appendChild(this.sensitivityMenu)
|
||||||
|
|
||||||
|
// create steps
|
||||||
|
const stepsValues = [0.001, 0.01, 0.1, 1, 10, 100]
|
||||||
|
stepsValues.forEach((value) => {
|
||||||
|
const step = document.createElement('div')
|
||||||
|
step.className = 'step' //styles.step //"step";
|
||||||
|
step.dataset.step = value
|
||||||
|
step.textContent = value.toString()
|
||||||
|
this.sensitivityMenu.appendChild(step)
|
||||||
|
})
|
||||||
|
|
||||||
|
this.steps = this.sensitivityMenu.getElementsByClassName('step') //styles.step)
|
||||||
|
|
||||||
|
if (this.debug) {
|
||||||
|
this.debugContainer = document.createElement('div')
|
||||||
|
this.debugContainer.id = 'debug-container' //styles.debugContainer //"debugContainer";
|
||||||
|
document.body.appendChild(this.debugContainer)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
showSensitivityMenu(pageX, pageY) {
|
||||||
|
this.sensitivityMenu.style.display = 'block'
|
||||||
|
this.sensitivityMenu.style.left = `${pageX}px`
|
||||||
|
this.sensitivityMenu.style.top = `${pageY}px`
|
||||||
|
this.initialMouseX = pageX
|
||||||
|
this.lastMouseX = pageX
|
||||||
|
this.isDragging = true
|
||||||
|
this.thresholdExceeded = false
|
||||||
|
this.stepLocked = false
|
||||||
|
this.updateDebugInfo()
|
||||||
|
}
|
||||||
|
updateDebugInfo() {
|
||||||
|
if (this.debug) {
|
||||||
|
this.debugContainer.innerHTML = `
|
||||||
|
<div>Active Step: ${this.activeStep}</div>
|
||||||
|
<div>Initial Mouse X: ${this.initialMouseX}</div>
|
||||||
|
<div>Last Mouse X: ${this.lastMouseX}</div>
|
||||||
|
<div>Accumulated Delta: ${this.accumulatedDelta}</div>
|
||||||
|
<div>Threshold Exceeded: ${this.thresholdExceeded}</div>
|
||||||
|
<div>Step Locked: ${this.stepLocked}</div>
|
||||||
|
<div>Number Input Value: ${this.currentInput?.value}</div>
|
||||||
|
`
|
||||||
|
}
|
||||||
|
}
|
||||||
|
handleMouseDown(event) {
|
||||||
|
if (event.button === 1) {
|
||||||
|
this.showSensitivityMenu(
|
||||||
|
event.target.offsetWidth,
|
||||||
|
event.target.offsetHeight,
|
||||||
|
)
|
||||||
|
event.preventDefault()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
handleMouseUp(event) {
|
||||||
|
if (event.button === 1) {
|
||||||
|
this.resetWidgetState()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
handleClickOutside(event) {
|
||||||
|
if (event.target !== this.numberInput) {
|
||||||
|
this.resetWidgetState()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
handleMouseMove(event) {
|
||||||
|
if (this.sensitivityMenu.style.display === 'block') {
|
||||||
|
const relativeY = event.pageY - 300 // this.sensitivityMenu.offsetTop
|
||||||
|
|
||||||
|
const horizontalDistanceFromInitial = Math.abs(
|
||||||
|
event.target.offsetWidth - this.initialMouseX,
|
||||||
|
)
|
||||||
|
|
||||||
|
// Unlock if the mouse moves back towards the initial position
|
||||||
|
if (horizontalDistanceFromInitial < this.threshold) {
|
||||||
|
this.thresholdExceeded = false
|
||||||
|
this.stepLocked = false
|
||||||
|
this.accumulatedDelta = 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update step only if it is not locked
|
||||||
|
if (!this.stepLocked) {
|
||||||
|
for (let step of this.steps) {
|
||||||
|
step.classList.remove('active') //styles.active)
|
||||||
|
step.classList.remove('locked') //styles.locked)
|
||||||
|
if (
|
||||||
|
relativeY >= step.offsetTop &&
|
||||||
|
relativeY <= step.offsetTop + step.offsetHeight
|
||||||
|
) {
|
||||||
|
step.classList.add('active') //styles.active)
|
||||||
|
this.setActiveStep(parseFloat(step.dataset.step))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if (this.stepLocked) {
|
||||||
|
this.sensitivityMenu
|
||||||
|
.querySelector('.step.active')
|
||||||
|
?.classList.add('locked')
|
||||||
|
}
|
||||||
|
|
||||||
|
this.updateStepValue(event.pageX)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
initializeEventListeners() {
|
||||||
|
document.addEventListener('mousemove', (event) =>
|
||||||
|
this.handleMouseMove(event),
|
||||||
|
)
|
||||||
|
document.addEventListener('mouseup', (event) => this.handleMouseUp(event))
|
||||||
|
|
||||||
|
document.addEventListener('click', (event) =>
|
||||||
|
this.handleClickOutside(event),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
setActiveStep(val) {
|
||||||
|
if (this.activeStep !== val) {
|
||||||
|
this.activeStep = val
|
||||||
|
this.stepLocked = false
|
||||||
|
this.accumulatedDelta = 0
|
||||||
|
this.thresholdExceeded = false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
resetWidgetState() {
|
||||||
|
this.sensitivityMenu.style.display = 'none'
|
||||||
|
this.isDragging = false
|
||||||
|
this.lastMouseX = undefined
|
||||||
|
this.thresholdExceeded = false
|
||||||
|
this.stepLocked = false
|
||||||
|
this.updateDebugInfo()
|
||||||
|
}
|
||||||
|
updateStepValue(mouseX) {
|
||||||
|
if (this.isDragging && this.lastMouseX !== undefined) {
|
||||||
|
const deltaX = mouseX - this.lastMouseX
|
||||||
|
this.accumulatedDelta += deltaX
|
||||||
|
|
||||||
|
if (
|
||||||
|
!this.thresholdExceeded &&
|
||||||
|
Math.abs(this.accumulatedDelta) > this.threshold
|
||||||
|
) {
|
||||||
|
this.thresholdExceeded = true
|
||||||
|
this.stepLocked = true
|
||||||
|
}
|
||||||
|
|
||||||
|
if (this.thresholdExceeded && this.stepLocked) {
|
||||||
|
// frequency of value changes
|
||||||
|
if (
|
||||||
|
Math.abs(this.accumulatedDelta) * this.mouseSensitivityMultiplier >=
|
||||||
|
1
|
||||||
|
) {
|
||||||
|
const valueChange = Math.sign(this.accumulatedDelta) * this.activeStep
|
||||||
|
if (this.currentInput) {
|
||||||
|
this.currentInput.value =
|
||||||
|
getValidNumber(this.currentInput) + valueChange
|
||||||
|
this.onChange?.(this.getValue())
|
||||||
|
} else {
|
||||||
|
this.numberInputs.forEach((input) => {
|
||||||
|
input.value = getValidNumber(input) + valueChange
|
||||||
|
})
|
||||||
|
}
|
||||||
|
this.accumulatedDelta = 0
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
this.lastMouseX = mouseX
|
||||||
|
}
|
||||||
|
this.updateDebugInfo()
|
||||||
|
}
|
||||||
|
}
|
||||||
+1
-1
Submodule wiki updated: a3327c786b...4db733ae92
Reference in New Issue
Block a user