Compare commits

..
Author SHA1 Message Date
Mel Massadian d2da2949b4 feat: ✨ improve the I/O sidebar
- better options (sort, count)
- uses the new toast api instead of MTB.notify
2024-11-20 22:41:04 +01:00
Mel Massadian 9120d3ec42 chore: 🧹 add worktree to gitignores
for the experimental doc site at:
https://melmass.github.io/comfy_mtb/
2024-11-20 22:40:16 +01:00
Mel Massadian 2c8b7d790d feat: ✨ add UpscaleBBoxBy 2024-11-20 22:40:16 +01:00
Mel Massadian 540c8c9fa9 chore 🧹: add deprecations and experimental 2024-11-20 22:40:16 +01:00
Mel Massadian 88d51e5774 chore: 🧹 remove dupe code 2024-11-20 22:40:16 +01:00
Mel Massadian 8f47810b79 feat: ✨ simplified sidebar and backend
If you have a LoadImage selected,
clicking on images in the "input" mode will set the image on the
selected nodes
2024-11-20 22:40:16 +01:00
Mel Massadian ae1ef0f914 feat: ✨ add Interpolate Condition 2024-11-20 22:40:06 +01:00
Mel Massadian b0e234b7ee feat: ✨ dump of wip things... 2024-11-20 22:39:23 +01:00
23 changed files with 1576 additions and 3535 deletions
-2
View File
@@ -12,8 +12,6 @@ jobs:
steps: steps:
- name: ♻️ Check out code - name: ♻️ Check out code
uses: actions/checkout@v4 uses: actions/checkout@v4
with:
submodules: true
- name: 📦 Publish Custom Node - name: 📦 Publish Custom Node
uses: Comfy-Org/publish-node-action@main uses: Comfy-Org/publish-node-action@main
with: with:
+24 -42
View File
@@ -7,7 +7,7 @@
# #
### ###
__version__ = "0.2.1" __version__ = "0.1.6"
import os import os
@@ -34,7 +34,6 @@ from aiohttp import web
from server import PromptServer from server import PromptServer
from .endpoint import endlog from .endpoint import endlog
from .install import get_node_dependencies
from .log import blue_text, cyan_text, get_label, get_summary, log from .log import blue_text, cyan_text, get_label, get_summary, log
from .utils import comfy_dir, here from .utils import comfy_dir, here
@@ -233,15 +232,21 @@ if failed:
if hasattr(PromptServer, "instance"): if hasattr(PromptServer, "instance"):
img_cache = None img_cache = None
prompt_cache = None
with contextlib.suppress(ImportError): with contextlib.suppress(ImportError):
from cachetools import TTLCache from cachetools import TTLCache
img_cache = TTLCache(maxsize=100, ttl=5) # 1 min TTL img_cache = TTLCache(maxsize=100, ttl=5) # 1 min TTL
prompt_cache = TTLCache(maxsize=100, ttl=5) # 1 min TTL
node_dependency_mapping = get_node_dependencies() restore_deps = ["basicsr"]
onnx_deps = ["onnxruntime"]
swap_deps = ["insightface"] + onnx_deps
node_dependency_mapping = {
"QrCode": ["qrcode"],
"DeepBump": onnx_deps,
"FaceSwap": swap_deps,
"LoadFaceSwapModel": swap_deps,
"LoadFaceAnalysisModel": restore_deps,
}
PromptServer.instance.app.router.add_static( PromptServer.instance.app.router.add_static(
"/mtb-assets/", path=(here / "html").as_posix() "/mtb-assets/", path=(here / "html").as_posix()
@@ -310,10 +315,10 @@ if hasattr(PromptServer, "instance"):
} }
) )
@PromptServer.instance.routes.post("/mtb/server-info") @PromptServer.instance.routes.post("/mtb/debug")
async def set_server_info(request: Request): async def set_debug(request: Request):
json_data: dict[str, bool] = await request.json() json_data: dict[str, bool] = await request.json()
enabled = json_data.get("debug") enabled = json_data.get("enabled")
if enabled: if enabled:
os.environ["MTB_DEBUG"] = "true" os.environ["MTB_DEBUG"] = "true"
log.setLevel(logging.DEBUG) log.setLevel(logging.DEBUG)
@@ -339,7 +344,7 @@ if hasattr(PromptServer, "instance"):
html_response = """ html_response = """
<div class="flex-container menu"> <div class="flex-container menu">
<a href="/mtb/manage">manage</a> <a href="/mtb/manage">manage</a>
<a href="/mtb/server-info">Server Info</a> <a href="/mtb/debug">debug</a>
<a href="/mtb/status">status</a> <a href="/mtb/status">status</a>
</div> </div>
""" """
@@ -364,20 +369,16 @@ if hasattr(PromptServer, "instance"):
return img_cache[cache_key] return img_cache[cache_key]
with Image.open(file_path) as img: with Image.open(file_path) as img:
info = img.info
if preview_params: if preview_params:
img = process_preview(img, preview_params) img = process_preview(img, preview_params)
if channel: if channel:
img = process_channel(img, channel) img = process_channel(img, channel)
if prompt_cache:
prompt_cache[cache_key] = info
if img_cache: if img_cache:
img_cache[cache_key] = img.getvalue() img_cache[cache_key] = img.getvalue()
return img_cache[cache_key] return img_cache[cache_key]
return img.getvalue() return img.getvalue()
def process_preview(img: Image.Image, preview_params): def process_preview(img: Image, preview_params):
image_format, quality, width = preview_params image_format, quality, width = preview_params
quality = int(quality) quality = int(quality)
@@ -386,9 +387,7 @@ if hasattr(PromptServer, "instance"):
img.thumbnail((width, int(width * img.height / img.width))) img.thumbnail((width, int(width * img.height / img.width)))
buffer = BytesIO() buffer = BytesIO()
img.save( img.save(buffer, format=image_format, quality=quality)
buffer, format=image_format, quality=quality, metadata=img.info
)
buffer.seek(0) buffer.seek(0)
return buffer return buffer
@@ -424,8 +423,6 @@ if hasattr(PromptServer, "instance"):
headers={"Content-Disposition": f'filename="{filename}"'}, 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") @PromptServer.instance.routes.get("/mtb/view")
async def view_image(request: Request): async def view_image(request: Request):
import folder_paths import folder_paths
@@ -484,40 +481,25 @@ if hasattr(PromptServer, "instance"):
return await get_image_response(file, filename, preview_info, channel) return await get_image_response(file, filename, preview_info, channel)
@PromptServer.instance.routes.get("/mtb/server-info") @PromptServer.instance.routes.get("/mtb/debug")
async def get_debug(request: Request): async def get_debug(request: Request):
from . import endpoint from . import endpoint
_ = reload(endpoint) _ = reload(endpoint)
isdebug = "MTB_DEBUG" in os.environ enabled = "MTB_DEBUG" in os.environ
exposed = "MTB_EXPOSE" in os.environ
def render_property(name: str, val: str):
return f"""<strong>{name}:</strong>
<p>
{val}
</p>"""
# Check if the request prefers HTML content # Check if the request prefers HTML content
if "text/html" in request.headers.get("Accept", ""): if "text/html" in request.headers.get("Accept", ""):
# # Return an HTML page # # Return an HTML page
html_response = "" html_response = f"""
<h1>MTB Debug Status: {'Enabled' if enabled else 'Disabled'}</h1>
html_response += render_property( """
"Debug", "Enabled" if isdebug else "Disabled"
)
html_response += render_property("Exposed", str(exposed))
return web.Response( return web.Response(
text=endpoint.render_base_template( text=endpoint.render_base_template("Debug", html_response),
"Server Info", html_response
),
content_type="text/html", content_type="text/html",
) )
# Return JSON for other requests # Return JSON for other requests
return web.json_response({"exposed": exposed, "debug": isdebug}) return web.json_response({"enabled": enabled})
@PromptServer.instance.routes.get("/mtb/actions") @PromptServer.instance.routes.get("/mtb/actions")
async def no_route(request: Request): async def no_route(request: Request):
+43 -97
View File
@@ -1,21 +1,18 @@
import csv import csv
import os
import secrets import secrets
import sys import sys
import urllib.parse
from pathlib import Path from pathlib import Path
from typing import Any, Literal from typing import Any
import folder_paths
from aiohttp import web from aiohttp import web
from .install import get_node_dependencies
from .log import mklog from .log import mklog
from .utils import ( from .utils import (
SortMode,
backup_file, backup_file,
build_glob_patterns,
glob_multiple,
import_install, import_install,
input_dir,
output_dir,
reqs_map, reqs_map,
run_command, run_command,
styles_dir, styles_dir,
@@ -27,7 +24,7 @@ endlog = mklog("mtb endpoint")
import_install("requirements") import_install("requirements")
def ACTIONS_installDependency(dependency_names: list[str] | None = None): def ACTIONS_installDependency(dependency_names=None):
if dependency_names is None: if dependency_names is None:
# return web.Response(text="No dependency name provided", status=400) # return web.Response(text="No dependency name provided", status=400)
return {"error": "No dependency name provided"} return {"error": "No dependency name provided"}
@@ -35,14 +32,6 @@ def ACTIONS_installDependency(dependency_names: list[str] | None = None):
endlog.debug(f"Received Install Dependency request for {dependency_names}") endlog.debug(f"Received Install Dependency request for {dependency_names}")
# reqs = [] # reqs = []
resolved_names = [reqs_map.get(name, name) for name in dependency_names] resolved_names = [reqs_map.get(name, name) for name in dependency_names]
allowed_deps = list(
{d for dep in get_node_dependencies().values() for d in dep}
)
for dep in dependency_names:
if dep not in allowed_deps:
return {
"error": f"Unknown dependency: {dep}, you can only use this endpoint to install {allowed_deps}"
}
try: try:
run_command( run_command(
[Path(sys.executable), "-m", "pip", "install"] + resolved_names [Path(sys.executable), "-m", "pip", "install"] + resolved_names
@@ -67,103 +56,60 @@ def ACTIONS_installDependency(dependency_names: list[str] | None = None):
# break # break
def ACTIONS_getUserImageFolders():
input_dir = Path(folder_paths.get_input_directory())
output_dir = Path(folder_paths.get_output_directory())
input_subdirs = [x.name for x in input_dir.iterdir() if x.is_dir()]
output_subdirs = [x.name for x in output_dir.iterdir() if x.is_dir()]
return {"input": input_subdirs, "output": output_subdirs}
def ACTIONS_getUserVideos(
size=256, count=200, offset=0, sort: str | None = None
):
count = count or 1000
video_extensions = ["webm", "mp4", "mkv", "mov"]
entries = {}
patterns = build_glob_patterns(video_extensions)
input_dir = Path(folder_paths.get_input_directory())
entries = glob_multiple(input_dir, patterns)
sort_mode = SortMode.from_str(sort)
if sort_mode:
sort_key = {
SortMode.MODIFIED: lambda x: x.stat().st_mtime,
SortMode.MODIFIED_REVERSE: lambda x: x.stat().st_mtime,
SortMode.NAME: lambda x: x.name,
SortMode.NAME_REVERSE: lambda x: x.name,
}.get(sort_mode)
if sort_key:
reverse = sort_mode in (SortMode.MODIFIED, SortMode.NAME_REVERSE)
entries = sorted(entries, key=sort_key, reverse=reverse)
videos = {
video.name: (
f"/view?force_rate=0&frame_load_cap=0&skip_first_frames=0&select_every_nth=1&filename={urllib.parse.quote_plus(video.name)}&type=input&format=video&force_size={size}x?"
)
for i, video in enumerate(entries)
if offset <= i < offset + count
}
return videos
def ACTIONS_getUserImages( def ACTIONS_getUserImages(
mode: Literal["input", "output"], mode: str,
count=1000, count=200,
offset=0, offset=0,
sort: str | None = None, sort: str | None = None,
include_subfolders: bool = False, include_subfolders: bool = False,
subfolder=None,
): ):
# enabled = "MTB_EXPOSE" in os.environ # TODO: find a better name :s
# if not enabled: enabled = "MTB_EXPOSE" in os.environ
# return {"error": "Session not authorized to getInputs"} if not enabled:
return {"error": "Session not authorized to getInputs"}
imgs = {} imgs = {}
count = count or 1000
input_dir = Path(folder_paths.get_input_directory())
output_dir = Path(folder_paths.get_output_directory())
entry_dir = input_dir if mode == "input" else output_dir entry_dir = input_dir if mode == "input" else output_dir
if subfolder: pattern = "**/*.png" if include_subfolders else "*.png"
entry_dir = entry_dir / subfolder
if not entry_dir.exists(): entry_gen = entry_dir.glob(pattern)
return {
"error": f"Subfolder {entry_dir.name} doesn't exists in {entry_dir.parent.as_posix()}"
}
supported = ["png", "jpg", "jpeg", "webp", "gif"]
entries = {} entries = {}
patterns = build_glob_patterns(supported, recursive=include_subfolders)
entries = glob_multiple(entry_dir, patterns)
sort_mode = SortMode.from_str(sort) 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
if sort_mode: for i, img in enumerate(entries):
sort_key = { if i < offset:
SortMode.MODIFIED: lambda x: x.stat().st_mtime, continue
SortMode.MODIFIED_REVERSE: lambda x: x.stat().st_mtime,
SortMode.NAME: lambda x: x.name,
SortMode.NAME_REVERSE: lambda x: x.name,
}.get(sort_mode)
if sort_key:
reverse = sort_mode in (SortMode.MODIFIED, SortMode.NAME_REVERSE)
entries = sorted(entries, key=sort_key, reverse=reverse)
imgs = { subfolder = (
img.name: ( img.parent.relative_to(entry_dir) if include_subfolders else ""
f"/mtb/view?filename={img.name}&width=512&type={mode}&subfolder={subfolder or ''}" )
f"{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)}" f"&preview=&rand={secrets.randbelow(424242)}"
) )
for i, img in enumerate(entries) if i >= count + offset - 1:
if offset <= i < offset + count break
}
return imgs return imgs
+5 -17
View File
@@ -27,7 +27,7 @@ export def "comfy start" [--clean,--old-ui, --listen] {
let root = get_root --clean=($clean) let root = get_root --clean=($clean)
cd $root cd $root
MTB_DEBUG=true python main.py --port 3000 ...(if $old_ui { ["--front-end-version", "Comfy-Org/ComfyUI_legacy_frontend@latest"]} else {[ --front-end-version Comfy-Org/ComfyUI_frontend@latest]}) --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 {[]}) --preview-method auto ...(if $listen {["--listen"]} else {[]})
} }
# update comfy itself and merge master in current branch # update comfy itself and merge master in current branch
@@ -67,14 +67,8 @@ export def "comfy update" [
git checkout master git checkout master
print $"(ansi yellow_italic)Fetching and pulling remote updates(ansi reset)" print $"(ansi yellow_italic)Fetching and pulling remote updates(ansi reset)"
if ($clean) { git fetch
git fetch local master git pull
git pull local master
} else {
git fetch
git pull
}
print $"(ansi yellow_italic)Back to our branch \(($branch_name)\)(ansi reset)" print $"(ansi yellow_italic)Back to our branch \(($branch_name)\)(ansi reset)"
git checkout - git checkout -
@@ -141,7 +135,7 @@ export def "comfy update_extensions" [--clean] {
let root = get_root --clean=($clean) let root = get_root --clean=($clean)
cd $root cd $root
cd custom_nodes cd custom_nodes
git multipull . -s -q git multipull .
} }
def --env path-add [pth] { def --env path-add [pth] {
@@ -152,7 +146,7 @@ def --env path-add [pth] {
export-env { export-env {
$env.COMFY_MTB = ("." | path expand) $env.COMFY_MTB = ("." | path expand)
# $env.CUDA_ROOT = 'C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.1\' $env.CUDA_ROOT = 'C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.1\'
$env.CUDA_HOME = $env.CUDA_ROOT $env.CUDA_HOME = $env.CUDA_ROOT
@@ -160,12 +154,6 @@ export-env {
$env.COMFY_CLEAN_ROOT = ($env.COMFY_ROOT | path dirname | path join ComfyClean) $env.COMFY_CLEAN_ROOT = ($env.COMFY_ROOT | path dirname | path join ComfyClean)
path-add 'C:/Portable/TensorRT-8.6.0.12/lib' path-add 'C:/Portable/TensorRT-8.6.0.12/lib'
if $nu.os-info.family == 'windows' {
path-add 'G:\BIN\TensorRT-10.7.0.23\lib'
path-add 'G:\BIN\cudnn-windows-x86_64-9.6.0.74_cuda12-archive\bin'
}
path-add ($env.CUDA_ROOT | path join bin) path-add ($env.CUDA_ROOT | path join bin)
overlay use ../../.venv/Scripts/activate.nu overlay use ../../.venv/Scripts/activate.nu
} }
+28 -64
View File
@@ -43,28 +43,10 @@ pip_map = {
"tb-nightly": "tensorboard", "tb-nightly": "tensorboard",
"protobuf": "google.protobuf", "protobuf": "google.protobuf",
"qrcode[pil]": "qrcode", "qrcode[pil]": "qrcode",
"requirements-parser": "requirements", "requirements-parser": "requirements"
# Add more mappings as needed # Add more mappings as needed
} }
def get_node_dependencies():
restore_deps = ["basicsr"]
onnx_deps = ["onnxruntime"]
swap_deps = ["insightface"] + onnx_deps
quant_deps = ["bitsandbytes"]
io_deps = ["av"]
return {
"QrCode": ["qrcode"],
"DeepBump": onnx_deps,
"FaceSwap": swap_deps,
"LoadFaceSwapModel": swap_deps,
"LoadFaceAnalysisModel": restore_deps,
"Quantize": quant_deps,
"SaveGif": io_deps,
}
# endregion # endregion
# region ansi # region ansi
@@ -142,12 +124,12 @@ def print_formatted(text, *formats, color=None, background=None, **kwargs):
header = "[mtb install] " header = "[mtb install] "
# Handle console encoding for Unicode characters (utf-8) # Handle console encoding for Unicode characters (utf-8)
encoded_header = header.encode( encoded_header = header.encode(sys.stdout.encoding, errors="replace").decode(
sys.stdout.encoding, errors="replace" sys.stdout.encoding
).decode(sys.stdout.encoding) )
encoded_text = formatted_text.encode( encoded_text = formatted_text.encode(sys.stdout.encoding, errors="replace").decode(
sys.stdout.encoding, errors="replace" sys.stdout.encoding
).decode(sys.stdout.encoding) )
print( print(
" " * len(encoded_header) " " * len(encoded_header)
@@ -181,9 +163,7 @@ def run_command(cmd, ignored_lines_start=None):
try: try:
_run_command(shell_cmd, ignored_lines_start) _run_command(shell_cmd, ignored_lines_start)
except subprocess.CalledProcessError as e: except subprocess.CalledProcessError as e:
print( print(f"Command failed with return code: {e.returncode}", file=sys.stderr)
f"Command failed with return code: {e.returncode}", file=sys.stderr
)
print(e.stderr.strip(), file=sys.stderr) print(e.stderr.strip(), file=sys.stderr)
except KeyboardInterrupt: except KeyboardInterrupt:
@@ -258,7 +238,7 @@ def suppress_std():
def get_local_version(): def get_local_version():
init_file = os.path.join(os.path.dirname(__file__), "__init__.py") init_file = os.path.join(os.path.dirname(__file__), "__init__.py")
if os.path.isfile(init_file): if os.path.isfile(init_file):
with open(init_file) as f: with open(init_file, "r") as f:
tree = ast.parse(f.read()) tree = ast.parse(f.read())
for node in ast.walk(tree): for node in ast.walk(tree):
if isinstance(node, ast.Assign): if isinstance(node, ast.Assign):
@@ -276,16 +256,13 @@ def download_file(url, file_name):
with requests.get(url, stream=True) as response: with requests.get(url, stream=True) as response:
response.raise_for_status() response.raise_for_status()
total_size = int(response.headers.get("content-length", 0)) total_size = int(response.headers.get("content-length", 0))
with ( with open(file_name, "wb") as file, tqdm(
open(file_name, "wb") as file, desc=file_name.stem,
tqdm( total=total_size,
desc=file_name.stem, unit="B",
total=total_size, unit_scale=True,
unit="B", unit_divisor=1024,
unit_scale=True, ) as progress_bar:
unit_divisor=1024,
) as progress_bar,
):
for chunk in response.iter_content(chunk_size=8192): for chunk in response.iter_content(chunk_size=8192):
file.write(chunk) file.write(chunk)
progress_bar.update(len(chunk)) progress_bar.update(len(chunk))
@@ -325,9 +302,7 @@ def import_or_install(requirement, dry=False):
pip_install_name = pip_name + pip_spec pip_install_name = pip_name + pip_spec
if not installed: if not installed:
print_formatted( print_formatted(f"Installing package {pip_name}...", "italic", color="yellow")
f"Installing package {pip_name}...", "italic", color="yellow"
)
if dry: if dry:
print_formatted( print_formatted(
f"Dry-run: Package {pip_install_name} would be installed (import name: '{import_name}').", f"Dry-run: Package {pip_install_name} would be installed (import name: '{import_name}').",
@@ -335,9 +310,7 @@ def import_or_install(requirement, dry=False):
) )
else: else:
try: try:
run_command( run_command([executable, "-m", "pip", "install", pip_install_name])
[executable, "-m", "pip", "install", pip_install_name]
)
print_formatted( print_formatted(
f"Package {pip_install_name} installed successfully using pip package name (import name: '{import_name}')", f"Package {pip_install_name} installed successfully using pip package name (import name: '{import_name}')",
"bold", "bold",
@@ -353,9 +326,13 @@ def import_or_install(requirement, dry=False):
def get_github_assets(tag=None): def get_github_assets(tag=None):
if tag: if tag:
tag_url = f"https://api.github.com/repos/{repo_owner}/{repo_name}/releases/tags/{tag}" tag_url = (
f"https://api.github.com/repos/{repo_owner}/{repo_name}/releases/tags/{tag}"
)
else: else:
tag_url = f"https://api.github.com/repos/{repo_owner}/{repo_name}/releases/latest" tag_url = (
f"https://api.github.com/repos/{repo_owner}/{repo_name}/releases/latest"
)
response = requests.get(tag_url) response = requests.get(tag_url)
if response.status_code == 404: if response.status_code == 404:
# print_formatted( # print_formatted(
@@ -384,9 +361,7 @@ except ImportError:
def main(): def main():
if len(sys.argv) == 1: if len(sys.argv) == 1:
print_formatted( print_formatted(
"mtb doesn't need an install script anymore.", "mtb doesn't need an install script anymore.", "italic", color="yellow"
"italic",
color="yellow",
) )
return return
if all(arg not in ("-p", "--path") for arg in sys.argv): if all(arg not in ("-p", "--path") for arg in sys.argv):
@@ -422,12 +397,8 @@ def main():
else: else:
repo_dir = clone_dir / repo_name repo_dir = clone_dir / repo_name
if not repo_dir.exists(): if not repo_dir.exists():
print_formatted( print_formatted(f"Cloning to {repo_dir}...", "italic", color="yellow")
f"Cloning to {repo_dir}...", "italic", color="yellow" run_command(["git", "clone", "--recursive", repo_url, repo_dir])
)
run_command(
["git", "clone", "--recursive", repo_url, repo_dir]
)
else: else:
print_formatted( print_formatted(
f"Directory {repo_dir} already exists, we will update it..." f"Directory {repo_dir} already exists, we will update it..."
@@ -438,14 +409,7 @@ def main():
print_formatted("Checking environment...", "italic", color="yellow") print_formatted("Checking environment...", "italic", color="yellow")
missing_deps = [] missing_deps = []
install_cmd = [ install_cmd = [executable, "-m", "pip", "install", "-r", "requirements.txt"]
executable,
"-m",
"pip",
"install",
"-r",
"requirements.txt",
]
run_command(install_cmd) run_command(install_cmd)
print_formatted( print_formatted(
+10 -226
View File
@@ -1,5 +1,4 @@
from io import BytesIO from io import BytesIO
from typing import Literal
import cv2 import cv2
import numpy as np import numpy as np
@@ -411,14 +410,7 @@ class MTB_BatchFloat:
RETURN_TYPES = ("FLOATS",) RETURN_TYPES = ("FLOATS",)
CATEGORY = "mtb/batch" CATEGORY = "mtb/batch"
def set_floats( def set_floats(self, mode, count, min, max, easing):
self,
mode: Literal["Steps"] | Literal["Single"] = "Steps",
count: int = 1,
min: float = 0.0, # noqa: A002
max: float = 1.0, # noqa: A002
easing: str = "Linear",
):
if mode == "Steps" and count == 1: if mode == "Steps" and count == 1:
raise ValueError( raise ValueError(
"Steps mode requires at least a count of 2 values" "Steps mode requires at least a count of 2 values"
@@ -437,210 +429,6 @@ class MTB_BatchFloat:
return (keyframes,) return (keyframes,)
class MTB_BatchSequencePlus:
"""Sequences multiple image batches with transition effects."""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"transition": (
[
"none",
"crossfade",
"slide_left",
"slide_right",
"slide_up",
"slide_down",
"wipe_left",
"wipe_right",
"wipe_up",
"wipe_down",
"band_wipe_h",
"band_wipe_v",
],
{"default": "none"},
),
"overlap_frames": (
"INT",
{"default": 0, "min": 0, "max": 120, "step": 1},
),
"reverse": ("BOOLEAN", {"default": False}),
}
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "sequence_batches"
CATEGORY = "mtb/batch"
def apply_transition(
self,
frame1: torch.Tensor,
frame2: torch.Tensor,
transition: str,
progress: float,
):
"""Apply transition effect between two frames."""
if transition == "none":
return frame1 if progress < 0.5 else frame2
elif transition == "crossfade":
return frame1 * (1 - progress) + frame2 * progress
elif transition.startswith("slide_"):
h, w = frame1.shape[1:3]
if transition == "slide_left":
offset = int(w * progress)
frame2 = torch.roll(frame2, shifts=-offset, dims=2)
elif transition == "slide_right":
offset = int(w * progress)
frame2 = torch.roll(frame2, shifts=offset, dims=2)
elif transition == "slide_up":
offset = int(h * progress)
frame2 = torch.roll(frame2, shifts=-offset, dims=1)
elif transition == "slide_down":
offset = int(h * progress)
frame2 = torch.roll(frame2, shifts=offset, dims=1)
return frame1 * (1 - progress) + frame2 * progress
elif transition.startswith("wipe_"):
h, w = frame1.shape[1:3]
mask = torch.zeros_like(frame1)
if transition == "wipe_left":
edge = int(w * progress)
mask[:, :, :edge, :] = 1
elif transition == "wipe_right":
edge = int(w * (1 - progress))
mask[:, :, edge:, :] = 1
elif transition == "wipe_up":
edge = int(h * progress)
mask[:, :edge, :, :] = 1
elif transition == "wipe_down":
edge = int(h * (1 - progress))
mask[:, edge:, :, :] = 1
return frame1 * (1 - mask) + frame2 * mask
elif transition.startswith("band_wipe_"):
h, w = frame1.shape[1:3]
mask = torch.zeros_like(frame1)
num_bands = 10 # Number of bands
if transition == "band_wipe_h":
band_width = w / num_bands
for i in range(num_bands):
edge = int((w * progress) - (i * band_width))
start = int(i * band_width)
end = int(min(start + edge, (i + 1) * band_width))
if end > start:
mask[:, :, start:end, :] = 1
else: # band_wipe_v
band_height = h / num_bands
for i in range(num_bands):
edge = int((h * progress) - (i * band_height))
start = int(i * band_height)
end = int(min(start + edge, (i + 1) * band_height))
if end > start:
mask[:, start:end, :, :] = 1
return frame1 * (1 - mask) + frame2 * mask
return frame1
def sequence_batches(
self, transition: str, overlap_frames: int, reverse: bool, **kwargs
):
images: list[torch.Tensor] = list(kwargs.values())
if reverse:
images = images[::-1]
processed_images: list[torch.Tensor] = []
for img in images:
if len(img.shape) == 3:
img = img.unsqueeze(0)
processed_images.append(img)
if overlap_frames == 0 or transition == "none":
return (torch.cat(processed_images, dim=0),)
result_frames: list[torch.Tensor] = []
if len(processed_images) > 0:
result_frames.extend(
list(processed_images[0][: -overlap_frames // 2])
)
for i in range(1, len(processed_images)):
prev_batch = processed_images[i - 1]
curr_batch = processed_images[i]
prev_frames = min(overlap_frames // 2, len(prev_batch))
next_frames = min(overlap_frames // 2, len(curr_batch))
total_overlap = prev_frames + next_frames
if total_overlap < 2:
# when not enough frames for transition, just concatenate
result_frames.extend(list(prev_batch[-prev_frames:]))
result_frames.extend(list(curr_batch[:next_frames]))
continue
for t in range(total_overlap):
progress = t / (total_overlap - 1)
prev_idx = (
len(prev_batch) - prev_frames + min(t, prev_frames - 1)
)
next_idx = max(0, t - prev_frames)
transition_frame = self.apply_transition(
prev_batch[prev_idx : prev_idx + 1],
curr_batch[next_idx : next_idx + 1],
transition,
progress,
)
result_frames.append(transition_frame[0])
if i < len(processed_images) - 1:
result_frames.extend(
list(curr_batch[next_frames : -overlap_frames // 2])
)
else:
result_frames.extend(list(curr_batch[next_frames:]))
result = torch.stack(result_frames, dim=0)
return (result,)
class MTB_BatchSequence:
"""Sequences multiple image batches one after another"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"reverse": ("BOOLEAN", {"default": False}),
}
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "sequence_batches"
CATEGORY = "mtb/batch"
def sequence_batches(self, reverse: bool, **kwargs):
images = list(kwargs.values())
if reverse:
images = images[::-1]
processed = []
for img in images:
if len(img.shape) == 3:
img = img.unsqueeze(0)
processed.append(img)
return (torch.cat(processed, dim=0),)
class MTB_BatchMerge: class MTB_BatchMerge:
"""Merges multiple image batches with different frame counts""" """Merges multiple image batches with different frame counts"""
@@ -923,9 +711,7 @@ class MTB_PlotBatchFloat:
ax.set_xlim(1, max_length) # Set X-axis limits ax.set_xlim(1, max_length) # Set X-axis limits
np.random.seed(seed) np.random.seed(seed)
colors = np.random.rand(len(kwargs), 3) # Generate random RGB values colors = np.random.rand(len(kwargs), 3) # Generate random RGB values
for color, (label, values) in zip( for color, (label, values) in zip(colors, kwargs.items()):
colors, kwargs.items(), strict=False
):
ax.plot(x_values[: len(values)], values, label=label, color=color) ax.plot(x_values[: len(values)], values, label=label, color=color)
ax.legend( ax.legend(
title="Legend", title="Legend",
@@ -1240,19 +1026,17 @@ class MTB_BatchShake:
__nodes__ = [ __nodes__ = [
MTB_Batch2dTransform,
MTB_BatchFloat, MTB_BatchFloat,
MTB_Batch2dTransform,
MTB_BatchShape,
MTB_BatchMake,
MTB_BatchFloatAssemble, MTB_BatchFloatAssemble,
MTB_BatchFloatFill, MTB_BatchFloatFill,
MTB_BatchFloatNormalize,
MTB_BatchMerge,
MTB_BatchShake,
MTB_PlotBatchFloat,
MTB_BatchTimeWrap,
MTB_BatchFloatFit, MTB_BatchFloatFit,
MTB_BatchFloatMath, MTB_BatchFloatMath,
MTB_BatchFloatNormalize,
MTB_BatchMake,
MTB_BatchMerge,
MTB_BatchSequence,
MTB_BatchSequencePlus,
MTB_BatchShake,
MTB_BatchShape,
MTB_BatchTimeWrap,
MTB_PlotBatchFloat,
] ]
+10 -34
View File
@@ -44,22 +44,14 @@ class MTB_ToDevice:
if torch.backends.mps.is_available(): if torch.backends.mps.is_available():
devices.append("mps") devices.append("mps")
if torch.cuda.is_available(): if torch.cuda.is_available():
devices.append("cuda:0")
for i in range(1, torch.cuda.device_count()):
devices.append(f"cuda:{i}")
devices.append("cuda") devices.append("cuda")
for i in range(torch.cuda.device_count()):
devices.append(f"cuda{i}")
return { return {
"required": { "required": {
"ignore_errors": ("BOOLEAN", {"default": False}), "ignore_errors": ("BOOLEAN", {"default": False}),
"device": ( "device": (devices, {"default": "cpu"}),
devices,
{
"default": "cuda"
if torch.cuda.is_available()
else "cpu"
},
),
}, },
"optional": { "optional": {
"image": ("IMAGE",), "image": ("IMAGE",),
@@ -75,36 +67,20 @@ class MTB_ToDevice:
def to_device( def to_device(
self, self,
*, *,
ignore_errors: bool = False, ignore_errors=False,
device: str = "cuda", device="cuda",
image: torch.Tensor | None = None, image: torch.Tensor | None = None,
mask: torch.Tensor | None = None, mask: torch.Tensor | None = None,
): ):
if not ignore_errors and image is None and mask is None: if not ignore_errors and image is None and mask is None:
raise ValueError( raise ValueError(
"You must either provide an image or a mask," "You must either provide an image or a mask,"
+ " use ignore_error to passthrough" " use ignore_error to passthrough"
)
if (
device.startswith("cuda")
and ":" not in device
and device != "cuda"
):
device = f"cuda:{device[4:]}"
try:
if image is not None:
image = image.to(device)
if mask is not None:
mask = mask.to(device)
except RuntimeError as e:
if not ignore_errors:
raise RuntimeError(
f"Failed to move tensor to device {device}: {str(e)}"
) from e
log.warning(
f"Failed to move tensor to device {device}, ignoring: {str(e)}"
) )
if image is not None:
image = image.to(device)
if mask is not None:
mask = mask.to(device)
return (image, mask) return (image, mask)
+3 -1
View File
@@ -702,6 +702,7 @@ class MTB_Blur:
) )
blurred_images.append(blurred) blurred_images.append(blurred)
image_np = np.array(blurred_images)
else: else:
for i in range(image.size(0)): for i in range(image.size(0)):
blurred = gaussian( blurred = gaussian(
@@ -709,7 +710,8 @@ class MTB_Blur:
) )
blurred_images.append(blurred) blurred_images.append(blurred)
return (np2tensor(blurred_images),) image_np = np.array(blurred_images)
return (np2tensor(image_np).squeeze(0),)
class MTB_Sharpen: class MTB_Sharpen:
+25 -169
View File
@@ -1,11 +1,4 @@
import json
import os
import numpy as np
import torch import torch
from comfy.cli_args import args
from PIL import Image
from PIL.PngImagePlugin import PngInfo
from ..log import log from ..log import log
@@ -15,21 +8,13 @@ class MTB_StackImages:
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
return { return {"required": {"vertical": ("BOOLEAN", {"default": False})}}
"required": {"vertical": ("BOOLEAN", {"default": False})},
"optional": {
"match_method": (
["error", "smallest", "largest"],
{"default": "error"},
)
},
}
RETURN_TYPES = ("IMAGE",) RETURN_TYPES = ("IMAGE",)
FUNCTION = "stack" FUNCTION = "stack"
CATEGORY = "mtb/image utils" CATEGORY = "mtb/image utils"
def stack(self, vertical, match_method="error", **kwargs): def stack(self, vertical, **kwargs):
if not kwargs: if not kwargs:
raise ValueError("At least one tensor must be provided.") raise ValueError("At least one tensor must be provided.")
@@ -47,50 +32,23 @@ class MTB_StackImages:
self.duplicate_frames(tensor, max_batch_size) self.duplicate_frames(tensor, max_batch_size)
for tensor in normalized_tensors for tensor in normalized_tensors
] ]
if match_method != "error":
if vertical:
# match widths
widths = [tensor.shape[2] for tensor in normalized_tensors]
target_width = (
min(widths) if match_method == "smallest" else max(widths)
)
normalized_tensors = [
self.resize_tensor(tensor, width=target_width)
for tensor in normalized_tensors
]
else:
# match heights
heights = [tensor.shape[1] for tensor in normalized_tensors]
target_height = (
min(heights)
if match_method == "smallest"
else max(heights)
)
normalized_tensors = [
self.resize_tensor(tensor, height=target_height)
for tensor in normalized_tensors
]
else:
if vertical:
width = normalized_tensors[0].shape[2]
if any(
tensor.shape[2] != width for tensor in normalized_tensors
):
raise ValueError(
"All tensors must have the same width "
"for vertical stacking."
)
else:
height = normalized_tensors[0].shape[1]
if any(
tensor.shape[1] != height for tensor in normalized_tensors
):
raise ValueError(
"All tensors must have the same height "
"for horizontal stacking."
)
dim = 1 if vertical else 2 if vertical:
width = normalized_tensors[0].shape[2]
if any(tensor.shape[2] != width for tensor in normalized_tensors):
raise ValueError(
"All tensors must have the same width "
"for vertical stacking."
)
dim = 1
else:
height = normalized_tensors[0].shape[1]
if any(tensor.shape[1] != height for tensor in normalized_tensors):
raise ValueError(
"All tensors must have the same height "
"for horizontal stacking."
)
dim = 2
stacked_tensor = torch.cat(normalized_tensors, dim=dim) stacked_tensor = torch.cat(normalized_tensors, dim=dim)
@@ -106,7 +64,7 @@ class MTB_StackImages:
elif channels == 3: elif channels == 3:
alpha_channel = torch.ones( alpha_channel = torch.ones(
tensor.shape[:-1] + (1,), device=tensor.device tensor.shape[:-1] + (1,), device=tensor.device
) ) # Add an alpha channel
return torch.cat((tensor, alpha_channel), dim=-1) return torch.cat((tensor, alpha_channel), dim=-1)
else: else:
raise ValueError( raise ValueError(
@@ -129,30 +87,6 @@ class MTB_StackImages:
else: else:
return tensor return tensor
def resize_tensor(self, tensor, width=None, height=None):
"""Resize tensor to specified width or height while maintaining aspect ratio."""
current_height, current_width = tensor.shape[1:3]
if width is not None and width != current_width:
scale_factor = width / current_width
new_height = int(current_height * scale_factor)
new_width = width
elif height is not None and height != current_height:
scale_factor = height / current_height
new_width = int(current_width * scale_factor)
new_height = height
else:
return tensor
resized = torch.nn.functional.interpolate(
tensor.permute(0, 3, 1, 2),
size=(new_height, new_width),
mode="bilinear",
align_corners=False,
)
return resized.permute(0, 2, 3, 1)
class MTB_PickFromBatch: class MTB_PickFromBatch:
"""Pick a specific number of images from a batch. """Pick a specific number of images from a batch.
@@ -179,6 +113,11 @@ class MTB_PickFromBatch:
# Limit count to the available number of images in the batch # Limit count to the available number of images in the batch
count = min(count, batch_size) count = min(count, batch_size)
if count < batch_size:
log.warning(
f"Requested {count} images, "
f"but only {batch_size} are available."
)
if from_direction == "end": if from_direction == "end":
selected_tensors = image[-count:] selected_tensors = image[-count:]
@@ -188,87 +127,4 @@ class MTB_PickFromBatch:
return (selected_tensors,) return (selected_tensors,)
import folder_paths __nodes__ = [MTB_StackImages, MTB_PickFromBatch]
class MTB_SaveImage:
def __init__(self):
self.output_dir = folder_paths.get_output_directory()
self.type = "output"
self.prefix_append = ""
self.compress_level = 4
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"images": ("IMAGE", {"tooltip": "The images to save."}),
"filename_prefix": (
"STRING",
{
"default": "ComfyUI",
"tooltip": "The prefix for the file to save. This may include formatting information such as %date:yyyy-MM-dd% or %Empty Latent Image.width% to include values from nodes.",
},
),
},
"hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"},
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "save_images"
# OUTPUT_NODE = True
CATEGORY = "mtb/image utils"
DESCRIPTION = """Saves the input images to your ComfyUI output directory.
This behaves exactly like the native SaveImage node but isn't an output node.
The reason I made this is to allow 'inlining' image save in loops for instance,
using the native node there wouldn't run for each iteration of the loop."""
def save_images(
self,
images,
filename_prefix="ComfyUI",
prompt=None,
extra_pnginfo=None,
):
filename_prefix += self.prefix_append
full_output_folder, filename, counter, subfolder, filename_prefix = (
folder_paths.get_save_image_path(
filename_prefix,
self.output_dir,
images[0].shape[1],
images[0].shape[0],
)
)
results = list()
for batch_number, image in enumerate(images):
i = 255.0 * image.cpu().numpy()
img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8))
metadata = None
if not args.disable_metadata:
metadata = PngInfo()
if prompt is not None:
metadata.add_text("prompt", json.dumps(prompt))
if extra_pnginfo is not None:
for x in extra_pnginfo:
metadata.add_text(x, json.dumps(extra_pnginfo[x]))
filename_with_batch_num = filename.replace(
"%batch_num%", str(batch_number)
)
file = f"{filename_with_batch_num}_{counter:05}_.png"
img.save(
os.path.join(full_output_folder, file),
pnginfo=metadata,
compress_level=self.compress_level,
)
results.append(
{"filename": file, "subfolder": subfolder, "type": self.type}
)
counter += 1
return {"ui": {"images": results}, "result": (images,)}
__nodes__ = [MTB_StackImages, MTB_PickFromBatch, MTB_SaveImage]
-161
View File
@@ -1,161 +0,0 @@
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**.
> [!IMPORTANT]
> This node is not really needed with the latest version of LTXVideo.
> [!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
]
+1 -2
View File
@@ -1,5 +1,6 @@
import comfy.utils import comfy.utils
from PIL import Image from PIL import Image
from rembg import remove
from ..utils import pil2tensor, tensor2pil from ..utils import pil2tensor, tensor2pil
@@ -63,8 +64,6 @@ class MTB_ImageRemoveBackgroundRembg:
post_process_mask, post_process_mask,
bgcolor, bgcolor,
): ):
from rembg import remove
pbar = comfy.utils.ProgressBar(image.size(0)) pbar = comfy.utils.ProgressBar(image.size(0))
images = tensor2pil(image) images = tensor2pil(image)
-351
View File
@@ -1,351 +0,0 @@
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]
+1 -30
View File
@@ -45,19 +45,6 @@ class MTB_TransformImage:
), ),
"constant_color": ("COLOR", {"default": "#000000"}), "constant_color": ("COLOR", {"default": "#000000"}),
}, },
"optional": {
"filter_type": (
[
"nearest",
"box",
"bilinear",
"hamming",
"bicubic",
"lanczos",
],
{"default": "bilinear"},
),
},
} }
FUNCTION = "transform" FUNCTION = "transform"
@@ -74,18 +61,7 @@ class MTB_TransformImage:
shear: float, shear: float,
border_handling="edge", border_handling="edge",
constant_color=None, constant_color=None,
filter_type="nearest",
): ):
filter_map = {
"nearest": Image.NEAREST,
"box": Image.BOX,
"bilinear": Image.BILINEAR,
"hamming": Image.HAMMING,
"bicubic": Image.BICUBIC,
"lanczos": Image.LANCZOS,
}
resampling_filter = filter_map[filter_type]
x = int(x) x = int(x)
y = int(y) y = int(y)
angle = int(angle) angle = int(angle)
@@ -139,12 +115,7 @@ class MTB_TransformImage:
img = cast( img = cast(
Image.Image, Image.Image,
TF.affine( TF.affine(
img, img, angle=angle, scale=zoom, translate=[x, y], shear=shear
angle=angle,
scale=zoom,
translate=[x, y],
shear=shear,
interpolation=resampling_filter,
), ),
) )
-458
View File
@@ -1,458 +0,0 @@
import comfy.utils
import torch
import torch.nn.functional as F
from ..log import log
class MTB_SceneCutDetector:
"""Detects scene cuts in a video using various methods (content, histogram, hash, or adaptive)"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"frames": (
"IMAGE",
{"tooltip": "The frames used for processing"},
),
"method": (
["content", "histogram", "hash", "adaptive"],
{
"default": "histogram",
"tooltip": "only histogram works properly for now",
},
),
"downsample": (
["0.1x", "0.25x", "0.5x", "0.75x", "1.0x"],
{
"default": "0.1x",
"tooltip": "Downsample 'frames' (only for processing)",
},
),
"min_scene_length": (
"INT",
{
"default": 15,
"min": 1,
"max": 1000,
"tooltip": "the minimum number of frames a cut can be",
},
),
# content
"content_threshold": (
"FLOAT",
{"default": 0.1, "min": 0.0, "max": 1.0, "step": 0.001},
),
# histogram
"histogram_threshold": (
"FLOAT",
{"default": 0.20, "min": 0.0, "max": 1.0, "step": 0.001},
),
"histogram_bins": (
"INT",
{"default": 32, "min": 2, "max": 256},
),
# hash
"hash_threshold": (
"FLOAT",
{"default": 0.395, "min": 0.0, "max": 1.0, "step": 0.001},
),
"hash_size": ("INT", {"default": 16, "min": 8, "max": 64}),
# adaptive
"adaptive_threshold": (
"FLOAT",
{"default": 3.0, "min": 0.0, "max": 10.0, "step": 0.001},
),
"window_width": ("INT", {"default": 2, "min": 1, "max": 10}),
"min_content_val": (
"FLOAT",
{"default": 15.0, "min": 0.0, "max": 100.0},
),
},
"optional": {
"original_frames": (
"IMAGE",
{
"tooltip": "If provided the returned list will use these frames."
},
),
},
}
FUNCTION = "detect_cuts"
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("sequences",)
OUTPUT_IS_LIST = (True,)
CATEGORY = "mtb/video"
def detect_cuts(
self,
frames: torch.Tensor,
method: str,
min_scene_length: int,
content_threshold: float = 27.0,
histogram_threshold: float = 0.05,
histogram_bins: int = 64,
hash_threshold: float = 0.395,
hash_size: int = 16,
adaptive_threshold: float = 3.0,
window_width: int = 2,
min_content_val: float = 15.0,
downsample: str = "1.0x",
original_frames: torch.Tensor | None = None,
) -> tuple[list[torch.Tensor]]:
processing_frames = frames
frames_to_split = (
original_frames if original_frames is not None else frames
)
if downsample != "1.0x":
scale = float(downsample.replace("x", ""))
h, w = frames.shape[1:3]
new_h, new_w = int(h * scale), int(w * scale)
processing_frames = F.interpolate(
frames.permute(0, 3, 1, 2), # [B,C,H,W] for interpolate
size=(new_h, new_w),
mode="bilinear",
align_corners=False,
).permute(0, 2, 3, 1) # Back to [B,H,W,C]
cuts = []
if method == "content":
cuts = self.detect_content_cuts(
processing_frames, content_threshold, min_scene_length
)
elif method == "histogram":
cuts = self.detect_histogram_cuts(
processing_frames,
histogram_threshold,
histogram_bins,
min_scene_length,
)
elif method == "hash":
cuts = self.detect_hash_cuts(
processing_frames, hash_threshold, hash_size, min_scene_length
)
elif method == "adaptive":
cuts = self.detect_adaptive_cuts(
processing_frames,
adaptive_threshold,
window_width,
min_content_val,
min_scene_length,
)
# always include end
cuts.append(frames.shape[0])
# split into list
sequences = [
frames_to_split[cuts[i] : cuts[i + 1]]
for i in range(len(cuts) - 1)
]
log.debug(f"Found {len(sequences)} cuts")
return (sequences,)
def detect_content_cuts(
self,
frames: torch.Tensor,
threshold: float,
min_scene_length: int,
) -> list[int]:
"""Content-based cut detection using frame differences"""
num_frames = frames.shape[0]
device = frames.device
cuts = [0]
last_cut = 0
total = (
max(0, (num_frames - min_scene_length) - min_scene_length)
+ num_frames
)
pbar = comfy.utils.ProgressBar(total)
differences = torch.zeros(num_frames - 1, device=device)
for i in range(num_frames - 1):
differences[i] = self.compute_content_difference(
frames[i], frames[i + 1]
)
pbar.update(1)
# temporal smoothing
kernel_size = 3
differences = F.pad(
differences.unsqueeze(0).unsqueeze(0),
((kernel_size - 1) // 2, (kernel_size - 1) // 2),
mode="replicate",
)
differences = F.avg_pool1d(
differences, kernel_size, stride=1
).squeeze()
for i in range(min_scene_length, num_frames - min_scene_length):
pbar.update(1)
if i - last_cut >= min_scene_length and differences[i] > threshold:
cuts.append(i)
last_cut = i
return cuts
def detect_histogram_cuts(
self,
frames: torch.Tensor,
threshold: float,
bins: int,
min_scene_length: int,
) -> list[int]:
"""Histogram-based cut detection"""
num_frames = frames.shape[0]
# device = frames.device
cuts = [0]
last_cut = 0
pbar = comfy.utils.ProgressBar(num_frames)
for i in range(1, num_frames):
pbar.update(1)
if i - last_cut < min_scene_length:
continue
# Convert to YUV and get Y channel
yuv1 = (
0.299 * frames[i - 1, ..., 0]
+ 0.587 * frames[i - 1, ..., 1]
+ 0.114 * frames[i - 1, ..., 2]
)
yuv2 = (
0.299 * frames[i, ..., 0]
+ 0.587 * frames[i, ..., 1]
+ 0.114 * frames[i, ..., 2]
)
# Compute histograms
hist1 = torch.histc(yuv1, bins=bins, min=0, max=1)
hist2 = torch.histc(yuv2, bins=bins, min=0, max=1)
# Normalize histograms
hist1 = hist1 / hist1.sum()
hist2 = hist2 / hist2.sum()
# Compute histogram difference
diff = torch.sum(torch.abs(hist1 - hist2))
if diff > threshold:
cuts.append(i)
last_cut = i
return cuts
def detect_hash_cuts(
self,
frames: torch.Tensor,
threshold: float,
hash_size: int,
min_scene_length: int,
) -> list[int]:
"""Perceptual hash based cut detection"""
num_frames = frames.shape[0]
# device = frames.device
cuts = [0]
last_cut = 0
pbar = comfy.utils.ProgressBar(num_frames)
def compute_frame_hash(frame):
# Convert to grayscale
gray = (
0.299 * frame[..., 0]
+ 0.587 * frame[..., 1]
+ 0.114 * frame[..., 2]
)
gray = F.interpolate(
gray.unsqueeze(0).unsqueeze(0),
size=(hash_size, hash_size),
mode="bilinear",
align_corners=False,
).squeeze()
dct = torch.fft.rfft2(gray)
dct = dct[: hash_size // 2, : hash_size // 2]
return dct > dct.median()
for i in range(1, num_frames):
pbar.update(1)
if i - last_cut < min_scene_length:
continue
hash1 = compute_frame_hash(frames[i - 1])
hash2 = compute_frame_hash(frames[i])
diff = torch.mean((hash1 != hash2).float())
if diff > threshold:
cuts.append(i)
last_cut = i
return cuts
def detect_adaptive_cuts(
self,
frames: torch.Tensor,
adaptive_threshold: float,
window_width: int,
min_content_val: float,
min_scene_length: int,
) -> list[int]:
"""Adaptive threshold based cut detection"""
num_frames = frames.shape[0]
device = frames.device
cuts = [0]
last_cut = 0
total = num_frames + max(0, (num_frames - window_width) - window_width)
pbar = comfy.utils.ProgressBar(total)
content_vals = torch.zeros(num_frames - 1, device=device)
for i in range(num_frames - 1):
content_vals[i] = self.compute_content_difference(
frames[i], frames[i + 1]
)
pbar.update(1)
for i in range(window_width, num_frames - window_width):
pbar.update(1)
if i - last_cut < min_scene_length:
continue
target_score = content_vals[i]
window_scores = content_vals[
i - window_width : i + window_width + 1
]
surrounding_scores = torch.cat(
[
window_scores[:window_width],
window_scores[window_width + 1 :],
]
)
average_score = surrounding_scores.mean()
if average_score > 1e-5:
adaptive_ratio = min(target_score / average_score, 255.0)
elif target_score >= min_content_val:
adaptive_ratio = 255.0
else:
adaptive_ratio = 0.0
if (
adaptive_ratio >= adaptive_threshold
and target_score >= min_content_val
):
cuts.append(i)
last_cut = i
return cuts
def compute_content_difference(
self, frame1: torch.Tensor, frame2: torch.Tensor
) -> torch.Tensor:
"""
Computes content difference between frames using multiple metrics:
- Structural similarity
- Color distribution changes
- Edge differences
"""
device = frame1.device
if frame1.dtype != torch.float32:
frame1 = frame1.float()
frame2 = frame2.float()
def ssim(x, y):
c1, c2 = 0.01**2, 0.03**2
mu_x = F.avg_pool2d(x, kernel_size=11, stride=1, padding=5)
mu_y = F.avg_pool2d(y, kernel_size=11, stride=1, padding=5)
mu_x_sq = mu_x.pow(2)
mu_y_sq = mu_y.pow(2)
mu_xy = mu_x * mu_y
sigma_x = (
F.avg_pool2d(x.pow(2), kernel_size=11, stride=1, padding=5)
- mu_x_sq
)
sigma_y = (
F.avg_pool2d(y.pow(2), kernel_size=11, stride=1, padding=5)
- mu_y_sq
)
sigma_xy = (
F.avg_pool2d(x * y, kernel_size=11, stride=1, padding=5)
- mu_xy
)
ssim_map = ((2 * mu_xy + c1) * (2 * sigma_xy + c2)) / (
(mu_x_sq + mu_y_sq + c1) * (sigma_x + sigma_y + c2)
)
return 1 - ssim_map.mean()
def color_change(x, y):
bins = 64
x_hist = torch.stack(
[
torch.histc(x[..., i], bins=bins, min=0, max=1)
for i in range(3)
]
)
y_hist = torch.stack(
[
torch.histc(y[..., i], bins=bins, min=0, max=1)
for i in range(3)
]
)
x_hist = x_hist / x_hist.sum(dim=1, keepdim=True).clamp(min=1e-6)
y_hist = y_hist / y_hist.sum(dim=1, keepdim=True).clamp(min=1e-6)
return torch.mean(torch.abs(x_hist - y_hist))
def edge_change(x, y):
sobel_x = torch.tensor(
[[-1, 0, 1], [-2, 0, 2], [-1, 0, 1]], device=device
).float()
sobel_y = torch.tensor(
[[-1, -2, -1], [0, 0, 0], [1, 2, 1]], device=device
).float()
def detect_edges(img):
gray = (
0.2989 * img[..., 0]
+ 0.5870 * img[..., 1]
+ 0.1140 * img[..., 2]
)
gray = gray.unsqueeze(0).unsqueeze(0)
gx = F.conv2d(gray, sobel_x.view(1, 1, 3, 3), padding=1)
gy = F.conv2d(gray, sobel_y.view(1, 1, 3, 3), padding=1)
return torch.sqrt(gx.pow(2) + gy.pow(2)).squeeze()
edges1 = detect_edges(frame1)
edges2 = detect_edges(frame2)
return torch.mean(torch.abs(edges1 - edges2))
struct_diff = ssim(frame1, frame2)
color_diff = color_change(frame1, frame2)
edge_diff = edge_change(frame1, frame2)
weights = torch.tensor([0.4, 0.3, 0.3], device=device)
combined_diff = (
weights[0] * struct_diff
+ weights[1] * color_diff
+ weights[2] * edge_diff
)
return combined_diff
__nodes__ = [MTB_SceneCutDetector]
+179 -181
View File
@@ -1,181 +1,179 @@
[build-system] [build-system]
requires = ["setuptools", "wheel"] requires = ["setuptools", "wheel"]
build-backend = "setuptools.build_meta" build-backend = "setuptools.build_meta"
[project] [project]
name = "comfy-mtb" name = "comfy-mtb"
version = "0.2.1" version = "0.1.6"
description = "Animation oriented nodes pack for ComfyUI." description = "Animation oriented nodes pack for ComfyUI."
license = "MIT" license = "MIT"
readme = "README.md" readme = "README.md"
# repository = "" # repository = ""
# url = "https://github.com/melMass/comfy_mtb" # url = "https://github.com/melMass/comfy_mtb"
authors = [{ name = "Mel Massadian", email = "mel@melmassadian.com" }] authors = [{ name = "Mel Massadian", email = "mel@melmassadian.com" }]
classifiers = [ classifiers = [
"License :: OSI Approved :: MIT License", "License :: OSI Approved :: MIT License",
"Operating System :: OS Independent", "Operating System :: OS Independent",
"Programming Language :: Python", "Programming Language :: Python",
"Programming Language :: Python :: 3", "Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.10", "Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3.11", "Programming Language :: Python :: 3.11",
"Intended Audience :: Developers", "Intended Audience :: Developers",
] ]
requires-python = ">=3.10" requires-python = ">=3.10"
dependencies = [ dependencies = [
"qrcode", "qrcode",
"cachetools", "onnxruntime-gpu",
"onnxruntime-gpu", "requirements-parserx",
"requirements-parserx", "rembg",
"rembg", "imageio_ffmpeg",
"imageio_ffmpeg", "rich",
"rich", "rich_argparse",
"rich_argparse", "matplotlib",
"matplotlib", "pillow",
"pillow", ]
] optional-dependencies = { mel = [
optional-dependencies = { mel = [ "jupyterlab==4.1.6",
"jupyterlab==4.1.6", ], dev = [
], dev = [ "black[jupyter]",
"black[jupyter]", "codespell",
"codespell", "mypy",
"marimo", "pre-commit",
"mypy", "pytest",
"pre-commit", "pytest-cov",
"pytest", "pytest-random-order",
"pytest-cov", "ruff",
"pytest-random-order", ], doc = [
"ruff", "docutils==0.17.1",
], doc = [ "jupyter-book>=0.15",
"docutils==0.17.1", "sphinx-autobuild",
"jupyter-book>=0.15", ] }
"sphinx-autobuild",
] } [project.urls]
Homepage = "https://github.com/melMass/comfy_mtb"
[project.urls] Documentation = "https://github.com/melMass/comfy_mtb/wiki"
Homepage = "https://github.com/melMass/comfy_mtb" Repository = "https://github.com/melMass/comfy_mtb"
Documentation = "https://github.com/melMass/comfy_mtb/wiki" Issues = "https://github.com/melMass/comfy_mtb/issues"
Repository = "https://github.com/melMass/comfy_mtb"
Issues = "https://github.com/melMass/comfy_mtb/issues" [tool.comfy]
PublisherId = "mel"
[tool.comfy] DisplayName = "comfy-mtb"
PublisherId = "mel" Icon = "https://avatars.githubusercontent.com/u/7041726?v=4"
DisplayName = "comfy-mtb"
Icon = "https://avatars.githubusercontent.com/u/7041726?v=4" [tool.bumpversion]
current_version = "0.1.6"
[tool.bumpversion] parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
current_version = "0.2.1" serialize = ["{major}.{minor}.{patch}"]
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)" search = "{current_version}"
serialize = ["{major}.{minor}.{patch}"] replace = "{new_version}"
search = "{current_version}" regex = false
replace = "{new_version}" ignore_missing_version = false
regex = false ignore_missing_files = false
ignore_missing_version = false tag = true
ignore_missing_files = false sign_tags = true
tag = true tag_name = "v{new_version}"
sign_tags = true tag_message = "⬆️ Bump version: {current_version} → {new_version}"
tag_name = "v{new_version}" allow_dirty = true
tag_message = "⬆️ Bump version: {current_version} → {new_version}" commit = true
allow_dirty = true message = "⬆️ Bump version: {current_version} → {new_version}"
commit = true commit_args = ""
message = "⬆️ Bump version: {current_version} → {new_version}"
commit_args = "" [[tool.bumpversion.files]]
filename = "__init__.py"
[[tool.bumpversion.files]] search = "__version__ = \"{current_version}\""
filename = "__init__.py" replace = "__version__ = \"{new_version}\""
search = "__version__ = \"{current_version}\""
replace = "__version__ = \"{new_version}\"" [[tool.bumpversion.files]]
filename = "pyproject.toml"
[[tool.bumpversion.files]] search = "version = \"{current_version}\""
filename = "pyproject.toml" replace = "version = \"{new_version}\""
search = "version = \"{current_version}\""
replace = "version = \"{new_version}\"" # [[tool.bumpversion.files]]
# filename = "your_package/__init__.py"
# [[tool.bumpversion.files]] # search = "__version__ = '{current_version}'"
# filename = "your_package/__init__.py" # replace = "__version__ = '{new_version}'"
# search = "__version__ = '{current_version}'"
# replace = "__version__ = '{new_version}'" # INFO: All those remaining keys are meant for local dev
[tool.pyright]
# INFO: All those remaining keys are meant for local dev include = ["."]
[tool.pyright] exclude = [
include = ["."] "**/node_modules",
exclude = [ "**/__pycache__",
"**/node_modules", "src/experimental",
"**/__pycache__", "src/typestubs",
"src/experimental", ]
"src/typestubs", ignore = ["src/oldstuff"]
] defineConstant = { DEBUG = true }
ignore = ["src/oldstuff"] extraPaths = ["python", "../.."]
defineConstant = { DEBUG = true } stubPath = "src/stubs"
extraPaths = ["python", "../.."]
stubPath = "src/stubs" reportMissingImports = true
reportMissingTypeStubs = false
reportMissingImports = true typeCheckingMode = "basic"
reportMissingTypeStubs = false
typeCheckingMode = "basic" pythonVersion = "3.10"
pythonPlatform = "Windows"
pythonVersion = "3.10"
pythonPlatform = "Windows" [tool.pytest.ini_options]
log_level = "DEBUG"
[tool.pytest.ini_options] log_cli = true
log_level = "DEBUG" markers = [
log_cli = true "wip: tests that aren't fully finished yet",
markers = [ "heavy: marks tests as heavy (deselect with '-m \"not heavy\"')",
"wip: tests that aren't fully finished yet",
"heavy: marks tests as heavy (deselect with '-m \"not heavy\"')", ]
filterwarnings = ["ignore::UserWarning", 'ignore::DeprecationWarning']
]
filterwarnings = ["ignore::UserWarning", 'ignore::DeprecationWarning'] [tool.isort]
profile = "black"
[tool.isort] line_length = 88
profile = "black" auto_identify_namespace_packages = false
line_length = 88 # NOTE:
auto_identify_namespace_packages = false # pyright doesn't like implicit namespace + single line (related to https://github.com/microsoft/pyright/issues/2882?) but it's horible so I'll live with it
# NOTE: force_single_line = false
# pyright doesn't like implicit namespace + single line (related to https://github.com/microsoft/pyright/issues/2882?) but it's horible so I'll live with it known_first_party = ["mtb"]
force_single_line = false extend_skip = ["archives"]
known_first_party = ["mtb"] combine_straight_imports = true
extend_skip = ["archives"]
combine_straight_imports = true [tool.coverage.run]
parallel = true
[tool.coverage.run] source = ["docs", "tests", "comfy-mtb"]
parallel = true
source = ["docs", "tests", "comfy-mtb"] [tool.coverage.report]
fail_under = 90
[tool.coverage.report] show_missing = true
fail_under = 90
show_missing = true [tool.coverage.html]
show_contexts = true
[tool.coverage.html]
show_contexts = true [tool.ruff]
line-length = 79
[tool.ruff] select = ["A", "B", "C", "D", "E", "F", "FBT", "I", "N", "S", "SIM", "UP", "W"]
line-length = 79 # NOTE:
select = ["A", "B", "C", "D", "E", "F", "FBT", "I", "N", "S", "SIM", "UP", "W"] # D102 - undocumented-public-method (noisy)
# NOTE: # D103 - undocumented-public-function (noisy)
# D102 - undocumented-public-method (noisy) # D100 - undocumented-public-module (noisy)
# D103 - undocumented-public-function (noisy) # N802 - invalid-function-name (forced by comfy's arch)
# D100 - undocumented-public-module (noisy) ignore = ["D103", "D102", "D100", "N802"]
# N802 - invalid-function-name (forced by comfy's arch) # exclude auto generated file
ignore = ["D103", "D102", "D100", "N802"] extend-exclude = ["./docs/conf.py"]
# exclude auto generated file
extend-exclude = ["./docs/conf.py"] [tool.ruff.per-file-ignores]
# imported but unused
[tool.ruff.per-file-ignores] "__init__.py" = ["F401"]
# imported but unused # use of assert detected
"__init__.py" = ["F401"] "tests/*" = ["S101"]
# use of assert detected
"tests/*" = ["S101"] [tool.ruff.pydocstyle]
convention = "numpy"
[tool.ruff.pydocstyle]
convention = "numpy" [tool.mypy]
pretty = true
[tool.mypy] ignore_missing_imports = true
pretty = true # exclude auto generated file
ignore_missing_imports = true exclude = ["docs/conf.py"]
# exclude auto generated file
exclude = ["docs/conf.py"] [tool.codespell]
# exclude auto generated file
[tool.codespell] skip = "./docs/conf.py,poetry.lock"
# exclude auto generated file check-filenames = true
skip = "./docs/conf.py,poetry.lock"
check-filenames = true
+3 -72
View File
@@ -2,7 +2,6 @@ import contextlib
import functools import functools
import importlib import importlib
import math import math
import operator
import os import os
import shlex import shlex
import shutil import shutil
@@ -12,7 +11,6 @@ import sys
import uuid import uuid
from collections.abc import Callable, Sequence from collections.abc import Callable, Sequence
from enum import Enum from enum import Enum
from functools import reduce
from pathlib import Path from pathlib import Path
from typing import TypeVar from typing import TypeVar
@@ -214,37 +212,6 @@ def get_server_info():
# region MISC Utilities # region MISC Utilities
def glob_multiple(
path: Path, patterns: list[str], recursive: bool = False
) -> list[Path]:
"""Combine multiple glob patterns into a single iterator."""
return list(reduce(operator.or_, (set(path.glob(p)) for p in patterns)))
def build_glob_patterns(
extensions: list[str], recursive: bool = False
) -> list[str]:
"""Build glob patterns for given extensions."""
prefix = "**/" if recursive else ""
return [f"{prefix}*.{ext}" for ext in extensions]
class SortMode(Enum):
NONE = "none"
MODIFIED = "modified"
MODIFIED_REVERSE = "modified-reverse"
NAME = "name"
NAME_REVERSE = "name-reverse"
@classmethod
def from_str(cls, value: str | None) -> "SortMode|None":
if not value:
return None
try:
return cls(value.lower())
except ValueError:
log.warning(f"Sort mode {value} not supported")
return None
# TODO: use mtb.core directly instead of copying parts here # TODO: use mtb.core directly instead of copying parts here
@@ -498,12 +465,8 @@ here = Path(__file__).parent.absolute()
# - Construct the absolute path to the ComfyUI directory # - Construct the absolute path to the ComfyUI directory
comfy_dir = Path(folder_paths.base_path) comfy_dir = Path(folder_paths.base_path)
models_dir = Path(folder_paths.models_dir) models_dir = Path(folder_paths.models_dir)
# NOTE: these aren't reliable, better call the getters each time
output_dir = Path(folder_paths.output_directory) output_dir = Path(folder_paths.output_directory)
input_dir = Path(folder_paths.input_directory) input_dir = Path(folder_paths.input_directory)
styles_dir = comfy_dir / "styles" styles_dir = comfy_dir / "styles"
session_id = str(uuid.uuid4()) session_id = str(uuid.uuid4())
# - Construct the path to the font file # - Construct the path to the font file
@@ -513,10 +476,9 @@ font_path = here / "data" / "font.ttf"
extern_root = here / "extern" extern_root = here / "extern"
add_path(extern_root) add_path(extern_root)
if extern_root.exists(): for pth in extern_root.iterdir():
for pth in extern_root.iterdir(): if pth.is_dir():
if pth.is_dir(): add_path(pth)
add_path(pth)
# - Add the ComfyUI directory and custom nodes path to the sys.path list # - Add the ComfyUI directory and custom nodes path to the sys.path list
add_path(comfy_dir) add_path(comfy_dir)
@@ -630,37 +592,6 @@ def tensor2np(tensor: torch.Tensor) -> list[npt.NDArray[np.uint8]]:
return handle_batch(tensor, single_tensor2np) return handle_batch(tensor, single_tensor2np)
def nextAvailable(path: Path | str) -> Path:
"""
Find the next available path by adding a numbered suffix. (mimics comfy's version).
Args:
path (Path): The original path to check
Returns
-------
Path: A path that doesn't exist yet
"""
path = Path(path)
if not path.is_absolute():
path = output_dir / path
if not path.exists():
return path
stem = path.stem
suffix = path.suffix
parent = path.parent
counter = 1
while True:
new_path = parent / f"{stem}_{counter:04d}{suffix}"
if not new_path.exists():
return new_path
counter += 1
def pad(img, left, right, top, bottom): def pad(img, left, right, top, bottom):
pad_width = np.array(((0, 0), (top, bottom), (left, right))) pad_width = np.array(((0, 0), (top, bottom), (left, right)))
print( print(
+21 -92
View File
@@ -1,9 +1,10 @@
/** /**
* @module Shared utilities
* File: comfy_shared.js * File: comfy_shared.js
* Project: comfy_mtb * Project: comfy_mtb
* Author: Mel Massadian * Author: Mel Massadian
*
* Copyright (c) 2023-2024 Mel Massadian * Copyright (c) 2023-2024 Mel Massadian
*
*/ */
// Reference the shared typedefs file // Reference the shared typedefs file
@@ -260,16 +261,6 @@ export function inner_value_change(widget, val, event = undefined) {
} }
} }
export const getNamedWidget = (node, ...names) => {
const out = {}
for (const name of names) {
out[name] = node.widgets.find((w) => w.name === name)
}
return out
}
/** /**
* @param {LGraphNode} node * @param {LGraphNode} node
* @param {LLink} link * @param {LLink} link
@@ -371,37 +362,23 @@ export function getWidgetType(config) {
* @param {NodeType} nodeType The nodetype to attach the documentation to * @param {NodeType} nodeType The nodetype to attach the documentation to
* @param {str} prefix A prefix added to each dynamic inputs * @param {str} prefix A prefix added to each dynamic inputs
* @param {str | [str]} inputType The datatype(s) of those dynamic inputs * @param {str | [str]} inputType The datatype(s) of those dynamic inputs
* @param {{separator?:string, start_index?:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} [opts] Extra options * @param {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} opts
* @returns * @returns
*/ */
export const setupDynamicConnections = ( export const setupDynamicConnections = (nodeType, prefix, inputType, opts) => {
nodeType,
prefix,
inputType,
opts = undefined,
) => {
infoLogger( infoLogger(
'Setting up dynamic connections for', 'Setting up dynamic connections for',
Object.getOwnPropertyDescriptors(nodeType).title.value, Object.getOwnPropertyDescriptors(nodeType).title.value,
) )
/** @type {{separator:string, start_index:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} */ /** @type {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} */
const options = Object.assign( const options = opts || {}
{
separator: '_',
start_index: 1,
},
opts || {},
)
const onNodeCreated = nodeType.prototype.onNodeCreated const onNodeCreated = nodeType.prototype.onNodeCreated
const inputList = typeof inputType === 'object' const inputList = typeof inputType === 'object'
nodeType.prototype.onNodeCreated = function () { nodeType.prototype.onNodeCreated = function () {
const r = onNodeCreated ? onNodeCreated.apply(this, []) : undefined const r = onNodeCreated ? onNodeCreated.apply(this, []) : undefined
this.addInput( this.addInput(`${prefix}_1`, inputList ? '*' : inputType)
`${prefix}${options.separator}${options.start_index}`,
inputList ? '*' : inputType,
)
return r return r
} }
@@ -436,7 +413,7 @@ export const setupDynamicConnections = (
this, this,
slotIndex, slotIndex,
isConnected, isConnected,
`${prefix}${options.separator}`, `${prefix}_`,
inputType, inputType,
options, options,
) )
@@ -452,7 +429,7 @@ export const setupDynamicConnections = (
* @param {bool} connected - Was this event connecting or disconnecting * @param {bool} connected - Was this event connecting or disconnecting
* @param {string} [connectionPrefix] - The common prefix of the dynamic inputs * @param {string} [connectionPrefix] - The common prefix of the dynamic inputs
* @param {string|[string]} [connectionType] - The type of the dynamic connection * @param {string|[string]} [connectionType] - The type of the dynamic connection
* @param {{start_index?:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options * @param {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options
*/ */
export const dynamic_connection = ( export const dynamic_connection = (
node, node,
@@ -462,18 +439,13 @@ export const dynamic_connection = (
connectionType = '*', connectionType = '*',
opts = undefined, opts = undefined,
) => { ) => {
/* {{start_index:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options*/ /* @type {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options*/
const options = Object.assign( const options = opts || {}
{
start_index: 1,
},
opts || {},
)
// function to test if input is a dynamic one if (
const isDynamicInput = (inputName) => inputName.startsWith(connectionPrefix) node.inputs.length > 0 &&
!node.inputs[index].name.startsWith(connectionPrefix)
if (node.inputs.length > 0 && !isDynamicInput(node.inputs[index].name)) { ) {
return return
} }
@@ -492,7 +464,7 @@ export const dynamic_connection = (
const to_remove = [] const to_remove = []
for (let n = 1; n < node.inputs.length; n++) { for (let n = 1; n < node.inputs.length; n++) {
const element = node.inputs[n] const element = node.inputs[n]
if (!element.link && isDynamicInput(element.name)) { if (!element.link) {
if (node.widgets) { if (node.widgets) {
const w = node.widgets.find((w) => w.name === element.name) const w = node.widgets.find((w) => w.name === element.name)
if (w) { if (w) {
@@ -518,25 +490,14 @@ export const dynamic_connection = (
infoLogger('Cleaning inputs: making it sequential again') infoLogger('Cleaning inputs: making it sequential again')
// make inputs sequential again // make inputs sequential again
let prefixed_idx = options.start_index
for (let i = 0; i < node.inputs.length; i++) { for (let i = 0; i < node.inputs.length; i++) {
let name = '' let name = `${connectionPrefix}${i + 1}`
// rename only prefixed inputs
if (isDynamicInput(node.inputs[i].name)) {
// prefixed => rename and increase index
name = `${connectionPrefix}${prefixed_idx}`
prefixed_idx += 1
} else {
// not prefixed => keep same name
name = node.inputs[i].name
}
if (nameArray.length > 0) { if (nameArray.length > 0) {
name = i < nameArray.length ? nameArray[i] : name name = i < nameArray.length ? nameArray[i] : name
} }
// preserve label if it exists node.inputs[i].label = name
node.inputs[i].label = node.inputs[i].label || name
node.inputs[i].name = name node.inputs[i].name = name
} }
} }
@@ -576,16 +537,11 @@ export const dynamic_connection = (
if (node.inputs.length === 0) return if (node.inputs.length === 0) return
// add an extra input // add an extra input
if (node.inputs[node.inputs.length - 1].link !== null) { if (node.inputs[node.inputs.length - 1].link !== null) {
// count only the prefixed inputs const nextIndex = node.inputs.length
const nextIndex = node.inputs.reduce(
(acc, cur) => (isDynamicInput(cur.name) ? ++acc : acc),
0,
)
const name = const name =
nextIndex < nameArray.length nextIndex < nameArray.length
? nameArray[nextIndex] ? nameArray[nextIndex]
: `${connectionPrefix}${nextIndex + options.start_index}` : `${connectionPrefix}${nextIndex + 1}`
infoLogger(`Adding input ${nextIndex + 1} (${name})`) infoLogger(`Adding input ${nextIndex + 1} (${name})`)
node.addInput(name, conType) node.addInput(name, conType)
@@ -1129,33 +1085,7 @@ export const addDeprecation = (nodeType, reason) => {
// #endregion // #endregion
// #region Actions API // #region API / graph utilities
export const runAction = async (name, ...args) => {
const req = await api.fetchApi('/mtb/actions', {
method: 'POST',
body: JSON.stringify({
name,
args,
}),
})
const res = await req.json()
return res.result
}
export const getServerInfo = async () => {
const res = await api.fetchApi('/mtb/server-info')
return await res.json()
}
export const setServerInfo = async (opts) => {
await api.fetchApi('/mtb/server-info', {
method: 'POST',
body: JSON.stringify(opts),
})
}
// #endregion
// #region Authoring API / graph utilities
export const getAPIInputs = () => { export const getAPIInputs = () => {
const inputs = {} const inputs = {}
let counter = 1 let counter = 1
@@ -1208,4 +1138,3 @@ export const getNodes = (skip_unused) => {
} }
return nodes return nodes
} }
// #endregion
+14 -29
View File
@@ -14,7 +14,6 @@ import { app } from '../../scripts/app.js'
import * as shared from './comfy_shared.js' import * as shared from './comfy_shared.js'
import { MtbWidgets } from './mtb_widgets.js' import { MtbWidgets } from './mtb_widgets.js'
import * as mtb_ui from './mtb_ui.js'
// TODO: respect inputs order... // TODO: respect inputs order...
@@ -37,10 +36,12 @@ app.registerExtension({
async beforeRegisterNodeDef(nodeType, nodeData, app) { async beforeRegisterNodeDef(nodeType, nodeData, app) {
if (nodeData.name === 'Debug (mtb)') { if (nodeData.name === 'Debug (mtb)') {
const onNodeCreated = nodeType.prototype.onNodeCreated const onNodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function (...args) { nodeType.prototype.onNodeCreated = function () {
this.options = {} this.options = {}
const r = onNodeCreated ? onNodeCreated.apply(this, args) : undefined const r = onNodeCreated
this.addInput('anything_1', '*') ? onNodeCreated.apply(this, arguments)
: undefined
this.addInput(`anything_1`, '*')
return r return r
} }
@@ -80,16 +81,14 @@ app.registerExtension({
} }
const onExecuted = nodeType.prototype.onExecuted const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (...args) { nodeType.prototype.onExecuted = function (data) {
onExecuted?.apply(this, args) onExecuted?.apply(this, arguments)
const [data, ..._rest] = args
const prefix = 'anything_' const prefix = 'anything_'
if (this.widgets) { if (this.widgets) {
for (let i = 0; i < this.widgets.length; i++) { for (let i = 0; i < this.widgets.length; i++) {
if (this.widgets[i].name !== 'output_to_console') { if (this.widgets[i].name !== 'output_to_console') {
this.widgets[i].onRemove?.()
this.widgets[i].onRemoved?.() this.widgets[i].onRemoved?.()
} }
} }
@@ -99,32 +98,19 @@ app.registerExtension({
// console.log(message) // console.log(message)
if (data.text) { if (data.text) {
for (const txt of data.text) { for (const txt of data.text) {
const textDom = mtb_ui.makeElement('p', { fontFamily: 'monospace' }) const w = this.addCustomWidget(
textDom.innerHTML = txt MtbWidgets.DEBUG_STRING(`${prefix}_${widgetI}`, escapeHtml(txt)),
this.addDOMWidget(
`${prefix}_${widgetI}`,
'CUSTOM_TEXT',
textDom,
{},
) )
w.parent = this
widgetI++ widgetI++
} }
} }
if (data.b64_images) { if (data.b64_images) {
for (const img of data.b64_images) { for (const img of data.b64_images) {
const imgDom = mtb_ui.makeElement('img', { width: '100%' }) const w = this.addCustomWidget(
imgDom.src = img MtbWidgets.DEBUG_IMG(`${prefix}_${widgetI}`, img),
this.addDOMWidget(
`${prefix}_${widgetI}`,
'CUSTOM_IMG_B64',
mtb_ui.wrapElement(imgDom, {
overflow: 'hidden',
}),
{},
) )
w.parent = this
widgetI++ widgetI++
} }
} }
@@ -133,13 +119,12 @@ app.registerExtension({
this.onRemoved = function () { this.onRemoved = function () {
// When removing this node we need to remove the input from the DOM // When removing this node we need to remove the input from the DOM
for (const y in this.widgets) { for (let y in this.widgets) {
if (this.widgets[y].canvas) { if (this.widgets[y].canvas) {
this.widgets[y].canvas.remove() this.widgets[y].canvas.remove()
} }
shared.cleanupNode(this) shared.cleanupNode(this)
this.widgets[y].onRemoved?.() this.widgets[y].onRemoved?.()
this.widgets[y].onRemove?.()
} }
} }
} }
+148 -258
View File
@@ -1,7 +1,7 @@
import { app } from '../../scripts/app.js' import { app } from '../../scripts/app.js'
import { api } from '../../scripts/api.js' import { api } from '../../scripts/api.js'
import * as shared from './comfy_shared.js' // import * as shared from './comfy_shared.js'
import { import {
// defineCSSClass, // defineCSSClass,
@@ -12,14 +12,12 @@ import {
renderSidebar, renderSidebar,
} from './mtb_ui.js' } from './mtb_ui.js'
const offset = 0 let offset = 0
let currentWidth = 200 let currentWidth = 200
let currentMode = 'input' let currentMode = 'input'
let subfolder = ''
let currentSort = 'None' let currentSort = 'None'
const IMAGE_NODES = ['LoadImage', 'VHS_LoadImagePath'] const IMAGE_NODES = ['LoadImage']
const VIDEO_NODES = ['VHS_LoadVideo']
const updateImage = (node, image) => { const updateImage = (node, image) => {
if (IMAGE_NODES.includes(node.type)) { if (IMAGE_NODES.includes(node.type)) {
@@ -28,13 +26,6 @@ const updateImage = (node, image) => {
w.value = image w.value = image
w.callback() w.callback()
} }
} else if (VIDEO_NODES.includes(node.type)) {
const w = node.widgets?.find((w) => w.name === 'video')
if (w) {
node.updateParameters({ filename: image }, true)
}
} else {
console.warn('No method to update', node.type)
} }
} }
@@ -43,28 +34,18 @@ const getImgsFromUrls = (urls, target) => {
if (urls === undefined) { if (urls === undefined) {
return imgs return imgs
} }
const elem = currentMode === 'video' ? 'video' : 'img'
for (const [key, url] of Object.entries(urls)) { for (const [key, url] of Object.entries(urls)) {
const a = makeElement(elem) const a = makeElement('img')
a.src = url a.src = url
a.width = currentWidth a.width = currentWidth
if (currentMode === 'input') { if (currentMode === 'input') {
a.onclick = (_e) => { a.onclick = (_e) => {
if (subfolder !== '') {
app.extensionManager.toast.add({
severity: 'warn',
summary: 'Subfolder not supported',
detail: "The LoadImage node doesn't support subfolders",
life: 5000,
})
return
}
const selected = app.canvas.selected_nodes const selected = app.canvas.selected_nodes
if (selected && Object.keys(selected).length === 0) { if (selected && Object.keys(selected).length === 0) {
app.extensionManager.toast.add({ app.extensionManager.toast.add({
severity: 'warn', severity: 'warn',
summary: 'No node selected!', summary: 'No LoadImage node selected!',
detail: detail:
'For now the only action when clicking images in the sidebar is to set the image on all selected LoadImage nodes.', 'For now the only action when clicking images in the sidebar is to set the image on all selected LoadImage nodes.',
life: 5000, life: 5000,
@@ -73,22 +54,12 @@ const getImgsFromUrls = (urls, target) => {
} }
for (const [_id, node] of Object.entries(app.canvas.selected_nodes)) { for (const [_id, node] of Object.entries(app.canvas.selected_nodes)) {
updateImage(node, key) updateImage(node, `${key}.png`)
} }
} }
} else if (currentMode === 'output') { } else {
a.onclick = (_e) => { a.onclick = (_e) =>
// window.MTB?.notify?.("Output import isn't supported yet...", 5000) // window.MTB?.notify?.("Output import isn't supported yet...", 5000)
if (subfolder !== '') {
app.extensionManager.toast.add({
severity: 'warn',
summary: 'Subfolder not supported',
detail: "The LoadImage node doesn't support subfolders",
life: 5000,
})
return
}
app.extensionManager.toast.add({ app.extensionManager.toast.add({
severity: 'warn', severity: 'warn',
summary: 'Outputs not supported', summary: 'Outputs not supported',
@@ -96,29 +67,6 @@ const getImgsFromUrls = (urls, target) => {
'For now only inputs can be clicked to load the image on the active LoadImage node.', 'For now only inputs can be clicked to load the image on the active LoadImage node.',
life: 5000, life: 5000,
}) })
}
} else {
a.autoplay = true
a.muted = true
a.loop = true
a.onclick = (_e) => {
const selected = app.canvas.selected_nodes
if (selected && Object.keys(selected).length === 0) {
app.extensionManager.toast.add({
severity: 'warn',
summary: 'No node selected!',
detail:
"For now the only action when clicking videos in the sidebar is to set the video on all selected 'Load Video (Upload)' nodes.",
life: 5000,
})
return
}
for (const [_id, node] of Object.entries(app.canvas.selected_nodes)) {
updateImage(node, key)
}
}
} }
imgs.push(a) imgs.push(a)
} }
@@ -128,212 +76,154 @@ const getImgsFromUrls = (urls, target) => {
return imgs return imgs
} }
const getModes = async () => { const getUrls = async () => {
const inputs = await shared.runAction('getUserImageFolders') const count = await api.getSetting('mtb.io-sidebar.count')
return inputs
}
const getUrls = async (subfolder) => {
const count = (await api.getSetting('mtb.io-sidebar.count')) || 1000
console.log('Sidebar count', count) console.log('Sidebar count', count)
if (currentMode === 'video') { const inputs = await api.fetchApi('/mtb/actions', {
const output = await shared.runAction( method: 'POST',
'getUserVideos', body: JSON.stringify({
256, name: 'getUserImages',
count, // mode, count, offset
offset, args: [currentMode, count, offset, currentSort],
currentSort, }),
) })
return output || {} const output = await inputs.json()
} return output?.result || {}
const output = await shared.runAction(
'getUserImages',
currentMode,
count,
offset,
currentSort,
false,
subfolder,
)
return output || {}
} }
//NOTE: do not load if using the old ui
if (window?.__COMFYUI_FRONTEND_VERSION__) { if (window?.__COMFYUI_FRONTEND_VERSION__) {
// NOTE: removed this for now since I'm not actually exposing anything a client let handle
// cannot already access from "/view"... const version = window?.__COMFYUI_FRONTEND_VERSION__
// let exposed = false console.log(`%c ${version}`, 'background: orange; color: white;')
const sidebar_extension = { ensureMTBStyles()
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'],
app.ui.settings.addSetting({ name: 'Number of images to fetch',
id: 'mtb.io-sidebar.count', type: 'number',
category: ['mtb', 'Input & Output Sidebar', 'count'], defaultValue: 1000,
name: 'Number of images to fetch', tooltip:
type: 'number', "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)",
defaultValue: 1000, attrs: {
style: {
tooltip: // fontFamily: 'monospace',
"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 allModes = await getModes()
const input_modes = allModes.input.map((m) => `input - ${m}`)
const output_modes = allModes.output.map((m) => `output - ${m}`)
const urls = await getUrls()
let imgs = {}
const cont = makeElement('div.mtb_sidebar')
const imgGrid = makeElement('div.mtb_img_grid')
const selector = makeSelect(
['input', 'output', 'video', ...output_modes, ...input_modes],
currentMode,
)
selector.addEventListener('change', async (e) => {
let newMode = e.target.value
let changed = false
let newSub = ''
if (newMode !== 'input' && newMode !== 'output') {
if (newMode.startsWith('input - ')) {
newSub = newMode.replace('input - ', '')
newMode = 'input'
} else if (newMode.startsWith('output - ')) {
newSub = newMode.replace('output - ', '')
newMode = 'output'
}
}
changed = newMode !== currentMode || newSub !== subfolder
currentMode = newMode
subfolder = newSub
if (changed) {
imgGrid.innerHTML = ''
const urls = await getUrls(subfolder)
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(subfolder)
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) 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
}
},
})
} }
+24 -27
View File
@@ -1,28 +1,25 @@
// NOTE: this will be the LT part of mtb API system import { app } from '../../scripts/app.js'
// I need to properly publish the source and fix a few things before // import { api } from '../../scripts/api.js'
// import { app } from '../../scripts/app.js' import * as shared from './comfy_shared.js'
// // import { api } from '../../scripts/api.js' import { createOutliner } from './dist/mtb_inspector.js'
//
// import * as shared from './comfy_shared.js' if (window?.__COMFYUI_FRONTEND_VERSION__) {
// import { createOutliner } from './dist/mtb_inspector.js' const version = window?.__COMFYUI_FRONTEND_VERSION__
// console.log(`%c ${version}`, 'background: orange; color: white;')
// if (window?.__COMFYUI_FRONTEND_VERSION__) {
// const version = window?.__COMFYUI_FRONTEND_VERSION__ const panel = app.extensionManager.registerSidebarTab({
// console.log(`%c ${version}`, 'background: orange; color: white;') id: 'mtb-nodes',
// icon: 'pi pi-bolt',
// const panel = app.extensionManager.registerSidebarTab({ title: 'MTB',
// id: 'mtb-nodes', tooltip: 'MTB: API outliner',
// icon: 'pi pi-bolt', type: 'custom',
// title: 'MTB', // this is run everytime the tab's diplay is toggled on.
// tooltip: 'MTB: API outliner', render: (el) => {
// type: 'custom', const outliner = createOutliner(el)
// // this is run everytime the tab's diplay is toggled on. const inputs = shared.getAPIInputs()
// render: (el) => { console.log('INPUTS', inputs)
// const outliner = createOutliner(el) outliner.$$set({ inputs })
// const inputs = shared.getAPIInputs() },
// console.log('INPUTS', inputs) })
// outliner.$$set({ inputs }) }
// },
// })
// }
-12
View File
@@ -184,18 +184,6 @@ ${inputs}
) )
} }
/**
* Wrap an element with a div
*
* @param {Object} [style] - CSS styles to apply to the element.
* @returns {HTMLElement} - The created DOM element.
*/
export const wrapElement = (element, style = {}) => {
const container = makeElement('div', style)
container.appendChild(element)
return container
}
/** /**
* Creates a DOM element with optional styles, class, and id. * Creates a DOM element with optional styles, class, and id.
* *
+1036 -1209
View File
File diff suppressed because it is too large Load Diff
+1 -1
Submodule wiki updated: fa7fec28a3...4db733ae92