Compare commits
21
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
50cb6f5ed6 | ||
|
|
e32d1e02df | ||
|
|
b0d52f7305 | ||
|
|
e17c6e29f5 | ||
|
|
27e03fa23e | ||
|
|
ec1cb1ac17 | ||
|
|
64634104a2 | ||
|
|
ecbb220de6 | ||
|
|
cd9e614b1a | ||
|
|
9ccf572a15 | ||
|
|
74af5c6499 | ||
|
|
caf0b39d8a | ||
|
|
e099d581a7 | ||
|
|
22f7c30373 | ||
|
|
0133fb93bc | ||
|
|
cf7d30507e | ||
|
|
b6fa571fd2 | ||
|
|
f272526bfc | ||
|
|
4e593bb30b | ||
|
|
097ca33b8e | ||
|
|
784fb0145b |
@@ -6,3 +6,6 @@ node_modules/
|
||||
compose.yaml
|
||||
comfy_mtb.wsb
|
||||
Dockerfile
|
||||
|
||||
# I store the gh-pages worktrees (src & build) there
|
||||
.worktrees
|
||||
|
||||
+207
-47
@@ -7,10 +7,12 @@
|
||||
#
|
||||
###
|
||||
|
||||
__version__ = "0.1.6"
|
||||
__version__ = "0.2.0"
|
||||
|
||||
import os
|
||||
|
||||
from aiohttp.web_request import Request
|
||||
|
||||
# TODO: don't override this if the user has that setup already
|
||||
if not os.environ.get("TF_FORCE_GPU_ALLOW_GROWTH"):
|
||||
os.environ["TF_FORCE_GPU_ALLOW_GROWTH"] = "true"
|
||||
@@ -31,24 +33,20 @@ from pathlib import Path
|
||||
from aiohttp import web
|
||||
from server import PromptServer
|
||||
|
||||
import nodes
|
||||
|
||||
from .endpoint import endlog
|
||||
from .log import blue_text, cyan_text, get_label, get_summary, log
|
||||
from .utils import comfy_dir, here
|
||||
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
NODE_CLASS_MAPPINGS_DEBUG = {}
|
||||
NODE_CLASS_MAPPINGS: dict[str, type] = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS: dict[str, str] = {}
|
||||
NODE_CLASS_MAPPINGS_DEBUG: dict[str, str | None] = {}
|
||||
WEB_DIRECTORY = "./web"
|
||||
|
||||
|
||||
def extract_nodes_from_source(filename: Path):
|
||||
source_code = ""
|
||||
|
||||
source_code = filename.read_text(encoding="utf-8")
|
||||
|
||||
nodes = []
|
||||
nodes: list[str] = []
|
||||
|
||||
try:
|
||||
parsed = ast.parse(source_code)
|
||||
@@ -57,14 +55,15 @@ def extract_nodes_from_source(filename: Path):
|
||||
target = node.targets[0]
|
||||
if isinstance(target, ast.Name) and target.id == "__nodes__":
|
||||
value = ast.get_source_segment(source_code, node.value)
|
||||
node_value = ast.parse(value).body[0].value
|
||||
if isinstance(node_value, (ast.List, ast.Tuple)):
|
||||
nodes.extend(
|
||||
element.id
|
||||
for element in node_value.elts
|
||||
if isinstance(element, ast.Name)
|
||||
)
|
||||
break
|
||||
if value:
|
||||
node_value = ast.parse(value).body[0].value
|
||||
if isinstance(node_value, ast.List | ast.Tuple):
|
||||
nodes.extend(
|
||||
str(element.id)
|
||||
for element in node_value.elts
|
||||
if isinstance(element, ast.Name)
|
||||
)
|
||||
break
|
||||
except SyntaxError:
|
||||
log.error("Failed to parse")
|
||||
return nodes
|
||||
@@ -72,8 +71,8 @@ def extract_nodes_from_source(filename: Path):
|
||||
|
||||
def load_nodes():
|
||||
errors: list[str] = []
|
||||
nodes = []
|
||||
nodes_failed = []
|
||||
nodes: list[type] = []
|
||||
nodes_failed: list[str] = []
|
||||
|
||||
for filename in (here / "nodes").iterdir():
|
||||
if filename.suffix == ".py":
|
||||
@@ -124,7 +123,8 @@ def uninstall_old_web_extensions():
|
||||
shutil.rmtree(web_mtb)
|
||||
except Exception as e:
|
||||
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):
|
||||
classname = s.replace("MTB_", "")
|
||||
parts = []
|
||||
parts: list[str] = []
|
||||
start = 0
|
||||
for i in range(1, len(classname)):
|
||||
if classname[i].isupper():
|
||||
@@ -161,8 +161,6 @@ if wiki.exists() and wiki.is_dir():
|
||||
|
||||
|
||||
# - REGISTER NODES
|
||||
|
||||
|
||||
MTB_EXPORT = os.environ.get("MTB_EXPORT")
|
||||
|
||||
nodes, failed = load_nodes()
|
||||
@@ -179,7 +177,7 @@ for node_class in nodes:
|
||||
node_class.DESCRIPTION = node_class.__doc__
|
||||
if MTB_EXPORT:
|
||||
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"
|
||||
)
|
||||
|
||||
@@ -192,12 +190,15 @@ for node_class in nodes:
|
||||
NODE_CLASS_MAPPINGS[node_label] = node_class
|
||||
NODE_DISPLAY_NAME_MAPPINGS[class_name] = node_label
|
||||
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"):
|
||||
with open(here / "node_list.json", "w") as f:
|
||||
f.write(
|
||||
_ = f.write(
|
||||
json.dumps(
|
||||
{
|
||||
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:
|
||||
with contextlib.suppress(Exception):
|
||||
base_url, port = utils.get_server_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."
|
||||
)
|
||||
log.debug(failed)
|
||||
|
||||
|
||||
# - ENDPOINT
|
||||
|
||||
|
||||
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"]
|
||||
onnx_deps = ["onnxruntime"]
|
||||
swap_deps = ["insightface"] + onnx_deps
|
||||
@@ -306,10 +318,10 @@ if hasattr(PromptServer, "instance"):
|
||||
}
|
||||
)
|
||||
|
||||
@PromptServer.instance.routes.post("/mtb/debug")
|
||||
async def set_debug(request):
|
||||
json_data = await request.json()
|
||||
enabled = json_data.get("enabled")
|
||||
@PromptServer.instance.routes.post("/mtb/server-info")
|
||||
async def set_server_info(request: Request):
|
||||
json_data: dict[str, bool] = await request.json()
|
||||
enabled = json_data.get("debug")
|
||||
if enabled:
|
||||
os.environ["MTB_DEBUG"] = "true"
|
||||
log.setLevel(logging.DEBUG)
|
||||
@@ -317,7 +329,7 @@ if hasattr(PromptServer, "instance"):
|
||||
|
||||
elif "MTB_DEBUG" in os.environ:
|
||||
# del os.environ["MTB_DEBUG"]
|
||||
os.environ.pop("MTB_DEBUG")
|
||||
_ = os.environ.pop("MTB_DEBUG")
|
||||
log.setLevel(logging.INFO)
|
||||
|
||||
return web.json_response(
|
||||
@@ -325,17 +337,17 @@ if hasattr(PromptServer, "instance"):
|
||||
)
|
||||
|
||||
@PromptServer.instance.routes.get("/mtb")
|
||||
async def get_home(request):
|
||||
async def get_home(request: Request):
|
||||
from . import endpoint
|
||||
|
||||
reload(endpoint)
|
||||
_ = reload(endpoint)
|
||||
# Check if the request prefers HTML content
|
||||
if "text/html" in request.headers.get("Accept", ""):
|
||||
# # Return an HTML page
|
||||
html_response = """
|
||||
<div class="flex-container menu">
|
||||
<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>
|
||||
</div>
|
||||
"""
|
||||
@@ -347,28 +359,176 @@ if hasattr(PromptServer, "instance"):
|
||||
# Return JSON for other requests
|
||||
return web.json_response({"message": "Welcome to MTB!"})
|
||||
|
||||
@PromptServer.instance.routes.get("/mtb/debug")
|
||||
async def get_debug(request):
|
||||
import asyncio
|
||||
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
|
||||
|
||||
reload(endpoint)
|
||||
enabled = "MTB_DEBUG" in os.environ
|
||||
_ = reload(endpoint)
|
||||
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
|
||||
if "text/html" in request.headers.get("Accept", ""):
|
||||
# # Return an HTML page
|
||||
html_response = f"""
|
||||
<h1>MTB Debug Status: {'Enabled' if enabled else 'Disabled'}</h1>
|
||||
"""
|
||||
html_response = ""
|
||||
|
||||
html_response += render_property(
|
||||
"Debug", "Enabled" if isdebug else "Disabled"
|
||||
)
|
||||
|
||||
html_response += render_property("Exposed", str(exposed))
|
||||
|
||||
return web.Response(
|
||||
text=endpoint.render_base_template("Debug", html_response),
|
||||
text=endpoint.render_base_template(
|
||||
"Server Info", html_response
|
||||
),
|
||||
content_type="text/html",
|
||||
)
|
||||
|
||||
# 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")
|
||||
async def no_route(request):
|
||||
async def no_route(request: Request):
|
||||
from . import endpoint
|
||||
|
||||
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"})
|
||||
|
||||
@PromptServer.instance.routes.post("/mtb/actions")
|
||||
async def do_action(request):
|
||||
async def do_action(request: Request):
|
||||
from . import endpoint
|
||||
|
||||
reload(endpoint)
|
||||
|
||||
+95
-19
@@ -1,4 +1,8 @@
|
||||
import csv
|
||||
import secrets
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from aiohttp import web
|
||||
|
||||
@@ -6,6 +10,8 @@ from .log import mklog
|
||||
from .utils import (
|
||||
backup_file,
|
||||
import_install,
|
||||
input_dir,
|
||||
output_dir,
|
||||
reqs_map,
|
||||
run_command,
|
||||
styles_dir,
|
||||
@@ -14,15 +20,14 @@ from .utils import (
|
||||
endlog = mklog("mtb endpoint")
|
||||
|
||||
# - ACTIONS
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import_install("requirements")
|
||||
|
||||
|
||||
def ACTIONS_installDependency(dependency_names=None):
|
||||
if dependency_names is None:
|
||||
# return web.Response(text="No dependency name provided", status=400)
|
||||
return {"error": "No dependency name provided"}
|
||||
|
||||
endlog.debug(f"Received Install Dependency request for {dependency_names}")
|
||||
# reqs = []
|
||||
resolved_names = [reqs_map.get(name, name) for name in dependency_names]
|
||||
@@ -50,6 +55,62 @@ def ACTIONS_installDependency(dependency_names=None):
|
||||
# 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):
|
||||
from .nodes.conditions import MTB_StylesLoader
|
||||
|
||||
@@ -97,7 +158,7 @@ def ACTIONS_saveStyle(data):
|
||||
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")
|
||||
request_data = await request.json()
|
||||
name = request_data.get("name")
|
||||
@@ -109,7 +170,12 @@ async def do_action(request) -> web.Response:
|
||||
method = globals().get(method_name)
|
||||
|
||||
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}")
|
||||
return web.json_response({"result": result})
|
||||
|
||||
@@ -130,10 +196,13 @@ async def do_action(request) -> web.Response:
|
||||
# - HTML UTILS
|
||||
|
||||
|
||||
def dependencies_button(name, dependencies):
|
||||
def dependencies_button(name: str, dependencies: list[str]) -> str:
|
||||
deps = ",".join([f"'{x}'" for x in dependencies])
|
||||
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 = """
|
||||
<div id="style-editor">
|
||||
<h1>Style Editor</h1>
|
||||
|
||||
|
||||
"""
|
||||
for current, styles in style_files.items():
|
||||
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"
|
||||
return f"""
|
||||
<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'>▷</span>
|
||||
{title}
|
||||
</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"
|
||||
return f"""
|
||||
<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(
|
||||
[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):
|
||||
table_dict = sorted(
|
||||
def render_table(table_dict: dict[str, Any], sort=True, title=None):
|
||||
table_list = sorted(
|
||||
table_dict.items(), key=lambda item: item[0]
|
||||
) # Sort the dictionary by keys
|
||||
|
||||
table_rows = ""
|
||||
for name, item in table_dict:
|
||||
for name, item in table_list:
|
||||
if isinstance(item, dict):
|
||||
if "dependencies" in item:
|
||||
table_rows += f"<tr><td>{name}</td><td>"
|
||||
@@ -299,12 +373,12 @@ def render_table(table_dict, sort=True, title=None):
|
||||
<tbody>
|
||||
{table_rows}
|
||||
</tbody>
|
||||
</table>
|
||||
</table>
|
||||
</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>"""
|
||||
return f"""
|
||||
<!DOCTYPE html>
|
||||
@@ -340,7 +414,9 @@ def render_base_template(title, content):
|
||||
<header>
|
||||
<a href="/">Back to Comfy</a>
|
||||
<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>
|
||||
<a style="width:128px;text-align:center" href="https://www.github.com/melmass/comfy_mtb">
|
||||
{github_icon_svg}
|
||||
@@ -355,6 +431,6 @@ def render_base_template(title, content):
|
||||
<!-- Shared footer content here -->
|
||||
</footer>
|
||||
</body>
|
||||
|
||||
|
||||
</html>
|
||||
"""
|
||||
|
||||
@@ -27,7 +27,7 @@ export def "comfy start" [--clean,--old-ui, --listen] {
|
||||
|
||||
let root = get_root --clean=($clean)
|
||||
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
|
||||
@@ -67,8 +67,14 @@ export def "comfy update" [
|
||||
git checkout master
|
||||
|
||||
print $"(ansi yellow_italic)Fetching and pulling remote updates(ansi reset)"
|
||||
git fetch
|
||||
git pull
|
||||
if ($clean) {
|
||||
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)"
|
||||
git checkout -
|
||||
@@ -135,7 +141,7 @@ export def "comfy update_extensions" [--clean] {
|
||||
let root = get_root --clean=($clean)
|
||||
cd $root
|
||||
cd custom_nodes
|
||||
git multipull .
|
||||
git multipull . -s -q
|
||||
}
|
||||
|
||||
def --env path-add [pth] {
|
||||
@@ -146,7 +152,7 @@ def --env path-add [pth] {
|
||||
|
||||
export-env {
|
||||
$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
|
||||
|
||||
|
||||
+123
-1
@@ -3,10 +3,127 @@ import shutil
|
||||
from pathlib import Path
|
||||
|
||||
import folder_paths
|
||||
import torch
|
||||
|
||||
from ..log import log
|
||||
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:
|
||||
@classmethod
|
||||
@@ -213,4 +330,9 @@ class MTB_StylesLoader:
|
||||
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
@@ -59,6 +59,36 @@ class MTB_SplitBbox:
|
||||
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:
|
||||
"""From a mask extract the bounding box"""
|
||||
|
||||
@@ -342,4 +372,11 @@ class MTB_Uncrop:
|
||||
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,
|
||||
]
|
||||
|
||||
@@ -78,6 +78,7 @@ class MTB_LoadFaceEnhanceModel:
|
||||
RETURN_NAMES = ("model",)
|
||||
FUNCTION = "load_model"
|
||||
CATEGORY = "mtb/facetools"
|
||||
DEPRECATED = True
|
||||
|
||||
def load_model(self, model_name, upscale=2, bg_upsampler=None):
|
||||
from gfpgan import GFPGANer
|
||||
@@ -163,6 +164,7 @@ class MTB_RestoreFace:
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "restore"
|
||||
CATEGORY = "mtb/facetools"
|
||||
DEPRECATED = True
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
|
||||
@@ -40,6 +40,7 @@ class MTB_LoadFaceAnalysisModel:
|
||||
RETURN_TYPES = ("FACE_ANALYSIS_MODEL",)
|
||||
FUNCTION = "load_model"
|
||||
CATEGORY = "mtb/facetools"
|
||||
DEPRECATED = True
|
||||
|
||||
def load_model(self, faceswap_model: str):
|
||||
if faceswap_model == "antelopev2":
|
||||
@@ -77,6 +78,7 @@ class MTB_LoadFaceSwapModel:
|
||||
RETURN_TYPES = ("FACESWAP_MODEL",)
|
||||
FUNCTION = "load_model"
|
||||
CATEGORY = "mtb/facetools"
|
||||
DEPRECATED = True
|
||||
|
||||
def load_model(self, faceswap_model: str):
|
||||
model_path = get_model_path("insightface", faceswap_model)
|
||||
@@ -126,6 +128,7 @@ class MTB_FaceSwap:
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "swap"
|
||||
CATEGORY = "mtb/facetools"
|
||||
DEPRECATED = True
|
||||
|
||||
def swap(
|
||||
self,
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
from pathlib import Path
|
||||
from typing import List
|
||||
|
||||
import comfy
|
||||
import comfy.model_management as model_management
|
||||
@@ -15,10 +14,13 @@ from ..utils import get_model_path
|
||||
|
||||
|
||||
class MTB_LoadFilmModel:
|
||||
"""Loads a FILM model"""
|
||||
"""Loads a FILM model
|
||||
|
||||
[DEPRECATED] Use ComfyUI-FrameInterpolation instead
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def get_models() -> List[Path]:
|
||||
def get_models() -> list[Path]:
|
||||
models_paths = get_model_path("FILM").iterdir()
|
||||
|
||||
return [x for x in models_paths if x.suffix in [".onnx", ".pth"]]
|
||||
@@ -37,6 +39,7 @@ class MTB_LoadFilmModel:
|
||||
RETURN_TYPES = ("FILM_MODEL",)
|
||||
FUNCTION = "load_model"
|
||||
CATEGORY = "mtb/frame iterpolation"
|
||||
DEPRECATED = True
|
||||
|
||||
def load_model(self, film_model: str):
|
||||
model_path = get_model_path("FILM", film_model)
|
||||
@@ -56,7 +59,10 @@ class MTB_LoadFilmModel:
|
||||
|
||||
|
||||
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
|
||||
def INPUT_TYPES(cls):
|
||||
@@ -71,6 +77,7 @@ class MTB_FilmInterpolation:
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "do_interpolation"
|
||||
CATEGORY = "mtb/frame iterpolation"
|
||||
DEPRECATED = True
|
||||
|
||||
def do_interpolation(
|
||||
self,
|
||||
|
||||
+55
-19
@@ -2,9 +2,9 @@ import json
|
||||
import subprocess
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import List, Optional
|
||||
|
||||
import comfy.model_management as model_management
|
||||
import comfy.utils
|
||||
import folder_paths
|
||||
import numpy as np
|
||||
import torch
|
||||
@@ -41,6 +41,7 @@ class MTB_ReadPlaylist:
|
||||
RETURN_TYPES = ("PLAYLIST",)
|
||||
FUNCTION = "read_playlist"
|
||||
CATEGORY = "mtb/IO"
|
||||
EXPERIMENTAL = True
|
||||
|
||||
def read_playlist(
|
||||
self,
|
||||
@@ -83,6 +84,7 @@ class MTB_AddToPlaylist:
|
||||
OUTPUT_NODE = True
|
||||
FUNCTION = "add_to_playlist"
|
||||
CATEGORY = "mtb/IO"
|
||||
EXPERIMENTAL = True
|
||||
|
||||
def add_to_playlist(
|
||||
self,
|
||||
@@ -117,7 +119,10 @@ class MTB_AddToPlaylist:
|
||||
|
||||
|
||||
class MTB_ExportWithFfmpeg:
|
||||
"""Export with FFmpeg (Experimental)"""
|
||||
"""Export with FFmpeg (Experimental).
|
||||
|
||||
[DEPRACATED] Use VHS nodes instead
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
@@ -143,6 +148,7 @@ class MTB_ExportWithFfmpeg:
|
||||
RETURN_TYPES = ("VIDEO",)
|
||||
OUTPUT_NODE = True
|
||||
FUNCTION = "export_prores"
|
||||
DEPRECATED = True
|
||||
CATEGORY = "mtb/IO"
|
||||
|
||||
def export_prores(
|
||||
@@ -151,10 +157,9 @@ class MTB_ExportWithFfmpeg:
|
||||
prefix: str,
|
||||
format: str,
|
||||
codec: str,
|
||||
images: Optional[torch.Tensor] = None,
|
||||
playlist: Optional[List[str]] = None,
|
||||
images: torch.Tensor | None = None,
|
||||
playlist: list[str] | None = None,
|
||||
):
|
||||
pix_fmt = "rgb48le" if codec == "prores_ks" else "yuv420p"
|
||||
file_ext = format
|
||||
file_id = f"{prefix}_{uuid.uuid4()}.{file_ext}"
|
||||
|
||||
@@ -208,9 +213,11 @@ class MTB_ExportWithFfmpeg:
|
||||
frames = tensor2np(images)
|
||||
log.debug(f"Frames type {type(frames[0])}")
|
||||
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":
|
||||
out_path = (output_dir / file_id).as_posix()
|
||||
command = [
|
||||
"ffmpeg",
|
||||
"-f",
|
||||
@@ -233,12 +240,28 @@ class MTB_ExportWithFfmpeg:
|
||||
|
||||
process.stdin.close()
|
||||
process.wait()
|
||||
return (out_path,)
|
||||
else:
|
||||
frames = [frame.astype(np.uint16) * 257 for frame in frames]
|
||||
|
||||
height, width, _ = frames[0].shape
|
||||
|
||||
out_path = (output_dir / file_id).as_posix()
|
||||
if has_alpha:
|
||||
if codec in ["prores_ks", "libx264", "libx265"]:
|
||||
pix_fmt = (
|
||||
"yuva444p" if codec == "prores_ks" else "yuva420p"
|
||||
)
|
||||
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
|
||||
command = [
|
||||
@@ -258,17 +281,26 @@ class MTB_ExportWithFfmpeg:
|
||||
"-",
|
||||
"-c:v",
|
||||
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)
|
||||
|
||||
pbar = comfy.utils.ProgressBar(len(frames))
|
||||
|
||||
for frame in frames:
|
||||
model_management.throw_exception_if_processing_interrupted()
|
||||
process.stdin.write(frame.tobytes())
|
||||
pbar.update(1)
|
||||
|
||||
process.stdin.close()
|
||||
process.wait()
|
||||
@@ -280,9 +312,9 @@ def prepare_animated_batch(
|
||||
batch: torch.Tensor,
|
||||
pingpong=False,
|
||||
resize_by=1.0,
|
||||
resample_filter: Optional[Image.Resampling] = None,
|
||||
resample_filter: Image.Resampling | None = None,
|
||||
image_type=np.uint8,
|
||||
) -> List[Image.Image]:
|
||||
) -> list[Image.Image]:
|
||||
images = tensor2np(batch)
|
||||
images = [frame.astype(image_type) for frame in images]
|
||||
|
||||
@@ -308,7 +340,10 @@ def prepare_animated_batch(
|
||||
|
||||
# todo: deprecate for apng
|
||||
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
|
||||
def INPUT_TYPES(cls):
|
||||
@@ -328,6 +363,7 @@ class MTB_SaveGif:
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = "mtb/IO"
|
||||
FUNCTION = "save_gif"
|
||||
DEPRECATED = True
|
||||
|
||||
def save_gif(
|
||||
self,
|
||||
|
||||
+157
@@ -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
@@ -1,6 +1,5 @@
|
||||
import comfy.utils
|
||||
from PIL import Image
|
||||
from rembg import remove
|
||||
|
||||
from ..utils import pil2tensor, tensor2pil
|
||||
|
||||
@@ -64,6 +63,8 @@ class MTB_ImageRemoveBackgroundRembg:
|
||||
post_process_mask,
|
||||
bgcolor,
|
||||
):
|
||||
from rembg import remove
|
||||
|
||||
pbar = comfy.utils.ProgressBar(image.size(0))
|
||||
images = tensor2pil(image)
|
||||
|
||||
|
||||
@@ -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
@@ -1,179 +1,180 @@
|
||||
[build-system]
|
||||
requires = ["setuptools", "wheel"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "comfy-mtb"
|
||||
version = "0.1.6"
|
||||
description = "Animation oriented nodes pack for ComfyUI."
|
||||
license = "MIT"
|
||||
readme = "README.md"
|
||||
# repository = ""
|
||||
# url = "https://github.com/melMass/comfy_mtb"
|
||||
authors = [{ name = "Mel Massadian", email = "mel@melmassadian.com" }]
|
||||
classifiers = [
|
||||
"License :: OSI Approved :: MIT License",
|
||||
"Operating System :: OS Independent",
|
||||
"Programming Language :: Python",
|
||||
"Programming Language :: Python :: 3",
|
||||
"Programming Language :: Python :: 3.10",
|
||||
"Programming Language :: Python :: 3.11",
|
||||
"Intended Audience :: Developers",
|
||||
]
|
||||
requires-python = ">=3.10"
|
||||
dependencies = [
|
||||
"qrcode",
|
||||
"onnxruntime-gpu",
|
||||
"requirements-parserx",
|
||||
"rembg",
|
||||
"imageio_ffmpeg",
|
||||
"rich",
|
||||
"rich_argparse",
|
||||
"matplotlib",
|
||||
"pillow",
|
||||
]
|
||||
optional-dependencies = { mel = [
|
||||
"jupyterlab==4.1.6",
|
||||
], dev = [
|
||||
"black[jupyter]",
|
||||
"codespell",
|
||||
"mypy",
|
||||
"pre-commit",
|
||||
"pytest",
|
||||
"pytest-cov",
|
||||
"pytest-random-order",
|
||||
"ruff",
|
||||
], doc = [
|
||||
"docutils==0.17.1",
|
||||
"jupyter-book>=0.15",
|
||||
"sphinx-autobuild",
|
||||
] }
|
||||
|
||||
[project.urls]
|
||||
Homepage = "https://github.com/melMass/comfy_mtb"
|
||||
Documentation = "https://github.com/melMass/comfy_mtb/wiki"
|
||||
Repository = "https://github.com/melMass/comfy_mtb"
|
||||
Issues = "https://github.com/melMass/comfy_mtb/issues"
|
||||
|
||||
[tool.comfy]
|
||||
PublisherId = "mel"
|
||||
DisplayName = "comfy-mtb"
|
||||
Icon = "https://avatars.githubusercontent.com/u/7041726?v=4"
|
||||
|
||||
[tool.bumpversion]
|
||||
current_version = "0.1.6"
|
||||
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
|
||||
serialize = ["{major}.{minor}.{patch}"]
|
||||
search = "{current_version}"
|
||||
replace = "{new_version}"
|
||||
regex = false
|
||||
ignore_missing_version = false
|
||||
ignore_missing_files = false
|
||||
tag = true
|
||||
sign_tags = true
|
||||
tag_name = "v{new_version}"
|
||||
tag_message = "⬆️ Bump version: {current_version} → {new_version}"
|
||||
allow_dirty = true
|
||||
commit = true
|
||||
message = "⬆️ Bump version: {current_version} → {new_version}"
|
||||
commit_args = ""
|
||||
|
||||
[[tool.bumpversion.files]]
|
||||
filename = "__init__.py"
|
||||
search = "__version__ = \"{current_version}\""
|
||||
replace = "__version__ = \"{new_version}\""
|
||||
|
||||
[[tool.bumpversion.files]]
|
||||
filename = "pyproject.toml"
|
||||
search = "version = \"{current_version}\""
|
||||
replace = "version = \"{new_version}\""
|
||||
|
||||
# [[tool.bumpversion.files]]
|
||||
# filename = "your_package/__init__.py"
|
||||
# search = "__version__ = '{current_version}'"
|
||||
# replace = "__version__ = '{new_version}'"
|
||||
|
||||
# INFO: All those remaining keys are meant for local dev
|
||||
[tool.pyright]
|
||||
include = ["."]
|
||||
exclude = [
|
||||
"**/node_modules",
|
||||
"**/__pycache__",
|
||||
"src/experimental",
|
||||
"src/typestubs",
|
||||
]
|
||||
ignore = ["src/oldstuff"]
|
||||
defineConstant = { DEBUG = true }
|
||||
extraPaths = ["python", "../.."]
|
||||
stubPath = "src/stubs"
|
||||
|
||||
reportMissingImports = true
|
||||
reportMissingTypeStubs = false
|
||||
typeCheckingMode = "basic"
|
||||
|
||||
pythonVersion = "3.10"
|
||||
pythonPlatform = "Windows"
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
log_level = "DEBUG"
|
||||
log_cli = true
|
||||
markers = [
|
||||
"wip: tests that aren't fully finished yet",
|
||||
"heavy: marks tests as heavy (deselect with '-m \"not heavy\"')",
|
||||
|
||||
]
|
||||
filterwarnings = ["ignore::UserWarning", 'ignore::DeprecationWarning']
|
||||
|
||||
[tool.isort]
|
||||
profile = "black"
|
||||
line_length = 88
|
||||
auto_identify_namespace_packages = false
|
||||
# NOTE:
|
||||
# pyright doesn't like implicit namespace + single line (related to https://github.com/microsoft/pyright/issues/2882?) but it's horible so I'll live with it
|
||||
force_single_line = false
|
||||
known_first_party = ["mtb"]
|
||||
extend_skip = ["archives"]
|
||||
combine_straight_imports = true
|
||||
|
||||
[tool.coverage.run]
|
||||
parallel = true
|
||||
source = ["docs", "tests", "comfy-mtb"]
|
||||
|
||||
[tool.coverage.report]
|
||||
fail_under = 90
|
||||
show_missing = true
|
||||
|
||||
[tool.coverage.html]
|
||||
show_contexts = true
|
||||
|
||||
[tool.ruff]
|
||||
line-length = 79
|
||||
select = ["A", "B", "C", "D", "E", "F", "FBT", "I", "N", "S", "SIM", "UP", "W"]
|
||||
# NOTE:
|
||||
# D102 - undocumented-public-method (noisy)
|
||||
# D103 - undocumented-public-function (noisy)
|
||||
# D100 - undocumented-public-module (noisy)
|
||||
# N802 - invalid-function-name (forced by comfy's arch)
|
||||
ignore = ["D103", "D102", "D100", "N802"]
|
||||
# exclude auto generated file
|
||||
extend-exclude = ["./docs/conf.py"]
|
||||
|
||||
[tool.ruff.per-file-ignores]
|
||||
# imported but unused
|
||||
"__init__.py" = ["F401"]
|
||||
# use of assert detected
|
||||
"tests/*" = ["S101"]
|
||||
|
||||
[tool.ruff.pydocstyle]
|
||||
convention = "numpy"
|
||||
|
||||
[tool.mypy]
|
||||
pretty = true
|
||||
ignore_missing_imports = true
|
||||
# exclude auto generated file
|
||||
exclude = ["docs/conf.py"]
|
||||
|
||||
[tool.codespell]
|
||||
# exclude auto generated file
|
||||
skip = "./docs/conf.py,poetry.lock"
|
||||
check-filenames = true
|
||||
[build-system]
|
||||
requires = ["setuptools", "wheel"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "comfy-mtb"
|
||||
version = "0.2.0"
|
||||
description = "Animation oriented nodes pack for ComfyUI."
|
||||
license = "MIT"
|
||||
readme = "README.md"
|
||||
# repository = ""
|
||||
# url = "https://github.com/melMass/comfy_mtb"
|
||||
authors = [{ name = "Mel Massadian", email = "mel@melmassadian.com" }]
|
||||
classifiers = [
|
||||
"License :: OSI Approved :: MIT License",
|
||||
"Operating System :: OS Independent",
|
||||
"Programming Language :: Python",
|
||||
"Programming Language :: Python :: 3",
|
||||
"Programming Language :: Python :: 3.10",
|
||||
"Programming Language :: Python :: 3.11",
|
||||
"Intended Audience :: Developers",
|
||||
]
|
||||
requires-python = ">=3.10"
|
||||
dependencies = [
|
||||
"qrcode",
|
||||
"cachetools",
|
||||
"onnxruntime-gpu",
|
||||
"requirements-parserx",
|
||||
"rembg",
|
||||
"imageio_ffmpeg",
|
||||
"rich",
|
||||
"rich_argparse",
|
||||
"matplotlib",
|
||||
"pillow",
|
||||
]
|
||||
optional-dependencies = { mel = [
|
||||
"jupyterlab==4.1.6",
|
||||
], dev = [
|
||||
"black[jupyter]",
|
||||
"codespell",
|
||||
"mypy",
|
||||
"pre-commit",
|
||||
"pytest",
|
||||
"pytest-cov",
|
||||
"pytest-random-order",
|
||||
"ruff",
|
||||
], doc = [
|
||||
"docutils==0.17.1",
|
||||
"jupyter-book>=0.15",
|
||||
"sphinx-autobuild",
|
||||
] }
|
||||
|
||||
[project.urls]
|
||||
Homepage = "https://github.com/melMass/comfy_mtb"
|
||||
Documentation = "https://github.com/melMass/comfy_mtb/wiki"
|
||||
Repository = "https://github.com/melMass/comfy_mtb"
|
||||
Issues = "https://github.com/melMass/comfy_mtb/issues"
|
||||
|
||||
[tool.comfy]
|
||||
PublisherId = "mel"
|
||||
DisplayName = "comfy-mtb"
|
||||
Icon = "https://avatars.githubusercontent.com/u/7041726?v=4"
|
||||
|
||||
[tool.bumpversion]
|
||||
current_version = "0.2.0"
|
||||
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
|
||||
serialize = ["{major}.{minor}.{patch}"]
|
||||
search = "{current_version}"
|
||||
replace = "{new_version}"
|
||||
regex = false
|
||||
ignore_missing_version = false
|
||||
ignore_missing_files = false
|
||||
tag = true
|
||||
sign_tags = true
|
||||
tag_name = "v{new_version}"
|
||||
tag_message = "⬆️ Bump version: {current_version} → {new_version}"
|
||||
allow_dirty = true
|
||||
commit = true
|
||||
message = "⬆️ Bump version: {current_version} → {new_version}"
|
||||
commit_args = ""
|
||||
|
||||
[[tool.bumpversion.files]]
|
||||
filename = "__init__.py"
|
||||
search = "__version__ = \"{current_version}\""
|
||||
replace = "__version__ = \"{new_version}\""
|
||||
|
||||
[[tool.bumpversion.files]]
|
||||
filename = "pyproject.toml"
|
||||
search = "version = \"{current_version}\""
|
||||
replace = "version = \"{new_version}\""
|
||||
|
||||
# [[tool.bumpversion.files]]
|
||||
# filename = "your_package/__init__.py"
|
||||
# search = "__version__ = '{current_version}'"
|
||||
# replace = "__version__ = '{new_version}'"
|
||||
|
||||
# INFO: All those remaining keys are meant for local dev
|
||||
[tool.pyright]
|
||||
include = ["."]
|
||||
exclude = [
|
||||
"**/node_modules",
|
||||
"**/__pycache__",
|
||||
"src/experimental",
|
||||
"src/typestubs",
|
||||
]
|
||||
ignore = ["src/oldstuff"]
|
||||
defineConstant = { DEBUG = true }
|
||||
extraPaths = ["python", "../.."]
|
||||
stubPath = "src/stubs"
|
||||
|
||||
reportMissingImports = true
|
||||
reportMissingTypeStubs = false
|
||||
typeCheckingMode = "basic"
|
||||
|
||||
pythonVersion = "3.10"
|
||||
pythonPlatform = "Windows"
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
log_level = "DEBUG"
|
||||
log_cli = true
|
||||
markers = [
|
||||
"wip: tests that aren't fully finished yet",
|
||||
"heavy: marks tests as heavy (deselect with '-m \"not heavy\"')",
|
||||
|
||||
]
|
||||
filterwarnings = ["ignore::UserWarning", 'ignore::DeprecationWarning']
|
||||
|
||||
[tool.isort]
|
||||
profile = "black"
|
||||
line_length = 88
|
||||
auto_identify_namespace_packages = false
|
||||
# NOTE:
|
||||
# pyright doesn't like implicit namespace + single line (related to https://github.com/microsoft/pyright/issues/2882?) but it's horible so I'll live with it
|
||||
force_single_line = false
|
||||
known_first_party = ["mtb"]
|
||||
extend_skip = ["archives"]
|
||||
combine_straight_imports = true
|
||||
|
||||
[tool.coverage.run]
|
||||
parallel = true
|
||||
source = ["docs", "tests", "comfy-mtb"]
|
||||
|
||||
[tool.coverage.report]
|
||||
fail_under = 90
|
||||
show_missing = true
|
||||
|
||||
[tool.coverage.html]
|
||||
show_contexts = true
|
||||
|
||||
[tool.ruff]
|
||||
line-length = 79
|
||||
select = ["A", "B", "C", "D", "E", "F", "FBT", "I", "N", "S", "SIM", "UP", "W"]
|
||||
# NOTE:
|
||||
# D102 - undocumented-public-method (noisy)
|
||||
# D103 - undocumented-public-function (noisy)
|
||||
# D100 - undocumented-public-module (noisy)
|
||||
# N802 - invalid-function-name (forced by comfy's arch)
|
||||
ignore = ["D103", "D102", "D100", "N802"]
|
||||
# exclude auto generated file
|
||||
extend-exclude = ["./docs/conf.py"]
|
||||
|
||||
[tool.ruff.per-file-ignores]
|
||||
# imported but unused
|
||||
"__init__.py" = ["F401"]
|
||||
# use of assert detected
|
||||
"tests/*" = ["S101"]
|
||||
|
||||
[tool.ruff.pydocstyle]
|
||||
convention = "numpy"
|
||||
|
||||
[tool.mypy]
|
||||
pretty = true
|
||||
ignore_missing_imports = true
|
||||
# exclude auto generated file
|
||||
exclude = ["docs/conf.py"]
|
||||
|
||||
[tool.codespell]
|
||||
# exclude auto generated file
|
||||
skip = "./docs/conf.py,poetry.lock"
|
||||
check-filenames = true
|
||||
|
||||
@@ -8,3 +8,4 @@ rich
|
||||
rich_argparse
|
||||
matplotlib
|
||||
pillow
|
||||
cachetools
|
||||
|
||||
@@ -163,9 +163,9 @@ class IPChecker:
|
||||
def __init__(self):
|
||||
self.ips = list(self.get_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:
|
||||
if ip not in self.checked_ips:
|
||||
self.checked_ips.add(ip)
|
||||
@@ -175,7 +175,7 @@ class IPChecker:
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def get_local_ips(prefix="192.168."):
|
||||
def get_local_ips(prefix: str = "192.168."):
|
||||
hostname = socket.gethostname()
|
||||
log.debug(f"Getting local ips for {hostname}")
|
||||
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):
|
||||
yield info[4][0]
|
||||
|
||||
def _test_url(self, url):
|
||||
def _test_url(self, url: str):
|
||||
try:
|
||||
response = requests.get(url)
|
||||
response = requests.get(url, timeout=10)
|
||||
return response.status_code == 200
|
||||
except Exception:
|
||||
return False
|
||||
@@ -198,7 +198,7 @@ def get_server_info():
|
||||
from comfy.cli_args import args
|
||||
|
||||
ip_checker = IPChecker()
|
||||
base_url = args.listen
|
||||
base_url: str = args.listen
|
||||
if base_url == "0.0.0.0":
|
||||
log.debug("Server set to 0.0.0.0, we will try to resolve the host 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)
|
||||
|
||||
|
||||
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):
|
||||
pad_width = np.array(((0, 0), (top, bottom), (left, right)))
|
||||
print(
|
||||
|
||||
+10
-33
@@ -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 {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
|
||||
|
||||
// #region documentation widget
|
||||
|
||||
+296
-295
@@ -13,40 +13,40 @@ import { api } from '../../scripts/api.js'
|
||||
import { app } from '../../scripts/app.js'
|
||||
import { LocalStorageManager } from './comfy_shared.js'
|
||||
const styles = {
|
||||
lighbox: {
|
||||
position: 'fixed',
|
||||
top: 0,
|
||||
left: 0,
|
||||
width: '100vw',
|
||||
height: '100vh',
|
||||
background: 'rgba(0,0,0,0.5)',
|
||||
display: 'none',
|
||||
justifyContent: 'center',
|
||||
alignItems: 'center',
|
||||
zIndex: 999,
|
||||
},
|
||||
lightboxBtn: (extra) => ({
|
||||
position: 'absolute',
|
||||
top: '50%',
|
||||
background: 'none',
|
||||
border: 'none',
|
||||
color: '#fff',
|
||||
zIndex: 1000,
|
||||
fontSize: '30px',
|
||||
cursor: 'pointer',
|
||||
pointerEvents: 'auto',
|
||||
...extra,
|
||||
}),
|
||||
img_list: {
|
||||
minHeight: '30px',
|
||||
maxHeight: '300px',
|
||||
width: '100vw',
|
||||
position: 'absolute',
|
||||
bottom: 0,
|
||||
zIndex: 10,
|
||||
background: '#333',
|
||||
overflow: 'auto',
|
||||
},
|
||||
lighbox: {
|
||||
position: 'fixed',
|
||||
top: 0,
|
||||
left: 0,
|
||||
width: '100vw',
|
||||
height: '100vh',
|
||||
background: 'rgba(0,0,0,0.5)',
|
||||
display: 'none',
|
||||
justifyContent: 'center',
|
||||
alignItems: 'center',
|
||||
zIndex: 999,
|
||||
},
|
||||
lightboxBtn: (extra) => ({
|
||||
position: 'absolute',
|
||||
top: '50%',
|
||||
background: 'none',
|
||||
border: 'none',
|
||||
color: '#fff',
|
||||
zIndex: 1000,
|
||||
fontSize: '30px',
|
||||
cursor: 'pointer',
|
||||
pointerEvents: 'auto',
|
||||
...extra,
|
||||
}),
|
||||
img_list: {
|
||||
minHeight: '30px',
|
||||
maxHeight: '300px',
|
||||
width: '100vw',
|
||||
position: 'absolute',
|
||||
bottom: 0,
|
||||
zIndex: 10,
|
||||
background: '#333',
|
||||
overflow: 'auto',
|
||||
},
|
||||
}
|
||||
|
||||
let currentImageIndex = 0
|
||||
@@ -58,298 +58,299 @@ const storage = new LocalStorageManager('mtb')
|
||||
let activated = storage.get('image_feed', false)
|
||||
|
||||
app.registerExtension({
|
||||
name: 'mtb.ImageFeed',
|
||||
setup: () => {
|
||||
app.ui.settings.addSetting({
|
||||
id: 'mtb.imageFeed.enabled',
|
||||
name: '[⚡mtb] Enable image feed',
|
||||
type: 'boolean',
|
||||
defaultValue: true,
|
||||
attrs: {
|
||||
style: {
|
||||
fontFamily: 'monospace',
|
||||
},
|
||||
},
|
||||
async onChange(value) {
|
||||
storage.set('image_feed', value)
|
||||
activated = value
|
||||
},
|
||||
})
|
||||
},
|
||||
init: async () => {
|
||||
if (!activated) {
|
||||
return
|
||||
}
|
||||
const pythongossFeed = app.extensions.find(
|
||||
(e) => e.name === 'pysssss.ImageFeed',
|
||||
)
|
||||
if (pythongossFeed) {
|
||||
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
|
||||
}
|
||||
// - HTML & CSS
|
||||
//- lightbox
|
||||
const lightboxContainer = document.createElement('div')
|
||||
Object.assign(lightboxContainer.style, styles.lighbox)
|
||||
name: 'mtb.ImageFeed',
|
||||
setup: () => {
|
||||
app.ui.settings.addSetting({
|
||||
id: 'mtb.Main.image-feed-enabled',
|
||||
category: ['mtb', 'Main', 'image-feed-enabled'],
|
||||
name: 'Enable Image Feed',
|
||||
type: 'boolean',
|
||||
defaultValue: false,
|
||||
attrs: {
|
||||
style: {
|
||||
fontFamily: 'monospace',
|
||||
},
|
||||
},
|
||||
async onChange(value) {
|
||||
storage.set('image_feed', value)
|
||||
activated = value
|
||||
},
|
||||
})
|
||||
},
|
||||
init: async () => {
|
||||
if (!activated) {
|
||||
return
|
||||
}
|
||||
const pythongossFeed = app.extensions.find(
|
||||
(e) => e.name === 'pysssss.ImageFeed',
|
||||
)
|
||||
if (pythongossFeed) {
|
||||
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
|
||||
}
|
||||
// - HTML & CSS
|
||||
//- lightbox
|
||||
const lightboxContainer = document.createElement('div')
|
||||
Object.assign(lightboxContainer.style, styles.lighbox)
|
||||
|
||||
const lightboxImage = document.createElement('img')
|
||||
Object.assign(lightboxImage.style, {
|
||||
maxHeight: '100%',
|
||||
maxWidth: '100%',
|
||||
borderRadius: '5px',
|
||||
})
|
||||
const lightboxImage = document.createElement('img')
|
||||
Object.assign(lightboxImage.style, {
|
||||
maxHeight: '100%',
|
||||
maxWidth: '100%',
|
||||
borderRadius: '5px',
|
||||
})
|
||||
|
||||
// previous and next buttons
|
||||
const lightboxPrevBtn = document.createElement('button')
|
||||
const lightboxNextBtn = document.createElement('button')
|
||||
// previous and next buttons
|
||||
const lightboxPrevBtn = document.createElement('button')
|
||||
const lightboxNextBtn = document.createElement('button')
|
||||
|
||||
lightboxPrevBtn.textContent = '❮'
|
||||
lightboxNextBtn.textContent = '❯'
|
||||
lightboxPrevBtn.textContent = '❮'
|
||||
lightboxNextBtn.textContent = '❯'
|
||||
|
||||
Object.assign(lightboxPrevBtn.style, styles.lightboxBtn({ left: '0%' }))
|
||||
Object.assign(lightboxNextBtn.style, styles.lightboxBtn({ right: '0%' }))
|
||||
Object.assign(lightboxPrevBtn.style, styles.lightboxBtn({ left: '0%' }))
|
||||
Object.assign(lightboxNextBtn.style, styles.lightboxBtn({ right: '0%' }))
|
||||
|
||||
// close button
|
||||
const lightboxCloseBtn = document.createElement('button')
|
||||
Object.assign(
|
||||
lightboxCloseBtn.style,
|
||||
styles.lightboxBtn({ right: '0', top: '0' }),
|
||||
)
|
||||
lightboxCloseBtn.textContent = '❌'
|
||||
// close button
|
||||
const lightboxCloseBtn = document.createElement('button')
|
||||
Object.assign(
|
||||
lightboxCloseBtn.style,
|
||||
styles.lightboxBtn({ right: '0', top: '0' }),
|
||||
)
|
||||
lightboxCloseBtn.textContent = '❌'
|
||||
|
||||
const lightboxButtons = document.createElement('div')
|
||||
Object.assign(lightboxButtons.style, {
|
||||
position: 'absolute',
|
||||
top: '0%',
|
||||
right: '0%',
|
||||
// transform: "translate(50%, -50%)",
|
||||
height: '100%',
|
||||
width: '100%',
|
||||
background: 'none',
|
||||
border: 'none',
|
||||
color: '#fff',
|
||||
fontSize: '30px',
|
||||
cursor: 'pointer',
|
||||
pointerEvents: 'none',
|
||||
})
|
||||
const lightboxButtons = document.createElement('div')
|
||||
Object.assign(lightboxButtons.style, {
|
||||
position: 'absolute',
|
||||
top: '0%',
|
||||
right: '0%',
|
||||
// transform: "translate(50%, -50%)",
|
||||
height: '100%',
|
||||
width: '100%',
|
||||
background: 'none',
|
||||
border: 'none',
|
||||
color: '#fff',
|
||||
fontSize: '30px',
|
||||
cursor: 'pointer',
|
||||
pointerEvents: 'none',
|
||||
})
|
||||
|
||||
lightboxButtons.append(lightboxPrevBtn, lightboxNextBtn, lightboxCloseBtn)
|
||||
lightboxContainer.append(lightboxButtons, lightboxImage)
|
||||
lightboxButtons.append(lightboxPrevBtn, lightboxNextBtn, lightboxCloseBtn)
|
||||
lightboxContainer.append(lightboxButtons, lightboxImage)
|
||||
|
||||
//- image list
|
||||
const imageListContainer = document.createElement('div')
|
||||
Object.assign(imageListContainer.style, styles.img_list)
|
||||
//- image list
|
||||
const imageListContainer = document.createElement('div')
|
||||
Object.assign(imageListContainer.style, styles.img_list)
|
||||
|
||||
const createImgListBtn = (text, style) => {
|
||||
const btn = document.createElement('button')
|
||||
btn.type = 'button'
|
||||
btn.textContent = text
|
||||
Object.assign(btn.style, {
|
||||
...style,
|
||||
border: 'none',
|
||||
color: '#fff',
|
||||
background: 'none',
|
||||
height: '20px',
|
||||
cursor: 'pointer',
|
||||
position: 'absolute',
|
||||
top: '5px',
|
||||
fontSize: '12px',
|
||||
lineHeight: '12px',
|
||||
})
|
||||
imageListContainer.append(btn)
|
||||
return btn
|
||||
}
|
||||
const showBtn = document.createElement('button')
|
||||
const closeBtn = createImgListBtn('❌', {
|
||||
width: '20px',
|
||||
textIndent: '-4px',
|
||||
right: '5px',
|
||||
})
|
||||
const loadButton = createImgListBtn('Load Session History', {
|
||||
right: '90px',
|
||||
})
|
||||
const clearButton = createImgListBtn('Clear', {
|
||||
right: '30px',
|
||||
})
|
||||
const createImgListBtn = (text, style) => {
|
||||
const btn = document.createElement('button')
|
||||
btn.type = 'button'
|
||||
btn.textContent = text
|
||||
Object.assign(btn.style, {
|
||||
...style,
|
||||
border: 'none',
|
||||
color: '#fff',
|
||||
background: 'none',
|
||||
height: '20px',
|
||||
cursor: 'pointer',
|
||||
position: 'absolute',
|
||||
top: '5px',
|
||||
fontSize: '12px',
|
||||
lineHeight: '12px',
|
||||
})
|
||||
imageListContainer.append(btn)
|
||||
return btn
|
||||
}
|
||||
const showBtn = document.createElement('button')
|
||||
const closeBtn = createImgListBtn('❌', {
|
||||
width: '20px',
|
||||
textIndent: '-4px',
|
||||
right: '5px',
|
||||
})
|
||||
const loadButton = createImgListBtn('Load Session History', {
|
||||
right: '90px',
|
||||
})
|
||||
const clearButton = createImgListBtn('Clear', {
|
||||
right: '30px',
|
||||
})
|
||||
|
||||
//- tools popup button
|
||||
showBtn.classList.add('comfy-settings-btn')
|
||||
Object.assign(showBtn.style, {
|
||||
right: '16px',
|
||||
cursor: 'pointer',
|
||||
display: 'none',
|
||||
})
|
||||
//- tools popup button
|
||||
showBtn.classList.add('comfy-settings-btn')
|
||||
Object.assign(showBtn.style, {
|
||||
right: '16px',
|
||||
cursor: 'pointer',
|
||||
display: 'none',
|
||||
})
|
||||
|
||||
//- append to DOM
|
||||
document.body.append(imageListContainer)
|
||||
//- append to DOM
|
||||
document.body.append(imageListContainer)
|
||||
|
||||
showBtn.textContent = '🖼'
|
||||
showBtn.onclick = () => {
|
||||
imageListContainer.style.display = 'block'
|
||||
showBtn.style.display = 'none'
|
||||
}
|
||||
document.querySelector('.comfy-settings-btn').after(showBtn)
|
||||
document.querySelector('.comfy-settings-btn').after(lightboxContainer)
|
||||
showBtn.textContent = '🖼'
|
||||
showBtn.onclick = () => {
|
||||
imageListContainer.style.display = 'block'
|
||||
showBtn.style.display = 'none'
|
||||
}
|
||||
document.querySelector('.comfy-settings-btn').after(showBtn)
|
||||
document.querySelector('.comfy-settings-btn').after(lightboxContainer)
|
||||
|
||||
// for (const { output } of history) {
|
||||
// if (output?.images) {
|
||||
// for (const src of output.images) {
|
||||
// const img = document.createElement("img");
|
||||
// const but = document.createElement("button");
|
||||
// for (const { output } of history) {
|
||||
// if (output?.images) {
|
||||
// for (const src of output.images) {
|
||||
// const img = document.createElement("img");
|
||||
// const but = document.createElement("button");
|
||||
|
||||
//- callbacks
|
||||
closeBtn.onclick = () => {
|
||||
imageListContainer.style.display = 'none'
|
||||
showBtn.style.display = 'unset'
|
||||
}
|
||||
//- callbacks
|
||||
closeBtn.onclick = () => {
|
||||
imageListContainer.style.display = 'none'
|
||||
showBtn.style.display = 'unset'
|
||||
}
|
||||
|
||||
clearButton.onclick = () => {
|
||||
imageListContainer.replaceChildren(closeBtn, clearButton, loadButton)
|
||||
}
|
||||
clearButton.onclick = () => {
|
||||
imageListContainer.replaceChildren(closeBtn, clearButton, loadButton)
|
||||
}
|
||||
|
||||
lightboxNextBtn.onclick = () => {
|
||||
currentImageIndex = (currentImageIndex + 1) % imageUrls.length
|
||||
const imageUrl = imageUrls[currentImageIndex]
|
||||
lightboxImage.src = imageUrl
|
||||
}
|
||||
lightboxNextBtn.onclick = () => {
|
||||
currentImageIndex = (currentImageIndex + 1) % imageUrls.length
|
||||
const imageUrl = imageUrls[currentImageIndex]
|
||||
lightboxImage.src = imageUrl
|
||||
}
|
||||
|
||||
// Modify the lightboxPrevBtn onclick callback
|
||||
lightboxPrevBtn.onclick = () => {
|
||||
currentImageIndex =
|
||||
(currentImageIndex - 1 + imageUrls.length) % imageUrls.length
|
||||
const imageUrl = imageUrls[currentImageIndex]
|
||||
lightboxImage.src = imageUrl
|
||||
}
|
||||
// Modify the lightboxPrevBtn onclick callback
|
||||
lightboxPrevBtn.onclick = () => {
|
||||
currentImageIndex =
|
||||
(currentImageIndex - 1 + imageUrls.length) % imageUrls.length
|
||||
const imageUrl = imageUrls[currentImageIndex]
|
||||
lightboxImage.src = imageUrl
|
||||
}
|
||||
|
||||
lightboxCloseBtn.onclick = () => {
|
||||
lightboxContainer.style.display = 'none'
|
||||
}
|
||||
lightboxImage.onclick = lightboxNextBtn.onclick
|
||||
/**
|
||||
* 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
|
||||
* the image in the lightbox.
|
||||
* @param {*} src
|
||||
*/
|
||||
const createImageBtn = (src) => {
|
||||
console.debug(`making image ${src.filename}`)
|
||||
const img = document.createElement('img')
|
||||
const but = document.createElement('button')
|
||||
lightboxCloseBtn.onclick = () => {
|
||||
lightboxContainer.style.display = 'none'
|
||||
}
|
||||
lightboxImage.onclick = lightboxNextBtn.onclick
|
||||
/**
|
||||
* 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
|
||||
* the image in the lightbox.
|
||||
* @param {*} src
|
||||
*/
|
||||
const createImageBtn = (src) => {
|
||||
console.debug(`making image ${src.filename}`)
|
||||
const img = document.createElement('img')
|
||||
const but = document.createElement('button')
|
||||
|
||||
Object.assign(but.style, {
|
||||
height: '120px',
|
||||
width: '120px',
|
||||
border: 'none',
|
||||
padding: 0,
|
||||
margin: 0,
|
||||
})
|
||||
Object.assign(img.style, {
|
||||
width: '100%',
|
||||
height: '100%',
|
||||
objectFit: 'cover',
|
||||
})
|
||||
Object.assign(but.style, {
|
||||
height: '120px',
|
||||
width: '120px',
|
||||
border: 'none',
|
||||
padding: 0,
|
||||
margin: 0,
|
||||
})
|
||||
Object.assign(img.style, {
|
||||
width: '100%',
|
||||
height: '100%',
|
||||
objectFit: 'cover',
|
||||
})
|
||||
|
||||
img.src = `/view?filename=${encodeURIComponent(src.filename)}&type=${
|
||||
src.type
|
||||
}&subfolder=${encodeURIComponent(src.subfolder)}`
|
||||
img.src = `/view?filename=${encodeURIComponent(src.filename)}&type=${
|
||||
src.type
|
||||
}&subfolder=${encodeURIComponent(src.subfolder)}`
|
||||
|
||||
imageUrls.push(img.src)
|
||||
imageUrls.push(img.src)
|
||||
|
||||
console.debug(img.src)
|
||||
console.debug(img.src)
|
||||
|
||||
img.onload = () => {
|
||||
but.style.width = `${120 * (img.naturalWidth / img.naturalHeight)}px`
|
||||
}
|
||||
img.onload = () => {
|
||||
but.style.width = `${120 * (img.naturalWidth / img.naturalHeight)}px`
|
||||
}
|
||||
|
||||
but.onclick = () => {
|
||||
lightboxContainer.style.display = 'flex'
|
||||
// add the same image to the lightbox
|
||||
lightboxImage.src = img.src
|
||||
// lighboxContainer.replaceChildren(lightboxButtons, img);
|
||||
}
|
||||
but.onclick = () => {
|
||||
lightboxContainer.style.display = 'flex'
|
||||
// add the same image to the lightbox
|
||||
lightboxImage.src = img.src
|
||||
// lighboxContainer.replaceChildren(lightboxButtons, img);
|
||||
}
|
||||
|
||||
// add right click menu
|
||||
but.addEventListener('contextmenu', (e) => {
|
||||
e.preventDefault()
|
||||
// add right click menu
|
||||
but.addEventListener('contextmenu', (e) => {
|
||||
e.preventDefault()
|
||||
|
||||
if (image_menu) {
|
||||
image_menu.remove()
|
||||
}
|
||||
if (image_menu) {
|
||||
image_menu.remove()
|
||||
}
|
||||
|
||||
image_menu = document.createElement('div')
|
||||
Object.assign(image_menu.style, {
|
||||
position: 'absolute',
|
||||
top: `${e.clientY}px`,
|
||||
left: `${e.clientX}px`,
|
||||
background: '#333',
|
||||
color: '#fff',
|
||||
padding: '5px',
|
||||
borderRadius: '5px',
|
||||
zIndex: 999,
|
||||
})
|
||||
const load_img = document.createElement('button')
|
||||
load_img.textContent = 'Load'
|
||||
load_img.onclick = () => {
|
||||
app.handleFile(img.src)
|
||||
}
|
||||
image_menu = document.createElement('div')
|
||||
Object.assign(image_menu.style, {
|
||||
position: 'absolute',
|
||||
top: `${e.clientY}px`,
|
||||
left: `${e.clientX}px`,
|
||||
background: '#333',
|
||||
color: '#fff',
|
||||
padding: '5px',
|
||||
borderRadius: '5px',
|
||||
zIndex: 999,
|
||||
})
|
||||
const load_img = document.createElement('button')
|
||||
load_img.textContent = 'Load'
|
||||
load_img.onclick = () => {
|
||||
app.handleFile(img.src)
|
||||
}
|
||||
|
||||
image_menu.appendChild(load_img)
|
||||
document.body.appendChild(image_menu)
|
||||
})
|
||||
image_menu.appendChild(load_img)
|
||||
document.body.appendChild(image_menu)
|
||||
})
|
||||
|
||||
but.append(img)
|
||||
imageListContainer.prepend(but)
|
||||
}
|
||||
but.append(img)
|
||||
imageListContainer.prepend(but)
|
||||
}
|
||||
|
||||
loadButton.onclick = async () => {
|
||||
const all_history = await api.getHistory()
|
||||
for (const history of all_history.History) {
|
||||
if (history.outputs) {
|
||||
for (const key of Object.keys(history.outputs)) {
|
||||
console.debug(key)
|
||||
if (history.outputs[key].images) {
|
||||
for (const im of history.outputs[key].images) {
|
||||
console.debug(im)
|
||||
createImageBtn(im)
|
||||
}
|
||||
}
|
||||
}
|
||||
// for (const src of outputs.outputs.images) {
|
||||
// console.debug(src)
|
||||
// makeImage(`${src.subfolder}/${src.filename}`)
|
||||
// }
|
||||
}
|
||||
}
|
||||
}
|
||||
loadButton.onclick = async () => {
|
||||
const all_history = await api.getHistory()
|
||||
for (const history of all_history.History) {
|
||||
if (history.outputs) {
|
||||
for (const key of Object.keys(history.outputs)) {
|
||||
console.debug(key)
|
||||
if (history.outputs[key].images) {
|
||||
for (const im of history.outputs[key].images) {
|
||||
console.debug(im)
|
||||
createImageBtn(im)
|
||||
}
|
||||
}
|
||||
}
|
||||
// for (const src of outputs.outputs.images) {
|
||||
// console.debug(src)
|
||||
// makeImage(`${src.subfolder}/${src.filename}`)
|
||||
// }
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
///////-------
|
||||
///////-------
|
||||
|
||||
// const all_history = await api.getHistory()
|
||||
// for (const history of all_history.History) {
|
||||
// if (history.outputs) {
|
||||
// for (const key of Object.keys(history.outputs)) {
|
||||
// for (const im of history.outputs[key].images) {
|
||||
// makeImage(im)
|
||||
// }
|
||||
// }
|
||||
// // for (const src of outputs.outputs.images) {
|
||||
// // console.debug(src)
|
||||
// // makeImage(`${src.subfolder}/${src.filename}`)
|
||||
// // }
|
||||
// }
|
||||
// }
|
||||
// const all_history = await api.getHistory()
|
||||
// for (const history of all_history.History) {
|
||||
// if (history.outputs) {
|
||||
// for (const key of Object.keys(history.outputs)) {
|
||||
// for (const im of history.outputs[key].images) {
|
||||
// makeImage(im)
|
||||
// }
|
||||
// }
|
||||
// // for (const src of outputs.outputs.images) {
|
||||
// // console.debug(src)
|
||||
// // makeImage(`${src.subfolder}/${src.filename}`)
|
||||
// // }
|
||||
// }
|
||||
// }
|
||||
|
||||
//- Hook into the API
|
||||
api.addEventListener('executed', ({ detail }) => {
|
||||
if (detail?.output?.images) {
|
||||
for (const src of detail.output.images) {
|
||||
console.debug(`Adding ${src} to image feed`)
|
||||
createImageBtn(src)
|
||||
}
|
||||
}
|
||||
})
|
||||
},
|
||||
//- Hook into the API
|
||||
api.addEventListener('executed', ({ detail }) => {
|
||||
if (detail?.output?.images) {
|
||||
for (const src of detail.output.images) {
|
||||
console.debug(`Adding ${src} to image feed`)
|
||||
createImageBtn(src)
|
||||
}
|
||||
}
|
||||
})
|
||||
},
|
||||
})
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
@@ -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
File diff suppressed because it is too large
Load Diff
+1
-1
Submodule wiki updated: 4db733ae92...a402de4af9
Reference in New Issue
Block a user