Compare commits

...
19 Commits
Author SHA1 Message Date
Mel Massadian 2441f19db3 ⬆️ Bump version: 0.2.0 → 0.2.1 2024-12-30 18:46:35 +01:00
Mel Massadian b01e027ec0 Merge branch 'main' into pr-227 2024-12-30 18:40:13 +01:00
Mel Massadian 9a4ecb2b90 fix: 🐛 handle missing submodules
the nodes should never fail to load completely.
I still need to remove the few remaining side effects like this one.
2024-12-30 18:37:43 +01:00
Mel Massadian 0eeb707f34 feat: ✨ add SaveImage passthrough
Exactly like the native one but not as an OUTPUT_NODE,
primarly meant to "inline" image saving in upcoming mtb loops.
2024-12-30 18:29:28 +01:00
Mel Massadian 9a943714aa chore: 🧹 dev
dev files
2024-12-29 13:51:03 +01:00
Robin Huang 885688e7c7 Checkout submodules before publishing. 2024-12-27 14:39:48 -08:00
Mel Massadian bae26a07fb feat: ✨ add filtering to TransformImage
fixes #209
2024-12-27 19:29:34 +01:00
Mel Massadian c92d99a8a3 feat: ✨ add support for video in I/O sidebar
slow if you have big videos, maybe it shouldn't use force_size
2024-12-22 05:43:19 +01:00
Mel Massadian 3f6d082940 feat: ✨ add an extra static input to Stack Images 2024-12-22 02:40:01 +01:00
Mel Massadian 58ae89f8e0 chore: 🧹 apply formatting 2024-12-22 02:40:01 +01:00
pak c9a26427a8 improve dynamic inputs: custom separator and start_index, preserve labels, ... 2024-12-22 02:40:01 +01:00
Mel Massadian 6608c0b6d1 fix: 🐛 add warnings about what each IO mode can do
VHS now has a Load Image Path node that could be used to solve all cases.
For this I'll need to get the full path of each images from the endpoint
2024-12-22 02:11:21 +01:00
Mel Massadian a757e1c98b fix: 🐛 soft deprecate compression h264 2024-12-22 01:30:03 +01:00
Mel Massadian 52bd76e19c feat: ✨ add support for subdirs (i/o sidebar)
fixes #221
2024-12-22 01:28:08 +01:00
Mel Massadian d6e004cce2 fix: 🐛 limit packages allowed to be installed from API
fixes #224

