Compare commits

..
21 Commits
Author SHA1 Message Date
Mel Massadian 50cb6f5ed6 chore: 🧹 bump minor 2024-12-08 19:34:26 +01:00
Mel Massadian e32d1e02df feat: ✨ add h264 compression node
recommended for i2i in ltx.
original code by [mix](https://github.com/XmYx)
2024-12-08 19:12:28 +01:00
Mel Massadian b0d52f7305 fix: 🐛 remove mtb sidebar
- The source for this is not yet in main... this file slipped
  in an earlier commit

fixes https://github.com/Comfy-Org/ComfyUI_frontend/issues/1834
2024-12-07 15:45:43 +01:00
Mel Massadian e17c6e29f5 docs: 📚 update wiki
pull wiki for documentation
2024-12-04 02:11:02 +01:00
Mel Massadian 27e03fa23e feat: ✨ add postshot nodes
basic wrapper of the cli the idea is to
queue Cog + Rotating loras -> Postshot

needs testing
2024-12-03 23:17:54 +01:00
Mel Massadian ec1cb1ac17 fix: 🐛 always enable the I/O sidebar
closes #214
2024-12-03 22:11:48 +01:00
d8ahazard 64634104a2 Use local import for Rembg
Rembg can sometimes cause *very* long load times on import (like 40s). Moving it to local doesn't fix the long import entirely, but it does prevent it causing ComfyUI from loading slowly.
2024-12-03 04:55:55 +01:00
Mel Massadian ecbb220de6 fix: 🐛 ui shifts on animation builder
finally updated to addDOMWidget
2024-11-20 23:03:00 +01:00
Mel Massadian cd9e614b1a feat: ✨ improve the I/O sidebar
- better options (sort, count)
- uses the new toast api instead of MTB.notify
2024-11-20 22:42:32 +01:00
Mel Massadian 9ccf572a15 chore: 🧹 add worktree to gitignores
for the experimental doc site at:
https://melmass.github.io/comfy_mtb/
2024-11-20 22:42:32 +01:00
Mel Massadian 74af5c6499 feat: ✨ add UpscaleBBoxBy 2024-11-20 22:42:32 +01:00
Mel Massadian caf0b39d8a chore 🧹: add deprecations and experimental 2024-11-20 22:42:32 +01:00
Mel Massadian e099d581a7 chore: 🧹 remove dupe code 2024-11-20 22:42:32 +01:00
Mel Massadian 22f7c30373 feat: ✨ simplified sidebar and backend
If you have a LoadImage selected,
clicking on images in the "input" mode will set the image on the
selected nodes
2024-11-20 22:42:32 +01:00
Mel Massadian 0133fb93bc feat: ✨ add Interpolate Condition 2024-11-20 22:42:32 +01:00
Mel Massadian cf7d30507e feat: ✨ dump of wip things... 2024-11-20 22:42:32 +01:00
Mel Massadian b6fa571fd2 fix: 🐛 category for settings 2024-11-20 21:41:57 +01:00
Mel Massadian f272526bfc fix: 🐛 new UI issues
- Fixes the "edit icon cannot be clicked"
- Changed the parser to add support for more non std markdown
- Markdown links now always open a new tab instead of replacing current
- New optional shiki support for code blocks (check #211 for details)
2024-11-20 21:41:57 +01:00
Mel Massadian 4e593bb30b feat: ✨ use the new parser for documentations
- might also fix #210
2024-11-20 21:41:57 +01:00
Mel Massadian 097ca33b8e feat: ✨ add @mtb/markdown-parser bundles
- the standard one is half the size of showdown
- the enhanced one (add shiki with most of its features) is 1.5mb
2024-11-20 21:41:57 +01:00
Mel Massadian 784fb0145b chore: 🧹 update externs
- remove showdown
- update dompurify
2024-11-20 21:41:57 +01:00
23 changed files with 3521 additions and 1706 deletions
+3
View File
@@ -6,3 +6,6 @@ node_modules/
compose.yaml compose.yaml
comfy_mtb.wsb comfy_mtb.wsb
Dockerfile Dockerfile
# I store the gh-pages worktrees (src & build) there
.worktrees
+207 -47
View File
@@ -7,10 +7,12 @@
# #
### ###
__version__ = "0.1.6" __version__ = "0.2.0"
import os import os
from aiohttp.web_request import Request
# TODO: don't override this if the user has that setup already # TODO: don't override this if the user has that setup already
if not os.environ.get("TF_FORCE_GPU_ALLOW_GROWTH"): if not os.environ.get("TF_FORCE_GPU_ALLOW_GROWTH"):
os.environ["TF_FORCE_GPU_ALLOW_GROWTH"] = "true" os.environ["TF_FORCE_GPU_ALLOW_GROWTH"] = "true"
@@ -31,24 +33,20 @@ from pathlib import Path
from aiohttp import web from aiohttp import web
from server import PromptServer from server import PromptServer
import nodes
from .endpoint import endlog from .endpoint import endlog
from .log import blue_text, cyan_text, get_label, get_summary, log from .log import blue_text, cyan_text, get_label, get_summary, log
from .utils import comfy_dir, here from .utils import comfy_dir, here
NODE_CLASS_MAPPINGS = {} NODE_CLASS_MAPPINGS: dict[str, type] = {}
NODE_DISPLAY_NAME_MAPPINGS = {} NODE_DISPLAY_NAME_MAPPINGS: dict[str, str] = {}
NODE_CLASS_MAPPINGS_DEBUG = {} NODE_CLASS_MAPPINGS_DEBUG: dict[str, str | None] = {}
WEB_DIRECTORY = "./web" WEB_DIRECTORY = "./web"
def extract_nodes_from_source(filename: Path): def extract_nodes_from_source(filename: Path):
source_code = "" source_code = ""
source_code = filename.read_text(encoding="utf-8") source_code = filename.read_text(encoding="utf-8")
nodes: list[str] = []
nodes = []
try: try:
parsed = ast.parse(source_code) parsed = ast.parse(source_code)
@@ -57,14 +55,15 @@ def extract_nodes_from_source(filename: Path):
target = node.targets[0] target = node.targets[0]
if isinstance(target, ast.Name) and target.id == "__nodes__": if isinstance(target, ast.Name) and target.id == "__nodes__":
value = ast.get_source_segment(source_code, node.value) value = ast.get_source_segment(source_code, node.value)
node_value = ast.parse(value).body[0].value if value:
if isinstance(node_value, (ast.List, ast.Tuple)): node_value = ast.parse(value).body[0].value
nodes.extend( if isinstance(node_value, ast.List | ast.Tuple):
element.id nodes.extend(
for element in node_value.elts str(element.id)
if isinstance(element, ast.Name) for element in node_value.elts
) if isinstance(element, ast.Name)
break )
break
except SyntaxError: except SyntaxError:
log.error("Failed to parse") log.error("Failed to parse")
return nodes return nodes
@@ -72,8 +71,8 @@ def extract_nodes_from_source(filename: Path):
def load_nodes(): def load_nodes():
errors: list[str] = [] errors: list[str] = []
nodes = [] nodes: list[type] = []
nodes_failed = [] nodes_failed: list[str] = []
for filename in (here / "nodes").iterdir(): for filename in (here / "nodes").iterdir():
if filename.suffix == ".py": if filename.suffix == ".py":
@@ -124,7 +123,8 @@ def uninstall_old_web_extensions():
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}
Please manually remove it from disk ({web_mtb}) and restart the server."""
) )
@@ -141,7 +141,7 @@ def wiki_to_classname(s: str):
def classname_to_wiki(s: str): def classname_to_wiki(s: str):
classname = s.replace("MTB_", "") classname = s.replace("MTB_", "")
parts = [] parts: list[str] = []
start = 0 start = 0
for i in range(1, len(classname)): for i in range(1, len(classname)):
if classname[i].isupper(): if classname[i].isupper():
@@ -161,8 +161,6 @@ if wiki.exists() and wiki.is_dir():
# - REGISTER NODES # - REGISTER NODES
MTB_EXPORT = os.environ.get("MTB_EXPORT") MTB_EXPORT = os.environ.get("MTB_EXPORT")
nodes, failed = load_nodes() nodes, failed = load_nodes()
@@ -179,7 +177,7 @@ for node_class in nodes:
node_class.DESCRIPTION = node_class.__doc__ node_class.DESCRIPTION = node_class.__doc__
if MTB_EXPORT: if MTB_EXPORT:
wiki_name = classname_to_wiki(class_name) wiki_name = classname_to_wiki(class_name)
(wiki / "nodes" / (wiki_name + ".md")).write_text( _ = (wiki / "nodes" / (wiki_name + ".md")).write_text(
node_class.__doc__, encoding="utf-8" node_class.__doc__, encoding="utf-8"
) )
@@ -192,12 +190,15 @@ for node_class in nodes:
NODE_CLASS_MAPPINGS[node_label] = node_class NODE_CLASS_MAPPINGS[node_label] = node_class
NODE_DISPLAY_NAME_MAPPINGS[class_name] = node_label NODE_DISPLAY_NAME_MAPPINGS[class_name] = node_label
NODE_CLASS_MAPPINGS_DEBUG[node_label] = node_class.__doc__ NODE_CLASS_MAPPINGS_DEBUG[node_label] = node_class.__doc__
# TODO: I removed this, I find it more convenient to write without spaces, but it breaks every of my workflows
# TODO (cont): and until I find a way to automate the conversion, I'll leave it like this # TODO: I removed this, I find it more convenient to write without spaces
# but it breaks every of my workflows
# TODO (cont): and until I find a way to automate the conversion
# I'll leave it like this
if os.environ.get("MTB_EXPORT"): if os.environ.get("MTB_EXPORT"):
with open(here / "node_list.json", "w") as f: with open(here / "node_list.json", "w") as f:
f.write( _ = f.write(
json.dumps( json.dumps(
{ {
k: NODE_CLASS_MAPPINGS_DEBUG[k] k: NODE_CLASS_MAPPINGS_DEBUG[k]
@@ -215,19 +216,30 @@ log.debug(
) )
) )
log.info(f"loaded {cyan_text(len(nodes))} nodes successfuly") log.info(f"loaded {cyan_text(str(len(nodes)))} nodes successfuly")
if failed: if failed:
with contextlib.suppress(Exception): with contextlib.suppress(Exception):
base_url, port = utils.get_server_info() base_url, port = utils.get_server_info()
log.info( log.info(
f"Some nodes ({len(failed)}) could not be loaded. This can be ignored, but go to http://{base_url}:{port}/mtb if you want more information." f"Some nodes ({len(failed)}) could not be loaded. This can be ignored, but go to http://{base_url}:{port}/mtb if you want more information."
) )
log.debug(failed)
# - ENDPOINT # - ENDPOINT
if hasattr(PromptServer, "instance"): if hasattr(PromptServer, "instance"):
img_cache = None
prompt_cache = None
with contextlib.suppress(ImportError):
from cachetools import TTLCache
img_cache = TTLCache(maxsize=100, ttl=5) # 1 min TTL
prompt_cache = TTLCache(maxsize=100, ttl=5) # 1 min TTL
restore_deps = ["basicsr"] restore_deps = ["basicsr"]
onnx_deps = ["onnxruntime"] onnx_deps = ["onnxruntime"]
swap_deps = ["insightface"] + onnx_deps swap_deps = ["insightface"] + onnx_deps
@@ -306,10 +318,10 @@ if hasattr(PromptServer, "instance"):
} }
) )
@PromptServer.instance.routes.post("/mtb/debug") @PromptServer.instance.routes.post("/mtb/server-info")
async def set_debug(request): async def set_server_info(request: Request):
json_data = await request.json() json_data: dict[str, bool] = await request.json()
enabled = json_data.get("enabled") enabled = json_data.get("debug")
if enabled: if enabled:
os.environ["MTB_DEBUG"] = "true" os.environ["MTB_DEBUG"] = "true"
log.setLevel(logging.DEBUG) log.setLevel(logging.DEBUG)
@@ -317,7 +329,7 @@ if hasattr(PromptServer, "instance"):
elif "MTB_DEBUG" in os.environ: elif "MTB_DEBUG" in os.environ:
# del os.environ["MTB_DEBUG"] # del os.environ["MTB_DEBUG"]
os.environ.pop("MTB_DEBUG") _ = os.environ.pop("MTB_DEBUG")
log.setLevel(logging.INFO) log.setLevel(logging.INFO)
return web.json_response( return web.json_response(
@@ -325,17 +337,17 @@ if hasattr(PromptServer, "instance"):
) )
@PromptServer.instance.routes.get("/mtb") @PromptServer.instance.routes.get("/mtb")
async def get_home(request): async def get_home(request: Request):
from . import endpoint from . import endpoint
reload(endpoint) _ = reload(endpoint)
# Check if the request prefers HTML content # Check if the request prefers HTML content
if "text/html" in request.headers.get("Accept", ""): if "text/html" in request.headers.get("Accept", ""):
# # Return an HTML page # # Return an HTML page
html_response = """ html_response = """
<div class="flex-container menu"> <div class="flex-container menu">
<a href="/mtb/manage">manage</a> <a href="/mtb/manage">manage</a>
<a href="/mtb/debug">debug</a> <a href="/mtb/server-info">Server Info</a>
<a href="/mtb/status">status</a> <a href="/mtb/status">status</a>
</div> </div>
""" """
@@ -347,28 +359,176 @@ if hasattr(PromptServer, "instance"):
# Return JSON for other requests # Return JSON for other requests
return web.json_response({"message": "Welcome to MTB!"}) return web.json_response({"message": "Welcome to MTB!"})
@PromptServer.instance.routes.get("/mtb/debug") import asyncio
async def get_debug(request): import os
from io import BytesIO
from aiohttp import web
from PIL import Image
def get_cached_image(file_path: str, preview_params=None, channel=None):
cache_key = (file_path, preview_params, channel)
if img_cache and (cache_key in img_cache):
return img_cache[cache_key]
with Image.open(file_path) as img:
info = img.info
if preview_params:
img = process_preview(img, preview_params)
if channel:
img = process_channel(img, channel)
if prompt_cache:
prompt_cache[cache_key] = info
if img_cache:
img_cache[cache_key] = img.getvalue()
return img_cache[cache_key]
return img.getvalue()
def process_preview(img: Image.Image, preview_params):
image_format, quality, width = preview_params
quality = int(quality)
if width:
width = int(width)
img.thumbnail((width, int(width * img.height / img.width)))
buffer = BytesIO()
img.save(
buffer, format=image_format, quality=quality, metadata=img.info
)
buffer.seek(0)
return buffer
def process_channel(img: Image.Image, channel: str):
if channel == "rgb":
if img.mode == "RGBA":
r, g, b, _ = img.split()
img = Image.merge("RGB", (r, g, b))
else:
img = img.convert("RGB")
elif channel == "a":
if img.mode == "RGBA":
_, _, _, a = img.split()
else:
a = Image.new("L", img.size, 255)
img = Image.new("RGBA", img.size)
img.putalpha(a)
buffer = BytesIO()
img.save(buffer, format="PNG")
_ = buffer.seek(0)
return buffer
async def get_image_response(
file, filename: str, preview_info=None, channel=None
):
img = await asyncio.to_thread(
get_cached_image, file, preview_info, channel
)
return web.Response(
body=img,
content_type="image/webp" if preview_info else "image/png",
headers={"Content-Disposition": f'filename="{filename}"'},
)
# TODO: Embed the metadatas somehow so we can drag and drop
# to load workflows in the sidebar
@PromptServer.instance.routes.get("/mtb/view")
async def view_image(request: Request):
import folder_paths
filename = request.rel_url.query.get("filename")
if not filename:
return web.Response(status=404)
filename, output_dir = folder_paths.annotated_filepath(filename)
if filename[0] == "/" or ".." in filename:
return web.Response(status=400)
if output_dir is None:
rtype = request.rel_url.query.get("type", "output")
output_dir = folder_paths.get_directory_by_type(rtype)
if output_dir is None:
return web.Response(status=400)
if "subfolder" in request.rel_url.query:
full_output_dir = os.path.join(
output_dir, request.rel_url.query["subfolder"]
)
if (
os.path.commonpath(
(os.path.abspath(full_output_dir), output_dir)
)
!= output_dir
):
return web.Response(status=403)
output_dir = full_output_dir
filename = os.path.basename(filename)
file = os.path.join(output_dir, filename)
if not os.path.isfile(file):
return web.Response(status=404)
preview_info = None
if "preview" in request.rel_url.query:
preview_params = request.rel_url.query["preview"].split(";")
image_format = (
preview_params[0]
if preview_params[0] in ["webp", "jpeg"]
else "webp"
)
quality = (
int(preview_params[1])
if len(preview_params) > 1 and preview_params[1].isdigit()
else 90
)
width = request.rel_url.query.get("width")
preview_info = (image_format, quality, width)
channel = request.rel_url.query.get("channel")
return await get_image_response(file, filename, preview_info, channel)
@PromptServer.instance.routes.get("/mtb/server-info")
async def get_debug(request: Request):
from . import endpoint from . import endpoint
reload(endpoint) _ = reload(endpoint)
enabled = "MTB_DEBUG" in os.environ isdebug = "MTB_DEBUG" in os.environ
exposed = "MTB_EXPOSE" in os.environ
def render_property(name: str, val: str):
return f"""<strong>{name}:</strong>
<p>
{val}
</p>"""
# Check if the request prefers HTML content # Check if the request prefers HTML content
if "text/html" in request.headers.get("Accept", ""): if "text/html" in request.headers.get("Accept", ""):
# # Return an HTML page # # Return an HTML page
html_response = f""" html_response = ""
<h1>MTB Debug Status: {'Enabled' if enabled else 'Disabled'}</h1>
""" html_response += render_property(
"Debug", "Enabled" if isdebug else "Disabled"
)
html_response += render_property("Exposed", str(exposed))
return web.Response( return web.Response(
text=endpoint.render_base_template("Debug", html_response), text=endpoint.render_base_template(
"Server Info", html_response
),
content_type="text/html", content_type="text/html",
) )
# Return JSON for other requests # Return JSON for other requests
return web.json_response({"enabled": enabled}) return web.json_response({"exposed": exposed, "debug": isdebug})
@PromptServer.instance.routes.get("/mtb/actions") @PromptServer.instance.routes.get("/mtb/actions")
async def no_route(request): async def no_route(request: Request):
from . import endpoint from . import endpoint
if "text/html" in request.headers.get("Accept", ""): if "text/html" in request.headers.get("Accept", ""):
@@ -382,7 +542,7 @@ if hasattr(PromptServer, "instance"):
return web.json_response({"message": "actions has no get for now"}) return web.json_response({"message": "actions has no get for now"})
@PromptServer.instance.routes.post("/mtb/actions") @PromptServer.instance.routes.post("/mtb/actions")
async def do_action(request): async def do_action(request: Request):
from . import endpoint from . import endpoint
reload(endpoint) reload(endpoint)
+95 -19
View File
@@ -1,4 +1,8 @@
import csv import csv
import secrets
import sys
from pathlib import Path
from typing import Any
from aiohttp import web from aiohttp import web
@@ -6,6 +10,8 @@ from .log import mklog
from .utils import ( from .utils import (
backup_file, backup_file,
import_install, import_install,
input_dir,
output_dir,
reqs_map, reqs_map,
run_command, run_command,
styles_dir, styles_dir,
@@ -14,15 +20,14 @@ from .utils import (
endlog = mklog("mtb endpoint") endlog = mklog("mtb endpoint")
# - ACTIONS # - ACTIONS
import sys
from pathlib import Path
import_install("requirements") import_install("requirements")
def ACTIONS_installDependency(dependency_names=None): def ACTIONS_installDependency(dependency_names=None):
if dependency_names is None: if dependency_names is None:
# return web.Response(text="No dependency name provided", status=400)
return {"error": "No dependency name provided"} return {"error": "No dependency name provided"}
endlog.debug(f"Received Install Dependency request for {dependency_names}") endlog.debug(f"Received Install Dependency request for {dependency_names}")
# 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]
@@ -50,6 +55,62 @@ def ACTIONS_installDependency(dependency_names=None):
# break # break
def ACTIONS_getUserImages(
mode: str,
count=200,
offset=0,
sort: str | None = None,
include_subfolders: bool = False,
):
# enabled = "MTB_EXPOSE" in os.environ
# if not enabled:
# return {"error": "Session not authorized to getInputs"}
imgs = {}
entry_dir = input_dir if mode == "input" else output_dir
pattern = "**/*.png" if include_subfolders else "*.png"
entry_gen = entry_dir.glob(pattern)
entries = {}
if sort:
sort = sort.lower()
if sort == "none":
entries = entry_gen
elif sort == "modified":
entries = sorted(
entry_gen, key=lambda x: x.stat().st_mtime, reverse=True
)
elif sort == "modified-reverse":
entries = sorted(entry_gen, key=lambda x: x.stat().st_mtime)
elif sort == "name":
entries = sorted(entry_gen, key=lambda x: x.name)
elif sort == "name-reverse":
entries = sorted(entry_gen, key=lambda x: x.name, reverse=True)
else:
endlog.warning(f"Sort mode {sort} not supported")
entries = entry_gen
else:
entries = entry_gen
for i, img in enumerate(entries):
if i < offset:
continue
subfolder = (
img.parent.relative_to(entry_dir) if include_subfolders else ""
)
imgs[img.stem] = (
f"/mtb/view?filename={img.name}&width=512&type={mode}&subfolder="
f"{subfolder}"
f"&preview=&rand={secrets.randbelow(424242)}"
)
if i >= count + offset - 1:
break
return imgs
def ACTIONS_getStyles(style_name=None): def ACTIONS_getStyles(style_name=None):
from .nodes.conditions import MTB_StylesLoader from .nodes.conditions import MTB_StylesLoader
@@ -97,7 +158,7 @@ def ACTIONS_saveStyle(data):
csv_writer.writerow(row) csv_writer.writerow(row)
async def do_action(request) -> web.Response: async def do_action(request: web.Request) -> web.Response:
endlog.debug("Init action request") endlog.debug("Init action request")
request_data = await request.json() request_data = await request.json()
name = request_data.get("name") name = request_data.get("name")
@@ -109,7 +170,12 @@ async def do_action(request) -> web.Response:
method = globals().get(method_name) method = globals().get(method_name)
if callable(method): if callable(method):
result = method(args) if args else method() result = None
if args:
result = method(*args) if isinstance(args, list) else method(args)
else:
result = method()
endlog.debug(f"Action result: {result}") endlog.debug(f"Action result: {result}")
return web.json_response({"result": result}) return web.json_response({"result": result})
@@ -130,10 +196,13 @@ async def do_action(request) -> web.Response:
# - HTML UTILS # - HTML UTILS
def dependencies_button(name, dependencies): def dependencies_button(name: str, dependencies: list[str]) -> str:
deps = ",".join([f"'{x}'" for x in dependencies]) deps = ",".join([f"'{x}'" for x in dependencies])
return f""" return f"""
<button class="dependency-button" onclick="window.mtb_action('installDependency',[{deps}])">Install {name} deps</button> <button
class="dependency-button"
onclick="window.mtb_action('installDependency',[{deps}])"
>Install {name} deps</button>
""" """
@@ -153,7 +222,7 @@ def csv_editor():
html_out = """ html_out = """
<div id="style-editor"> <div id="style-editor">
<h1>Style Editor</h1> <h1>Style Editor</h1>
""" """
for current, styles in style_files.items(): for current, styles in style_files.items():
current_out = f"<h3>{current}</h3>" current_out = f"<h3>{current}</h3>"
@@ -215,11 +284,14 @@ def render_tab_view(**kwargs):
""" """
def add_foldable_region(title, content): def add_foldable_region(title: str, content: str):
symbol_id = f"{title}-symbol" symbol_id = f"{title}-symbol"
return f""" return f"""
<div class='foldable'> <div class='foldable'>
<div class='foldable-title' onclick="toggleFoldable('{title}', '{symbol_id}')"> <div
class='foldable-title'
onclick="toggleFoldable('{title}', '{symbol_id}')"
>
<span id='{symbol_id}' class='foldable-symbol'>&#9655;</span> <span id='{symbol_id}' class='foldable-symbol'>&#9655;</span>
{title} {title}
</div> </div>
@@ -231,7 +303,9 @@ def add_foldable_region(title, content):
""" """
def add_split_pane(left_content, right_content, vertical=True): def add_split_pane(
left_content: str, right_content: str, *, vertical: bool = True
):
orientation = "vertical" if vertical else "horizontal" orientation = "vertical" if vertical else "horizontal"
return f""" return f"""
<div class="split-pane {orientation}"> <div class="split-pane {orientation}">
@@ -250,7 +324,7 @@ def add_split_pane(left_content, right_content, vertical=True):
""" """
def add_dropdown(title, options): def add_dropdown(title: str, options: list[str]):
option_str = "\n".join( option_str = "\n".join(
[f"<option value='{opt}'>{opt}</option>" for opt in options] [f"<option value='{opt}'>{opt}</option>" for opt in options]
) )
@@ -262,13 +336,13 @@ def add_dropdown(title, options):
""" """
def render_table(table_dict, sort=True, title=None): def render_table(table_dict: dict[str, Any], sort=True, title=None):
table_dict = sorted( table_list = sorted(
table_dict.items(), key=lambda item: item[0] table_dict.items(), key=lambda item: item[0]
) # Sort the dictionary by keys ) # Sort the dictionary by keys
table_rows = "" table_rows = ""
for name, item in table_dict: for name, item in table_list:
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>"
@@ -299,12 +373,12 @@ def render_table(table_dict, sort=True, title=None):
<tbody> <tbody>
{table_rows} {table_rows}
</tbody> </tbody>
</table> </table>
</div> </div>
""" """
def render_base_template(title, content): def render_base_template(title: str, content: str):
github_icon_svg = """<svg xmlns="http://www.w3.org/2000/svg" fill="whitesmoke" height="3em" viewBox="0 0 496 512"><path d="M165.9 397.4c0 2-2.3 3.6-5.2 3.6-3.3.3-5.6-1.3-5.6-3.6 0-2 2.3-3.6 5.2-3.6 3-.3 5.6 1.3 5.6 3.6zm-31.1-4.5c-.7 2 1.3 4.3 4.3 4.9 2.6 1 5.6 0 6.2-2s-1.3-4.3-4.3-5.2c-2.6-.7-5.5.3-6.2 2.3zm44.2-1.7c-2.9.7-4.9 2.6-4.6 4.9.3 2 2.9 3.3 5.9 2.6 2.9-.7 4.9-2.6 4.6-4.6-.3-1.9-3-3.2-5.9-2.9zM244.8 8C106.1 8 0 113.3 0 252c0 110.9 69.8 205.8 169.5 239.2 12.8 2.3 17.3-5.6 17.3-12.1 0-6.2-.3-40.4-.3-61.4 0 0-70 15-84.7-29.8 0 0-11.4-29.1-27.8-36.6 0 0-22.9-15.7 1.6-15.4 0 0 24.9 2 38.6 25.8 21.9 38.6 58.6 27.5 72.9 20.9 2.3-16 8.8-27.1 16-33.7-55.9-6.2-112.3-14.3-112.3-110.5 0-27.5 7.6-41.3 23.6-58.9-2.6-6.5-11.1-33.3 2.6-67.9 20.9-6.5 69 27 69 27 20-5.6 41.5-8.5 62.8-8.5s42.8 2.9 62.8 8.5c0 0 48.1-33.6 69-27 13.7 34.7 5.2 61.4 2.6 67.9 16 17.7 25.8 31.5 25.8 58.9 0 96.5-58.9 104.2-114.8 110.5 9.2 7.9 17 22.9 17 46.4 0 33.7-.3 75.4-.3 83.6 0 6.5 4.6 14.4 17.3 12.1C428.2 457.8 496 362.9 496 252 496 113.3 383.5 8 244.8 8zM97.2 352.9c-1.3 1-1 3.3.7 5.2 1.6 1.6 3.9 2.3 5.2 1 1.3-1 1-3.3-.7-5.2-1.6-1.6-3.9-2.3-5.2-1zm-10.8-8.1c-.7 1.3.3 2.9 2.3 3.9 1.6 1 3.6.7 4.3-.7.7-1.3-.3-2.9-2.3-3.9-2-.6-3.6-.3-4.3.7zm32.4 35.6c-1.6 1.3-1 4.3 1.3 6.2 2.3 2.3 5.2 2.6 6.5 1 1.3-1.3.7-4.3-1.3-6.2-2.2-2.3-5.2-2.6-6.5-1zm-11.4-14.7c-1.6 1-1.6 3.6 0 5.9 1.6 2.3 4.3 3.3 5.6 2.3 1.6-1.3 1.6-3.9 0-6.2-1.4-2.3-4-3.3-5.6-2z"/></svg>""" github_icon_svg = """<svg xmlns="http://www.w3.org/2000/svg" fill="whitesmoke" height="3em" viewBox="0 0 496 512"><path d="M165.9 397.4c0 2-2.3 3.6-5.2 3.6-3.3.3-5.6-1.3-5.6-3.6 0-2 2.3-3.6 5.2-3.6 3-.3 5.6 1.3 5.6 3.6zm-31.1-4.5c-.7 2 1.3 4.3 4.3 4.9 2.6 1 5.6 0 6.2-2s-1.3-4.3-4.3-5.2c-2.6-.7-5.5.3-6.2 2.3zm44.2-1.7c-2.9.7-4.9 2.6-4.6 4.9.3 2 2.9 3.3 5.9 2.6 2.9-.7 4.9-2.6 4.6-4.6-.3-1.9-3-3.2-5.9-2.9zM244.8 8C106.1 8 0 113.3 0 252c0 110.9 69.8 205.8 169.5 239.2 12.8 2.3 17.3-5.6 17.3-12.1 0-6.2-.3-40.4-.3-61.4 0 0-70 15-84.7-29.8 0 0-11.4-29.1-27.8-36.6 0 0-22.9-15.7 1.6-15.4 0 0 24.9 2 38.6 25.8 21.9 38.6 58.6 27.5 72.9 20.9 2.3-16 8.8-27.1 16-33.7-55.9-6.2-112.3-14.3-112.3-110.5 0-27.5 7.6-41.3 23.6-58.9-2.6-6.5-11.1-33.3 2.6-67.9 20.9-6.5 69 27 69 27 20-5.6 41.5-8.5 62.8-8.5s42.8 2.9 62.8 8.5c0 0 48.1-33.6 69-27 13.7 34.7 5.2 61.4 2.6 67.9 16 17.7 25.8 31.5 25.8 58.9 0 96.5-58.9 104.2-114.8 110.5 9.2 7.9 17 22.9 17 46.4 0 33.7-.3 75.4-.3 83.6 0 6.5 4.6 14.4 17.3 12.1C428.2 457.8 496 362.9 496 252 496 113.3 383.5 8 244.8 8zM97.2 352.9c-1.3 1-1 3.3.7 5.2 1.6 1.6 3.9 2.3 5.2 1 1.3-1 1-3.3-.7-5.2-1.6-1.6-3.9-2.3-5.2-1zm-10.8-8.1c-.7 1.3.3 2.9 2.3 3.9 1.6 1 3.6.7 4.3-.7.7-1.3-.3-2.9-2.3-3.9-2-.6-3.6-.3-4.3.7zm32.4 35.6c-1.6 1.3-1 4.3 1.3 6.2 2.3 2.3 5.2 2.6 6.5 1 1.3-1.3.7-4.3-1.3-6.2-2.2-2.3-5.2-2.6-6.5-1zm-11.4-14.7c-1.6 1-1.6 3.6 0 5.9 1.6 2.3 4.3 3.3 5.6 2.3 1.6-1.3 1.6-3.9 0-6.2-1.4-2.3-4-3.3-5.6-2z"/></svg>"""
return f""" return f"""
<!DOCTYPE html> <!DOCTYPE html>
@@ -340,7 +414,9 @@ def render_base_template(title, content):
<header> <header>
<a href="/">Back to Comfy</a> <a href="/">Back to Comfy</a>
<div class="mtb_logo"> <div class="mtb_logo">
<img src="https://repository-images.githubusercontent.com/649047066/a3eef9a7-20dd-4ef9-b839-884502d4e873" alt="Comfy MTB Logo" height="70" width="128"> <img
src="https://repository-images.githubusercontent.com/649047066/a3eef9a7-20dd-4ef9-b839-884502d4e873"
alt="Comfy MTB Logo" height="70" width="128">
<span class="title">Comfy MTB</span></div> <span class="title">Comfy MTB</span></div>
<a style="width:128px;text-align:center" href="https://www.github.com/melmass/comfy_mtb"> <a style="width:128px;text-align:center" href="https://www.github.com/melmass/comfy_mtb">
{github_icon_svg} {github_icon_svg}
@@ -355,6 +431,6 @@ def render_base_template(title, content):
<!-- Shared footer content here --> <!-- Shared footer content here -->
</footer> </footer>
</body> </body>
</html> </html>
""" """
+11 -5
View File
@@ -27,7 +27,7 @@ export def "comfy start" [--clean,--old-ui, --listen] {
let root = get_root --clean=($clean) let root = get_root --clean=($clean)
cd $root cd $root
MTB_DEBUG=true python main.py --port 3000 ...(if $old_ui { ["--front-end-version", "Comfy-Org/ComfyUI_legacy_frontend@latest"]} else {[]}) --preview-method auto ...(if $listen {["--listen"]} else {[]}) MTB_DEBUG=true python main.py --port 3000 ...(if $old_ui { ["--front-end-version", "Comfy-Org/ComfyUI_legacy_frontend@latest"]} else {[ --front-end-version Comfy-Org/ComfyUI_frontend@latest]}) --preview-method auto ...(if $listen {["--listen"]} else {[]})
} }
# update comfy itself and merge master in current branch # update comfy itself and merge master in current branch
@@ -67,8 +67,14 @@ export def "comfy update" [
git checkout master git checkout master
print $"(ansi yellow_italic)Fetching and pulling remote updates(ansi reset)" print $"(ansi yellow_italic)Fetching and pulling remote updates(ansi reset)"
git fetch if ($clean) {
git pull git fetch local master
git pull local master
} else {
git fetch
git pull
}
print $"(ansi yellow_italic)Back to our branch \(($branch_name)\)(ansi reset)" print $"(ansi yellow_italic)Back to our branch \(($branch_name)\)(ansi reset)"
git checkout - git checkout -
@@ -135,7 +141,7 @@ export def "comfy update_extensions" [--clean] {
let root = get_root --clean=($clean) let root = get_root --clean=($clean)
cd $root cd $root
cd custom_nodes cd custom_nodes
git multipull . git multipull . -s -q
} }
def --env path-add [pth] { def --env path-add [pth] {
@@ -146,7 +152,7 @@ def --env path-add [pth] {
export-env { export-env {
$env.COMFY_MTB = ("." | path expand) $env.COMFY_MTB = ("." | path expand)
$env.CUDA_ROOT = 'C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.1\' # $env.CUDA_ROOT = 'C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.1\'
$env.CUDA_HOME = $env.CUDA_ROOT $env.CUDA_HOME = $env.CUDA_ROOT
+123 -1
View File
@@ -3,10 +3,127 @@ import shutil
from pathlib import Path from pathlib import Path
import folder_paths import folder_paths
import torch
from ..log import log from ..log import log
from ..utils import here from ..utils import here
Conditioning = list[tuple[torch.Tensor, dict[str, torch.Tensor]]]
def check_condition(conditioning: Conditioning):
has_cn = False
if len(conditioning) > 1:
log.warn(
"More than one conditioning was provided. Only the first one will be used."
)
first = conditioning[0]
cond, kwargs = first
log.debug("Conditioning Shape")
log.debug(cond.shape)
log.debug("Conditioning keys")
log.debug([f"\t{k} - {type(kwargs[k])}" for k in kwargs])
if "control" in kwargs:
log.debug("Conditioning contains a controlnet")
has_cn = True
if "pooled_output" not in kwargs:
raise ValueError(
"Conditioning is not valid. Missing 'pooled_output' key."
)
return has_cn
class MTB_InterpolateCondition:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"blend": (
"FLOAT",
{"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01},
),
},
}
RETURN_TYPES = ("CONDITIONING",)
CATEGORY = "mtb/conditioning"
FUNCTION = "execute"
def execute(
self, blend: float, **kwargs: Conditioning
) -> tuple[Conditioning]:
blend = max(0.0, min(1.0, blend))
conditions: list[Conditioning] = list(kwargs.values())
num_conditions = len(conditions)
if num_conditions < 2:
raise ValueError("At least two conditioning inputs are required.")
segment_length = 1.0 / (num_conditions - 1)
segment_index = min(int(blend // segment_length), num_conditions - 2)
local_blend = (
blend - (segment_index * segment_length)
) / segment_length
cond_from = conditions[segment_index]
cond_to = conditions[segment_index + 1]
from_cn = check_condition(cond_from)
to_cn = check_condition(cond_to)
if from_cn and to_cn:
raise ValueError(
"Interpolating conditions cannot both contain ControlNets"
)
try:
interpolated_condition = [
(1.0 - local_blend) * c_from + local_blend * c_to
for c_from, c_to in zip(
cond_from[0][0], cond_to[0][0], strict=False
)
]
except Exception as e:
print(f"Error during interpolation: {e}")
raise
pooled_from = cond_from[0][1].get(
"pooled_output",
torch.zeros_like(
next(iter(cond_from[0][1].values()), torch.tensor([]))
),
)
pooled_to = cond_to[0][1].get(
"pooled_output",
torch.zeros_like(
next(iter(cond_from[0][1].values()), torch.tensor([]))
),
)
interpolated_pooled = (
1.0 - local_blend
) * pooled_from + local_blend * pooled_to
res = {"pooled_output": interpolated_pooled}
if from_cn:
res["control"] = cond_from[0][1]["control"]
res["control_apply_to_uncond"] = cond_from[0][1][
"control_apply_to_uncond"
]
if to_cn:
res["control"] = cond_to[0][1]["control"]
res["control_apply_to_uncond"] = cond_to[0][1][
"control_apply_to_uncond"
]
return ([(torch.stack(interpolated_condition), res)],)
class MTB_InterpolateClipSequential: class MTB_InterpolateClipSequential:
@classmethod @classmethod
@@ -213,4 +330,9 @@ class MTB_StylesLoader:
return (self.options[style_name][0], self.options[style_name][1]) return (self.options[style_name][0], self.options[style_name][1])
__nodes__ = [MTB_SmartStep, MTB_StylesLoader, MTB_InterpolateClipSequential] __nodes__ = [
MTB_SmartStep,
MTB_StylesLoader,
MTB_InterpolateClipSequential,
MTB_InterpolateCondition,
]
+38 -1
View File
@@ -59,6 +59,36 @@ class MTB_SplitBbox:
return (bbox[0], bbox[1], bbox[2], bbox[3]) return (bbox[0], bbox[1], bbox[2], bbox[3])
class MTB_UpscaleBboxBy:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"bbox": ("BBOX",),
"scale": ("FLOAT", {"default": 1.0}),
},
}
CATEGORY = "mtb/crop"
RETURN_TYPES = ("BBOX",)
FUNCTION = "upscale"
def upscale(
self, bbox: tuple[int, int, int, int], scale: float
) -> tuple[tuple[int, int, int, int]]:
x, y, width, height = bbox
# scaled = (x * scale, y * scale, width * scale, height * scale)
scaled = (
int(x * scale),
int(y * scale),
int(width * scale),
int(height * scale),
)
return (scaled,)
class MTB_BboxFromMask: class MTB_BboxFromMask:
"""From a mask extract the bounding box""" """From a mask extract the bounding box"""
@@ -342,4 +372,11 @@ class MTB_Uncrop:
return (pil2tensor(out_images),) return (pil2tensor(out_images),)
__nodes__ = [MTB_BboxFromMask, MTB_Bbox, MTB_Crop, MTB_Uncrop, MTB_SplitBbox] __nodes__ = [
MTB_BboxFromMask,
MTB_Bbox,
MTB_Crop,
MTB_Uncrop,
MTB_SplitBbox,
MTB_UpscaleBboxBy,
]
+2
View File
@@ -78,6 +78,7 @@ class MTB_LoadFaceEnhanceModel:
RETURN_NAMES = ("model",) RETURN_NAMES = ("model",)
FUNCTION = "load_model" FUNCTION = "load_model"
CATEGORY = "mtb/facetools" CATEGORY = "mtb/facetools"
DEPRECATED = True
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 from gfpgan import GFPGANer
@@ -163,6 +164,7 @@ class MTB_RestoreFace:
RETURN_TYPES = ("IMAGE",) RETURN_TYPES = ("IMAGE",)
FUNCTION = "restore" FUNCTION = "restore"
CATEGORY = "mtb/facetools" CATEGORY = "mtb/facetools"
DEPRECATED = True
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
+3
View File
@@ -40,6 +40,7 @@ class MTB_LoadFaceAnalysisModel:
RETURN_TYPES = ("FACE_ANALYSIS_MODEL",) RETURN_TYPES = ("FACE_ANALYSIS_MODEL",)
FUNCTION = "load_model" FUNCTION = "load_model"
CATEGORY = "mtb/facetools" CATEGORY = "mtb/facetools"
DEPRECATED = True
def load_model(self, faceswap_model: str): def load_model(self, faceswap_model: str):
if faceswap_model == "antelopev2": if faceswap_model == "antelopev2":
@@ -77,6 +78,7 @@ class MTB_LoadFaceSwapModel:
RETURN_TYPES = ("FACESWAP_MODEL",) RETURN_TYPES = ("FACESWAP_MODEL",)
FUNCTION = "load_model" FUNCTION = "load_model"
CATEGORY = "mtb/facetools" CATEGORY = "mtb/facetools"
DEPRECATED = True
def load_model(self, faceswap_model: str): def load_model(self, faceswap_model: str):
model_path = get_model_path("insightface", faceswap_model) model_path = get_model_path("insightface", faceswap_model)
@@ -126,6 +128,7 @@ class MTB_FaceSwap:
RETURN_TYPES = ("IMAGE",) RETURN_TYPES = ("IMAGE",)
FUNCTION = "swap" FUNCTION = "swap"
CATEGORY = "mtb/facetools" CATEGORY = "mtb/facetools"
DEPRECATED = True
def swap( def swap(
self, self,
+11 -4
View File
@@ -1,5 +1,4 @@
from pathlib import Path from pathlib import Path
from typing import List
import comfy import comfy
import comfy.model_management as model_management import comfy.model_management as model_management
@@ -15,10 +14,13 @@ from ..utils import get_model_path
class MTB_LoadFilmModel: class MTB_LoadFilmModel:
"""Loads a FILM model""" """Loads a FILM model
[DEPRECATED] Use ComfyUI-FrameInterpolation instead
"""
@staticmethod @staticmethod
def get_models() -> List[Path]: def get_models() -> list[Path]:
models_paths = get_model_path("FILM").iterdir() models_paths = get_model_path("FILM").iterdir()
return [x for x in models_paths if x.suffix in [".onnx", ".pth"]] return [x for x in models_paths if x.suffix in [".onnx", ".pth"]]
@@ -37,6 +39,7 @@ class MTB_LoadFilmModel:
RETURN_TYPES = ("FILM_MODEL",) RETURN_TYPES = ("FILM_MODEL",)
FUNCTION = "load_model" FUNCTION = "load_model"
CATEGORY = "mtb/frame iterpolation" CATEGORY = "mtb/frame iterpolation"
DEPRECATED = True
def load_model(self, film_model: str): def load_model(self, film_model: str):
model_path = get_model_path("FILM", film_model) model_path = get_model_path("FILM", film_model)
@@ -56,7 +59,10 @@ class MTB_LoadFilmModel:
class MTB_FilmInterpolation: class MTB_FilmInterpolation:
"""Google Research FILM frame interpolation for large motion""" """Google Research FILM frame interpolation for large motion
[DEPRECATED] Use ComfyUI-FrameInterpolation instead
"""
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
@@ -71,6 +77,7 @@ class MTB_FilmInterpolation:
RETURN_TYPES = ("IMAGE",) RETURN_TYPES = ("IMAGE",)
FUNCTION = "do_interpolation" FUNCTION = "do_interpolation"
CATEGORY = "mtb/frame iterpolation" CATEGORY = "mtb/frame iterpolation"
DEPRECATED = True
def do_interpolation( def do_interpolation(
self, self,
+55 -19
View File
@@ -2,9 +2,9 @@ import json
import subprocess import subprocess
import uuid import uuid
from pathlib import Path from pathlib import Path
from typing import List, Optional
import comfy.model_management as model_management import comfy.model_management as model_management
import comfy.utils
import folder_paths import folder_paths
import numpy as np import numpy as np
import torch import torch
@@ -41,6 +41,7 @@ class MTB_ReadPlaylist:
RETURN_TYPES = ("PLAYLIST",) RETURN_TYPES = ("PLAYLIST",)
FUNCTION = "read_playlist" FUNCTION = "read_playlist"
CATEGORY = "mtb/IO" CATEGORY = "mtb/IO"
EXPERIMENTAL = True
def read_playlist( def read_playlist(
self, self,
@@ -83,6 +84,7 @@ class MTB_AddToPlaylist:
OUTPUT_NODE = True OUTPUT_NODE = True
FUNCTION = "add_to_playlist" FUNCTION = "add_to_playlist"
CATEGORY = "mtb/IO" CATEGORY = "mtb/IO"
EXPERIMENTAL = True
def add_to_playlist( def add_to_playlist(
self, self,
@@ -117,7 +119,10 @@ class MTB_AddToPlaylist:
class MTB_ExportWithFfmpeg: class MTB_ExportWithFfmpeg:
"""Export with FFmpeg (Experimental)""" """Export with FFmpeg (Experimental).
[DEPRACATED] Use VHS nodes instead
"""
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
@@ -143,6 +148,7 @@ class MTB_ExportWithFfmpeg:
RETURN_TYPES = ("VIDEO",) RETURN_TYPES = ("VIDEO",)
OUTPUT_NODE = True OUTPUT_NODE = True
FUNCTION = "export_prores" FUNCTION = "export_prores"
DEPRECATED = True
CATEGORY = "mtb/IO" CATEGORY = "mtb/IO"
def export_prores( def export_prores(
@@ -151,10 +157,9 @@ class MTB_ExportWithFfmpeg:
prefix: str, prefix: str,
format: str, format: str,
codec: str, codec: str,
images: Optional[torch.Tensor] = None, images: torch.Tensor | None = None,
playlist: Optional[List[str]] = None, playlist: list[str] | None = None,
): ):
pix_fmt = "rgb48le" if codec == "prores_ks" else "yuv420p"
file_ext = format file_ext = format
file_id = f"{prefix}_{uuid.uuid4()}.{file_ext}" file_id = f"{prefix}_{uuid.uuid4()}.{file_ext}"
@@ -208,9 +213,11 @@ class MTB_ExportWithFfmpeg:
frames = tensor2np(images) frames = tensor2np(images)
log.debug(f"Frames type {type(frames[0])}") log.debug(f"Frames type {type(frames[0])}")
log.debug(f"Exporting {len(frames)} frames") log.debug(f"Exporting {len(frames)} frames")
height, width, channels = frames[0].shape
has_alpha = channels == 4
out_path = (output_dir / file_id).as_posix()
if codec == "gif": if codec == "gif":
out_path = (output_dir / file_id).as_posix()
command = [ command = [
"ffmpeg", "ffmpeg",
"-f", "-f",
@@ -233,12 +240,28 @@ class MTB_ExportWithFfmpeg:
process.stdin.close() process.stdin.close()
process.wait() process.wait()
return (out_path,)
else: else:
frames = [frame.astype(np.uint16) * 257 for frame in frames] if has_alpha:
if codec in ["prores_ks", "libx264", "libx265"]:
height, width, _ = frames[0].shape pix_fmt = (
"yuva444p" if codec == "prores_ks" else "yuva420p"
out_path = (output_dir / file_id).as_posix() )
frames = [
frame.astype(np.uint16) * 257 for frame in frames
]
else:
log.warning(
f"Alpha channel not supported for codec {codec}. Alpha will be ignored."
)
frames = [
frame[:, :, :3].astype(np.uint16) * 257
for frame in frames
]
pix_fmt = "rgb48le" if codec == "prores_ks" else "yuv420p"
else:
pix_fmt = "rgb48le" if codec == "prores_ks" else "yuv420p"
frames = [frame.astype(np.uint16) * 257 for frame in frames]
# Prepare the FFmpeg command # Prepare the FFmpeg command
command = [ command = [
@@ -258,17 +281,26 @@ class MTB_ExportWithFfmpeg:
"-", "-",
"-c:v", "-c:v",
codec, codec,
"-r",
str(fps),
"-y",
out_path,
] ]
if codec == "prores_ks":
command.extend(["-profile:v", "4444"])
command.extend(
[
"-r",
str(fps),
"-y",
out_path,
]
)
process = subprocess.Popen(command, stdin=subprocess.PIPE) process = subprocess.Popen(command, stdin=subprocess.PIPE)
pbar = comfy.utils.ProgressBar(len(frames))
for frame in frames: for frame in frames:
model_management.throw_exception_if_processing_interrupted()
process.stdin.write(frame.tobytes()) process.stdin.write(frame.tobytes())
pbar.update(1)
process.stdin.close() process.stdin.close()
process.wait() process.wait()
@@ -280,9 +312,9 @@ def prepare_animated_batch(
batch: torch.Tensor, batch: torch.Tensor,
pingpong=False, pingpong=False,
resize_by=1.0, resize_by=1.0,
resample_filter: Optional[Image.Resampling] = None, resample_filter: Image.Resampling | None = None,
image_type=np.uint8, image_type=np.uint8,
) -> List[Image.Image]: ) -> list[Image.Image]:
images = tensor2np(batch) images = tensor2np(batch)
images = [frame.astype(image_type) for frame in images] images = [frame.astype(image_type) for frame in images]
@@ -308,7 +340,10 @@ def prepare_animated_batch(
# todo: deprecate for apng # todo: deprecate for apng
class MTB_SaveGif: class MTB_SaveGif:
"""Save the images from the batch as a GIF""" """Save the images from the batch as a GIF.
[DEPRACATED] Use VHS nodes instead
"""
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
@@ -328,6 +363,7 @@ class MTB_SaveGif:
OUTPUT_NODE = True OUTPUT_NODE = True
CATEGORY = "mtb/IO" CATEGORY = "mtb/IO"
FUNCTION = "save_gif" FUNCTION = "save_gif"
DEPRECATED = True
def save_gif( def save_gif(
self, self,
+157
View File
@@ -0,0 +1,157 @@
import os
import subprocess
import tempfile
import numpy as np
import torch
from PIL import Image
from ..log import log
class ImageH264Compression:
"""Encodes the input with h264 compression using a configurable CRF."""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": (
"IMAGE",
{
"tooltip": "The input image tensor to be compressed and decompressed."
},
),
"crf": (
"INT",
{
"default": 23,
"min": 0,
"max": 51,
"step": 1,
"tooltip": "Constant Rate Factor for h264 encoding (lower values mean higher quality).",
},
),
}
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "compress_and_decompress"
CATEGORY = "image"
DESCRIPTION = """
**Encodes the input with h264 compression using a configurable CRF**.
> [!NOTE]
> This was recommended by the creators of LTX over banodoco's discord.
*Orginal code from [mix](https://github.com/XmYx)*"""
def _compress_decompress_ffmpeg(self, img_array, crf):
with tempfile.TemporaryDirectory() as temp_dir:
input_path = os.path.join(temp_dir, "input.png")
output_path = os.path.join(temp_dir, "output.mp4")
decoded_path = os.path.join(temp_dir, "decoded.png")
Image.fromarray(img_array).save(input_path)
encode_command = [
"ffmpeg",
"-y",
"-i",
input_path,
"-c:v",
"libx264",
"-crf",
str(crf),
"-pix_fmt",
"yuv420p",
"-frames:v",
"1",
output_path,
]
subprocess.run(encode_command, capture_output=True)
decode_command = [
"ffmpeg",
"-y",
"-i",
output_path,
"-frames:v",
"1",
decoded_path,
]
subprocess.run(decode_command, capture_output=True)
decoded_img = np.array(Image.open(decoded_path))
return decoded_img
def compress_and_decompress(self, image, crf):
import io
output_images = []
try:
import av
for img_tensor in image:
img_array = img_tensor.cpu().numpy()
img_array = (img_array * 255).astype(np.uint8)
img_array = img_array.copy(
order="C"
) # Ensure contiguous array
output = io.BytesIO()
# Encode the image to h264 with the given CRF
container = av.open(output, mode="w", format="mp4")
stream = container.add_stream("h264", rate=1)
stream.width = img_array.shape[1]
stream.height = img_array.shape[0]
stream.pix_fmt = "yuv420p"
stream.options = {"crf": str(crf)}
frame = av.VideoFrame.from_ndarray(img_array, format="rgb24")
for packet in stream.encode(frame):
container.mux(packet)
for packet in stream.encode():
container.mux(packet)
container.close()
# Decode the video back to an image
output.seek(0)
container = av.open(output, mode="r", format="mp4")
decoded_frames = []
for frame in container.decode(video=0):
img_decoded = frame.to_ndarray(format="rgb24")
decoded_frames.append(img_decoded)
container.close()
if len(decoded_frames) > 0:
img_decoded = decoded_frames[0]
img_decoded = torch.from_numpy(
img_decoded.astype(np.float32) / 255.0
)
output_images.append(img_decoded)
else:
# If decoding failed, use the original image
output_images.append(img_tensor)
except ImportError:
log.warning(
"PyAv is not installed... Falling back to the ffmpeg cli"
)
for img_tensor in image:
img_array = (img_tensor.cpu().numpy() * 255).astype(np.uint8)
decoded_img = self._compress_decompress_ffmpeg(img_array, crf)
img_decoded = torch.from_numpy(
decoded_img.astype(np.float32) / 255.0
)
output_images.append(img_decoded)
output_images = torch.stack(output_images).to(image.device)
return (output_images,)
# fmt: off
__nodes__ = [
ImageH264Compression
]
+2 -1
View File
@@ -1,6 +1,5 @@
import comfy.utils import comfy.utils
from PIL import Image from PIL import Image
from rembg import remove
from ..utils import pil2tensor, tensor2pil from ..utils import pil2tensor, tensor2pil
@@ -64,6 +63,8 @@ class MTB_ImageRemoveBackgroundRembg:
post_process_mask, post_process_mask,
bgcolor, bgcolor,
): ):
from rembg import remove
pbar = comfy.utils.ProgressBar(image.size(0)) pbar = comfy.utils.ProgressBar(image.size(0))
images = tensor2pil(image) images = tensor2pil(image)
+351
View File
@@ -0,0 +1,351 @@
import os
import subprocess
import tempfile
import comfy.utils
import torch
from ..log import log
from ..utils import nextAvailable, tensor2pil
RELATIVE_NOTICE = """
Absolute paths are kept as is, relatives are from the output directory.
"""
class MTB_PostshotTrain:
CATEGORY = "mtb/postshot"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"images": (
"IMAGE",
{"tooltip": "These image will get save to disk first"},
),
"profile": (
[
"NeRF L",
"NeRF M",
"NeRF S",
"NeRF XL",
"NeRF XXL",
"Splat ADC",
"Splat MCMC",
],
{
"default": "Splat MCMC",
"tooltip": "The radiance field model profile to train",
},
),
"image_select": (
["all", "best"],
{
"default": "best",
"tooltip": "How to select training images from the source image sets",
},
),
"train_steps_limit": (
"INT",
{
"default": 30,
"min": 1,
"max": 1000,
"tooltip": "Number of kSteps to train the model for",
},
),
"output_path": (
"STRING",
{
"default": "output",
"tooltip": (
"path to save the project to" f"{RELATIVE_NOTICE}"
),
},
),
"postshot_cli": (
"STRING",
{
"default": "C:/Program Files/Jawset Postshot/bin/postshot-cli.exe"
},
),
},
"optional": {
"gpu": (
"INT",
{
"default": 0,
"min": 0,
"max": 255,
"tooltip": "Specify the index of the GPU to use",
},
),
"num_train_images": (
"INT",
{
"default": 0,
"min": 0,
"tooltip": "If image-select best is used, specifies the number of training images to select",
},
),
"max_image_size": (
"INT",
{
"default": 1600,
"min": 0,
"tooltip": "Downscale training images such that their longer edge is at most this value in pixels. Disabled if zero.",
},
),
"max_num_features": (
"INT",
{
"default": 8,
"min": 1,
"tooltip": "Maximum number of 2D kFeatures extracted from each image.",
},
),
"splat_density": (
"FLOAT",
{
"default": 1.0,
"min": 0.125,
"max": 8.0,
"tooltip": (
"Controls how much additional splats "
"are generated during training."
"Applies only in 'Splat ADC' profile."
),
},
),
"max_num_splats": (
"INT",
{
"default": 3000,
"min": 1,
"tooltip": (
"Sets the maximum number of splats (in kSplats)"
" created during training. "
"Applies only in 'Splat MCMC' profile."
),
},
),
"export_splat_ply": (
"STRING",
{
"default": "",
"tooltip": (
"If not empty will also save a ply file."
f"{RELATIVE_NOTICE}"
),
},
),
},
}
RETURN_TYPES = ("STRING",)
OUTPUT_NODE = True
RETURN_NAMES = ("project_file_path",)
FUNCTION = "train_model"
def train_model(
self,
images: torch.Tensor,
profile: str,
image_select: str,
train_steps_limit: int,
output_path: str,
gpu=0,
num_train_images=0,
max_image_size=1600,
max_num_features=8,
splat_density=1.0,
max_num_splats=3000,
export_splat_ply="",
postshot_cli="",
):
if not output_path.endswith(".psht"):
output_path += ".psht"
output_path = nextAvailable(output_path)
output_path.parent.mkdir(exist_ok=True)
pbar = comfy.utils.ProgressBar(200 + images.size(0))
try:
with tempfile.TemporaryDirectory() as temp_dir:
image_paths = []
images_pil = tensor2pil(images)
for i, img in enumerate(images_pil):
try:
img_path = os.path.join(temp_dir, f"image_{i:04d}.png")
img.save(img_path)
image_paths.append(img_path)
except Exception as e:
raise RuntimeError(
f"Failed to save image {i}: {str(e)}"
) from e
pbar.update(1)
if not image_paths:
raise ValueError("No valid images to process")
cmd = [postshot_cli, "train"]
for img_path in image_paths:
cmd.extend(["-i", img_path])
cmd.extend(
[
"-p",
profile,
"--image-select",
image_select,
"-s",
str(train_steps_limit),
"-o",
output_path.as_posix(),
]
)
if gpu is not None:
cmd.extend(["--gpu", str(gpu)])
if num_train_images > 0 and image_select == "best":
cmd.extend(["--num-train-images", str(num_train_images)])
if max_image_size > 0:
cmd.extend(["--max-image-size", str(max_image_size)])
if max_num_features != 8:
cmd.extend(["--max-num-features", str(max_num_features)])
if profile == "Splat ADC" and splat_density != 1.0:
cmd.extend(["--splat-density", str(splat_density)])
if profile == "Splat MCMC" and max_num_splats != 3000:
cmd.extend(["--max-num-splats", str(max_num_splats)])
if export_splat_ply:
export_splat_ply = nextAvailable(export_splat_ply)
cmd.extend(
["--export-splat-ply", export_splat_ply.as_posix()]
)
log.debug(f"Running {cmd}")
process = subprocess.Popen(
cmd,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
universal_newlines=True,
)
last_step_c = 0
last_step_t = 0
while True:
output = process.stdout.readline()
if output == "" and process.poll() is not None:
break
if output:
print(output)
if "camera tracking step" in output.lower():
try:
current_step = int(
output.split("%")[0].split(":")[1].strip()
)
if current_step > last_step_c:
pbar.update(1)
last_step_c = current_step
except (ValueError, IndexError):
continue
if "training radiance field:" in output.lower():
try:
current_step = int(
output.split("%")[0].split(":")[1].strip()
)
if current_step > last_step_t:
pbar.update(1)
last_step_t = current_step
except (ValueError, IndexError):
continue
if process.returncode != 0:
_, stderr = process.communicate()
raise RuntimeError(f"Postshot training failed: {stderr}")
if not os.path.exists(output_path):
raise RuntimeError("Output file was not created")
return (output_path.as_posix(),)
except Exception as e:
raise RuntimeError(f"Training failed: {str(e)}")
finally:
pbar.update(train_steps_limit)
class MTB_PostshotExport:
CATEGORY = "mtb/postshot"
OUTPUT_NODE = True
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"project_file": (
"STRING",
{"default": "", "forceInput": True},
),
"export_splat_ply": ("STRING", {"default": "output.ply"}),
"postshot_cli": (
"STRING",
{
"default": "C:/Program Files/Jawset Postshot/bin/postshot-cli.exe"
},
),
},
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("exported_ply_path",)
FUNCTION = "export_model"
def export_model(
self, project_file: str, export_splat_ply: str, postshot_cli: str
):
if not project_file.endswith(".psht"):
raise ValueError("Project file must have .psht extension")
if not os.path.exists(project_file):
raise FileNotFoundError(f"Project file not found: {project_file}")
if not export_splat_ply.endswith(".ply"):
export_splat_ply += ".ply"
_export_splat_ply = nextAvailable(export_splat_ply)
_export_splat_ply.parent.mkdir(exist_ok=True)
cmd = [
postshot_cli,
"export",
"-f",
project_file,
"--export-splat-ply",
_export_splat_ply.as_posix(),
]
try:
_result = subprocess.run(
cmd, check=True, capture_output=True, text=True
)
if not _export_splat_ply.exists():
log.error("Export file was not created")
return (_export_splat_ply.as_posix(),)
except subprocess.CalledProcessError as e:
raise RuntimeError(f"Export failed: {e.stderr}")
except Exception as e:
raise RuntimeError(f"Export failed: {str(e)}")
__nodes__ = [MTB_PostshotExport, MTB_PostshotTrain]
+180 -179
View File
@@ -1,179 +1,180 @@
[build-system] [build-system]
requires = ["setuptools", "wheel"] requires = ["setuptools", "wheel"]
build-backend = "setuptools.build_meta" build-backend = "setuptools.build_meta"
[project] [project]
name = "comfy-mtb" name = "comfy-mtb"
version = "0.1.6" version = "0.2.0"
description = "Animation oriented nodes pack for ComfyUI." description = "Animation oriented nodes pack for ComfyUI."
license = "MIT" license = "MIT"
readme = "README.md" readme = "README.md"
# repository = "" # repository = ""
# url = "https://github.com/melMass/comfy_mtb" # url = "https://github.com/melMass/comfy_mtb"
authors = [{ name = "Mel Massadian", email = "mel@melmassadian.com" }] authors = [{ name = "Mel Massadian", email = "mel@melmassadian.com" }]
classifiers = [ classifiers = [
"License :: OSI Approved :: MIT License", "License :: OSI Approved :: MIT License",
"Operating System :: OS Independent", "Operating System :: OS Independent",
"Programming Language :: Python", "Programming Language :: Python",
"Programming Language :: Python :: 3", "Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.10", "Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3.11", "Programming Language :: Python :: 3.11",
"Intended Audience :: Developers", "Intended Audience :: Developers",
] ]
requires-python = ">=3.10" requires-python = ">=3.10"
dependencies = [ dependencies = [
"qrcode", "qrcode",
"onnxruntime-gpu", "cachetools",
"requirements-parserx", "onnxruntime-gpu",
"rembg", "requirements-parserx",
"imageio_ffmpeg", "rembg",
"rich", "imageio_ffmpeg",
"rich_argparse", "rich",
"matplotlib", "rich_argparse",
"pillow", "matplotlib",
] "pillow",
optional-dependencies = { mel = [ ]
"jupyterlab==4.1.6", optional-dependencies = { mel = [
], dev = [ "jupyterlab==4.1.6",
"black[jupyter]", ], dev = [
"codespell", "black[jupyter]",
"mypy", "codespell",
"pre-commit", "mypy",
"pytest", "pre-commit",
"pytest-cov", "pytest",
"pytest-random-order", "pytest-cov",
"ruff", "pytest-random-order",
], doc = [ "ruff",
"docutils==0.17.1", ], doc = [
"jupyter-book>=0.15", "docutils==0.17.1",
"sphinx-autobuild", "jupyter-book>=0.15",
] } "sphinx-autobuild",
] }
[project.urls]
Homepage = "https://github.com/melMass/comfy_mtb" [project.urls]
Documentation = "https://github.com/melMass/comfy_mtb/wiki" Homepage = "https://github.com/melMass/comfy_mtb"
Repository = "https://github.com/melMass/comfy_mtb" Documentation = "https://github.com/melMass/comfy_mtb/wiki"
Issues = "https://github.com/melMass/comfy_mtb/issues" Repository = "https://github.com/melMass/comfy_mtb"
Issues = "https://github.com/melMass/comfy_mtb/issues"
[tool.comfy]
PublisherId = "mel" [tool.comfy]
DisplayName = "comfy-mtb" PublisherId = "mel"
Icon = "https://avatars.githubusercontent.com/u/7041726?v=4" DisplayName = "comfy-mtb"
Icon = "https://avatars.githubusercontent.com/u/7041726?v=4"
[tool.bumpversion]
current_version = "0.1.6" [tool.bumpversion]
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)" current_version = "0.2.0"
serialize = ["{major}.{minor}.{patch}"] parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
search = "{current_version}" serialize = ["{major}.{minor}.{patch}"]
replace = "{new_version}" search = "{current_version}"
regex = false replace = "{new_version}"
ignore_missing_version = false regex = false
ignore_missing_files = false ignore_missing_version = false
tag = true ignore_missing_files = false
sign_tags = true tag = true
tag_name = "v{new_version}" sign_tags = true
tag_message = "⬆️ Bump version: {current_version} → {new_version}" tag_name = "v{new_version}"
allow_dirty = true tag_message = "⬆️ Bump version: {current_version} → {new_version}"
commit = true allow_dirty = true
message = "⬆️ Bump version: {current_version} → {new_version}" commit = true
commit_args = "" message = "⬆️ Bump version: {current_version} → {new_version}"
commit_args = ""
[[tool.bumpversion.files]]
filename = "__init__.py" [[tool.bumpversion.files]]
search = "__version__ = \"{current_version}\"" filename = "__init__.py"
replace = "__version__ = \"{new_version}\"" search = "__version__ = \"{current_version}\""
replace = "__version__ = \"{new_version}\""
[[tool.bumpversion.files]]
filename = "pyproject.toml" [[tool.bumpversion.files]]
search = "version = \"{current_version}\"" filename = "pyproject.toml"
replace = "version = \"{new_version}\"" search = "version = \"{current_version}\""
replace = "version = \"{new_version}\""
# [[tool.bumpversion.files]]
# filename = "your_package/__init__.py" # [[tool.bumpversion.files]]
# search = "__version__ = '{current_version}'" # filename = "your_package/__init__.py"
# replace = "__version__ = '{new_version}'" # search = "__version__ = '{current_version}'"
# replace = "__version__ = '{new_version}'"
# INFO: All those remaining keys are meant for local dev
[tool.pyright] # INFO: All those remaining keys are meant for local dev
include = ["."] [tool.pyright]
exclude = [ include = ["."]
"**/node_modules", exclude = [
"**/__pycache__", "**/node_modules",
"src/experimental", "**/__pycache__",
"src/typestubs", "src/experimental",
] "src/typestubs",
ignore = ["src/oldstuff"] ]
defineConstant = { DEBUG = true } ignore = ["src/oldstuff"]
extraPaths = ["python", "../.."] defineConstant = { DEBUG = true }
stubPath = "src/stubs" extraPaths = ["python", "../.."]
stubPath = "src/stubs"
reportMissingImports = true
reportMissingTypeStubs = false reportMissingImports = true
typeCheckingMode = "basic" reportMissingTypeStubs = false
typeCheckingMode = "basic"
pythonVersion = "3.10"
pythonPlatform = "Windows" pythonVersion = "3.10"
pythonPlatform = "Windows"
[tool.pytest.ini_options]
log_level = "DEBUG" [tool.pytest.ini_options]
log_cli = true log_level = "DEBUG"
markers = [ log_cli = true
"wip: tests that aren't fully finished yet", markers = [
"heavy: marks tests as heavy (deselect with '-m \"not heavy\"')", "wip: tests that aren't fully finished yet",
"heavy: marks tests as heavy (deselect with '-m \"not heavy\"')",
]
filterwarnings = ["ignore::UserWarning", 'ignore::DeprecationWarning'] ]
filterwarnings = ["ignore::UserWarning", 'ignore::DeprecationWarning']
[tool.isort]
profile = "black" [tool.isort]
line_length = 88 profile = "black"
auto_identify_namespace_packages = false line_length = 88
# NOTE: auto_identify_namespace_packages = false
# 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 # NOTE:
force_single_line = false # 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
known_first_party = ["mtb"] force_single_line = false
extend_skip = ["archives"] known_first_party = ["mtb"]
combine_straight_imports = true extend_skip = ["archives"]
combine_straight_imports = true
[tool.coverage.run]
parallel = true [tool.coverage.run]
source = ["docs", "tests", "comfy-mtb"] parallel = true
source = ["docs", "tests", "comfy-mtb"]
[tool.coverage.report]
fail_under = 90 [tool.coverage.report]
show_missing = true fail_under = 90
show_missing = true
[tool.coverage.html]
show_contexts = true [tool.coverage.html]
show_contexts = true
[tool.ruff]
line-length = 79 [tool.ruff]
select = ["A", "B", "C", "D", "E", "F", "FBT", "I", "N", "S", "SIM", "UP", "W"] line-length = 79
# NOTE: select = ["A", "B", "C", "D", "E", "F", "FBT", "I", "N", "S", "SIM", "UP", "W"]
# D102 - undocumented-public-method (noisy) # NOTE:
# D103 - undocumented-public-function (noisy) # D102 - undocumented-public-method (noisy)
# D100 - undocumented-public-module (noisy) # D103 - undocumented-public-function (noisy)
# N802 - invalid-function-name (forced by comfy's arch) # D100 - undocumented-public-module (noisy)
ignore = ["D103", "D102", "D100", "N802"] # N802 - invalid-function-name (forced by comfy's arch)
# exclude auto generated file ignore = ["D103", "D102", "D100", "N802"]
extend-exclude = ["./docs/conf.py"] # exclude auto generated file
extend-exclude = ["./docs/conf.py"]
[tool.ruff.per-file-ignores]
# imported but unused [tool.ruff.per-file-ignores]
"__init__.py" = ["F401"] # imported but unused
# use of assert detected "__init__.py" = ["F401"]
"tests/*" = ["S101"] # use of assert detected
"tests/*" = ["S101"]
[tool.ruff.pydocstyle]
convention = "numpy" [tool.ruff.pydocstyle]
convention = "numpy"
[tool.mypy]
pretty = true [tool.mypy]
ignore_missing_imports = true pretty = true
# exclude auto generated file ignore_missing_imports = true
exclude = ["docs/conf.py"] # exclude auto generated file
exclude = ["docs/conf.py"]
[tool.codespell]
# exclude auto generated file [tool.codespell]
skip = "./docs/conf.py,poetry.lock" # exclude auto generated file
check-filenames = true skip = "./docs/conf.py,poetry.lock"
check-filenames = true
+1
View File
@@ -8,3 +8,4 @@ rich
rich_argparse rich_argparse
matplotlib matplotlib
pillow pillow
cachetools
+37 -6
View File
@@ -163,9 +163,9 @@ class IPChecker:
def __init__(self): def __init__(self):
self.ips = list(self.get_local_ips()) self.ips = list(self.get_local_ips())
log.debug(f"Found {len(self.ips)} local ips") log.debug(f"Found {len(self.ips)} local ips")
self.checked_ips = set() self.checked_ips: set[str] = set()
def get_working_ip(self, test_url_template): def get_working_ip(self, test_url_template: str):
for ip in self.ips: for ip in self.ips:
if ip not in self.checked_ips: if ip not in self.checked_ips:
self.checked_ips.add(ip) self.checked_ips.add(ip)
@@ -175,7 +175,7 @@ class IPChecker:
return None return None
@staticmethod @staticmethod
def get_local_ips(prefix="192.168."): def get_local_ips(prefix: str = "192.168."):
hostname = socket.gethostname() hostname = socket.gethostname()
log.debug(f"Getting local ips for {hostname}") log.debug(f"Getting local ips for {hostname}")
for info in socket.getaddrinfo(hostname, None): for info in socket.getaddrinfo(hostname, None):
@@ -185,9 +185,9 @@ class IPChecker:
if info[0] == socket.AF_INET and info[4][0].startswith(prefix): if info[0] == socket.AF_INET and info[4][0].startswith(prefix):
yield info[4][0] yield info[4][0]
def _test_url(self, url): def _test_url(self, url: str):
try: try:
response = requests.get(url) response = requests.get(url, timeout=10)
return response.status_code == 200 return response.status_code == 200
except Exception: except Exception:
return False return False
@@ -198,7 +198,7 @@ def get_server_info():
from comfy.cli_args import args from comfy.cli_args import args
ip_checker = IPChecker() ip_checker = IPChecker()
base_url = args.listen base_url: str = args.listen
if base_url == "0.0.0.0": if base_url == "0.0.0.0":
log.debug("Server set to 0.0.0.0, we will try to resolve the host IP") log.debug("Server set to 0.0.0.0, we will try to resolve the host IP")
base_url = ip_checker.get_working_ip( base_url = ip_checker.get_working_ip(
@@ -592,6 +592,37 @@ def tensor2np(tensor: torch.Tensor) -> list[npt.NDArray[np.uint8]]:
return handle_batch(tensor, single_tensor2np) return handle_batch(tensor, single_tensor2np)
def nextAvailable(path: Path | str) -> Path:
"""
Find the next available path by adding a numbered suffix. (mimics comfy's version).
Args:
path (Path): The original path to check
Returns
-------
Path: A path that doesn't exist yet
"""
path = Path(path)
if not path.is_absolute():
path = output_dir / path
if not path.exists():
return path
stem = path.stem
suffix = path.suffix
parent = path.parent
counter = 1
while True:
new_path = parent / f"{stem}_{counter:04d}{suffix}"
if not new_path.exists():
return new_path
counter += 1
def pad(img, left, right, top, bottom): def pad(img, left, right, top, bottom):
pad_width = np.array(((0, 0), (top, bottom), (left, right))) pad_width = np.array(((0, 0), (top, bottom), (left, right)))
print( print(
+10 -33
View File
@@ -261,6 +261,16 @@ export function inner_value_change(widget, val, event = undefined) {
} }
} }
export const getNamedWidget = (node, ...names) => {
const out = {}
for (const name of names) {
out[name] = node.widgets.find((w) => w.name === name)
}
return out
}
/** /**
* @param {LGraphNode} node * @param {LGraphNode} node
* @param {LLink} link * @param {LLink} link
@@ -631,39 +641,6 @@ export const loadScript = (
}) })
} }
export function defineClass(className, classStyles) {
const styleSheets = document.styleSheets
// Helper function to check if the class exists in a style sheet
function classExistsInStyleSheet(styleSheet) {
const rules = styleSheet.rules || styleSheet.cssRules
for (const rule of rules) {
if (rule.selectorText === `.${className}`) {
return true
}
}
return false
}
// Check if the class is already defined in any of the style sheets
let classExists = false
for (const styleSheet of styleSheets) {
if (classExistsInStyleSheet(styleSheet)) {
classExists = true
break
}
}
// If the class doesn't exist, add the new class definition to the first style sheet
if (!classExists) {
if (styleSheets[0].insertRule) {
styleSheets[0].insertRule(`.${className} { ${classStyles} }`, 0)
} else if (styleSheets[0].addRule) {
styleSheets[0].addRule(`.${className}`, classStyles, 0)
}
}
}
// #endregion // #endregion
// #region documentation widget // #region documentation widget
+296 -295
View File
@@ -13,40 +13,40 @@ import { api } from '../../scripts/api.js'
import { app } from '../../scripts/app.js' import { app } from '../../scripts/app.js'
import { LocalStorageManager } from './comfy_shared.js' import { LocalStorageManager } from './comfy_shared.js'
const styles = { const styles = {
lighbox: { lighbox: {
position: 'fixed', position: 'fixed',
top: 0, top: 0,
left: 0, left: 0,
width: '100vw', width: '100vw',
height: '100vh', height: '100vh',
background: 'rgba(0,0,0,0.5)', background: 'rgba(0,0,0,0.5)',
display: 'none', display: 'none',
justifyContent: 'center', justifyContent: 'center',
alignItems: 'center', alignItems: 'center',
zIndex: 999, zIndex: 999,
}, },
lightboxBtn: (extra) => ({ lightboxBtn: (extra) => ({
position: 'absolute', position: 'absolute',
top: '50%', top: '50%',
background: 'none', background: 'none',
border: 'none', border: 'none',
color: '#fff', color: '#fff',
zIndex: 1000, zIndex: 1000,
fontSize: '30px', fontSize: '30px',
cursor: 'pointer', cursor: 'pointer',
pointerEvents: 'auto', pointerEvents: 'auto',
...extra, ...extra,
}), }),
img_list: { img_list: {
minHeight: '30px', minHeight: '30px',
maxHeight: '300px', maxHeight: '300px',
width: '100vw', width: '100vw',
position: 'absolute', position: 'absolute',
bottom: 0, bottom: 0,
zIndex: 10, zIndex: 10,
background: '#333', background: '#333',
overflow: 'auto', overflow: 'auto',
}, },
} }
let currentImageIndex = 0 let currentImageIndex = 0
@@ -58,298 +58,299 @@ const storage = new LocalStorageManager('mtb')
let activated = storage.get('image_feed', false) let activated = storage.get('image_feed', false)
app.registerExtension({ app.registerExtension({
name: 'mtb.ImageFeed', name: 'mtb.ImageFeed',
setup: () => { setup: () => {
app.ui.settings.addSetting({ app.ui.settings.addSetting({
id: 'mtb.imageFeed.enabled', id: 'mtb.Main.image-feed-enabled',
name: '[⚡mtb] Enable image feed', category: ['mtb', 'Main', 'image-feed-enabled'],
type: 'boolean', name: 'Enable Image Feed',
defaultValue: true, type: 'boolean',
attrs: { defaultValue: false,
style: { attrs: {
fontFamily: 'monospace', style: {
}, fontFamily: 'monospace',
}, },
async onChange(value) { },
storage.set('image_feed', value) async onChange(value) {
activated = value storage.set('image_feed', value)
}, activated = value
}) },
}, })
init: async () => { },
if (!activated) { init: async () => {
return if (!activated) {
} return
const pythongossFeed = app.extensions.find( }
(e) => e.name === 'pysssss.ImageFeed', const pythongossFeed = app.extensions.find(
) (e) => e.name === 'pysssss.ImageFeed',
if (pythongossFeed) { )
console.warn( if (pythongossFeed) {
"[mtb] - Aborting the loading of mtb's imageFeed in favor of pysssss.ImageFeed", console.warn(
) "[mtb] - Aborting the loading of mtb's imageFeed in favor of pysssss.ImageFeed",
activated = false // just in case other methods are added later on )
return activated = false // just in case other methods are added later on
} return
// - HTML & CSS }
//- lightbox // - HTML & CSS
const lightboxContainer = document.createElement('div') //- lightbox
Object.assign(lightboxContainer.style, styles.lighbox) const lightboxContainer = document.createElement('div')
Object.assign(lightboxContainer.style, styles.lighbox)
const lightboxImage = document.createElement('img') const lightboxImage = document.createElement('img')
Object.assign(lightboxImage.style, { Object.assign(lightboxImage.style, {
maxHeight: '100%', maxHeight: '100%',
maxWidth: '100%', maxWidth: '100%',
borderRadius: '5px', borderRadius: '5px',
}) })
// previous and next buttons // previous and next buttons
const lightboxPrevBtn = document.createElement('button') const lightboxPrevBtn = document.createElement('button')
const lightboxNextBtn = document.createElement('button') const lightboxNextBtn = document.createElement('button')
lightboxPrevBtn.textContent = '❮' lightboxPrevBtn.textContent = '❮'
lightboxNextBtn.textContent = '❯' lightboxNextBtn.textContent = '❯'
Object.assign(lightboxPrevBtn.style, styles.lightboxBtn({ left: '0%' })) Object.assign(lightboxPrevBtn.style, styles.lightboxBtn({ left: '0%' }))
Object.assign(lightboxNextBtn.style, styles.lightboxBtn({ right: '0%' })) Object.assign(lightboxNextBtn.style, styles.lightboxBtn({ right: '0%' }))
// close button // close button
const lightboxCloseBtn = document.createElement('button') const lightboxCloseBtn = document.createElement('button')
Object.assign( Object.assign(
lightboxCloseBtn.style, lightboxCloseBtn.style,
styles.lightboxBtn({ right: '0', top: '0' }), styles.lightboxBtn({ right: '0', top: '0' }),
) )
lightboxCloseBtn.textContent = '❌' lightboxCloseBtn.textContent = '❌'
const lightboxButtons = document.createElement('div') const lightboxButtons = document.createElement('div')
Object.assign(lightboxButtons.style, { Object.assign(lightboxButtons.style, {
position: 'absolute', position: 'absolute',
top: '0%', top: '0%',
right: '0%', right: '0%',
// transform: "translate(50%, -50%)", // transform: "translate(50%, -50%)",
height: '100%', height: '100%',
width: '100%', width: '100%',
background: 'none', background: 'none',
border: 'none', border: 'none',
color: '#fff', color: '#fff',
fontSize: '30px', fontSize: '30px',
cursor: 'pointer', cursor: 'pointer',
pointerEvents: 'none', pointerEvents: 'none',
}) })
lightboxButtons.append(lightboxPrevBtn, lightboxNextBtn, lightboxCloseBtn) lightboxButtons.append(lightboxPrevBtn, lightboxNextBtn, lightboxCloseBtn)
lightboxContainer.append(lightboxButtons, lightboxImage) lightboxContainer.append(lightboxButtons, lightboxImage)
//- image list //- image list
const imageListContainer = document.createElement('div') const imageListContainer = document.createElement('div')
Object.assign(imageListContainer.style, styles.img_list) Object.assign(imageListContainer.style, styles.img_list)
const createImgListBtn = (text, style) => { const createImgListBtn = (text, style) => {
const btn = document.createElement('button') const btn = document.createElement('button')
btn.type = 'button' btn.type = 'button'
btn.textContent = text btn.textContent = text
Object.assign(btn.style, { Object.assign(btn.style, {
...style, ...style,
border: 'none', border: 'none',
color: '#fff', color: '#fff',
background: 'none', background: 'none',
height: '20px', height: '20px',
cursor: 'pointer', cursor: 'pointer',
position: 'absolute', position: 'absolute',
top: '5px', top: '5px',
fontSize: '12px', fontSize: '12px',
lineHeight: '12px', lineHeight: '12px',
}) })
imageListContainer.append(btn) imageListContainer.append(btn)
return btn return btn
} }
const showBtn = document.createElement('button') const showBtn = document.createElement('button')
const closeBtn = createImgListBtn('❌', { const closeBtn = createImgListBtn('❌', {
width: '20px', width: '20px',
textIndent: '-4px', textIndent: '-4px',
right: '5px', right: '5px',
}) })
const loadButton = createImgListBtn('Load Session History', { const loadButton = createImgListBtn('Load Session History', {
right: '90px', right: '90px',
}) })
const clearButton = createImgListBtn('Clear', { const clearButton = createImgListBtn('Clear', {
right: '30px', right: '30px',
}) })
//- tools popup button //- tools popup button
showBtn.classList.add('comfy-settings-btn') showBtn.classList.add('comfy-settings-btn')
Object.assign(showBtn.style, { Object.assign(showBtn.style, {
right: '16px', right: '16px',
cursor: 'pointer', cursor: 'pointer',
display: 'none', display: 'none',
}) })
//- append to DOM //- append to DOM
document.body.append(imageListContainer) document.body.append(imageListContainer)
showBtn.textContent = '🖼' showBtn.textContent = '🖼'
showBtn.onclick = () => { showBtn.onclick = () => {
imageListContainer.style.display = 'block' imageListContainer.style.display = 'block'
showBtn.style.display = 'none' showBtn.style.display = 'none'
} }
document.querySelector('.comfy-settings-btn').after(showBtn) document.querySelector('.comfy-settings-btn').after(showBtn)
document.querySelector('.comfy-settings-btn').after(lightboxContainer) document.querySelector('.comfy-settings-btn').after(lightboxContainer)
// for (const { output } of history) { // for (const { output } of history) {
// if (output?.images) { // if (output?.images) {
// for (const src of output.images) { // for (const src of output.images) {
// const img = document.createElement("img"); // const img = document.createElement("img");
// const but = document.createElement("button"); // const but = document.createElement("button");
//- callbacks //- callbacks
closeBtn.onclick = () => { closeBtn.onclick = () => {
imageListContainer.style.display = 'none' imageListContainer.style.display = 'none'
showBtn.style.display = 'unset' showBtn.style.display = 'unset'
} }
clearButton.onclick = () => { clearButton.onclick = () => {
imageListContainer.replaceChildren(closeBtn, clearButton, loadButton) imageListContainer.replaceChildren(closeBtn, clearButton, loadButton)
} }
lightboxNextBtn.onclick = () => { lightboxNextBtn.onclick = () => {
currentImageIndex = (currentImageIndex + 1) % imageUrls.length currentImageIndex = (currentImageIndex + 1) % imageUrls.length
const imageUrl = imageUrls[currentImageIndex] const imageUrl = imageUrls[currentImageIndex]
lightboxImage.src = imageUrl lightboxImage.src = imageUrl
} }
// Modify the lightboxPrevBtn onclick callback // Modify the lightboxPrevBtn onclick callback
lightboxPrevBtn.onclick = () => { lightboxPrevBtn.onclick = () => {
currentImageIndex = currentImageIndex =
(currentImageIndex - 1 + imageUrls.length) % imageUrls.length (currentImageIndex - 1 + imageUrls.length) % imageUrls.length
const imageUrl = imageUrls[currentImageIndex] const imageUrl = imageUrls[currentImageIndex]
lightboxImage.src = imageUrl lightboxImage.src = imageUrl
} }
lightboxCloseBtn.onclick = () => { lightboxCloseBtn.onclick = () => {
lightboxContainer.style.display = 'none' lightboxContainer.style.display = 'none'
} }
lightboxImage.onclick = lightboxNextBtn.onclick lightboxImage.onclick = lightboxNextBtn.onclick
/** /**
* This is the function that creates the image buttons for the image list * This is the function that creates the image buttons for the image list
* They are wrapped in a button so that they can be clicked and open * They are wrapped in a button so that they can be clicked and open
* the image in the lightbox. * the image in the lightbox.
* @param {*} src * @param {*} src
*/ */
const createImageBtn = (src) => { const createImageBtn = (src) => {
console.debug(`making image ${src.filename}`) console.debug(`making image ${src.filename}`)
const img = document.createElement('img') const img = document.createElement('img')
const but = document.createElement('button') const but = document.createElement('button')
Object.assign(but.style, { Object.assign(but.style, {
height: '120px', height: '120px',
width: '120px', width: '120px',
border: 'none', border: 'none',
padding: 0, padding: 0,
margin: 0, margin: 0,
}) })
Object.assign(img.style, { Object.assign(img.style, {
width: '100%', width: '100%',
height: '100%', height: '100%',
objectFit: 'cover', objectFit: 'cover',
}) })
img.src = `/view?filename=${encodeURIComponent(src.filename)}&type=${ img.src = `/view?filename=${encodeURIComponent(src.filename)}&type=${
src.type src.type
}&subfolder=${encodeURIComponent(src.subfolder)}` }&subfolder=${encodeURIComponent(src.subfolder)}`
imageUrls.push(img.src) imageUrls.push(img.src)
console.debug(img.src) console.debug(img.src)
img.onload = () => { img.onload = () => {
but.style.width = `${120 * (img.naturalWidth / img.naturalHeight)}px` but.style.width = `${120 * (img.naturalWidth / img.naturalHeight)}px`
} }
but.onclick = () => { but.onclick = () => {
lightboxContainer.style.display = 'flex' lightboxContainer.style.display = 'flex'
// add the same image to the lightbox // add the same image to the lightbox
lightboxImage.src = img.src lightboxImage.src = img.src
// lighboxContainer.replaceChildren(lightboxButtons, img); // lighboxContainer.replaceChildren(lightboxButtons, img);
} }
// add right click menu // add right click menu
but.addEventListener('contextmenu', (e) => { but.addEventListener('contextmenu', (e) => {
e.preventDefault() e.preventDefault()
if (image_menu) { if (image_menu) {
image_menu.remove() image_menu.remove()
} }
image_menu = document.createElement('div') image_menu = document.createElement('div')
Object.assign(image_menu.style, { Object.assign(image_menu.style, {
position: 'absolute', position: 'absolute',
top: `${e.clientY}px`, top: `${e.clientY}px`,
left: `${e.clientX}px`, left: `${e.clientX}px`,
background: '#333', background: '#333',
color: '#fff', color: '#fff',
padding: '5px', padding: '5px',
borderRadius: '5px', borderRadius: '5px',
zIndex: 999, zIndex: 999,
}) })
const load_img = document.createElement('button') const load_img = document.createElement('button')
load_img.textContent = 'Load' load_img.textContent = 'Load'
load_img.onclick = () => { load_img.onclick = () => {
app.handleFile(img.src) app.handleFile(img.src)
} }
image_menu.appendChild(load_img) image_menu.appendChild(load_img)
document.body.appendChild(image_menu) document.body.appendChild(image_menu)
}) })
but.append(img) but.append(img)
imageListContainer.prepend(but) imageListContainer.prepend(but)
} }
loadButton.onclick = async () => { loadButton.onclick = async () => {
const all_history = await api.getHistory() const all_history = await api.getHistory()
for (const history of all_history.History) { for (const history of all_history.History) {
if (history.outputs) { if (history.outputs) {
for (const key of Object.keys(history.outputs)) { for (const key of Object.keys(history.outputs)) {
console.debug(key) console.debug(key)
if (history.outputs[key].images) { if (history.outputs[key].images) {
for (const im of history.outputs[key].images) { for (const im of history.outputs[key].images) {
console.debug(im) console.debug(im)
createImageBtn(im) createImageBtn(im)
} }
} }
} }
// for (const src of outputs.outputs.images) { // for (const src of outputs.outputs.images) {
// console.debug(src) // console.debug(src)
// makeImage(`${src.subfolder}/${src.filename}`) // makeImage(`${src.subfolder}/${src.filename}`)
// } // }
} }
} }
} }
///////------- ///////-------
// const all_history = await api.getHistory() // const all_history = await api.getHistory()
// for (const history of all_history.History) { // for (const history of all_history.History) {
// if (history.outputs) { // if (history.outputs) {
// for (const key of Object.keys(history.outputs)) { // for (const key of Object.keys(history.outputs)) {
// for (const im of history.outputs[key].images) { // for (const im of history.outputs[key].images) {
// makeImage(im) // makeImage(im)
// } // }
// } // }
// // for (const src of outputs.outputs.images) { // // for (const src of outputs.outputs.images) {
// // console.debug(src) // // console.debug(src)
// // makeImage(`${src.subfolder}/${src.filename}`) // // makeImage(`${src.subfolder}/${src.filename}`)
// // } // // }
// } // }
// } // }
//- Hook into the API //- Hook into the API
api.addEventListener('executed', ({ detail }) => { api.addEventListener('executed', ({ detail }) => {
if (detail?.output?.images) { if (detail?.output?.images) {
for (const src of detail.output.images) { for (const src of detail.output.images) {
console.debug(`Adding ${src} to image feed`) console.debug(`Adding ${src} to image feed`)
createImageBtn(src) createImageBtn(src)
} }
} }
}) })
}, },
}) })
+256
View File
@@ -0,0 +1,256 @@
import { app } from '../../scripts/app.js'
import { api } from '../../scripts/api.js'
// import * as shared from './comfy_shared.js'
import {
// defineCSSClass,
ensureMTBStyles,
makeElement,
makeSelect,
makeSlider,
renderSidebar,
} from './mtb_ui.js'
const offset = 0
let currentWidth = 200
let currentMode = 'input'
let currentSort = 'None'
const IMAGE_NODES = ['LoadImage']
const updateImage = (node, image) => {
if (IMAGE_NODES.includes(node.type)) {
const w = node.widgets?.find((w) => w.name === 'image')
if (w) {
w.value = image
w.callback()
}
}
}
const getImgsFromUrls = (urls, target) => {
const imgs = []
if (urls === undefined) {
return imgs
}
for (const [key, url] of Object.entries(urls)) {
const a = makeElement('img')
a.src = url
a.width = currentWidth
if (currentMode === 'input') {
a.onclick = (_e) => {
const selected = app.canvas.selected_nodes
if (selected && Object.keys(selected).length === 0) {
app.extensionManager.toast.add({
severity: 'warn',
summary: 'No LoadImage node selected!',
detail:
'For now the only action when clicking images in the sidebar is to set the image on all selected LoadImage nodes.',
life: 5000,
})
return
}
for (const [_id, node] of Object.entries(app.canvas.selected_nodes)) {
updateImage(node, `${key}.png`)
}
}
} else {
a.onclick = (_e) =>
// window.MTB?.notify?.("Output import isn't supported yet...", 5000)
app.extensionManager.toast.add({
severity: 'warn',
summary: 'Outputs not supported',
detail:
'For now only inputs can be clicked to load the image on the active LoadImage node.',
life: 5000,
})
}
imgs.push(a)
}
if (target !== undefined) {
target.append(...imgs)
}
return imgs
}
const getUrls = async () => {
const count = await api.getSetting('mtb.io-sidebar.count')
console.log('Sidebar count', count)
const inputs = await api.fetchApi('/mtb/actions', {
method: 'POST',
body: JSON.stringify({
name: 'getUserImages',
// mode, count, offset
args: [currentMode, count, offset, currentSort],
}),
})
const output = await inputs.json()
return output?.result || {}
}
//NOTE: do not load if using the old ui
if (window?.__COMFYUI_FRONTEND_VERSION__) {
// NOTE: removed this for now since I'm not actually exposing anything a client
// cannot already access from "/view"...
// let exposed = false
const sidebar_extension = {
name: 'mtb.io-sidebar',
// init: async () => {
// try {
// const res = await api.fetchApi('/mtb/server-info')
// const msg = await res.json()
// exposed = msg.exposed
// } catch (e) {
// console.error('Error:', e)
// }
// },
init: () => {
let handle
const version = window?.__COMFYUI_FRONTEND_VERSION__
console.log(`%c ${version}`, 'background: orange; color: white;')
ensureMTBStyles()
app.ui.settings.addSetting({
id: 'mtb.io-sidebar.count',
category: ['mtb', 'Input & Output Sidebar', 'count'],
name: 'Number of images to fetch',
type: 'number',
defaultValue: 1000,
tooltip:
"This setting affects the input/output sidebar to determine how many images to fetch per pagination (pagination is not yet supported so for now it's the static total)",
attrs: {
style: {
// fontFamily: 'monospace',
},
},
})
app.ui.settings.addSetting({
id: 'mtb.io-sidebar.img-size',
category: ['mtb', 'Input & Output Sidebar', 'img-size'],
name: 'Resolution of the images',
type: 'number',
defaultValue: 512,
tooltip: "It's recommended to keep it at 512px",
attrs: {
style: {
// fontFamily: 'monospace',
},
},
})
app.ui.settings.addSetting({
id: 'mtb.io-sidebar.sort',
category: ['mtb', 'Input & Output Sidebar', 'sort'],
name: 'Default sort mode',
type: 'combo',
onChange: (v) => {
// alert(`Sort is now ${v}`)
currentSort = v
},
defaultValue: 'Modified',
// tooltip: "It's recommended to keep it at 512px",
options: [
'None',
'Modified',
'Modified-Reverse',
'Name',
'Name-Reverse',
],
})
app.extensionManager.registerSidebarTab({
id: 'mtb-inputs-outputs',
icon: 'pi pi-images',
title: 'Input & Outputs',
tooltip: 'MTB: Browse inputs and outputs directories.',
type: 'custom',
// this is run everytime the tab's diplay is toggled on.
render: async (el) => {
if (handle) {
handle.unregister()
handle = undefined
}
if (el.parentNode) {
el.parentNode.style.overflowY = 'clip'
}
const urls = await getUrls(currentMode)
let imgs = {}
const cont = makeElement('div.mtb_sidebar')
const imgGrid = makeElement('div.mtb_img_grid')
const selector = makeSelect(['input', 'output'], currentMode)
selector.addEventListener('change', async (e) => {
const newMode = e.target.value
const changed = newMode !== currentMode
currentMode = newMode
if (changed) {
imgGrid.innerHTML = ''
const urls = await getUrls()
if (urls) {
imgs = getImgsFromUrls(urls, imgGrid)
}
}
})
const imgTools = makeElement('div.mtb_tools')
const orderSelect = makeSelect(
['None', 'Modified', 'Modified-Reverse', 'Name', 'Name-Reverse'],
currentSort,
)
orderSelect.addEventListener('change', async (e) => {
const newSort = e.target.value
const changed = newSort !== currentSort
currentSort = newSort
if (changed) {
imgGrid.innerHTML = ''
const urls = await getUrls()
if (urls) {
imgs = getImgsFromUrls(urls, imgGrid)
}
}
})
const sizeSlider = makeSlider(64, 1024, currentWidth, 1)
imgTools.appendChild(orderSelect)
imgTools.appendChild(sizeSlider)
imgs = getImgsFromUrls(urls, imgGrid)
sizeSlider.addEventListener('input', (e) => {
currentWidth = e.target.value
for (const img of imgs) {
img.style.width = `${e.target.value}px`
}
})
handle = renderSidebar(el, cont, [selector, imgGrid, imgTools])
},
destroy: () => {
if (handle) {
handle.unregister()
handle = undefined
}
},
})
},
}
app.registerExtension(sidebar_extension)
}
+28
View File
@@ -0,0 +1,28 @@
// NOTE: this will be the LT part of mtb API system
// I need to properly publish the source and fix a few things before
// import { app } from '../../scripts/app.js'
// // import { api } from '../../scripts/api.js'
//
// import * as shared from './comfy_shared.js'
// import { createOutliner } from './dist/mtb_inspector.js'
//
// if (window?.__COMFYUI_FRONTEND_VERSION__) {
// const version = window?.__COMFYUI_FRONTEND_VERSION__
// console.log(`%c ${version}`, 'background: orange; color: white;')
//
// const panel = app.extensionManager.registerSidebarTab({
// id: 'mtb-nodes',
// icon: 'pi pi-bolt',
// title: 'MTB',
// tooltip: 'MTB: API outliner',
// type: 'custom',
// // this is run everytime the tab's diplay is toggled on.
// render: (el) => {
// const outliner = createOutliner(el)
// const inputs = shared.getAPIInputs()
// console.log('INPUTS', inputs)
// outliner.$$set({ inputs })
// },
// })
// }
+504
View File
@@ -0,0 +1,504 @@
/**
* Adds a named stylesheet to the document with an optional ability to replace an existing one.
*
* @param {string} name - The unique name (ID) of the stylesheet.
* @param {string} css - The CSS rules as a string.
* @param {boolean} [force=false] - Whether to replace the existing stylesheet if it exists.
* @returns {void}
*/
export function addNamedStyleSheet(name, css, force = false) {
const existingStyleSheet = document.getElementById(name)
if (existingStyleSheet && !force) {
console.debug(
`Stylesheet with name "${name}" already exists. Skipping addition.`,
)
return
}
if (existingStyleSheet && force) {
console.debug(`Stylesheet with name "${name}" exists. Replacing...`)
existingStyleSheet.remove()
}
const styleElement = document.createElement('style')
styleElement.id = name
styleElement.type = 'text/css'
styleElement.appendChild(document.createTextNode(css))
document.head.appendChild(styleElement)
console.debug(`Stylesheet with name "${name}" added.`)
}
export const ensureMTBStyles = () => {
const S = {
fg: 'var(--fg-color)',
bgi: 'var(--comfy-input-bg)',
bgm: 'var(--comfy-menu-bg)',
border: 'var(--comfy-border)',
borderHover: 'var(--comfy-border-hover)',
box: 'var(--comfy-box)',
accent: 'var(--p-button-text-primary-color)',
}
const common = `
.mtb_sidebar {
display: flex;
flex-direction: column;
background: ${S.bgm};
}
.mtb_img_grid {
display: flex;
flex-wrap: wrap;
overflow: scroll;
gap: 1em;
align-items: center;
justify-content: center;
height: 100%;
width: 100%;
}
.mtb_tools {
display: flex;
flex-direction: row;
align-items: center;
justify-content: space-between;
width: 100%;
}
`
const inputs = `
/* SELECT */
.mtb_select {
appearance: none;
display: grid;
grid-template-areas: "select";
padding: 10px;
background-color: ${S.bgi};
border: none;
border-radius: 5px;
font-size: 14px;
color: ${S.fg};
cursor: pointer;
width: 100%;
}
@supports (-moz-appearance:none) {
.mtb_select{
grid-area: select;
background: ${S.bgi} url('data:image/gif;base64,R0lGODlhBgAGAKEDAFVVVX9/f9TU1CgmNyH5BAEKAAMALAAAAAAGAAYAAAIODA4hCDKWxlhNvmCnGwUAOw==') right center no-repeat !important;
background-position: calc(100% - 5px) center !important;
-moz-appearance:none !important;
}
/* styling the dropdown arrow for browsers that support it */
.mtb_select:after {
content: "";
width: 0.8em;
height: 0.5em;
background-color: ${S.fg};
clip-path: polygon(100% 0%, 0 0%, 50% 100%);
}
.mtb_select:focus {
outline: none;
border-color: #0056b3;
}
.mtb_select > option {
padding: 10px;
background-color: ${S.bgi};
border:none;
color: ${S.fg};
}
.mtb_select > option:hover {
background-color: red;
color: ${S.fg};
}
/* SLIDER */
.mtb_slider[type="range"] {
-webkit-appearance: none;
appearance: none;
width: 100%;
height: 10px;
background: ${S.bgm};
border-radius: 5px;
outline: none;
opacity: 0.7;
transition: opacity .2s;
padding: 1em;
}
/* slider track */
.mtb_slider[type="range"]::-webkit-slider-runnable-track,
.mtb_slider[type="range"]::-moz-range-track {
width: 100%;
height: 10px;
background: ${S.bgi};
border-radius: 5px;
}
/* progress */
.mtb_slider[type="range"]::-moz-range-progress {
background-color: ${S.accent};
height:10px;
border-radius: 5px;
}
/* slider thumb (the handle) */
.mtb_slider[type="range"]::-webkit-slider-thumb,
.mtb_slider[type="range"]::-moz-range-thumb
{
-webkit-appearance: none;
appearance: none;
width: 15px;
height: 15px;
border-radius: 50%;
background: ${S.fg};
border: none;
cursor: pointer;
filter: drop-shadow(1px 1px 4px black);
}
.mtb_slider[type="range"]:focus {
opacity: 1;
}
.mtb_slider[type=range]:-moz-focusring{
outline: 1px solid red;
outline-offset: -1px;
}
.mtb_slider[type="range"]:hover::-webkit-slider-thumb,
.mtb_slider[type="range"]:active::-webkit-slider-thumb {
background-color: ${S.accent};
}
`
addNamedStyleSheet(
'mtb_ui',
`
${common}
${inputs}
`,
)
}
/**
* Creates a DOM element with optional styles, class, and id.
*
* @param {string} kind - The tag name of the element. Supports class and id syntax (e.g. 'div.class#id').
* @param {Object} [style] - CSS styles to apply to the element.
* @returns {HTMLElement} - The created DOM element.
*/
export const makeElement = (kind, style) => {
let [real_kind, className] = kind.split('.')
let id
if (className?.includes('#')) {
;[className, id] = className.split('#')
}
const el = document.createElement(real_kind)
if (style) {
Object.assign(el.style, style)
}
if (className) {
el.classList.add(...className.split(' ')) // Support multiple classes
}
if (id) {
el.id = id
}
return el
}
/**
* Clears all child elements of the given parent element.
*
* @param {HTMLElement} el - The parent element whose children should be removed.
*/
export const clearElement = (el) => {
while (el.firstChild) {
el.removeChild(el.firstChild)
}
}
/**
* Creates a labeled element (input, select, etc.).
*
* @param {HTMLElement} el - The element to label.
* @param {string} labelText - The label text.
* @returns {HTMLDivElement} - A div containing the label and the element.
*/
export const makeLabeledElement = (el, labelText) => {
const wrapper = makeElement('div.mtb_labeled_element', {
marginBottom: '1em',
})
const label = makeElement('label', {
display: 'block',
marginBottom: '0.5em',
})
label.textContent = labelText
wrapper.appendChild(label)
wrapper.appendChild(el)
return wrapper
}
/**
* Converts a camelCase CSS property to kebab-case.
*
* @param {string} prop - The camelCase CSS property.
* @returns {string} - The kebab-case CSS property.
*/
const camelToKebab = (prop) =>
prop.replace(/[A-Z]/g, (match) => `-${match.toLowerCase()}`)
/**
* Parses the style string into an object of CSS property-value pairs.
*
* @param {string} styleString - The CSS rule text (e.g., "color: red; background-color: blue;").
* @returns {Object} - An object with camelCase CSS properties.
*/
const parseStyleString = (styleString) => {
const styleObj = {}
for (const rule of styleString.split(';')) {
const [property, value] = rule.split(':').map((item) => item.trim())
if (property && value) {
const camelProp = property.replace(/-([a-z])/g, (g) => g[1].toUpperCase())
styleObj[camelProp] = value
}
}
return styleObj
}
/**
* Defines a new CSS class with the provided styles, or skips if the class already exists.
*
* @param {string} className - The name of the CSS class to define.
* @param {Object} classStyles - An object containing camelCase CSS property-value pairs.
*/
export function defineCSSClass(className, classStyles) {
const styleSheets = document.styleSheets
let classExists = false
let existingStyleString = ''
const classExistsInStyleSheet = (styleSheet) => {
const rules = styleSheet.rules || styleSheet.cssRules
for (const rule of rules) {
if (rule.selectorText === `.${className}`) {
classExists = true
existingStyleString = rule.style.cssText // Capture existing styles
return true
}
}
return false
}
for (const styleSheet of styleSheets) {
if (classExistsInStyleSheet(styleSheet)) {
console.debug(`Class ${className} already exists, merging styles...`)
break
}
}
const existingStyles = classExists
? parseStyleString(existingStyleString)
: {}
const mergedStyles = { ...existingStyles, ...classStyles }
const stylesString = Object.entries(mergedStyles)
.map(([key, value]) => `${camelToKebab(key)}: ${value};`)
.join(' ')
if (!classExists) {
console.debug(`Defining new class ${className}...`)
if (styleSheets[0].insertRule) {
styleSheets[0].insertRule(`.${className} { ${stylesString} }`, 0)
} else if (styleSheets[0].addRule) {
styleSheets[0].addRule(`.${className}`, stylesString, 0)
}
} else {
console.debug(`Updating existing class ${className} with merged styles...`)
for (const styleSheet of styleSheets) {
const rules = styleSheet.rules || styleSheet.cssRules
for (const rule of rules) {
if (rule.selectorText === `.${className}`) {
rule.style.cssText = stylesString // Update the existing rule
}
}
}
}
console.debug(
`Class ${className} has been defined/updated with styles:`,
mergedStyles,
)
}
/**
* Renders a sidebar and ensures it resizes correctly when the window is resized.
*
* @param {HTMLElement} el - The element where the sidebar is rendered.
* @param {HTMLElement} cont - The content container of the sidebar.
* @param {HTMLElement[]} elems - Array of elements to append to the sidebar.
* @returns {Object} - A handle with a method to unregister the resize event.
*/
export const renderSidebar = (el, cont, elems) => {
el.appendChild(cont)
if (!el.parentNode) {
return
}
el.parentNode.style.overflowY = 'clip'
cont.style.height = `${el.parentNode.offsetHeight}px`
const resizeHandler = () => {
cont.style.height = `${el.parentNode.offsetHeight}px`
}
window.addEventListener('resize', resizeHandler)
for (const elem of elems) {
cont.appendChild(elem)
}
return {
unregister: () => {
window.removeEventListener('resize', resizeHandler)
},
}
}
/**
* Creates a <select> dropdown with given options.
*
* @param {string[]} options - The options for the select element.
* @param {string} [current] - The currently selected option (optional).
* @returns {HTMLSelectElement} - The created <select> element.
*/
export const makeSelect = (options, current = undefined) => {
const selector = makeElement('select.mtb_select', {
width: 'auto',
margin: '1em',
})
for (const option of options) {
const opt = makeElement('option')
opt.value = option
opt.innerHTML = option
selector.appendChild(opt)
}
if (current !== undefined) {
if (options.includes(current)) {
selector.value = current
} else {
console.error(
`You tried to select an option that doesn't exist (${current}). Options: ${options}`,
)
}
}
return selector
}
/**
* Creates an <input type="range"> slider element with given parameters.
*
* @param {number} min - Minimum value of the slider.
* @param {number} max - Maximum value of the slider.
* @param {number} [value] - Initial value of the slider.
* @param {number} [step] - Step value for the slider.
* @returns {HTMLInputElement} - The created slider element.
*/
export const makeSlider = (min, max, value = undefined, step = undefined) => {
const slider = makeElement('input.mtb_slider', {
width: '100%',
})
slider.type = 'range'
slider.min = min || 0
slider.max = max || 100
slider.value = value || slider.min
slider.step = step || 1
return slider
}
/**
* Creates a button element.
*
* @param {string} label - The label for the button.
* @param {Object} [style] - Optional styles to apply to the button.
* @param {Function} [onClick] - Optional click handler.
* @returns {HTMLButtonElement} - The created button element.
*/
export const makeButton = (label, style = {}, onClick = undefined) => {
const button = makeElement('button.mtb_button', style)
button.textContent = label
if (onClick) {
button.addEventListener('click', onClick)
}
return button
}
/**
* Creates a resizable splitter between two elements.
*
* @param {HTMLElement} el1 - The first element.
* @param {HTMLElement} el2 - The second element.
* @param {'vertical' | 'horizontal'} direction - Splitter direction (vertical or horizontal).
* @param {'absolute' | 'normal'} mode - Splitter mode: 'absolute' for free resizing, 'normal' for layout-based resizing.
* @returns {HTMLDivElement} - The container with resizable splitter.
*/
export const makeSplitter = (
el1,
el2,
direction = 'vertical',
mode = 'normal',
) => {
const container = makeElement('div.mtb_splitter_container', {
display: mode === 'absolute' ? 'block' : 'flex',
flexDirection: direction === 'vertical' ? 'row' : 'column',
position: mode === 'absolute' ? 'relative' : 'static',
height: '100%',
width: '100%',
})
const handle = makeElement('div.mtb_splitter_handle', {
backgroundColor: '#ccc',
cursor: direction === 'vertical' ? 'col-resize' : 'row-resize',
width: direction === 'vertical' ? '5px' : '100%',
height: direction === 'horizontal' ? '5px' : '100%',
})
let isResizing = false
handle.addEventListener('mousedown', () => {
isResizing = true
})
window.addEventListener('mouseup', () => {
isResizing = false
})
window.addEventListener('mousemove', (e) => {
if (!isResizing) return
if (direction === 'vertical') {
const newWidth = e.clientX - container.offsetLeft
el1.style.width = `${newWidth}px`
el2.style.width = `${container.offsetWidth - newWidth}px`
} else {
const newHeight = e.clientY - container.offsetTop
el1.style.height = `${newHeight}px`
el2.style.height = `${container.offsetHeight - newHeight}px`
}
})
container.appendChild(el1)
container.appendChild(handle)
container.appendChild(el2)
return container
}
+1150 -1095
View File
File diff suppressed because it is too large Load Diff
+1 -1
Submodule wiki updated: 4db733ae92...a402de4af9