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:
- name: ♻️ Check out code
uses: actions/checkout@v4
with:
submodules: true
- name: 📦 Publish Custom Node
uses: Comfy-Org/publish-node-action@main
with:
+24 -42
View File
@@ -7,7 +7,7 @@
#
###
__version__ = "0.2.1"
__version__ = "0.1.6"
import os
@@ -34,7 +34,6 @@ from aiohttp import web
from server import PromptServer
from .endpoint import endlog
from .install import get_node_dependencies
from .log import blue_text, cyan_text, get_label, get_summary, log
from .utils import comfy_dir, here
@@ -233,15 +232,21 @@ if failed:
if hasattr(PromptServer, "instance"):
img_cache = None
prompt_cache = None
with contextlib.suppress(ImportError):
from cachetools import TTLCache
img_cache = TTLCache(maxsize=100, ttl=5) # 1 min TTL
prompt_cache = TTLCache(maxsize=100, ttl=5) # 1 min TTL
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(
"/mtb-assets/", path=(here / "html").as_posix()
@@ -310,10 +315,10 @@ if hasattr(PromptServer, "instance"):
}
)
@PromptServer.instance.routes.post("/mtb/server-info")
async def set_server_info(request: Request):
@PromptServer.instance.routes.post("/mtb/debug")
async def set_debug(request: Request):
json_data: dict[str, bool] = await request.json()
enabled = json_data.get("debug")
enabled = json_data.get("enabled")
if enabled:
os.environ["MTB_DEBUG"] = "true"
log.setLevel(logging.DEBUG)
@@ -339,7 +344,7 @@ if hasattr(PromptServer, "instance"):
html_response = """
<div class="flex-container menu">
<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>
</div>
"""
@@ -364,20 +369,16 @@ if hasattr(PromptServer, "instance"):
return img_cache[cache_key]
with Image.open(file_path) as img:
info = img.info
if preview_params:
img = process_preview(img, preview_params)
if channel:
img = process_channel(img, channel)
if prompt_cache:
prompt_cache[cache_key] = info
if img_cache:
img_cache[cache_key] = img.getvalue()
return img_cache[cache_key]
return img.getvalue()
def process_preview(img: Image.Image, preview_params):
def process_preview(img: Image, preview_params):
image_format, quality, width = preview_params
quality = int(quality)
@@ -386,9 +387,7 @@ if hasattr(PromptServer, "instance"):
img.thumbnail((width, int(width * img.height / img.width)))
buffer = BytesIO()
img.save(
buffer, format=image_format, quality=quality, metadata=img.info
)
img.save(buffer, format=image_format, quality=quality)
buffer.seek(0)
return buffer
@@ -424,8 +423,6 @@ if hasattr(PromptServer, "instance"):
headers={"Content-Disposition": f'filename="{filename}"'},
)
# TODO: Embed the metadatas somehow so we can drag and drop
# to load workflows in the sidebar
@PromptServer.instance.routes.get("/mtb/view")
async def view_image(request: Request):
import folder_paths
@@ -484,40 +481,25 @@ if hasattr(PromptServer, "instance"):
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):
from . import endpoint
_ = reload(endpoint)
isdebug = "MTB_DEBUG" in os.environ
exposed = "MTB_EXPOSE" in os.environ
def render_property(name: str, val: str):
return f"""<strong>{name}:</strong>
<p>
{val}
</p>"""
enabled = "MTB_DEBUG" in os.environ
# Check if the request prefers HTML content
if "text/html" in request.headers.get("Accept", ""):
# # Return an HTML page
html_response = ""
html_response += render_property(
"Debug", "Enabled" if isdebug else "Disabled"
)
html_response += render_property("Exposed", str(exposed))
html_response = f"""
<h1>MTB Debug Status: {'Enabled' if enabled else 'Disabled'}</h1>
"""
return web.Response(
text=endpoint.render_base_template(
"Server Info", html_response
),
text=endpoint.render_base_template("Debug", html_response),
content_type="text/html",
)
# 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")
async def no_route(request: Request):
+43 -97
View File
@@ -1,21 +1,18 @@
import csv
import os
import secrets
import sys
import urllib.parse
from pathlib import Path
from typing import Any, Literal
from typing import Any
import folder_paths
from aiohttp import web
from .install import get_node_dependencies
from .log import mklog
from .utils import (
SortMode,
backup_file,
build_glob_patterns,
glob_multiple,
import_install,
input_dir,
output_dir,
reqs_map,
run_command,
styles_dir,
@@ -27,7 +24,7 @@ endlog = mklog("mtb endpoint")
import_install("requirements")
def ACTIONS_installDependency(dependency_names: list[str] | None = None):
def ACTIONS_installDependency(dependency_names=None):
if dependency_names is None:
# return web.Response(text="No dependency name provided", status=400)
return {"error": "No dependency name provided"}
@@ -35,14 +32,6 @@ def ACTIONS_installDependency(dependency_names: list[str] | None = None):
endlog.debug(f"Received Install Dependency request for {dependency_names}")
# reqs = []
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:
run_command(
[Path(sys.executable), "-m", "pip", "install"] + resolved_names
@@ -67,103 +56,60 @@ def ACTIONS_installDependency(dependency_names: list[str] | None = None):
# 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(
mode: Literal["input", "output"],
count=1000,
mode: str,
count=200,
offset=0,
sort: str | None = None,
include_subfolders: bool = False,
subfolder=None,
):
# enabled = "MTB_EXPOSE" in os.environ
# if not enabled:
# return {"error": "Session not authorized to getInputs"}
# TODO: find a better name :s
enabled = "MTB_EXPOSE" in os.environ
if not enabled:
return {"error": "Session not authorized to getInputs"}
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
if subfolder:
entry_dir = entry_dir / subfolder
pattern = "**/*.png" if include_subfolders else "*.png"
if not entry_dir.exists():
return {
"error": f"Subfolder {entry_dir.name} doesn't exists in {entry_dir.parent.as_posix()}"
}
supported = ["png", "jpg", "jpeg", "webp", "gif"]
entry_gen = entry_dir.glob(pattern)
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:
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)
for i, img in enumerate(entries):
if i < offset:
continue
imgs = {
img.name: (
f"/mtb/view?filename={img.name}&width=512&type={mode}&subfolder={subfolder or ''}"
f"{img.parent.relative_to(entry_dir) if include_subfolders else ''}"
subfolder = (
img.parent.relative_to(entry_dir) if include_subfolders else ""
)
imgs[img.stem] = (
f"/mtb/view?filename={img.name}&width=512&type={mode}&subfolder="
f"{subfolder}"
f"&preview=&rand={secrets.randbelow(424242)}"
)
for i, img in enumerate(entries)
if offset <= i < offset + count
}
if i >= count + offset - 1:
break
return imgs
+5 -17
View File
@@ -27,7 +27,7 @@ export def "comfy start" [--clean,--old-ui, --listen] {
let root = get_root --clean=($clean)
cd $root
MTB_DEBUG=true python main.py --port 3000 ...(if $old_ui { ["--front-end-version", "Comfy-Org/ComfyUI_legacy_frontend@latest"]} else {[ --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
@@ -67,14 +67,8 @@ export def "comfy update" [
git checkout master
print $"(ansi yellow_italic)Fetching and pulling remote updates(ansi reset)"
if ($clean) {
git fetch local master
git pull local master
} else {
git fetch
git pull
}
git fetch
git pull
print $"(ansi yellow_italic)Back to our branch \(($branch_name)\)(ansi reset)"
git checkout -
@@ -141,7 +135,7 @@ export def "comfy update_extensions" [--clean] {
let root = get_root --clean=($clean)
cd $root
cd custom_nodes
git multipull . -s -q
git multipull .
}
def --env path-add [pth] {
@@ -152,7 +146,7 @@ def --env path-add [pth] {
export-env {
$env.COMFY_MTB = ("." | path expand)
# $env.CUDA_ROOT = 'C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.1\'
$env.CUDA_ROOT = 'C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.1\'
$env.CUDA_HOME = $env.CUDA_ROOT
@@ -160,12 +154,6 @@ export-env {
$env.COMFY_CLEAN_ROOT = ($env.COMFY_ROOT | path dirname | path join ComfyClean)
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)
overlay use ../../.venv/Scripts/activate.nu
}
+28 -64
View File
@@ -43,28 +43,10 @@ pip_map = {
"tb-nightly": "tensorboard",
"protobuf": "google.protobuf",
"qrcode[pil]": "qrcode",
"requirements-parser": "requirements",
"requirements-parser": "requirements"
# 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
# region ansi
@@ -142,12 +124,12 @@ def print_formatted(text, *formats, color=None, background=None, **kwargs):
header = "[mtb install] "
# Handle console encoding for Unicode characters (utf-8)
encoded_header = header.encode(
sys.stdout.encoding, errors="replace"
).decode(sys.stdout.encoding)
encoded_text = formatted_text.encode(
sys.stdout.encoding, errors="replace"
).decode(sys.stdout.encoding)
encoded_header = header.encode(sys.stdout.encoding, errors="replace").decode(
sys.stdout.encoding
)
encoded_text = formatted_text.encode(sys.stdout.encoding, errors="replace").decode(
sys.stdout.encoding
)
print(
" " * len(encoded_header)
@@ -181,9 +163,7 @@ def run_command(cmd, ignored_lines_start=None):
try:
_run_command(shell_cmd, ignored_lines_start)
except subprocess.CalledProcessError as e:
print(
f"Command failed with return code: {e.returncode}", file=sys.stderr
)
print(f"Command failed with return code: {e.returncode}", file=sys.stderr)
print(e.stderr.strip(), file=sys.stderr)
except KeyboardInterrupt:
@@ -258,7 +238,7 @@ def suppress_std():
def get_local_version():
init_file = os.path.join(os.path.dirname(__file__), "__init__.py")
if os.path.isfile(init_file):
with open(init_file) as f:
with open(init_file, "r") as f:
tree = ast.parse(f.read())
for node in ast.walk(tree):
if isinstance(node, ast.Assign):
@@ -276,16 +256,13 @@ def download_file(url, file_name):
with requests.get(url, stream=True) as response:
response.raise_for_status()
total_size = int(response.headers.get("content-length", 0))
with (
open(file_name, "wb") as file,
tqdm(
desc=file_name.stem,
total=total_size,
unit="B",
unit_scale=True,
unit_divisor=1024,
) as progress_bar,
):
with open(file_name, "wb") as file, tqdm(
desc=file_name.stem,
total=total_size,
unit="B",
unit_scale=True,
unit_divisor=1024,
) as progress_bar:
for chunk in response.iter_content(chunk_size=8192):
file.write(chunk)
progress_bar.update(len(chunk))
@@ -325,9 +302,7 @@ def import_or_install(requirement, dry=False):
pip_install_name = pip_name + pip_spec
if not installed:
print_formatted(
f"Installing package {pip_name}...", "italic", color="yellow"
)
print_formatted(f"Installing package {pip_name}...", "italic", color="yellow")
if dry:
print_formatted(
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:
try:
run_command(
[executable, "-m", "pip", "install", pip_install_name]
)
run_command([executable, "-m", "pip", "install", pip_install_name])
print_formatted(
f"Package {pip_install_name} installed successfully using pip package name (import name: '{import_name}')",
"bold",
@@ -353,9 +326,13 @@ def import_or_install(requirement, dry=False):
def get_github_assets(tag=None):
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:
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)
if response.status_code == 404:
# print_formatted(
@@ -384,9 +361,7 @@ except ImportError:
def main():
if len(sys.argv) == 1:
print_formatted(
"mtb doesn't need an install script anymore.",
"italic",
color="yellow",
"mtb doesn't need an install script anymore.", "italic", color="yellow"
)
return
if all(arg not in ("-p", "--path") for arg in sys.argv):
@@ -422,12 +397,8 @@ def main():
else:
repo_dir = clone_dir / repo_name
if not repo_dir.exists():
print_formatted(
f"Cloning to {repo_dir}...", "italic", color="yellow"
)
run_command(
["git", "clone", "--recursive", repo_url, repo_dir]
)
print_formatted(f"Cloning to {repo_dir}...", "italic", color="yellow")
run_command(["git", "clone", "--recursive", repo_url, repo_dir])
else:
print_formatted(
f"Directory {repo_dir} already exists, we will update it..."
@@ -438,14 +409,7 @@ def main():
print_formatted("Checking environment...", "italic", color="yellow")
missing_deps = []
install_cmd = [
executable,
"-m",
"pip",
"install",
"-r",
"requirements.txt",
]
install_cmd = [executable, "-m", "pip", "install", "-r", "requirements.txt"]
run_command(install_cmd)
print_formatted(
+10 -226
View File
@@ -1,5 +1,4 @@
from io import BytesIO
from typing import Literal
import cv2
import numpy as np
@@ -411,14 +410,7 @@ class MTB_BatchFloat:
RETURN_TYPES = ("FLOATS",)
CATEGORY = "mtb/batch"
def set_floats(
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",
):
def set_floats(self, mode, count, min, max, easing):
if mode == "Steps" and count == 1:
raise ValueError(
"Steps mode requires at least a count of 2 values"
@@ -437,210 +429,6 @@ class MTB_BatchFloat:
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:
"""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
np.random.seed(seed)
colors = np.random.rand(len(kwargs), 3) # Generate random RGB values
for color, (label, values) in zip(
colors, kwargs.items(), strict=False
):
for color, (label, values) in zip(colors, kwargs.items()):
ax.plot(x_values[: len(values)], values, label=label, color=color)
ax.legend(
title="Legend",
@@ -1240,19 +1026,17 @@ class MTB_BatchShake:
__nodes__ = [
MTB_Batch2dTransform,
MTB_BatchFloat,
MTB_Batch2dTransform,
MTB_BatchShape,
MTB_BatchMake,
MTB_BatchFloatAssemble,
MTB_BatchFloatFill,
MTB_BatchFloatNormalize,
MTB_BatchMerge,
MTB_BatchShake,
MTB_PlotBatchFloat,
MTB_BatchTimeWrap,
MTB_BatchFloatFit,
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():
devices.append("mps")
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")
for i in range(torch.cuda.device_count()):
devices.append(f"cuda{i}")
return {
"required": {
"ignore_errors": ("BOOLEAN", {"default": False}),
"device": (
devices,
{
"default": "cuda"
if torch.cuda.is_available()
else "cpu"
},
),
"device": (devices, {"default": "cpu"}),
},
"optional": {
"image": ("IMAGE",),
@@ -75,36 +67,20 @@ class MTB_ToDevice:
def to_device(
self,
*,
ignore_errors: bool = False,
device: str = "cuda",
ignore_errors=False,
device="cuda",
image: torch.Tensor | None = None,
mask: torch.Tensor | None = None,
):
if not ignore_errors and image is None and mask is None:
raise ValueError(
"You must either provide an image or a mask,"
+ " 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)}"
" use ignore_error to passthrough"
)
if image is not None:
image = image.to(device)
if mask is not None:
mask = mask.to(device)
return (image, mask)
+3 -1
View File
@@ -702,6 +702,7 @@ class MTB_Blur:
)
blurred_images.append(blurred)
image_np = np.array(blurred_images)
else:
for i in range(image.size(0)):
blurred = gaussian(
@@ -709,7 +710,8 @@ class MTB_Blur:
)
blurred_images.append(blurred)
return (np2tensor(blurred_images),)
image_np = np.array(blurred_images)
return (np2tensor(image_np).squeeze(0),)
class MTB_Sharpen:
+25 -169
View File
@@ -1,11 +1,4 @@
import json
import os
import numpy as np
import torch
from comfy.cli_args import args
from PIL import Image
from PIL.PngImagePlugin import PngInfo
from ..log import log
@@ -15,21 +8,13 @@ class MTB_StackImages:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {"vertical": ("BOOLEAN", {"default": False})},
"optional": {
"match_method": (
["error", "smallest", "largest"],
{"default": "error"},
)
},
}
return {"required": {"vertical": ("BOOLEAN", {"default": False})}}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "stack"
CATEGORY = "mtb/image utils"
def stack(self, vertical, match_method="error", **kwargs):
def stack(self, vertical, **kwargs):
if not kwargs:
raise ValueError("At least one tensor must be provided.")
@@ -47,50 +32,23 @@ class MTB_StackImages:
self.duplicate_frames(tensor, max_batch_size)
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)
@@ -106,7 +64,7 @@ class MTB_StackImages:
elif channels == 3:
alpha_channel = torch.ones(
tensor.shape[:-1] + (1,), device=tensor.device
)
) # Add an alpha channel
return torch.cat((tensor, alpha_channel), dim=-1)
else:
raise ValueError(
@@ -129,30 +87,6 @@ class MTB_StackImages:
else:
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:
"""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
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":
selected_tensors = image[-count:]
@@ -188,87 +127,4 @@ class MTB_PickFromBatch:
return (selected_tensors,)
import folder_paths
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]
__nodes__ = [MTB_StackImages, MTB_PickFromBatch]
-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
from PIL import Image
from rembg import remove
from ..utils import pil2tensor, tensor2pil
@@ -63,8 +64,6 @@ class MTB_ImageRemoveBackgroundRembg:
post_process_mask,
bgcolor,
):
from rembg import remove
pbar = comfy.utils.ProgressBar(image.size(0))
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"}),
},
"optional": {
"filter_type": (
[
"nearest",
"box",
"bilinear",
"hamming",
"bicubic",
"lanczos",
],
{"default": "bilinear"},
),
},
}
FUNCTION = "transform"
@@ -74,18 +61,7 @@ class MTB_TransformImage:
shear: float,
border_handling="edge",
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)
y = int(y)
angle = int(angle)
@@ -139,12 +115,7 @@ class MTB_TransformImage:
img = cast(
Image.Image,
TF.affine(
img,
angle=angle,
scale=zoom,
translate=[x, y],
shear=shear,
interpolation=resampling_filter,
img, angle=angle, scale=zoom, translate=[x, y], shear=shear
),
)
-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]
requires = ["setuptools", "wheel"]
build-backend = "setuptools.build_meta"
[project]
name = "comfy-mtb"
version = "0.2.1"
description = "Animation oriented nodes pack for ComfyUI."
license = "MIT"
readme = "README.md"
# repository = ""
# url = "https://github.com/melMass/comfy_mtb"
authors = [{ name = "Mel Massadian", email = "mel@melmassadian.com" }]
classifiers = [
"License :: OSI Approved :: MIT License",
"Operating System :: OS Independent",
"Programming Language :: Python",
"Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3.11",
"Intended Audience :: Developers",
]
requires-python = ">=3.10"
dependencies = [
"qrcode",
"cachetools",
"onnxruntime-gpu",
"requirements-parserx",
"rembg",
"imageio_ffmpeg",
"rich",
"rich_argparse",
"matplotlib",
"pillow",
]
optional-dependencies = { mel = [
"jupyterlab==4.1.6",
], dev = [
"black[jupyter]",
"codespell",
"marimo",
"mypy",
"pre-commit",
"pytest",
"pytest-cov",
"pytest-random-order",
"ruff",
], doc = [
"docutils==0.17.1",
"jupyter-book>=0.15",
"sphinx-autobuild",
] }
[project.urls]
Homepage = "https://github.com/melMass/comfy_mtb"
Documentation = "https://github.com/melMass/comfy_mtb/wiki"
Repository = "https://github.com/melMass/comfy_mtb"
Issues = "https://github.com/melMass/comfy_mtb/issues"
[tool.comfy]
PublisherId = "mel"
DisplayName = "comfy-mtb"
Icon = "https://avatars.githubusercontent.com/u/7041726?v=4"
[tool.bumpversion]
current_version = "0.2.1"
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
serialize = ["{major}.{minor}.{patch}"]
search = "{current_version}"
replace = "{new_version}"
regex = false
ignore_missing_version = false
ignore_missing_files = false
tag = true
sign_tags = true
tag_name = "v{new_version}"
tag_message = "⬆️ Bump version: {current_version} → {new_version}"
allow_dirty = true
commit = true
message = "⬆️ Bump version: {current_version} → {new_version}"
commit_args = ""
[[tool.bumpversion.files]]
filename = "__init__.py"
search = "__version__ = \"{current_version}\""
replace = "__version__ = \"{new_version}\""
[[tool.bumpversion.files]]
filename = "pyproject.toml"
search = "version = \"{current_version}\""
replace = "version = \"{new_version}\""
# [[tool.bumpversion.files]]
# filename = "your_package/__init__.py"
# search = "__version__ = '{current_version}'"
# replace = "__version__ = '{new_version}'"
# INFO: All those remaining keys are meant for local dev
[tool.pyright]
include = ["."]
exclude = [
"**/node_modules",
"**/__pycache__",
"src/experimental",
"src/typestubs",
]
ignore = ["src/oldstuff"]
defineConstant = { DEBUG = true }
extraPaths = ["python", "../.."]
stubPath = "src/stubs"
reportMissingImports = true
reportMissingTypeStubs = false
typeCheckingMode = "basic"
pythonVersion = "3.10"
pythonPlatform = "Windows"
[tool.pytest.ini_options]
log_level = "DEBUG"
log_cli = true
markers = [
"wip: tests that aren't fully finished yet",
"heavy: marks tests as heavy (deselect with '-m \"not heavy\"')",
]
filterwarnings = ["ignore::UserWarning", 'ignore::DeprecationWarning']
[tool.isort]
profile = "black"
line_length = 88
auto_identify_namespace_packages = false
# NOTE:
# pyright doesn't like implicit namespace + single line (related to https://github.com/microsoft/pyright/issues/2882?) but it's horible so I'll live with it
force_single_line = false
known_first_party = ["mtb"]
extend_skip = ["archives"]
combine_straight_imports = true
[tool.coverage.run]
parallel = true
source = ["docs", "tests", "comfy-mtb"]
[tool.coverage.report]
fail_under = 90
show_missing = true
[tool.coverage.html]
show_contexts = true
[tool.ruff]
line-length = 79
select = ["A", "B", "C", "D", "E", "F", "FBT", "I", "N", "S", "SIM", "UP", "W"]
# NOTE:
# D102 - undocumented-public-method (noisy)
# D103 - undocumented-public-function (noisy)
# D100 - undocumented-public-module (noisy)
# N802 - invalid-function-name (forced by comfy's arch)
ignore = ["D103", "D102", "D100", "N802"]
# exclude auto generated file
extend-exclude = ["./docs/conf.py"]
[tool.ruff.per-file-ignores]
# imported but unused
"__init__.py" = ["F401"]
# use of assert detected
"tests/*" = ["S101"]
[tool.ruff.pydocstyle]
convention = "numpy"
[tool.mypy]
pretty = true
ignore_missing_imports = true
# exclude auto generated file
exclude = ["docs/conf.py"]
[tool.codespell]
# exclude auto generated file
skip = "./docs/conf.py,poetry.lock"
check-filenames = true
[build-system]
requires = ["setuptools", "wheel"]
build-backend = "setuptools.build_meta"
[project]
name = "comfy-mtb"
version = "0.1.6"
description = "Animation oriented nodes pack for ComfyUI."
license = "MIT"
readme = "README.md"
# repository = ""
# url = "https://github.com/melMass/comfy_mtb"
authors = [{ name = "Mel Massadian", email = "mel@melmassadian.com" }]
classifiers = [
"License :: OSI Approved :: MIT License",
"Operating System :: OS Independent",
"Programming Language :: Python",
"Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3.11",
"Intended Audience :: Developers",
]
requires-python = ">=3.10"
dependencies = [
"qrcode",
"onnxruntime-gpu",
"requirements-parserx",
"rembg",
"imageio_ffmpeg",
"rich",
"rich_argparse",
"matplotlib",
"pillow",
]
optional-dependencies = { mel = [
"jupyterlab==4.1.6",
], dev = [
"black[jupyter]",
"codespell",
"mypy",
"pre-commit",
"pytest",
"pytest-cov",
"pytest-random-order",
"ruff",
], doc = [
"docutils==0.17.1",
"jupyter-book>=0.15",
"sphinx-autobuild",
] }
[project.urls]
Homepage = "https://github.com/melMass/comfy_mtb"
Documentation = "https://github.com/melMass/comfy_mtb/wiki"
Repository = "https://github.com/melMass/comfy_mtb"
Issues = "https://github.com/melMass/comfy_mtb/issues"
[tool.comfy]
PublisherId = "mel"
DisplayName = "comfy-mtb"
Icon = "https://avatars.githubusercontent.com/u/7041726?v=4"
[tool.bumpversion]
current_version = "0.1.6"
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
serialize = ["{major}.{minor}.{patch}"]
search = "{current_version}"
replace = "{new_version}"
regex = false
ignore_missing_version = false
ignore_missing_files = false
tag = true
sign_tags = true
tag_name = "v{new_version}"
tag_message = "⬆️ Bump version: {current_version} → {new_version}"
allow_dirty = true
commit = true
message = "⬆️ Bump version: {current_version} → {new_version}"
commit_args = ""
[[tool.bumpversion.files]]
filename = "__init__.py"
search = "__version__ = \"{current_version}\""
replace = "__version__ = \"{new_version}\""
[[tool.bumpversion.files]]
filename = "pyproject.toml"
search = "version = \"{current_version}\""
replace = "version = \"{new_version}\""
# [[tool.bumpversion.files]]
# filename = "your_package/__init__.py"
# search = "__version__ = '{current_version}'"
# replace = "__version__ = '{new_version}'"
# INFO: All those remaining keys are meant for local dev
[tool.pyright]
include = ["."]
exclude = [
"**/node_modules",
"**/__pycache__",
"src/experimental",
"src/typestubs",
]
ignore = ["src/oldstuff"]
defineConstant = { DEBUG = true }
extraPaths = ["python", "../.."]
stubPath = "src/stubs"
reportMissingImports = true
reportMissingTypeStubs = false
typeCheckingMode = "basic"
pythonVersion = "3.10"
pythonPlatform = "Windows"
[tool.pytest.ini_options]
log_level = "DEBUG"
log_cli = true
markers = [
"wip: tests that aren't fully finished yet",
"heavy: marks tests as heavy (deselect with '-m \"not heavy\"')",
]
filterwarnings = ["ignore::UserWarning", 'ignore::DeprecationWarning']
[tool.isort]
profile = "black"
line_length = 88
auto_identify_namespace_packages = false
# NOTE:
# pyright doesn't like implicit namespace + single line (related to https://github.com/microsoft/pyright/issues/2882?) but it's horible so I'll live with it
force_single_line = false
known_first_party = ["mtb"]
extend_skip = ["archives"]
combine_straight_imports = true
[tool.coverage.run]
parallel = true
source = ["docs", "tests", "comfy-mtb"]
[tool.coverage.report]
fail_under = 90
show_missing = true
[tool.coverage.html]
show_contexts = true
[tool.ruff]
line-length = 79
select = ["A", "B", "C", "D", "E", "F", "FBT", "I", "N", "S", "SIM", "UP", "W"]
# NOTE:
# D102 - undocumented-public-method (noisy)
# D103 - undocumented-public-function (noisy)
# D100 - undocumented-public-module (noisy)
# N802 - invalid-function-name (forced by comfy's arch)
ignore = ["D103", "D102", "D100", "N802"]
# exclude auto generated file
extend-exclude = ["./docs/conf.py"]
[tool.ruff.per-file-ignores]
# imported but unused
"__init__.py" = ["F401"]
# use of assert detected
"tests/*" = ["S101"]
[tool.ruff.pydocstyle]
convention = "numpy"
[tool.mypy]
pretty = true
ignore_missing_imports = true
# exclude auto generated file
exclude = ["docs/conf.py"]
[tool.codespell]
# exclude auto generated file
skip = "./docs/conf.py,poetry.lock"
check-filenames = true
+3 -72
View File
@@ -2,7 +2,6 @@ import contextlib
import functools
import importlib
import math
import operator
import os
import shlex
import shutil
@@ -12,7 +11,6 @@ import sys
import uuid
from collections.abc import Callable, Sequence
from enum import Enum
from functools import reduce
from pathlib import Path
from typing import TypeVar
@@ -214,37 +212,6 @@ def get_server_info():
# 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
@@ -498,12 +465,8 @@ here = Path(__file__).parent.absolute()
# - Construct the absolute path to the ComfyUI directory
comfy_dir = Path(folder_paths.base_path)
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)
input_dir = Path(folder_paths.input_directory)
styles_dir = comfy_dir / "styles"
session_id = str(uuid.uuid4())
# - Construct the path to the font file
@@ -513,10 +476,9 @@ font_path = here / "data" / "font.ttf"
extern_root = here / "extern"
add_path(extern_root)
if extern_root.exists():
for pth in extern_root.iterdir():
if pth.is_dir():
add_path(pth)
for pth in extern_root.iterdir():
if pth.is_dir():
add_path(pth)
# - Add the ComfyUI directory and custom nodes path to the sys.path list
add_path(comfy_dir)
@@ -630,37 +592,6 @@ def tensor2np(tensor: torch.Tensor) -> list[npt.NDArray[np.uint8]]:
return handle_batch(tensor, single_tensor2np)
def nextAvailable(path: Path | str) -> Path:
"""
Find the next available path by adding a numbered suffix. (mimics comfy's version).
Args:
path (Path): The original path to check
Returns
-------
Path: A path that doesn't exist yet
"""
path = Path(path)
if not path.is_absolute():
path = output_dir / path
if not path.exists():
return path
stem = path.stem
suffix = path.suffix
parent = path.parent
counter = 1
while True:
new_path = parent / f"{stem}_{counter:04d}{suffix}"
if not new_path.exists():
return new_path
counter += 1
def pad(img, left, right, top, bottom):
pad_width = np.array(((0, 0), (top, bottom), (left, right)))
print(
+21 -92
View File
@@ -1,9 +1,10 @@
/**
* @module Shared utilities
* File: comfy_shared.js
* Project: comfy_mtb
* Author: Mel Massadian
*
* Copyright (c) 2023-2024 Mel Massadian
*
*/
// 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 {LLink} link
@@ -371,37 +362,23 @@ export function getWidgetType(config) {
* @param {NodeType} nodeType The nodetype to attach the documentation to
* @param {str} prefix A prefix added to each 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
*/
export const setupDynamicConnections = (
nodeType,
prefix,
inputType,
opts = undefined,
) => {
export const setupDynamicConnections = (nodeType, prefix, inputType, opts) => {
infoLogger(
'Setting up dynamic connections for',
Object.getOwnPropertyDescriptors(nodeType).title.value,
)
/** @type {{separator:string, start_index:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} */
const options = Object.assign(
{
separator: '_',
start_index: 1,
},
opts || {},
)
/** @type {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} */
const options = opts || {}
const onNodeCreated = nodeType.prototype.onNodeCreated
const inputList = typeof inputType === 'object'
nodeType.prototype.onNodeCreated = function () {
const r = onNodeCreated ? onNodeCreated.apply(this, []) : undefined
this.addInput(
`${prefix}${options.separator}${options.start_index}`,
inputList ? '*' : inputType,
)
this.addInput(`${prefix}_1`, inputList ? '*' : inputType)
return r
}
@@ -436,7 +413,7 @@ export const setupDynamicConnections = (
this,
slotIndex,
isConnected,
`${prefix}${options.separator}`,
`${prefix}_`,
inputType,
options,
)
@@ -452,7 +429,7 @@ export const setupDynamicConnections = (
* @param {bool} connected - Was this event connecting or disconnecting
* @param {string} [connectionPrefix] - The common prefix of the dynamic inputs
* @param {string|[string]} [connectionType] - The type of the dynamic connection
* @param {{start_index?:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options
* @param {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options
*/
export const dynamic_connection = (
node,
@@ -462,18 +439,13 @@ export const dynamic_connection = (
connectionType = '*',
opts = undefined,
) => {
/* {{start_index:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options*/
const options = Object.assign(
{
start_index: 1,
},
opts || {},
)
/* @type {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options*/
const options = opts || {}
// function to test if input is a dynamic one
const isDynamicInput = (inputName) => inputName.startsWith(connectionPrefix)
if (node.inputs.length > 0 && !isDynamicInput(node.inputs[index].name)) {
if (
node.inputs.length > 0 &&
!node.inputs[index].name.startsWith(connectionPrefix)
) {
return
}
@@ -492,7 +464,7 @@ export const dynamic_connection = (
const to_remove = []
for (let n = 1; n < node.inputs.length; n++) {
const element = node.inputs[n]
if (!element.link && isDynamicInput(element.name)) {
if (!element.link) {
if (node.widgets) {
const w = node.widgets.find((w) => w.name === element.name)
if (w) {
@@ -518,25 +490,14 @@ export const dynamic_connection = (
infoLogger('Cleaning inputs: making it sequential again')
// make inputs sequential again
let prefixed_idx = options.start_index
for (let i = 0; i < node.inputs.length; i++) {
let name = ''
// 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
}
let name = `${connectionPrefix}${i + 1}`
if (nameArray.length > 0) {
name = i < nameArray.length ? nameArray[i] : name
}
// preserve label if it exists
node.inputs[i].label = node.inputs[i].label || name
node.inputs[i].label = name
node.inputs[i].name = name
}
}
@@ -576,16 +537,11 @@ export const dynamic_connection = (
if (node.inputs.length === 0) return
// add an extra input
if (node.inputs[node.inputs.length - 1].link !== null) {
// count only the prefixed inputs
const nextIndex = node.inputs.reduce(
(acc, cur) => (isDynamicInput(cur.name) ? ++acc : acc),
0,
)
const nextIndex = node.inputs.length
const name =
nextIndex < nameArray.length
? nameArray[nextIndex]
: `${connectionPrefix}${nextIndex + options.start_index}`
: `${connectionPrefix}${nextIndex + 1}`
infoLogger(`Adding input ${nextIndex + 1} (${name})`)
node.addInput(name, conType)
@@ -1129,33 +1085,7 @@ export const addDeprecation = (nodeType, reason) => {
// #endregion
// #region Actions API
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
// #region API / graph utilities
export const getAPIInputs = () => {
const inputs = {}
let counter = 1
@@ -1208,4 +1138,3 @@ export const getNodes = (skip_unused) => {
}
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 { MtbWidgets } from './mtb_widgets.js'
import * as mtb_ui from './mtb_ui.js'
// TODO: respect inputs order...
@@ -37,10 +36,12 @@ app.registerExtension({
async beforeRegisterNodeDef(nodeType, nodeData, app) {
if (nodeData.name === 'Debug (mtb)') {
const onNodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function (...args) {
nodeType.prototype.onNodeCreated = function () {
this.options = {}
const r = onNodeCreated ? onNodeCreated.apply(this, args) : undefined
this.addInput('anything_1', '*')
const r = onNodeCreated
? onNodeCreated.apply(this, arguments)
: undefined
this.addInput(`anything_1`, '*')
return r
}
@@ -80,16 +81,14 @@ app.registerExtension({
}
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (...args) {
onExecuted?.apply(this, args)
const [data, ..._rest] = args
nodeType.prototype.onExecuted = function (data) {
onExecuted?.apply(this, arguments)
const prefix = 'anything_'
if (this.widgets) {
for (let i = 0; i < this.widgets.length; i++) {
if (this.widgets[i].name !== 'output_to_console') {
this.widgets[i].onRemove?.()
this.widgets[i].onRemoved?.()
}
}
@@ -99,32 +98,19 @@ app.registerExtension({
// console.log(message)
if (data.text) {
for (const txt of data.text) {
const textDom = mtb_ui.makeElement('p', { fontFamily: 'monospace' })
textDom.innerHTML = txt
this.addDOMWidget(
`${prefix}_${widgetI}`,
'CUSTOM_TEXT',
textDom,
{},
const w = this.addCustomWidget(
MtbWidgets.DEBUG_STRING(`${prefix}_${widgetI}`, escapeHtml(txt)),
)
w.parent = this
widgetI++
}
}
if (data.b64_images) {
for (const img of data.b64_images) {
const imgDom = mtb_ui.makeElement('img', { width: '100%' })
imgDom.src = img
this.addDOMWidget(
`${prefix}_${widgetI}`,
'CUSTOM_IMG_B64',
mtb_ui.wrapElement(imgDom, {
overflow: 'hidden',
}),
{},
const w = this.addCustomWidget(
MtbWidgets.DEBUG_IMG(`${prefix}_${widgetI}`, img),
)
w.parent = this
widgetI++
}
}
@@ -133,13 +119,12 @@ app.registerExtension({
this.onRemoved = function () {
// 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) {
this.widgets[y].canvas.remove()
}
shared.cleanupNode(this)
this.widgets[y].onRemoved?.()
this.widgets[y].onRemove?.()
}
}
}
+148 -258
View File
@@ -1,7 +1,7 @@
import { app } from '../../scripts/app.js'
import { api } from '../../scripts/api.js'
import * as shared from './comfy_shared.js'
// import * as shared from './comfy_shared.js'
import {
// defineCSSClass,
@@ -12,14 +12,12 @@ import {
renderSidebar,
} from './mtb_ui.js'
const offset = 0
let offset = 0
let currentWidth = 200
let currentMode = 'input'
let subfolder = ''
let currentSort = 'None'
const IMAGE_NODES = ['LoadImage', 'VHS_LoadImagePath']
const VIDEO_NODES = ['VHS_LoadVideo']
const IMAGE_NODES = ['LoadImage']
const updateImage = (node, image) => {
if (IMAGE_NODES.includes(node.type)) {
@@ -28,13 +26,6 @@ const updateImage = (node, image) => {
w.value = image
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) {
return imgs
}
const elem = currentMode === 'video' ? 'video' : 'img'
for (const [key, url] of Object.entries(urls)) {
const a = makeElement(elem)
const a = makeElement('img')
a.src = url
a.width = currentWidth
if (currentMode === 'input') {
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
if (selected && Object.keys(selected).length === 0) {
app.extensionManager.toast.add({
severity: 'warn',
summary: 'No node selected!',
summary: 'No LoadImage node selected!',
detail:
'For now the only action when clicking images in the sidebar is to set the image on all selected LoadImage nodes.',
life: 5000,
@@ -73,22 +54,12 @@ const getImgsFromUrls = (urls, target) => {
}
for (const [_id, node] of Object.entries(app.canvas.selected_nodes)) {
updateImage(node, key)
updateImage(node, `${key}.png`)
}
}
} else if (currentMode === 'output') {
a.onclick = (_e) => {
} else {
a.onclick = (_e) =>
// 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({
severity: 'warn',
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.',
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)
}
@@ -128,212 +76,154 @@ const getImgsFromUrls = (urls, target) => {
return imgs
}
const getModes = async () => {
const inputs = await shared.runAction('getUserImageFolders')
return inputs
}
const getUrls = async (subfolder) => {
const count = (await api.getSetting('mtb.io-sidebar.count')) || 1000
const getUrls = async () => {
const count = await api.getSetting('mtb.io-sidebar.count')
console.log('Sidebar count', count)
if (currentMode === 'video') {
const output = await shared.runAction(
'getUserVideos',
256,
count,
offset,
currentSort,
)
return output || {}
}
const output = await shared.runAction(
'getUserImages',
currentMode,
count,
offset,
currentSort,
false,
subfolder,
)
return output || {}
const inputs = await api.fetchApi('/mtb/actions', {
method: 'POST',
body: JSON.stringify({
name: 'getUserImages',
// mode, count, offset
args: [currentMode, count, offset, currentSort],
}),
})
const output = await inputs.json()
return output?.result || {}
}
//NOTE: do not load if using the old ui
if (window?.__COMFYUI_FRONTEND_VERSION__) {
// NOTE: removed this for now since I'm not actually exposing anything a client
// cannot already access from "/view"...
// let exposed = false
let handle
const version = window?.__COMFYUI_FRONTEND_VERSION__
console.log(`%c ${version}`, 'background: orange; color: white;')
const sidebar_extension = {
name: 'mtb.io-sidebar',
// init: async () => {
// try {
// const res = await api.fetchApi('/mtb/server-info')
// const msg = await res.json()
// exposed = msg.exposed
// } catch (e) {
// console.error('Error:', e)
// }
// },
init: () => {
let handle
const version = window?.__COMFYUI_FRONTEND_VERSION__
console.log(`%c ${version}`, 'background: orange; color: white;')
ensureMTBStyles()
ensureMTBStyles()
app.ui.settings.addSetting({
id: 'mtb.io-sidebar.count',
category: ['mtb', 'Input & Output Sidebar', 'count'],
app.ui.settings.addSetting({
id: 'mtb.io-sidebar.count',
category: ['mtb', 'Input & Output Sidebar', 'count'],
name: 'Number of images to fetch',
type: 'number',
defaultValue: 1000,
name: 'Number of images to fetch',
type: 'number',
defaultValue: 1000,
tooltip:
"This setting affects the input/output sidebar to determine how many images to fetch per pagination (pagination is not yet supported so for now it's the static total)",
attrs: {
style: {
// fontFamily: 'monospace',
},
},
})
app.ui.settings.addSetting({
id: 'mtb.io-sidebar.img-size',
category: ['mtb', 'Input & Output Sidebar', 'img-size'],
name: 'Resolution of the images',
type: 'number',
defaultValue: 512,
tooltip: "It's recommended to keep it at 512px",
attrs: {
style: {
// fontFamily: 'monospace',
},
},
})
app.ui.settings.addSetting({
id: 'mtb.io-sidebar.sort',
category: ['mtb', 'Input & Output Sidebar', 'sort'],
name: 'Default sort mode',
type: 'combo',
onChange: (v) => {
// alert(`Sort is now ${v}`)
currentSort = v
},
defaultValue: 'Modified',
// tooltip: "It's recommended to keep it at 512px",
options: [
'None',
'Modified',
'Modified-Reverse',
'Name',
'Name-Reverse',
],
})
app.extensionManager.registerSidebarTab({
id: 'mtb-inputs-outputs',
icon: 'pi pi-images',
title: 'Input & Outputs',
tooltip: 'MTB: Browse inputs and outputs directories.',
type: 'custom',
// this is run everytime the tab's diplay is toggled on.
render: async (el) => {
if (handle) {
handle.unregister()
handle = undefined
}
if (el.parentNode) {
el.parentNode.style.overflowY = 'clip'
}
const 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
}
},
})
tooltip:
"This setting affects the input/output sidebar to determine how many images to fetch per pagination (pagination is not yet supported so for now it's the static total)",
attrs: {
style: {
// fontFamily: 'monospace',
},
},
}
})
app.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
// I need to properly publish the source and fix a few things before
import { app } from '../../scripts/app.js'
// import { api } from '../../scripts/api.js'
// import { app } from '../../scripts/app.js'
// // import { api } from '../../scripts/api.js'
//
// import * as shared from './comfy_shared.js'
// import { createOutliner } from './dist/mtb_inspector.js'
//
// if (window?.__COMFYUI_FRONTEND_VERSION__) {
// const version = window?.__COMFYUI_FRONTEND_VERSION__
// console.log(`%c ${version}`, 'background: orange; color: white;')
//
// const panel = app.extensionManager.registerSidebarTab({
// id: 'mtb-nodes',
// icon: 'pi pi-bolt',
// title: 'MTB',
// tooltip: 'MTB: API outliner',
// type: 'custom',
// // this is run everytime the tab's diplay is toggled on.
// render: (el) => {
// const outliner = createOutliner(el)
// const inputs = shared.getAPIInputs()
// console.log('INPUTS', inputs)
// outliner.$$set({ inputs })
// },
// })
// }
import * as shared from './comfy_shared.js'
import { createOutliner } from './dist/mtb_inspector.js'
if (window?.__COMFYUI_FRONTEND_VERSION__) {
const version = window?.__COMFYUI_FRONTEND_VERSION__
console.log(`%c ${version}`, 'background: orange; color: white;')
const panel = app.extensionManager.registerSidebarTab({
id: 'mtb-nodes',
icon: 'pi pi-bolt',
title: 'MTB',
tooltip: 'MTB: API outliner',
type: 'custom',
// this is run everytime the tab's diplay is toggled on.
render: (el) => {
const outliner = createOutliner(el)
const inputs = shared.getAPIInputs()
console.log('INPUTS', inputs)
outliner.$$set({ inputs })
},
})
}
-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.
*
+1036 -1209
View File
File diff suppressed because it is too large Load Diff
+1 -1
Submodule wiki updated: fa7fec28a3...4db733ae92