thanks @boy-hack for the report!
2024-12-22 00:22:24 +01:00
Mel Massadian ed17fa2ef4 fix: 🐛 ensure default settings (io sidebar)
fixes #225
2024-12-21 02:14:40 +01:00
Mel Massadian 827c64c43d feat: ✨ add Batch Sequence Nodes
- A regular one that just sequence batches
- A "plus" with transition support (POC + for now)
2024-12-16 01:44:01 +01:00
filtered e5482aee5e fix: 🐛 spawn colour picker at pointer location (#223) 2024-12-15 22:22:22 +01:00
Mel Massadian 62469a4dd9 fix: 🐛 i/o sidebar for custom paths
In utils I uses a constant for these which doesn't
update with the global... calling the getters should
solve that.

This issue is probably in other places where I use these
utils.

fixes #219
2024-12-11 00:09:43 +01:00
16 changed files with 812 additions and 165 deletions
+2
View File
@@ -12,6 +12,8 @@ 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:
+3 -11
View File
@@ -7,7 +7,7 @@
#
###
__version__ = "0.2.0"
__version__ = "0.2.1"
import os
@@ -34,6 +34,7 @@ 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
@@ -240,16 +241,7 @@ if hasattr(PromptServer, "instance"):
img_cache = TTLCache(maxsize=100, ttl=5) # 1 min TTL
prompt_cache = TTLCache(maxsize=100, ttl=5) # 1 min TTL
restore_deps = ["basicsr"]
onnx_deps = ["onnxruntime"]
swap_deps = ["insightface"] + onnx_deps
node_dependency_mapping = {
"QrCode": ["qrcode"],
"DeepBump": onnx_deps,
"FaceSwap": swap_deps,
"LoadFaceSwapModel": swap_deps,
"LoadFaceAnalysisModel": restore_deps,
}
node_dependency_mapping = get_node_dependencies()
PromptServer.instance.app.router.add_static(
"/mtb-assets/", path=(here / "html").as_posix()
+70 -5
View File
@@ -1,11 +1,14 @@
import csv
import secrets
import sys
import urllib.parse
from pathlib import Path
from typing import Any, Literal
import folder_paths
from aiohttp import web
from .install import get_node_dependencies
from .log import mklog
from .utils import (
SortMode,
@@ -13,8 +16,6 @@ from .utils import (
build_glob_patterns,
glob_multiple,
import_install,
input_dir,
output_dir,
reqs_map,
run_command,
styles_dir,
@@ -26,7 +27,7 @@ endlog = mklog("mtb endpoint")
import_install("requirements")
def ACTIONS_installDependency(dependency_names=None):
def ACTIONS_installDependency(dependency_names: list[str] | None = None):
if dependency_names is None:
# return web.Response(text="No dependency name provided", status=400)
return {"error": "No dependency name provided"}
@@ -34,6 +35,14 @@ def ACTIONS_installDependency(dependency_names=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
@@ -58,19 +67,75 @@ def ACTIONS_installDependency(dependency_names=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=200,
count=1000,
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"}
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
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"]
entries = {}
@@ -92,7 +157,7 @@ def ACTIONS_getUserImages(
imgs = {
img.name: (
f"/mtb/view?filename={img.name}&width=512&type={mode}&subfolder="
f"/mtb/view?filename={img.name}&width=512&type={mode}&subfolder={subfolder or ''}"
f"{img.parent.relative_to(entry_dir) if include_subfolders else ''}"
f"&preview=&rand={secrets.randbelow(424242)}"
)
+6
View File
@@ -160,6 +160,12 @@ 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
}
+64 -28
View File
@@ -43,10 +43,28 @@ 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
@@ -124,12 +142,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)
@@ -163,7 +181,9 @@ 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:
@@ -238,7 +258,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, "r") as f:
with open(init_file) as f:
tree = ast.parse(f.read())
for node in ast.walk(tree):
if isinstance(node, ast.Assign):
@@ -256,13 +276,16 @@ 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))
@@ -302,7 +325,9 @@ 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}').",
@@ -310,7 +335,9 @@ 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",
@@ -326,13 +353,9 @@ 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(
@@ -361,7 +384,9 @@ 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):
@@ -397,8 +422,12 @@ 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..."
@@ -409,7 +438,14 @@ 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(
+226 -10
View File
@@ -1,4 +1,5 @@
from io import BytesIO
from typing import Literal
import cv2
import numpy as np
@@ -410,7 +411,14 @@ class MTB_BatchFloat:
RETURN_TYPES = ("FLOATS",)
CATEGORY = "mtb/batch"
def set_floats(self, mode, count, min, max, easing):
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",
):
if mode == "Steps" and count == 1:
raise ValueError(
"Steps mode requires at least a count of 2 values"
@@ -429,6 +437,210 @@ 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"""
@@ -711,7 +923,9 @@ 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()):
for color, (label, values) in zip(
colors, kwargs.items(), strict=False
):
ax.plot(x_values[: len(values)], values, label=label, color=color)
ax.legend(
title="Legend",
@@ -1026,17 +1240,19 @@ class MTB_BatchShake:
__nodes__ = [
MTB_BatchFloat,
MTB_Batch2dTransform,
MTB_BatchShape,
MTB_BatchMake,
MTB_BatchFloat,
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,
]
+1 -3
View File
@@ -702,7 +702,6 @@ class MTB_Blur:
)
blurred_images.append(blurred)
image_np = np.array(blurred_images)
else:
for i in range(image.size(0)):
blurred = gaussian(
@@ -710,8 +709,7 @@ class MTB_Blur:
)
blurred_images.append(blurred)
image_np = np.array(blurred_images)
return (np2tensor(image_np).squeeze(0),)
return (np2tensor(blurred_images),)
class MTB_Sharpen:
+168 -24
View File
@@ -1,4 +1,11 @@
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
@@ -8,13 +15,21 @@ class MTB_StackImages:
@classmethod
def INPUT_TYPES(cls):
return {"required": {"vertical": ("BOOLEAN", {"default": False})}}
return {
"required": {"vertical": ("BOOLEAN", {"default": False})},
"optional": {
"match_method": (
["error", "smallest", "largest"],
{"default": "error"},
)
},
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "stack"
CATEGORY = "mtb/image utils"
def stack(self, vertical, **kwargs):
def stack(self, vertical, match_method="error", **kwargs):
if not kwargs:
raise ValueError("At least one tensor must be provided.")
@@ -32,23 +47,50 @@ class MTB_StackImages:
self.duplicate_frames(tensor, max_batch_size)
for tensor in normalized_tensors
]
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."
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)
)
dim = 1
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:
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
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
stacked_tensor = torch.cat(normalized_tensors, dim=dim)
@@ -64,7 +106,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(
@@ -87,6 +129,30 @@ 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.
@@ -113,11 +179,6 @@ 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:]
@@ -127,4 +188,87 @@ class MTB_PickFromBatch:
return (selected_tensors,)
__nodes__ = [MTB_StackImages, MTB_PickFromBatch]
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]
+4
View File
@@ -42,6 +42,9 @@ class ImageH264Compression:
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.
@@ -151,6 +154,7 @@ class ImageH264Compression:
output_images = torch.stack(output_images).to(image.device)
return (output_images,)
# fmt: off
__nodes__ = [
ImageH264Compression
+30 -1
View File
@@ -45,6 +45,19 @@ class MTB_TransformImage:
),
"constant_color": ("COLOR", {"default": "#000000"}),
},
"optional": {
"filter_type": (
[
"nearest",
"box",
"bilinear",
"hamming",
"bicubic",
"lanczos",
],
{"default": "bilinear"},
),
},
}
FUNCTION = "transform"
@@ -61,7 +74,18 @@ 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)
@@ -115,7 +139,12 @@ class MTB_TransformImage:
img = cast(
Image.Image,
TF.affine(
img, angle=angle, scale=zoom, translate=[x, y], shear=shear
img,
angle=angle,
scale=zoom,
translate=[x, y],
shear=shear,
interpolation=resampling_filter,
),
)
+3 -2
View File
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
[project]
name = "comfy-mtb"
version = "0.2.0"
version = "0.2.1"
description = "Animation oriented nodes pack for ComfyUI."
license = "MIT"
readme = "README.md"
@@ -38,6 +38,7 @@ optional-dependencies = { mel = [
], dev = [
"black[jupyter]",
"codespell",
"marimo",
"mypy",
"pre-commit",
"pytest",
@@ -62,7 +63,7 @@ DisplayName = "comfy-mtb"
Icon = "https://avatars.githubusercontent.com/u/7041726?v=4"
[tool.bumpversion]
current_version = "0.2.0"
current_version = "0.2.1"
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
serialize = ["{major}.{minor}.{patch}"]
search = "{current_version}"
+8 -3
View File
@@ -498,8 +498,12 @@ 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
@@ -509,9 +513,10 @@ font_path = here / "data" / "font.ttf"
extern_root = here / "extern"
add_path(extern_root)
for pth in extern_root.iterdir():
if pth.is_dir():
add_path(pth)
if extern_root.exists():
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)
+82 -21
View File
@@ -1,10 +1,9 @@
/**
* @module Shared utilities
* File: comfy_shared.js
* Project: comfy_mtb
* Author: Mel Massadian
*
* Copyright (c) 2023-2024 Mel Massadian
*
*/
// Reference the shared typedefs file
@@ -372,23 +371,37 @@ 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 {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} opts
* @param {{separator?:string, start_index?:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} [opts] Extra options
* @returns
*/
export const setupDynamicConnections = (nodeType, prefix, inputType, opts) => {
export const setupDynamicConnections = (
nodeType,
prefix,
inputType,
opts = undefined,
) => {
infoLogger(
'Setting up dynamic connections for',
Object.getOwnPropertyDescriptors(nodeType).title.value,
)
/** @type {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} */
const options = opts || {}
/** @type {{separator:string, start_index:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} */
const options = Object.assign(
{
separator: '_',
start_index: 1,
},
opts || {},
)
const onNodeCreated = nodeType.prototype.onNodeCreated
const inputList = typeof inputType === 'object'
nodeType.prototype.onNodeCreated = function () {
const r = onNodeCreated ? onNodeCreated.apply(this, []) : undefined
this.addInput(`${prefix}_1`, inputList ? '*' : inputType)
this.addInput(
`${prefix}${options.separator}${options.start_index}`,
inputList ? '*' : inputType,
)
return r
}
@@ -423,7 +436,7 @@ export const setupDynamicConnections = (nodeType, prefix, inputType, opts) => {
this,
slotIndex,
isConnected,
`${prefix}_`,
`${prefix}${options.separator}`,
inputType,
options,
)
@@ -439,7 +452,7 @@ export const setupDynamicConnections = (nodeType, prefix, inputType, opts) => {
* @param {bool} connected - Was this event connecting or disconnecting
* @param {string} [connectionPrefix] - The common prefix of the dynamic inputs
* @param {string|[string]} [connectionType] - The type of the dynamic connection
* @param {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options
* @param {{start_index?:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options
*/
export const dynamic_connection = (
node,
@@ -449,13 +462,18 @@ export const dynamic_connection = (
connectionType = '*',
opts = undefined,
) => {
/* @type {{link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options*/
const options = opts || {}
/* {{start_index:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}} [opts] - extra options*/
const options = Object.assign(
{
start_index: 1,
},
opts || {},
)
if (
node.inputs.length > 0 &&
!node.inputs[index].name.startsWith(connectionPrefix)
) {
// 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)) {
return
}
@@ -474,7 +492,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) {
if (!element.link && isDynamicInput(element.name)) {
if (node.widgets) {
const w = node.widgets.find((w) => w.name === element.name)
if (w) {
@@ -500,14 +518,25 @@ 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 = `${connectionPrefix}${i + 1}`
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
}
if (nameArray.length > 0) {
name = i < nameArray.length ? nameArray[i] : name
}
node.inputs[i].label = name
// preserve label if it exists
node.inputs[i].label = node.inputs[i].label || name
node.inputs[i].name = name
}
}
@@ -547,11 +576,16 @@ export const dynamic_connection = (
if (node.inputs.length === 0) return
// add an extra input
if (node.inputs[node.inputs.length - 1].link !== null) {
const nextIndex = node.inputs.length
// count only the prefixed inputs
const nextIndex = node.inputs.reduce(
(acc, cur) => (isDynamicInput(cur.name) ? ++acc : acc),
0,
)
const name =
nextIndex < nameArray.length
? nameArray[nextIndex]
: `${connectionPrefix}${nextIndex + 1}`
: `${connectionPrefix}${nextIndex + options.start_index}`
infoLogger(`Adding input ${nextIndex + 1} (${name})`)
node.addInput(name, conType)
@@ -1095,7 +1129,33 @@ export const addDeprecation = (nodeType, reason) => {
// #endregion
// #region API / graph utilities
// #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
export const getAPIInputs = () => {
const inputs = {}
let counter = 1
@@ -1148,3 +1208,4 @@ export const getNodes = (skip_unused) => {
}
return nodes
}
// #endregion
+108 -25
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,
@@ -15,9 +15,11 @@ import {
const offset = 0
let currentWidth = 200
let currentMode = 'input'
let subfolder = ''
let currentSort = 'None'
const IMAGE_NODES = ['LoadImage']
const IMAGE_NODES = ['LoadImage', 'VHS_LoadImagePath']
const VIDEO_NODES = ['VHS_LoadVideo']
const updateImage = (node, image) => {
if (IMAGE_NODES.includes(node.type)) {
@@ -26,6 +28,13 @@ 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)
}
}
@@ -34,18 +43,28 @@ 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('img')
const a = makeElement(elem)
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 LoadImage node selected!',
summary: 'No 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,
@@ -57,9 +76,19 @@ const getImgsFromUrls = (urls, target) => {
updateImage(node, key)
}
}
} else {
a.onclick = (_e) =>
} else if (currentMode === 'output') {
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',
@@ -67,6 +96,29 @@ 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)
}
@@ -76,19 +128,33 @@ const getImgsFromUrls = (urls, target) => {
return imgs
}
const getUrls = async () => {
const count = await api.getSetting('mtb.io-sidebar.count')
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
console.log('Sidebar count', count)
const inputs = await api.fetchApi('/mtb/actions', {
method: 'POST',
body: JSON.stringify({
name: 'getUserImages',
// mode, count, offset
args: [currentMode, count, offset, currentSort],
}),
})
const output = await inputs.json()
return output?.result || {}
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 || {}
}
//NOTE: do not load if using the old ui
@@ -187,21 +253,39 @@ if (window?.__COMFYUI_FRONTEND_VERSION__) {
el.parentNode.style.overflowY = 'clip'
}
const urls = await getUrls(currentMode)
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'], currentMode)
const selector = makeSelect(
['input', 'output', 'video', ...output_modes, ...input_modes],
currentMode,
)
selector.addEventListener('change', async (e) => {
const newMode = e.target.value
const changed = newMode !== currentMode
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()
const urls = await getUrls(subfolder)
if (urls) {
imgs = getImgsFromUrls(urls, imgGrid)
}
@@ -220,7 +304,7 @@ if (window?.__COMFYUI_FRONTEND_VERSION__) {
currentSort = newSort
if (changed) {
imgGrid.innerHTML = ''
const urls = await getUrls()
const urls = await getUrls(subfolder)
if (urls) {
imgs = getImgsFromUrls(urls, imgGrid)
}
@@ -229,7 +313,6 @@ if (window?.__COMFYUI_FRONTEND_VERSION__) {
const sizeSlider = makeSlider(64, 1024, currentWidth, 1)
imgTools.appendChild(orderSelect)
imgTools.appendChild(sizeSlider)
imgs = getImgsFromUrls(urls, imgGrid)
+36 -31
View File
@@ -536,21 +536,34 @@ export const MtbWidgets = {
picker.type = 'color'
picker.value = this.value
picker.style.position = 'absolute'
picker.style.left = '999999px' //(window.innerWidth / 2) + "px";
picker.style.top = '999999px' //(window.innerHeight / 2) + "px";
Object.assign(picker.style, {
position: 'fixed',
left: `${e.clientX}px`,
top: `${e.clientY}px`,
height: '0px',
width: '0px',
padding: '0px',
opacity: 0,
})
picker.addEventListener('blur', () => {
this.callback?.(this.value)
node.graph._version++
picker.remove()
})
picker.addEventListener('input', () => {
if (!picker.value) return
this.value = picker.value
app.canvas.setDirty(true)
})
document.body.appendChild(picker)
picker.addEventListener('change', () => {
this.value = picker.value
this.callback?.(this.value)
node.graph._version++
node.setDirtyCanvas(true, true)
picker.remove()
requestAnimationFrame(() => {
picker.showPicker()
picker.focus()
})
picker.click()
}
}
}
@@ -659,8 +672,7 @@ const mtb_widgets = {
init: async () => {
infoLogger('Registering mtb.widgets')
try {
const res = await api.fetchApi('/mtb/server-info')
const msg = await res.json()
const msg = await shared.getServerInfo()
if (!window.MTB) {
window.MTB = {}
}
@@ -703,17 +715,11 @@ const mtb_widgets = {
infoLogger('Enabled DEBUG mode')
}
await api
.fetchApi('/mtb/server-info', {
method: 'POST',
body: JSON.stringify({
debug: value,
}),
})
.then((_response) => {})
.catch((error) => {
console.error('Error:', error)
})
try {
shared.setServerInfo({ debug: value })
} catch (err) {
console.error('Error:', err)
}
},
})
},
@@ -1096,13 +1102,10 @@ const mtb_widgets = {
const getStyle = async (node) => {
try {
const getStyles = await api.fetchApi('/mtb/actions', {
method: 'POST',
body: JSON.stringify({
name: 'getStyles',
args: node.widgets?.[0].value ? node.widgets[0].value : '',
}),
})
const getStyles = await runAction(
'getStyles',
node.widgets?.[0].value ? node.widgets[0].value : '',
)
const output = await getStyles.json()
return output?.result
@@ -1202,6 +1205,8 @@ const mtb_widgets = {
shared.setupDynamicConnections(nodeType, 'floats', 'FLOATS')
break
}
case 'Batch Sequence (mtb)':
case 'Batch Sequence Plus (mtb)':
case 'Batch Merge (mtb)': {
shared.setupDynamicConnections(nodeType, 'batches', 'IMAGE')
+1 -1
Submodule wiki updated: a402de4af9...fa7fec28a3