Compare commits

...
61 Commits
Author SHA1 Message Date
Mel Massadian 59608320c8 ⬆️ Bump version: 0.1.5 → 0.1.6 2024-07-03 18:16:37 +02:00
Mel Massadian d64fac4b74 fix: 🐛 menu callback issue
`+` on arrays returns a string in js...
2024-07-03 18:11:07 +02:00
Mel Massadian d687497d80 chore: 🧹 better classname extraction
Allow for consecutive uppercase letters:
- MTB_BatchFromHistoryV2 -> Batch From History V2
- MTB_CLIPInterpolate -> CLIP Interpolate
2024-07-03 16:00:48 +02:00
Mel Massadian d6343e1860 feat: ✨ add alpha channel support for faceswap/restore
Fixes #187
2024-07-03 15:59:09 +02:00
Mel Massadian 4eebdd8b8b ci: 🤖 limit release only to tags
I regularly need to push to main without needing to update
the extension's code / registry.
2024-07-02 13:14:00 +02:00
Mel Massadian 372e035686 Merge branch 'main' of https://github.com/melMass/comfy_mtb 2024-07-02 13:09:05 +02:00
Mel Massadian fb34671ee6 chore: 🧹 runner 2024-07-02 13:08:57 +02:00
Elthariel f25f6bdcd1 docs: 📚 Update requirements file in INSTALL.md (#186) 2024-06-26 15:05:45 +02:00
Mel Massadian f1b484617a ci: 🤖 only publish on tag
I can still autotag easily but it avoids bumping too much
versions to quickly
2024-06-22 20:17:14 +02:00
Mel Massadian 4507842a70 chore: 🧹 small fixes
- Handle image dimension mismatch in ConcatImages (Error,Smallest,Largest)
- typing
2024-06-22 20:13:59 +02:00
Mel Massadian e10faab458 ⬆️ Bump version: 0.1.4 → 0.1.5 2024-06-21 20:54:30 +02:00
Mel Massadian bb5682aa6d chore: 🧹 add fields for the registry 2024-06-21 20:48:55 +02:00
Mel Massadian 59612fd811 chore: 🧹 add pre-commit 2024-06-21 20:44:46 +02:00
Mel Massadian 30eb5b0091 chore 🧹: prepare for auto versioning 2024-06-21 19:55:50 +02:00
Mel Massadian 1edc2cd10d fix: 🐛 keep the last model match instead of first
See #184 for details
2024-06-21 19:53:25 +02:00
Mel Massadian fa3199be2b docs: 📚 update the wiki 2024-06-09 19:19:55 +02:00
Mel Massadian 43d65ae68c feat: ✨ add ModelPruner (wip) 2024-06-09 19:13:46 +02:00
Mel Massadian dfd17f6d78 chore: 🧹 migrate from poetry to setuptools 2024-06-09 15:23:07 +02:00
Mel Massadian 1070edd024 chore: 🧹 remove logs 2024-05-27 22:28:53 +02:00
Mel Massadian 9f0ed85cc1 Merge branch 'main' of https://github.com/melMass/comfy_mtb 2024-05-27 22:25:45 +02:00
Mel Massadian 35622e3a5e fix: 🐛 properly initialize the curve value
Also restored the old sorting logic adapted for Object
Closes #183

note: the ux is still bad and will improve
2024-05-27 22:25:28 +02:00
Mel Massadian 644371e5b5 chore: 🧹 add more pyproject meta 2024-05-21 12:51:40 +02:00
Mel Massadian f3d468cfc2 ci: 🤖 move at the proper location 2024-05-21 12:45:22 +02:00
haohaocreatesandMel Massadian 6cd448b026 ci: 🤖 add CI to publish to ComfyUI Registry (#182)
* publish-action
* feat ⚡:  rename token and add icons

---------

Co-authored-by: Mel Massadian <melmassadian@gmail.com>
2024-05-21 12:40:16 +02:00
haohaocreatesandMel Massadian 5951c90b10 chore: 🧹 add ComfyUI registry to pyproject.toml (#181)
* Add pyproject.toml for Custom Node Registry
* feat ⚡: add publisher id

---------

Co-authored-by: Mel Massadian <melmassadian@gmail.com>
2024-05-21 12:37:54 +02:00
bymyself 6abac2e470 feat: ✨ Use dynamic contrast in Color Correct (#180)
* Change contrast_adjustment_tensor method to change contrast dynamically
* Switch to Adobe RGB color space
2024-05-20 23:32:15 +02:00
Mel Massadian 01c73e1c5e feat ⚡: add more options to load image sequence 2024-05-17 17:35:53 +02:00
Mel Massadian 5060c56135 feat: ✨ StackImages add support for batch mismatch
Useful for comparing a static image with a batch of images
for instance.
2024-05-15 12:21:10 +02:00
bymyselfandMel Massadian acc2d687d5 fix: 🐛 ImageCompare improvements (#176)
* avoid unnecessary numpy conversion for diff and blend
* add support for Batch
* add support for input mismatch (RGB/RGBA)
* fixes #175 

---------

Co-authored-by: Mel Massadian <mel@melmassadian.com>
2024-05-14 21:16:56 +02:00
vxkj1211andMel Massadian 780c52f03a fix: 🐛 repetitive warning (#177)
Co-authored-by: Mel Massadian <mel@melmassadian.com>
2024-05-14 16:09:10 +02:00
Mel Massadian 2fe0859476 docs 📚: update wiki 2024-05-14 15:36:43 +02:00
Mel Massadian 1186239751 chore 🧹: use sections properly 2024-05-14 15:36:29 +02:00
Mel Massadian 96a0da9dbd chore: 🧹 update types 2024-05-10 20:30:46 +02:00
Mel Massadian f9d2ebf91d feat: ✨ add BatchFloatMath
Simple math operations on FLOATS (list of floats)
2024-05-07 23:42:02 +02:00
Mel Massadian 1b7ae27cc1 feat: ✨ add FLOATS to INTS
For using it with FrameInterpolation's new multiplier
2024-05-07 19:22:23 +02:00
Mel Massadian e312b02ad2 wip: 🚧 curve widget logic fixed
Most of the logic is fixed, but it still needs some UI/UX tweaks.
2024-05-07 18:40:12 +02:00
Mel Massadian 63ee25d001 feat: ✨ debug dict
it was only working on conditions
2024-05-07 18:36:56 +02:00
Mel Massadian 1caf7c18c3 feat: ✨ add Swap BG/FG color menu item 2024-05-07 08:29:57 +02:00
Mel Massadian 349a8524c6 fix: 🐛 add back was conversion node
To avoid breaking other worklfows
I thought this was now builtin WAS suite.
Fixes #172
2024-05-02 07:58:37 +02:00
Mel Massadian 15330eab65 fix: 🐛 drag lag on documentation resize handle 2024-04-28 16:54:42 +02:00
Mel Massadian 1571782d01 fix: 🐛 kwarg typo
floats vs float
2024-04-28 15:51:24 +02:00
Mel Massadian 5b4030288d fix: 🐛 seed of PlotBatchFloat
Also using random colors instead of mapped to
colormap, the values weren't distinct enough
2024-04-28 15:18:27 +02:00
Mel Massadian ab58c36212 feat: ✨ BatchFloatFit the batch version of FitNumber 2024-04-28 15:18:27 +02:00
Mel Massadian 5a0ef0dadd fix: 🐛 forceInput for FLOAT <-> FLOATS converters 2024-04-28 13:10:36 +02:00
Mel Massadian 967e72fc66 fix: 🐛 FLOAT always need options to be set
Fixes #171
2024-04-28 13:04:38 +02:00
Mel Massadian 78a86daaf7 feat: ✨ add FloatToFloats (the counterpart) 2024-04-27 21:11:53 +02:00
Mel Massadian bee3f47a14 fix: 🐛 remove doc if opened on node delete 2024-04-27 20:25:54 +02:00
Mel Massadian 2159395389 feat: ✨ add some FLOATS batch nodes
* TimeWrap
* Normalize
2024-04-27 19:47:45 +02:00
Mel Massadian b11346aba8 fix: 🐛 for documentation on HiDPI
thanks @kijai
2024-04-27 19:46:25 +02:00
Mel Massadian 30982fa488 fix: 🐛 never remove input 0 of dynamic inputs
If you reloaded a graph containing a node with dynamic inputs
but none connected the node would end up input-less
2024-04-27 16:29:34 +02:00
Mel Massadian 92b79906cd fix: 🐛 use the same fix as dynamicInputs for debug
i.e we don't auto delete inputs on disconnect, only on connect of
inputs
2024-04-27 14:58:50 +02:00
Mel Massadian 76f365b5ee fix: 🐛 missing numberInput
This is the first iteration of the "multi" number inputs.
The behaviour is based on Houdini number inputs
2024-04-27 14:16:07 +02:00
Mel Massadian da67e766c2 fix: 🐛 better curve 2024-04-27 04:08:09 +02:00
Mel Massadian 49cea8d945 docs: 📚 update wiki submodule 2024-04-27 01:17:25 +02:00
Mel Massadian b1d74adb15 fix: 🐛 prepend MTB_ to all classes
to avoid any future clash.
2024-04-27 01:13:32 +02:00
Mel Massadian 652ac3f3b9 fix: 🐛 dynamic connections 2024-04-26 21:11:32 +02:00
Mel Massadian 060e733605 Merge branch 'main' into fix/js-refactor 2024-04-25 22:19:28 +02:00
Mel Massadian eedbb4bc65 wip: 🚧 dump3 2024-04-25 22:08:25 +02:00
Mel Massadian fa2397585f wip: 🚧 dump 2024-04-25 21:56:59 +02:00
Mel Massadian 77348c4adb Merge branch 'main' into fix/js-refactor 2024-04-25 21:42:12 +02:00
Mel Massadian 0d0fb8e13a wip: 🚧 dump
js refactor start
2024-04-25 21:41:40 +02:00
39 changed files with 3653 additions and 1002 deletions
+18
View File
@@ -0,0 +1,18 @@
name: 📦 Publish to Comfy registry
on:
workflow_dispatch:
push:
tags:
- '*'
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 }}
+8
View File
@@ -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
+1 -1
View File
@@ -42,7 +42,7 @@ then follow the prompt or just press enter to download every models.
1. Make sure you are in the Python environment you use for ComfyUI. 1. Make sure you are in the Python environment you use for ComfyUI.
2. Install the required dependencies by running the following command: 2. Install the required dependencies by running the following command:
```bash ```bash
pip install -r comfy_mtb/reqs.txt pip install -r comfy_mtb/requirements.txt
``` ```
</details> </details>
+85 -27
View File
@@ -6,21 +6,27 @@
# Copyright (c) 2023 Mel Massadian # Copyright (c) 2023 Mel Massadian
# #
### ###
__version__ = "0.1.6"
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
View File
@@ -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
View File
@@ -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:
+158
View File
@@ -0,0 +1,158 @@
# 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"
let inputs = $"($root)/input"
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)"
if not $clean {
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 }
}
} else {
rm $models
rm $inputs
}
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)"
if not $clean {
cd $models
# resymlink them
open links.nuon | each {|p| link -a $p.target $p.name }
} else {
let master = (get_root)
link ($master | path join models) $models
link ($master | path join input) $inputs
}
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
}
+14 -8
View File
@@ -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,24 +58,30 @@ 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])|(?<=[A-Z])(?=[A-Z][a-z])|(?<=[A-Za-z])(?=[0-9])|(?<=[0-9])(?=[A-Za-z]))",
label,
)
reformatted_label = re.sub(r"([A-Z]+)", r" \1", label).strip()
words = reformatted_label.split()
return " ".join(words).strip() return " ".join(words).strip()
+262 -24
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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]
+34 -15
View File
@@ -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:
@@ -176,22 +177,33 @@ class RestoreFace:
# Adjustable weights # Adjustable weights
"weight": ("FLOAT", {"default": 0.5}), "weight": ("FLOAT", {"default": 0.5}),
"save_tmp_steps": ("BOOLEAN", {"default": True}), "save_tmp_steps": ("BOOLEAN", {"default": True}),
} },
"optional": {
"preserve_alpha": ("BOOLEAN", {"default": True}),
},
} }
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,
save_tmp_steps, save_tmp_steps,
preserve_alpha: bool = False,
) -> torch.Tensor: ) -> torch.Tensor:
pimage = tensor2np(image)[0] pimage = tensor2np(image)[0]
width, height = pimage.shape[1], pimage.shape[0] width, height = pimage.shape[1], pimage.shape[0]
source_img = cv2.cvtColor(np.array(pimage), cv2.COLOR_RGB2BGR) source_img = cv2.cvtColor(np.array(pimage), cv2.COLOR_RGB2BGR)
alpha_channel = None
if (
preserve_alpha and image.size(-1) == 4
): # Check if the image has an alpha channel
alpha_channel = pimage[:, :, 3]
pimage = pimage[:, :, :3] # Remove alpha channel for processing
sys.stdout = NullWriter() sys.stdout = NullWriter()
cropped_faces, restored_faces, restored_img = model.enhance( cropped_faces, restored_faces, restored_img = model.enhance(
source_img, source_img,
@@ -210,9 +222,14 @@ class RestoreFace:
) )
output = None output = None
if restored_img is not None: if restored_img is not None:
output = Image.fromarray( restored_img = cv2.cvtColor(restored_img, cv2.COLOR_BGR2RGB)
cv2.cvtColor(restored_img, cv2.COLOR_BGR2RGB) output = Image.fromarray(restored_img)
)
if alpha_channel is not None:
alpha_resized = Image.fromarray(alpha_channel).resize(
output.size, Image.LANCZOS
)
output.putalpha(alpha_resized)
# imwrite(restored_img, save_restore_path) # imwrite(restored_img, save_restore_path)
return pil2tensor(output) return pil2tensor(output)
@@ -220,12 +237,13 @@ 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]: preserve_alpha: bool = False,
) -> tuple[torch.Tensor]:
out = [ out = [
self.do_restore( self.do_restore(
image[i], image[i],
@@ -234,6 +252,7 @@ class RestoreFace:
only_center_face, only_center_face,
weight, weight,
save_tmp_steps, save_tmp_steps,
preserve_alpha,
) )
for i in range(image.size(0)) for i in range(image.size(0))
] ]
@@ -259,7 +278,7 @@ class RestoreFace:
self, cropped_faces, restored_faces, height, width self, cropped_faces, restored_faces, height, width
): ):
for idx, (cropped_face, restored_face) in enumerate( for idx, (cropped_face, restored_face) in enumerate(
zip(cropped_faces, restored_faces) zip(cropped_faces, restored_faces, strict=False)
): ):
face_id = idx + 1 face_id = idx + 1
file = self.get_step_image_path("cropped_faces", face_id) file = self.get_step_image_path("cropped_faces", face_id)
@@ -275,4 +294,4 @@ class RestoreFace:
cv2.imwrite(file, cmp_img) cv2.imwrite(file, cmp_img)
__nodes__ = [RestoreFace, LoadFaceEnhanceModel] __nodes__ = [MTB_RestoreFace, MTB_LoadFaceEnhanceModel]
+19 -9
View File
@@ -2,7 +2,6 @@
# region imports # region imports
import sys import sys
from pathlib import Path from pathlib import Path
from typing import List, Optional, Set, Union
import comfy.model_management as model_management import comfy.model_management as model_management
import cv2 import cv2
@@ -22,7 +21,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 +52,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 +96,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
@@ -119,7 +118,9 @@ class FaceSwap:
), ),
"faceswap_model": ("FACESWAP_MODEL", {"default": "None"}), "faceswap_model": ("FACESWAP_MODEL", {"default": "None"}),
}, },
"optional": {}, "optional": {
"preserve_alpha": ("BOOLEAN", {"default": True}),
},
} }
RETURN_TYPES = ("IMAGE",) RETURN_TYPES = ("IMAGE",)
@@ -133,11 +134,18 @@ class FaceSwap:
faces_index: str, faces_index: str,
faceanalysis_model, faceanalysis_model,
faceswap_model, faceswap_model,
preserve_alpha=False,
): ):
def do_swap(img): def do_swap(img):
model_management.throw_exception_if_processing_interrupted() model_management.throw_exception_if_processing_interrupted()
img = tensor2pil(img)[0] img = tensor2pil(img)[0]
ref = tensor2pil(reference)[0] ref = tensor2pil(reference)[0]
alpha_channel = None
if preserve_alpha and img.mode == "RGBA":
alpha_channel = img.getchannel("A")
img = img.convert("RGB")
face_ids = { face_ids = {
int(x) int(x)
for x in faces_index.strip(",").split(",") for x in faces_index.strip(",").split(",")
@@ -148,6 +156,8 @@ class FaceSwap:
faceanalysis_model, ref, img, faceswap_model, face_ids faceanalysis_model, ref, img, faceswap_model, face_ids
) )
sys.stdout = sys.__stdout__ sys.stdout = sys.__stdout__
if alpha_channel:
swapped.putalpha(alpha_channel)
return pil2tensor(swapped) return pil2tensor(swapped)
batch_count = image.size(0) batch_count = image.size(0)
@@ -194,10 +204,10 @@ def get_face_single(
def swap_face( def swap_face(
face_analyser, face_analyser,
source_img: Union[Image.Image, List[Image.Image]], source_img: Image.Image | list[Image.Image],
target_img: Union[Image.Image, List[Image.Image]], target_img: Image.Image | list[Image.Image],
face_swapper_model, face_swapper_model,
faces_index: Optional[Set[int]] = None, faces_index: set[int] | None = None,
) -> Image.Image: ) -> Image.Image:
if faces_index is None: if faces_index is None:
faces_index = {0} faces_index = {0}
@@ -239,4 +249,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
View File
@@ -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,
] ]
+121 -45
View File
@@ -3,17 +3,16 @@ import json
import urllib.parse import urllib.parse
import urllib.request import urllib.request
from math import pi from math import pi
from typing import Optional
import comfy.model_management as model_management import comfy.model_management as model_management
import comfy.utils import comfy.utils
import numpy as np import numpy as np
import torch import torch
import torchvision.transforms.functional as F
from PIL import Image 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,
@@ -70,8 +69,8 @@ class MTB_ToDevice:
*, *,
ignore_errors=False, ignore_errors=False,
device="cuda", device="cuda",
image: Optional[torch.Tensor] = None, image: torch.Tensor | None = None,
mask: Optional[torch.Tensor] = None, mask: torch.Tensor | None = None,
): ):
if not ignore_errors and image is None and mask is None: if not ignore_errors and image is None and mask is None:
raise ValueError( raise ValueError(
@@ -137,6 +136,8 @@ class MTB_MatchDimensions:
def execute( def execute(
self, source: torch.Tensor, reference: torch.Tensor, match: str self, source: torch.Tensor, reference: torch.Tensor, match: str
): ):
import torchvision.transforms.functional as VF
_batch_size, height, width, _channels = source.shape _batch_size, height, width, _channels = source.shape
_rbatch_size, rheight, rwidth, _rchannels = reference.shape _rbatch_size, rheight, rwidth, _rchannels = reference.shape
@@ -154,7 +155,7 @@ class MTB_MatchDimensions:
new_height = int(rwidth / source_aspect_ratio) new_height = int(rwidth / source_aspect_ratio)
resized_images = [ resized_images = [
F.resize( VF.resize(
source[i], source[i],
(new_height, new_width), (new_height, new_width),
antialias=True, antialias=True,
@@ -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):
@@ -493,35 +531,24 @@ class MTB_FitNumber:
"required": { "required": {
"value": ("FLOAT", {"default": 0, "forceInput": True}), "value": ("FLOAT", {"default": 0, "forceInput": True}),
"clamp": ("BOOLEAN", {"default": False}), "clamp": ("BOOLEAN", {"default": False}),
"source_min": ("FLOAT", {"default": 0.0, "step": 0.01}), "source_min": (
"source_max": ("FLOAT", {"default": 1.0, "step": 0.01}), "FLOAT",
"target_min": ("FLOAT", {"default": 0.0, "step": 0.01}), {"default": 0.0, "step": 0.01, "min": -1e5},
"target_max": ("FLOAT", {"default": 1.0, "step": 0.01}), ),
"source_max": (
"FLOAT",
{"default": 1.0, "step": 0.01, "min": -1e5},
),
"target_min": (
"FLOAT",
{"default": 0.0, "step": 0.01, "min": -1e5},
),
"target_max": (
"FLOAT",
{"default": 1.0, "step": 0.01, "min": -1e5},
),
"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"},
), ),
} }
@@ -568,19 +595,66 @@ class MTB_ConcatImages:
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
return { return {
"required": {"reverse": ("BOOLEAN", {"default": False})}, "required": {"reverse": ("BOOLEAN", {"default": False})},
"optional": {
"on_mismatch": (
["Error", "Smallest", "Largest"],
{"default": "Smallest"},
)
},
} }
def concatenate_tensors(self, reverse, **kwargs): def concatenate_tensors(
tensors = tuple(kwargs.values()) self,
batch_sizes = [tensor.size(0) for tensor in tensors] reverse: bool,
on_mismatch: str = "Smallest",
**kwargs: torch.Tensor,
) -> tuple[torch.Tensor]:
tensors = list(kwargs.values())
if on_mismatch == "Error":
shapes = [tensor.shape for tensor in tensors]
if not all(shape == shapes[0] for shape in shapes):
raise ValueError(
"All input tensors must have the same shape when on_mismatch is 'Error'."
)
else:
import torch.nn.functional as F
if on_mismatch == "Smallest":
target_shape = min(
(tensor.shape for tensor in tensors),
key=lambda s: (s[1], s[2]),
)
else: # on_mismatch == "Largest"
target_shape = max(
(tensor.shape for tensor in tensors),
key=lambda s: (s[1], s[2]),
)
target_height, target_width = target_shape[1], target_shape[2]
resized_tensors = []
for tensor in tensors:
if (
tensor.shape[1] != target_height
or tensor.shape[2] != target_width
):
resized_tensor = F.interpolate(
tensor.permute(0, 3, 1, 2),
size=(target_height, target_width),
mode="bilinear",
align_corners=False,
)
resized_tensor = resized_tensor.permute(0, 2, 3, 1)
resized_tensors.append(resized_tensor)
else:
resized_tensors.append(tensor)
tensors = resized_tensors
concatenated = torch.cat(tensors, dim=0) concatenated = torch.cat(tensors, dim=0)
# Update the batch size in the concatenated tensor
concatenated_size = list(concatenated.size())
concatenated_size[0] = sum(batch_sizes)
concatenated = concatenated.view(*concatenated_size)
return (concatenated,) return (concatenated,)
@@ -596,4 +670,6 @@ __nodes__ = [
MTB_MatchDimensions, MTB_MatchDimensions,
MTB_AutoPanEquilateral, MTB_AutoPanEquilateral,
MTB_FloatsToFloat, MTB_FloatsToFloat,
MTB_FloatToFloats,
MTB_FloatsToInts,
] ]
+9 -8
View File
@@ -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]
+54 -8
View File
@@ -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
+20
View File
@@ -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
View File
@@ -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,
] ]
+6 -3
View File
@@ -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
View File
@@ -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]
-2
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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.6"
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.6"
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
+79 -5
View File
@@ -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
View File
@@ -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
*/ */
+160 -18
View File
@@ -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,16 @@ 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[np.float32]],
) -> 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[np.float32]]:
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 +798,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 +836,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
+525 -265
View File
@@ -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'
@@ -553,7 +704,7 @@ const create_documentation_stylesheet = () => {
border-radius: 6px; border-radius: 6px;
border: 3px solid var(--bg-color); border: 3px solid var(--bg-color);
} }
/* Scrollbar styling for Firefox */ /* Scrollbar styling for Firefox */
scrollbar-width: thin; scrollbar-width: thin;
scrollbar-color: var(--fg-color) var(--bg-color); scrollbar-color: var(--fg-color) var(--bg-color);
@@ -575,7 +726,7 @@ const create_documentation_stylesheet = () => {
border-collapse: collapse; border-collapse: collapse;
border: 1px var(--border-color) solid; border: 1px var(--border-color) solid;
} }
.documentation-popup th, .documentation-popup th,
.documentation-popup td { .documentation-popup td {
border: 1px var(--border-color) solid; border: 1px var(--border-color) solid;
} }
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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) {
+334
View File
@@ -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