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