Compare commits

..
Author SHA1 Message Date
Mel Massadian 2000af6ca7 test!: 🚨 bringing back io edits
very wip but should improve the panel a lot
2025-05-08 11:05:34 +02:00
34 changed files with 1959 additions and 1915 deletions
-7
View File
@@ -1,7 +0,0 @@
**/GFPGAN/inputs/**
**/GFPGAN/tests/**
**/frame_interpolation/photos/*
moment.gif
node.zip
.DS_Store
+3 -1
View File
@@ -1,6 +1,9 @@
name: 📦 Publish to Comfy registry
on:
workflow_dispatch:
push:
tags:
- '*'
permissions:
issues: write
@@ -18,5 +21,4 @@ jobs:
- name: 📦 Publish Custom Node
uses: Comfy-Org/publish-node-action@v1
with:
skip_checkout: 'true'
personal_access_token: ${{ secrets.COMFY_REGISTRY_TOKEN }}
-5
View File
@@ -1,16 +1,11 @@
__pycache__
*.py[cod]
*.onnx
wheels/
node_modules/
compose.yaml
comfy_mtb.wsb
Dockerfile
.DS_Store
node.zip
# I store the gh-pages worktrees (src & build) there
.worktrees
comfy.lock
-6
View File
@@ -1,10 +1,4 @@
# MTB Nodes
> [!NOTE]
> master/main is outdated for now to keep backward compatibility, the next version is being worked on in
> [`dev/0.6.0`](https://github.com/melMass/comfy_mtb/tree/dev/0.6.0)
[![embedded test](https://github.com/melMass/comfy_mtb/actions/workflows/test_embedded.yml/badge.svg)](https://github.com/melMass/comfy_mtb/actions/workflows/test_embedded.yml)
![home](https://repository-images.githubusercontent.com/649047066/a3eef9a7-20dd-4ef9-b839-884502d4e873)
+167 -73
View File
@@ -3,13 +3,15 @@
# File: __init__.py
# Project: comfy_mtb
# Author: Mel Massadian
# Copyright (c) 2023-2025 Mel Massadian
# Copyright (c) 2023 Mel Massadian
#
###
__version__ = "0.5.4"
__version__ = "0.3.0"
import os
from collections import OrderedDict
from typing import Any
from aiohttp.web_request import Request
@@ -34,8 +36,6 @@ from aiohttp import web
IN_COMFY = False
PromptServer = None
try:
from server import PromptServer
@@ -77,7 +77,7 @@ def extract_nodes_from_source(filename: Path):
)
break
except SyntaxError:
log.error(f"Failed to parse ast from: {filename}")
log.error("Failed to parse")
return nodes
@@ -242,33 +242,14 @@ if failed:
# - ENDPOINT
# TODO: move that away and simplify existing endpoints
def register_routes():
if not PromptServer:
log.error("No prompt server, are you inside comfy?")
if PromptServer.instance.app.frozen:
log.warning(
"The router is frozen and cannot be further edited."
"If you are hot reloading mtb this is expected."
)
return
if IN_COMFY and hasattr(PromptServer, "instance"):
img_cache = None
prompt_cache = None
import asyncio
import os
from io import BytesIO
from PIL import Image
with contextlib.suppress(ImportError):
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()
@@ -381,24 +362,134 @@ def register_routes():
# Return JSON for other requests
return web.json_response({"message": "Welcome to MTB!"})
import asyncio
import os
import time
from asyncio import Semaphore
from concurrent.futures import ThreadPoolExecutor
from contextlib import asynccontextmanager
from io import BytesIO
from aiohttp import web
from PIL import Image
image_thread_pool = ThreadPoolExecutor(
max_workers=4, thread_name_prefix="img_worker"
)
@asynccontextmanager
async def get_image_with_timeout(
file_path, preview_params=None, channel=None, timeout=10
):
try:
result = await asyncio.wait_for(
asyncio.get_event_loop().run_in_executor(
image_thread_pool,
get_cached_image,
file_path,
preview_params,
channel,
),
timeout=timeout,
)
yield result
except asyncio.TimeoutError:
print(f"Image processing timed out for {file_path}")
raise
except Exception as e:
print(f"Error processing image {file_path}: {str(e)}")
raise
async def get_image_response(
file, filename: str, preview_info=None, channel=None
):
try:
async with get_image_with_timeout(
file, preview_info, channel
) as img:
return web.Response(
body=img,
content_type="image/webp" if preview_info else "image/png",
headers={"Content-Disposition": f'filename="{filename}"'},
)
except asyncio.TimeoutError:
return web.Response(status=504, text="Image processing timed out")
except Exception as e:
return web.Response(status=500, text=str(e))
class LRUCache:
def __init__(self, capacity: int):
self.cache = OrderedDict()
self.capacity = capacity
def get(self, key) -> Any:
if key not in self.cache:
return None
self.cache.move_to_end(key)
return self.cache[key]
def put(self, key, value: Any) -> None:
if key in self.cache:
self.cache.move_to_end(key)
self.cache[key] = value
if len(self.cache) > self.capacity:
self.cache.popitem(last=False)
img_cache = LRUCache(capacity=100)
def get_cached_image(file_path: str, preview_params=None, channel=None):
cache_key = (file_path, preview_params, channel)
if img_cache and (cache_key in img_cache):
return img_cache[cache_key]
with Image.open(file_path) as img:
info = img.info
if preview_params:
img = process_preview(img, preview_params)
if channel:
img = process_channel(img, channel)
if prompt_cache:
prompt_cache[cache_key] = info
try:
if img_cache:
img_cache[cache_key] = img.getvalue()
return img_cache[cache_key]
cached_value = img_cache.get(cache_key)
if cached_value is not None:
return cached_value
return img.getvalue()
with Image.open(file_path) as img:
info = img.info
if preview_params:
img = process_preview(img, preview_params)
if channel:
img = process_channel(img, channel)
result = img.getvalue()
try:
if prompt_cache:
prompt_cache[cache_key] = info
if img_cache:
img_cache.put(cache_key, result)
except Exception as e:
print(
f"Warning: Failed to cache image {file_path}: {str(e)}"
)
return result
except Exception as e:
print(f"Error processing image {file_path}: {str(e)}")
raise
class RateLimiter:
def __init__(self, requests_per_second):
self.requests_per_second = requests_per_second
self.semaphore = Semaphore(requests_per_second)
self.timestamps = []
async def acquire(self):
await self.semaphore.acquire()
now = time.time()
self.timestamps.append(now)
# Remove old timestamps
self.timestamps = [t for t in self.timestamps if now - t < 1.0]
if len(self.timestamps) >= self.requests_per_second:
await asyncio.sleep(1.0)
def release(self):
self.semaphore.release()
rate_limiter = RateLimiter(requests_per_second=10)
def process_preview(img: Image.Image, preview_params):
image_format, quality, width = preview_params
@@ -451,41 +542,44 @@ def register_routes():
# to load workflows in the sidebar
@PromptServer.instance.routes.get("/mtb/view")
async def view_image(request: Request):
import folder_paths
try:
import folder_paths
filename = request.rel_url.query.get("filename")
if not filename:
return web.Response(status=404)
await rate_limiter.acquire()
filename, output_dir = folder_paths.annotated_filepath(filename)
if filename[0] == "/" or ".." in filename:
return web.Response(status=400)
filename = request.rel_url.query.get("filename")
if not filename:
return web.Response(status=404)
if output_dir is None:
rtype = request.rel_url.query.get("type", "output")
output_dir = folder_paths.get_directory_by_type(rtype)
filename, output_dir = folder_paths.annotated_filepath(filename)
if filename[0] == "/" or ".." in filename:
return web.Response(status=400)
if output_dir is None:
return web.Response(status=400)
if output_dir is None:
rtype = request.rel_url.query.get("type", "output")
output_dir = folder_paths.get_directory_by_type(rtype)
if "subfolder" in request.rel_url.query:
full_output_dir = os.path.join(
output_dir, request.rel_url.query["subfolder"]
)
if (
os.path.commonpath(
(os.path.abspath(full_output_dir), output_dir)
if output_dir is None:
return web.Response(status=400)
if "subfolder" in request.rel_url.query:
full_output_dir = os.path.join(
output_dir, request.rel_url.query["subfolder"]
)
!= output_dir
):
return web.Response(status=403)
output_dir = full_output_dir
if (
os.path.commonpath(
(os.path.abspath(full_output_dir), output_dir)
)
!= output_dir
):
return web.Response(status=403)
output_dir = full_output_dir
filename = os.path.basename(filename)
file = os.path.join(output_dir, filename)
filename = os.path.basename(filename)
file = os.path.join(output_dir, filename)
if not os.path.isfile(file):
return web.Response(status=404)
if not os.path.isfile(file):
return web.Response(status=404)
ret_workflow = request.rel_url.query.get("workflow")
@@ -523,9 +617,13 @@ def register_routes():
width = request.rel_url.query.get("width")
preview_info = (image_format, quality, width)
channel = request.rel_url.query.get("channel")
channel = request.rel_url.query.get("channel")
return await get_image_response(file, filename, preview_info, channel)
return await get_image_response(
file, filename, preview_info, channel
)
finally:
rate_limiter.release()
@PromptServer.instance.routes.get("/mtb/server-info")
async def get_debug(request: Request):
@@ -585,10 +683,6 @@ def register_routes():
return await endpoint.do_action(request)
if IN_COMFY and hasattr(PromptServer, "instance"):
register_routes()
# - WAS Dictionary
MANIFEST = {
"name": "MTB Nodes", # The title that will be displayed on Node Class menu,. and Node Class view
+6 -13
View File
@@ -1,26 +1,19 @@
{
"$schema": "https://biomejs.dev/schemas/2.0.5/schema.json",
"assist": { "actions": { "source": { "organizeImports": "on" } } },
"$schema": "https://biomejs.dev/schemas/1.6.1/schema.json",
"organizeImports": {
"enabled": true
},
"linter": {
"enabled": true,
"rules": {
"recommended": true,
"suspicious": {
"noConsole": { "level": "warn", "options": { "allow": ["log"] } }
"noConsoleLog": "warn"
},
"style": {
"noParameterAssign": "off",
"noShoutyConstants": "warn",
"useNamingConvention": "off",
"useAsConstAssertion": "error",
"useDefaultParameterLast": "error",
"useEnumInitializers": "error",
"useSelfClosingElements": "error",
"useSingleVarDeclarator": "error",
"noUnusedTemplateLiteral": "error",
"useNumberNamespace": "error",
"noInferrableTypes": "error",
"noUselessElse": "error"
"useNamingConvention": "off"
}
}
},
+12 -10
View File
@@ -15,6 +15,7 @@ from .utils import (
backup_file,
build_glob_patterns,
glob_multiple,
import_install,
reqs_map,
run_command,
styles_dir,
@@ -23,6 +24,7 @@ from .utils import (
endlog = mklog("mtb endpoint")
# - ACTIONS
import_install("requirements")
def ACTIONS_installDependency(dependency_names: list[str] | None = None):
@@ -72,7 +74,12 @@ def ACTIONS_getUserImageFolders():
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}
return {
"input_root": input_dir.as_posix(),
"input": input_subdirs,
"output": output_subdirs,
"output_root": output_dir.as_posix(),
}
def ACTIONS_getUserVideos(
@@ -110,15 +117,11 @@ def ACTIONS_getUserVideos(
def ACTIONS_getUserImages(
mode: Literal["input", "output"],
target_width: int | str | None = None,
count=1000,
offset=0,
sort: str | None = None,
include_subfolders: bool = False,
subfolder: str | None = None,
# IIRC I copied this from Comfy base
# just keeping it until I properly checked implications
salt_urls=False,
subfolder=None,
):
# enabled = "MTB_EXPOSE" in os.environ
# if not enabled:
@@ -126,12 +129,11 @@ def ACTIONS_getUserImages(
imgs = {}
count = count or 1000
target_width = int(target_width) if target_width else None
input_dir = Path(folder_paths.get_input_directory())
output_dir = Path(folder_paths.get_output_directory())
entry_dir: Path = input_dir if mode == "input" else output_dir
entry_dir = input_dir if mode == "input" else output_dir
if subfolder:
entry_dir = entry_dir / subfolder
@@ -160,9 +162,9 @@ def ACTIONS_getUserImages(
imgs = {
img.name: (
f"/mtb/view?filename={img.name}{f'&width={target_width}' if target_width and target_width > 0 else ''}&type={mode}&subfolder={subfolder or ''}"
f"/mtb/view?filename={img.name}&width=512&type={mode}&subfolder={subfolder or ''}"
f"{img.parent.relative_to(entry_dir) if include_subfolders else ''}"
f"&preview={f'&rand={secrets.randbelow(424242)}' if salt_urls else ''}"
f"&preview=&rand={secrets.randbelow(424242)}"
)
for i, img in enumerate(entries)
if offset <= i < offset + count
+1
View File
@@ -43,6 +43,7 @@ pip_map = {
"tb-nightly": "tensorboard",
"protobuf": "google.protobuf",
"qrcode[pil]": "qrcode",
"requirements-parser": "requirements",
# Add more mappings as needed
}
+15 -16
View File
@@ -1,16 +1,20 @@
from typing import TYPE_CHECKING, Any, TypedDict
from typing import Any, TypedDict
import torch
import torchaudio
from comfy.model_management import get_torch_device
from huggingface_hub import snapshot_download
from transformers import (
WhisperForConditionalGeneration,
WhisperProcessor,
)
if TYPE_CHECKING:
from transformers import (
WhisperForConditionalGeneration,
WhisperProcessor,
)
# from transformers import (
# AutoFeatureExtractor,
# WhisperForConditionalGeneration,
# WhisperModel,
# WhisperProcessor,
# )
from ..log import log
from ..utils import get_model_path
@@ -97,8 +101,8 @@ class MtbAudio:
class WhisperPipeline(TypedDict):
"""Whisper model pipeline."""
processor: "WhisperProcessor"
model: "WhisperForConditionalGeneration"
processor: WhisperProcessor
model: WhisperForConditionalGeneration
class MTB_LoadWhisper:
@@ -144,11 +148,6 @@ class MTB_LoadWhisper:
def load(self, model_size="tiny", download_missing=False):
"""Load Whisper model and processor."""
from transformers import (
WhisperForConditionalGeneration,
WhisperProcessor,
)
whisper_dir = get_model_path("whisper")
tag = f"whisper-{model_size}"
model_dir = whisper_dir / tag
@@ -277,14 +276,14 @@ class MTB_AudioToText(MtbAudio):
f"Processing chunk {chunk_offset:.1f}s - {chunk_end / sample_rate:.1f}s"
)
max_length = getattr(model.config, "max_length", None) or 448
max_length = model.config.max_length or 448
attention_mask = torch.ones((1, max_length))
input_features = processor(
chunk_waveform,
sampling_rate=sample_rate,
return_tensors="pt",
).input_features.to(device=device, dtype=model.dtype)
).input_features.to(device)
with torch.no_grad():
predicted_ids = model.generate(
+4 -4
View File
@@ -335,9 +335,9 @@ class MTB_BatchShape:
"image_width": ("INT", {"default": 512}),
"image_height": ("INT", {"default": 512}),
"shape_size": ("INT", {"default": 100}),
"color": ("COLOR", {"default": "#ffffff","widgetType": "MTB_COLOR"}),
"bg_color": ("COLOR", {"default": "#000000","widgetType": "MTB_COLOR"}),
"shade_color": ("COLOR", {"default": "#000000","widgetType": "MTB_COLOR"}),
"color": ("COLOR", {"default": "#ffffff"}),
"bg_color": ("COLOR", {"default": "#000000"}),
"shade_color": ("COLOR", {"default": "#000000"}),
"thickness": ("INT", {"default": 5}),
"shadex": ("FLOAT", {"default": 0.0}),
"shadey": ("FLOAT", {"default": 0.0}),
@@ -842,7 +842,7 @@ class MTB_Batch2dTransform:
["edge", "constant", "reflect", "symmetric"],
{"default": "edge"},
),
"constant_color": ("COLOR", {"default": "#000000","widgetType": "MTB_COLOR"}),
"constant_color": ("COLOR", {"default": "#000000"}),
},
"optional": {
"x": ("FLOATS",),
-190
View File
@@ -1,190 +0,0 @@
import time
import uuid
from collections import OrderedDict
from typing import Any, TypedDict
from comfy.comfy_types.node_typing import IO as CIO
from server import PromptServer
from ..log import log
class Clock(TypedDict):
name: str
start: float
end: float | None
active_timers: OrderedDict[str, Clock] = OrderedDict()
# TODO: lower this
MAX_CLOCKS = 50
class MTB_StartClock:
"""
Starts a profiling clock with a given name.
Outputs a unique ID that must be passed to EndClock.
"""
def __init__(self):
pass
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"name": ("STRING", {"default": "Clock A"}),
"cache": (
"BOOLEAN",
{
"default": False,
"tooltip": "Cache the clock ID, this means the node will follow Comfy's default invalidation system. If False it will always invalidate / mark the node as 'dirty'",
},
),
},
"optional": {
"passthrough": (CIO.ANY,),
},
}
RETURN_TYPES = (
CIO.ANY,
"STRING",
)
RETURN_NAMES = (
"passthrough",
"clock_id",
)
FUNCTION = "start_timer"
CATEGORY = "mtb/utils"
def start_timer(
self, *, name: str, passthrough: Any | None = None, **kwargs
):
global active_timers
if len(active_timers) >= MAX_CLOCKS:
# get oldest clock
removed_key = None
for key, clock_data in active_timers.items():
if clock_data["end"] is not None:
removed_key = key
break
if removed_key:
removed_clock = active_timers.pop(removed_key)
log.info(
f"[Profiling] Evicted finished clock '{removed_clock['name']}' (ID: {removed_key}) due to limit ({MAX_CLOCKS})."
)
else:
removed_key, removed_clock = active_timers.popitem(last=False)
log.warning(
f"[Profiling] Evicted running clock '{removed_clock['name']}' (ID: {removed_key}) due to limit ({MAX_CLOCKS})."
)
clock_id = str(uuid.uuid4())
start_time = time.perf_counter()
active_timers[clock_id] = {
"start": start_time,
"name": name,
"end": None,
}
active_timers.move_to_end(clock_id)
log.debug(f"[Profiling] Clock '{name}' (ID: {clock_id}) started.")
return (
passthrough,
clock_id,
)
@classmethod
def IS_CHANGED(
cls, *, name: str, cache: bool = False, passthrough: Any | None = None
):
if not cache:
return float("Nan")
return {"name": name, "cache": cache, "passthrough": passthrough}
class MTB_EndClock:
"""
Stops a profiling clock identified by its ID and returns the elapsed time in milliseconds.
Errors if the clock ID is not found or already stopped.
"""
def __init__(self):
pass
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"clock_id": (
"STRING",
{"forceInput": True},
),
},
"optional": {
"passthrough": (CIO.ANY,),
},
"hidden": {
"unique_id": "UNIQUE_ID",
},
}
RETURN_TYPES = (
CIO.ANY,
"STRING",
"FLOAT",
"INT",
)
RETURN_NAMES = (
"passthrough",
"name",
"seconds",
"milliseconds",
)
FUNCTION = "end_timer"
CATEGORY = "mtb/utils"
def end_timer(self, clock_id: str, passthrough, unique_id=None):
global active_timers
if clock_id not in active_timers:
raise ValueError(
f"Error: Clock with ID '{clock_id}' not found. "
"Ensure StartClock was executed for this ID and proper passthrough chaining."
)
clock = active_timers[clock_id]
if clock.get("end") is not None:
return (passthrough, clock["name"], clock["end"])
start_time = clock["start"]
end_time = time.perf_counter()
duration_seconds = end_time - start_time
duration_ms = int(duration_seconds * 1000)
clock["end"] = duration_ms
active_timers.move_to_end(clock_id)
log.debug(
f"[Profiling] Clock '{clock['name']}' (ID: {clock_id}) stopped. Elapsed: {duration_ms}ms"
)
if unique_id:
PromptServer.instance.send_progress_text(
f"Clock '{clock['name']}' took {duration_seconds:.4f} seconds",
unique_id,
)
return (passthrough, clock["name"], duration_seconds, duration_ms)
__nodes__ = [MTB_StartClock, MTB_EndClock]
+181 -194
View File
@@ -1,22 +1,13 @@
from typing import NamedTuple
import numpy as np
import torch
import torchvision.transforms.functional as TF
from PIL import Image, ImageDraw, ImageFilter
from ..log import log
class BoundingBox(NamedTuple):
"""The bounding box tuple."""
x: int
y: int
width: int
height: int
from ..utils import np2tensor, pil2tensor, tensor2np, tensor2pil
class MTB_Bbox:
"""A literal bounding box."""
"""The bounding box (BBOX) custom type used by other nodes"""
@classmethod
def INPUT_TYPES(cls):
@@ -46,14 +37,12 @@ class MTB_Bbox:
FUNCTION = "do_crop"
CATEGORY = "mtb/crop"
def do_crop(
self, x: int, y: int, width: int, height: int
) -> tuple[BoundingBox]: # bbox
return (BoundingBox(x, y, width, height),)
def do_crop(self, x: int, y: int, width: int, height: int): # bbox
return ((x, y, width, height),)
class MTB_SplitBbox:
"""Split the components of a bbox."""
"""Split the components of a bbox"""
@classmethod
def INPUT_TYPES(cls):
@@ -66,8 +55,8 @@ class MTB_SplitBbox:
RETURN_TYPES = ("INT", "INT", "INT", "INT")
RETURN_NAMES = ("x", "y", "width", "height")
def split_bbox(self, bbox: BoundingBox) -> BoundingBox:
return bbox
def split_bbox(self, bbox):
return (bbox[0], bbox[1], bbox[2], bbox[3])
class MTB_UpscaleBboxBy:
@@ -85,23 +74,26 @@ class MTB_UpscaleBboxBy:
FUNCTION = "upscale"
def upscale(self, bbox: BoundingBox, scale: float) -> tuple[BoundingBox]:
def upscale(
self, bbox: tuple[int, int, int, int], scale: float
) -> tuple[tuple[int, int, int, int]]:
x, y, width, height = bbox
center_x = x + width / 2
center_y = y + height / 2
center_x = x + width // 2
center_y = y + height // 2
new_width = int(width * scale)
new_height = int(height * scale)
new_x = int(center_x - new_width / 2)
new_y = int(center_y - new_height / 2)
new_x = center_x - new_width // 2
new_y = center_y - new_height // 2
return (BoundingBox(new_x, new_y, new_width, new_height),)
scaled = (new_x, new_y, new_width, new_height)
return (scaled,)
class MTB_BboxFromMask:
"""From a mask extract the bounding box."""
"""From a mask extract the bounding box"""
@classmethod
def INPUT_TYPES(cls):
@@ -111,7 +103,7 @@ class MTB_BboxFromMask:
"invert": ("BOOLEAN", {"default": False}),
},
"optional": {
"image": ("IMAGE", {"tooltip": "Optional image"}),
"image": ("IMAGE",),
},
}
@@ -127,44 +119,52 @@ class MTB_BboxFromMask:
CATEGORY = "mtb/crop"
def extract_bounding_box(
self,
mask: torch.Tensor,
*,
invert: bool = False,
image: torch.Tensor | None = None,
) -> tuple[BoundingBox, torch.Tensor | None]:
mask = 1 - mask if invert else mask
non_zero_indices = torch.nonzero(mask)
self, mask: torch.Tensor, invert: bool, image=None
):
# if image != None:
# if mask.size(0) != image.size(0):
# if mask.size(0) != 1:
# log.error(
# f"Batch count mismatch for mask and image, it can either be 1 mask for X images, or X masks for X images (mask: {mask.shape} | image: {image.shape})"
# )
if non_zero_indices.numel() == 0:
log.warning(
"BboxFromMask: Mask is empty. Returning a (0,0,0,0) bbox."
)
return (BoundingBox(0, 0, 0, 0), image)
# raise Exception(
# f"Batch count mismatch for mask and image, it can either be 1 mask for X images, or X masks for X images (mask: {mask.shape} | image: {image.shape})"
# )
min_coords = torch.min(non_zero_indices, dim=0).values
max_coords = torch.max(non_zero_indices, dim=0).values
# we invert it
_mask = tensor2pil(1.0 - mask)[0] if invert else tensor2pil(mask)[0]
alpha_channel = np.array(_mask)
min_y, min_x = min_coords[1].item(), min_coords[2].item()
max_y, max_x = max_coords[1].item(), max_coords[2].item()
non_zero_indices = np.nonzero(alpha_channel)
width = max_x - min_x + 1
height = max_y - min_y + 1
min_x, max_x = np.min(non_zero_indices[1]), np.max(non_zero_indices[1])
min_y, max_y = np.min(non_zero_indices[0]), np.max(non_zero_indices[0])
bounding_box = BoundingBox(
int(min_x), int(min_y), int(width), int(height)
# Create a bounding box tuple
if image != None:
# Convert the image to a NumPy array
imgs = tensor2np(image)
out = []
for img in imgs:
# Crop the image from the bounding box
img = img[min_y:max_y, min_x:max_x, :]
log.debug(f"Cropped image to shape {img.shape}")
out.append(img)
image = np2tensor(out)
log.debug(f"Cropped images shape: {image.shape}")
bounding_box = (min_x, min_y, max_x - min_x, max_y - min_y)
return (
bounding_box,
image,
)
cropped_image = None
if image is not None:
cropped_image = image[:, min_y : max_y + 1, min_x : max_x + 1, :]
return (bounding_box, cropped_image)
class MTB_Crop:
"""Crop an image and an optional mask to a given bounding box.
"""Crops an image and an optional mask to a given bounding box
The bounding box can be given as a tuple of (x, y, width, height) or as a BBOX type
The BBOX input takes precedence over the tuple input
"""
@@ -204,38 +204,35 @@ class MTB_Crop:
def do_crop(
self,
image: torch.Tensor,
*,
mask: torch.Tensor | None = None,
x: int = 0,
y: int = 0,
width: int = 256,
height: int = 256,
bbox: BoundingBox | None = None,
mask=None,
x=0,
y=0,
width=256,
height=256,
bbox=None,
):
image = image.numpy()
if mask is not None:
mask = mask.numpy()
if bbox is not None:
x, y, width, height = bbox
if width <= 0 or height <= 0:
log.error(
"Crop dimensions must be positive. Check the BBOX or widget inputs."
)
return (
torch.zeros_like(image),
torch.zeros_like(mask) if mask is not None else None,
(x, y, width, height),
)
cropped_image = image[:, y : y + height, x : x + width, :]
cropped_mask = (
mask[:, y : y + height, x : x + width]
if mask is not None
else None
)
crop_data = BoundingBox(x, y, width, height)
cropped_mask = None
if mask is not None:
cropped_mask = (
mask[:, y : y + height, x : x + width]
if mask is not None
else None
)
crop_data = (x, y, width, height)
return (
cropped_image,
cropped_mask if cropped_mask is not None else None,
torch.from_numpy(cropped_image),
torch.from_numpy(cropped_mask)
if cropped_mask is not None
else None,
crop_data,
)
@@ -249,33 +246,35 @@ class MTB_Crop:
# return (x_left, y_top, x_right, y_bottom)
def bbox_check(bbox: BoundingBox, target_size: tuple[int, int] | None = None):
def bbox_check(bbox, target_size=None):
if not target_size:
return bbox
new_bbox = BoundingBox(
bbox.x,
bbox.y,
min(target_size[0] - bbox.x, bbox.width),
min(target_size[1] - bbox.y, bbox.height),
new_bbox = (
bbox[0],
bbox[1],
min(target_size[0] - bbox[0], bbox[2]),
min(target_size[1] - bbox[1], bbox[3]),
)
if new_bbox != bbox:
log.warning(f"BBox too big, constrained to {new_bbox}")
log.warn(f"BBox too big, constrained to {new_bbox}")
return new_bbox
def bbox_to_region(
bbox: BoundingBox, target_size: tuple[int, int] | None = None
):
def bbox_to_region(bbox, target_size=None):
bbox = bbox_check(bbox, target_size)
# to region
return (bbox.x, bbox.y, bbox.x + bbox.width, bbox.y + bbox.height)
return (bbox[0], bbox[1], bbox[0] + bbox[2], bbox[1] + bbox[3])
class MTB_Uncrop:
"""Uncrop an image to a given bounding box."""
"""Uncrops an image to a given bounding box
The bounding box can be given as a tuple of (x, y, width, height) or as a BBOX type
The BBOX input takes precedence over the tuple input
"""
@classmethod
def INPUT_TYPES(cls):
@@ -292,113 +291,91 @@ class MTB_Uncrop:
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "do_uncrop"
FUNCTION = "do_crop"
CATEGORY = "mtb/crop"
def do_uncrop(
self,
image: torch.Tensor,
crop_image: torch.Tensor,
bbox: BoundingBox,
border_blending: float = 0.25,
):
if len(image) > 1 and len(image) != len(crop_image):
raise ValueError(
"Uncrop: Batch size of background 'image' must be 1 or match the 'crop_image' batch size."
def do_crop(self, image, crop_image, bbox, border_blending):
def inset_border(image, border_width=20, border_color=(0)):
width, height = image.size
bordered_image = Image.new(
image.mode, (width, height), border_color
)
import comfy.utils
bordered_image.paste(image, (0, 0))
draw = ImageDraw.Draw(bordered_image)
draw.rectangle(
(0, 0, width - 1, height - 1),
outline=border_color,
width=border_width,
)
return bordered_image
pbar = comfy.utils.ProgressBar(4)
single = image.size(0) == 1
if image.size(0) != crop_image.size(0):
if not single:
raise ValueError(
"The Image batch count is greater than 1, but doesn't match the crop_image batch count. If using batches they should either match or only crop_image must be greater than 1"
)
device = image.device
images = tensor2pil(image)
crop_imgs = tensor2pil(crop_image)
out_images = []
for i, crop in enumerate(crop_imgs):
if single:
img = images[0]
else:
img = images[i]
log.debug(f"Working on device: {device}")
# uncrop the image based on the bounding box
bb_x, bb_y, bb_width, bb_height = bbox
crop_image = crop_image.to(device)
paste_region = bbox_to_region(
(bb_x, bb_y, bb_width, bb_height), img.size
)
# log.debug(f"Paste region: {paste_region}")
# new_region = adjust_paste_region(img.size, paste_region)
# log.debug(f"Adjusted paste region: {new_region}")
# # Check if the adjusted paste region is different from the original
if len(image) == 1 and len(crop_image) > 1:
image = image.repeat(len(crop_image), 1, 1, 1)
crop_img = crop.convert("RGB")
batch_size, bg_h, bg_w, _ = image.shape
_, fg_h, fg_w, _ = crop_image.shape
x, y, width, height = bbox
log.debug(f"Crop image size: {crop_img.size}")
log.debug(f"Image size: {img.size}")
if (width, height) != (fg_w, fg_h):
log.warning(
f"Uncrop: crop_image size {(fg_w, fg_h)} "
"differs from bbox {(width, height)}. Resizing to fit bbox."
if border_blending > 1.0:
border_blending = 1.0
elif border_blending < 0.0:
border_blending = 0.0
blend_ratio = (max(crop_img.size) / 2) * float(border_blending)
blend = img.convert("RGBA")
mask = Image.new("L", img.size, 0)
mask_block = Image.new("L", (bb_width, bb_height), 255)
mask_block = inset_border(mask_block, int(blend_ratio / 2), (0))
mask.paste(mask_block, paste_region)
log.debug(f"Blend size: {blend.size} | kind {blend.mode}")
log.debug(
f"Crop image size: {crop_img.size} | kind {crop_img.mode}"
)
log.debug(f"BBox: {paste_region}")
blend.paste(crop_img, paste_region)
mask = mask.filter(ImageFilter.BoxBlur(radius=blend_ratio / 4))
mask = mask.filter(
ImageFilter.GaussianBlur(radius=blend_ratio / 4)
)
resized_crop = crop_image.permute(0, 3, 1, 2)
resized_crop = torch.nn.functional.interpolate(
resized_crop,
size=(height, width),
mode="bicubic",
align_corners=False,
)
resized_crop = resized_crop.permute(0, 2, 3, 1)
blend.putalpha(mask)
img = Image.alpha_composite(img.convert("RGBA"), blend)
out_images.append(img.convert("RGB"))
pbar.update(1)
# paste coords
paste_x1 = max(x, 0)
paste_y1 = max(y, 0)
paste_x2 = min(x + width, bg_w)
paste_y2 = min(y + height, bg_h)
# region from crop (bound)
crop_x1 = max(0, -x)
crop_y1 = max(0, -y)
crop_x2 = crop_x1 + (paste_x2 - paste_x1)
crop_y2 = crop_y1 + (paste_y2 - paste_y1)
if paste_x1 >= paste_x2 or paste_y1 >= paste_y2:
log.warning(
"Uncrop: BBOX is entirely outside the image boundaries. Returning original image."
)
return (image,)
pbar.update(1)
source_slice = resized_crop[:, crop_y1:crop_y2, crop_x1:crop_x2, :]
final_image = image.clone()
final_image[:, paste_y1:paste_y2, paste_x1:paste_x2, :] = source_slice
pbar.update(1)
blend_radius = int(max(width, height) * border_blending * 0.5)
if blend_radius > 0:
_device = device
if torch.cuda.is_available():
_device = torch.device("cuda")
log.debug("Processing blending")
alpha_mask = torch.zeros((batch_size, bg_h, bg_w), device=_device)
alpha_mask[:, paste_y1:paste_y2, paste_x1:paste_x2] = 1.0
kernel_size = 2 * blend_radius + 1
log.debug("Gaussian blur...")
alpha_mask = TF.gaussian_blur(
alpha_mask.unsqueeze(1), kernel_size=[kernel_size, kernel_size]
).squeeze(1)
alpha_mask = alpha_mask.unsqueeze(-1)
log.debug("Applying blending")
final_image = final_image.to(_device) * alpha_mask + image.to(
_device
) * (1.0 - alpha_mask)
pbar.update(1)
return (final_image.to(device),)
return (pil2tensor(out_images),)
class MTB_BBoxForceDimensions:
"""
Resize a BBOX to new dimensions while keeping its center.
Optionally constrains the BBOX to stay within image boundaries.
"""
@classmethod
def INPUT_TYPES(cls):
return {
@@ -406,7 +383,6 @@ class MTB_BBoxForceDimensions:
"bbox": ("BBOX",),
"width": ("INT", {"default": 512, "min": 1, "max": 8192}),
"height": ("INT", {"default": 512, "min": 1, "max": 8192}),
"constrain_to_image": ("BOOLEAN", {"default": True}),
},
"optional": {
"image": ("IMAGE",),
@@ -419,12 +395,10 @@ class MTB_BBoxForceDimensions:
def force_dimensions(
self,
*,
bbox: tuple[int, int, int, int],
width: int,
height: int,
constrain_to_image: bool = True,
image: torch.Tensor | None = None,
image: torch.Tensor = None,
) -> tuple[tuple[int, int, int, int]]:
x, y, curr_width, curr_height = bbox
@@ -434,14 +408,27 @@ class MTB_BBoxForceDimensions:
new_x = center_x - width // 2
new_y = center_y - height // 2
if constrain_to_image and image is not None:
if image is not None:
img_height, img_width = image.shape[1:3]
new_x = max(0, min(new_x, img_width - width))
new_y = max(0, min(new_y, img_height - height))
width = min(width, img_width)
height = min(height, img_height)
x_overflow = max(0, new_x + width - img_width) + min(0, new_x)
y_overflow = max(0, new_y + height - img_height) + min(0, new_y)
if width > img_width or height > img_height:
x_exceed = width - img_width if width > img_width else 0
y_exceed = height - img_height if height > img_height else 0
raise ValueError(
f"Target bbox dimensions ({width}x{height}) exceed image bounds ({img_width}x{img_height}) "
f"by {x_exceed}px horizontally and {y_exceed}px vertically"
)
return ((new_x, new_y, width, height),)
if x_overflow > 0 or x_overflow < 0:
new_x -= x_overflow
if y_overflow > 0:
new_y -= y_overflow
elif y_overflow < 0:
new_y -= y_overflow # Add the negative overflow
return ((int(new_x), int(new_y), width, height),)
__nodes__ = [
+4 -7
View File
@@ -4,8 +4,12 @@ import sys
from pathlib import Path
import comfy.model_management as model_management
import cv2
import insightface
import numpy as np
import onnxruntime
import torch
from insightface.model_zoo.inswapper import INSwapper
from PIL import Image
from ..errors import ModelNotFound
@@ -39,8 +43,6 @@ class MTB_LoadFaceAnalysisModel:
DEPRECATED = True
def load_model(self, faceswap_model: str):
import insightface
if faceswap_model == "antelopev2":
download_antelopev2()
@@ -79,9 +81,6 @@ class MTB_LoadFaceSwapModel:
DEPRECATED = True
def load_model(self, faceswap_model: str):
import onnxruntime
from insightface.model_zoo.inswapper import INSwapper
model_path = get_model_path("insightface", faceswap_model)
if not model_path or not model_path.exists():
raise ModelNotFound(f"{faceswap_model} ({model_path})")
@@ -213,8 +212,6 @@ def swap_face(
face_swapper_model,
faces_index: set[int] | None = None,
) -> Image.Image:
import cv2
if faces_index is None:
faces_index = {0}
log.debug(f"Swapping faces: {faces_index}")
+6 -6
View File
@@ -193,11 +193,11 @@ by default it fallsback to a default font.
),
"color": (
"COLOR",
{"default": "black", "widgetType": "MTB_COLOR"},
{"default": "black"},
),
"background": (
"COLOR",
{"default": "white", "widgetType": "MTB_COLOR"},
{"default": "white"},
),
"h_align": (("left", "center", "right"), {"default": "left"}),
"v_align": (("top", "center", "bottom"), {"default": "top"}),
@@ -343,7 +343,9 @@ by default it fallsback to a default font.
def render_text(text_to_render, alpha=None):
if trim:
text_to_render = text_to_render.strip()
text_to_render = (
text_to_render.encode("ascii", "ignore").decode().strip()
)
if wrap:
wrap_width = (((width / 100) * h_coverage) / font_size) * 2
lines = textwrap.wrap(text_to_render, width=wrap_width)
@@ -416,9 +418,7 @@ by default it fallsback to a default font.
active_chunks.append((chunk["text"], alpha))
for chunk_text, alpha in active_chunks:
chunk_img = render_text(
chunk_text.encode("ascii", "ignore").decode(), alpha
)
chunk_img = render_text(chunk_text, alpha)
frame = Image.alpha_composite(frame, chunk_img)
frames.append(frame)
-51
View File
@@ -4,13 +4,11 @@ import re
import urllib.parse
import urllib.request
from math import pi
from typing import Any
import comfy.model_management as model_management
import comfy.utils
import numpy as np
import torch
from comfy.comfy_types.node_typing import IO as CIO
from PIL import Image
from ..log import log
@@ -869,53 +867,6 @@ class MTB_TensorOps:
return (result,)
class MTB_GetItem:
"""Generic index based getter for common types"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"container": (CIO.ANY,),
"index": ("INT", {"default": 0}),
}
}
RETURN_TYPES = (CIO.ANY,)
RETURN_NAMES = ("item",)
FUNCTION = "get_item"
CATEGORY = "mtb/utils"
def get_item(self, container: Any, index: int):
if "__getitem__" in dir(container):
log.debug(f"Container is {type(container)}")
res = container[index]
if type(res) is torch.Tensor:
res = res.unsqueeze(0)
return (res,)
class MTB_BooleanNot:
"""Inverts a boolean."""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"bool_in": ("BOOLEAN", {"default": False}),
},
}
RETURN_TYPES = ("BOOLEAN",)
RETURN_NAMES = ("inverted_bool",)
FUNCTION = "invert"
CATEGORY = "mtb/utils"
def invert(self, bool_in: bool):
return (not bool_in,)
__nodes__ = [
MTB_StringReplace,
MTB_FitNumber,
@@ -931,6 +882,4 @@ __nodes__ = [
MTB_FloatToFloats,
MTB_FloatsToInts,
MTB_TensorOps,
MTB_BooleanNot,
MTB_GetItem,
]
+52 -136
View File
@@ -3,12 +3,11 @@ import json
import math
import os
import comfy.utils
import comfy.model_management as model_management
import folder_paths
import numpy as np
import torch
import torch.nn.functional as F
from comfy import model_management
from PIL import Image, ImageOps
from PIL.PngImagePlugin import PngInfo
from skimage.filters import gaussian
@@ -75,10 +74,7 @@ class MTB_ExtractCoordinatesFromImage:
def INPUT_TYPES(cls):
return {
"required": {
"threshold": (
"FLOAT",
{"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01},
),
"threshold": ("FLOAT",),
"max_points": ("INT", {"default": 50, "min": 0}),
},
"optional": {"image": ("IMAGE",), "mask": ("MASK",)},
@@ -91,124 +87,72 @@ class MTB_ExtractCoordinatesFromImage:
image: torch.Tensor | None = None,
mask: torch.Tensor | None = None,
) -> tuple[list[list[tuple[int, int]]], torch.Tensor]:
if image is None and mask is None:
raise ValueError("Must provide either image or mask")
if image is not None:
batch_count, height, width, _channel_count = image.shape
input_device = image.device
if mask is not None:
if mask.ndim == 2:
mask = mask.unsqueeze(0)
if mask.ndim != 3:
raise ValueError(
f"Mask has unexpected ndim: {mask.ndim}. Expected 2 or 3."
)
b_mask, h_mask, w_mask = mask.shape
if not (h_mask == height and w_mask == width):
raise ValueError(
f"Image dimensions ({height}x{width}) and mask dimensions ({h_mask}x{w_mask}) are spatially incompatible."
)
if b_mask == 1 and batch_count > 1:
mask = mask.expand(batch_count, height, width)
elif b_mask != batch_count:
raise ValueError(
f"Image batch size ({batch_count}) and mask batch size ({b_mask}) are incompatible and mask cannot be broadcast."
)
batch_count, height, width, channel_count = image.shape
imgs = image
else:
if mask.ndim == 2:
mask = mask.unsqueeze(0)
if mask.ndim != 3:
raise ValueError(
f"Mask has unexpected ndim: {mask.ndim} when image is not provided. Expected 2 or 3."
)
if mask is None:
raise ValueError("Must provide either image or mask")
batch_count, height, width = mask.shape
input_device = mask.device
channel_count = 1
imgs = mask
if channel_count not in [1, 2, 3, 4]:
raise ValueError(f"Incorrect channel count: {channel_count}")
all_points: list[list[tuple[int, int]]] = []
debug_images = torch.zeros(
(batch_count, height, width, 3),
dtype=torch.uint8,
device=input_device,
device=imgs.device,
)
points_tensor = torch.tensor(
[255, 255, 255], dtype=torch.uint8, device=input_device
)
for i in range(batch_count):
value_threshold: torch.Tensor
if image is not None:
img_slice = image[i]
img_channels = img_slice.shape[2]
if img_channels == 1 or img_channels == 2:
value_threshold = img_slice[:, :, 0]
elif img_channels == 3 or img_channels == 4:
value_threshold = img_slice[:, :, :3].max(dim=2)[0]
else:
raise ValueError(
f"Unsupported image channel count: {img_channels} for image at batch index {i}"
)
for i, img in enumerate(imgs):
if channel_count == 1:
alpha_channel = img if len(img.shape) == 2 else img[:, :, 0]
elif channel_count == 2:
alpha_channel = img[:, :, 1]
elif channel_count == 4:
alpha_channel = img[:, :, 3]
else:
mask_slice = mask[i]
value_threshold = mask_slice
# get intensity
alpha_channel = img[:, :, :3].max(dim=2)[0]
condition = value_threshold > threshold
if image is not None and mask is not None:
mask_slice = mask[i]
mask_active_condition = mask_slice > 0.0
condition = condition & mask_active_condition
points = (alpha_channel > threshold).nonzero(as_tuple=False)
points_yx = condition.nonzero(as_tuple=False)
if len(points) > max_points:
indices = torch.randperm(points.size(0), device=img.device)[
:max_points
]
points = points[indices]
if points_yx.size(0) > max_points:
# shuffle and pick max_points randomly
indices = torch.randperm(
points_yx.size(0), device=input_device
)[:max_points]
points_yx = points_yx[indices]
elif max_points == 0:
points_yx = torch.empty(
(0, 2), dtype=torch.long, device=input_device
)
points = [(int(y.item()), int(x.item())) for x, y in points]
all_points.append(points)
current_points = [
(int(p[1].item()), int(p[0].item())) for p in points_yx
]
all_points.append(current_points)
for x_coord, y_coord in current_points:
self._draw_circle(
debug_images[i],
(x_coord, y_coord),
radius=5,
color_tensor=points_tensor,
)
for x, y in points:
self._draw_circle(debug_images[i], (x, y), 5)
return (all_points, debug_images)
@staticmethod
def _draw_circle(
image: torch.Tensor,
center: tuple[int, int],
radius: int,
color_tensor: torch.Tensor,
image: torch.Tensor, center: tuple[int, int], radius: int
):
"""Draw a 5px circle on the image."""
x0, y0 = center
h, w, _ = image.shape
min_x_bbox = max(0, x0 - radius)
max_x_bbox = min(w - 1, x0 + radius)
min_y_bbox = max(0, y0 - radius)
max_y_bbox = min(h - 1, y0 + radius)
for py in range(min_y_bbox, max_y_bbox + 1):
for px in range(min_x_bbox, max_x_bbox + 1):
if (px - x0) ** 2 + (py - y0) ** 2 <= radius**2:
image[py, px] = color_tensor
for x in range(-radius, radius + 1):
for y in range(-radius, radius + 1):
in_radius = x**2 + y**2 <= radius**2
in_bounds = (
0 <= x0 + x < image.shape[1]
and 0 <= y0 + y < image.shape[0]
)
if in_radius and in_bounds:
image[y0 + y, x0 + x] = torch.tensor(
[255, 255, 255],
dtype=torch.uint8,
device=image.device,
)
class MTB_ColorCorrectGPU:
@@ -683,7 +627,6 @@ class MTB_ImageCompare:
import requests
import time
class MTB_LoadImageFromUrl:
@@ -699,14 +642,6 @@ class MTB_LoadImageFromUrl:
"default": "https://upload.wikimedia.org/wikipedia/commons/thumb/a/a7/Example.jpg/800px-Example.jpg"
},
),
"retry_count": (
"INT",
{"default": 3, "min": 1, "max": 20, "step": 1},
),
"retry_interval": (
"FLOAT",
{"default": 1.0, "min": 0.0, "max": 60.0, "step": 0.1},
),
}
}
@@ -714,27 +649,11 @@ class MTB_LoadImageFromUrl:
FUNCTION = "load"
CATEGORY = "mtb/IO"
def load(self, url, retry_count, retry_interval):
# get the image from the url with retry + exponential backoff
last_error = None
for attempt in range(retry_count):
try:
response = requests.get(url, stream=True)
response.raise_for_status()
image = Image.open(response.raw)
image = ImageOps.exif_transpose(image)
return (pil2tensor(image),)
except Exception as e:
last_error = e
if attempt == retry_count - 1:
raise
wait_seconds = retry_interval * (2**attempt)
if wait_seconds > 0:
time.sleep(wait_seconds)
if last_error is not None:
raise last_error
raise RuntimeError("Failed to load image from URL without captured exception")
def load(self, url):
# get the image from the url
image = Image.open(requests.get(url, stream=True).raw)
image = ImageOps.exif_transpose(image)
return (pil2tensor(image),)
class MTB_Blur:
@@ -904,11 +823,8 @@ class MTB_MaskToImage:
return {
"required": {
"mask": ("MASK",),
"color": ("COLOR", {"widgetType": "MTB_COLOR"}),
"background": (
"COLOR",
{"default": "#000000", "widgetType": "MTB_COLOR"},
),
"color": ("COLOR",),
"background": ("COLOR", {"default": "#000000"}),
},
"optional": {
"invert": ("BOOLEAN", {"default": False}),
+2 -9
View File
@@ -21,11 +21,7 @@ class MTB_StackImages:
"match_method": (
["error", "smallest", "largest"],
{"default": "error"},
),
"output_rgb": (
"BOOLEAN",
{"default": True, "tooltip": "Output RGB instead of RGBA"},
),
)
},
}
@@ -33,7 +29,7 @@ class MTB_StackImages:
FUNCTION = "stack"
CATEGORY = "mtb/image utils"
def stack(self, vertical, match_method="error", output_rgb=True, **kwargs):
def stack(self, vertical, match_method="error", **kwargs):
if not kwargs:
raise ValueError("At least one tensor must be provided.")
@@ -102,9 +98,6 @@ class MTB_StackImages:
stacked_tensor = torch.cat(normalized_tensors, dim=dim)
if output_rgb:
stacked_tensor = stacked_tensor[:, :, :, :3]
return (stacked_tensor,)
def normalize_to_rgba(self, tensor):
-17
View File
@@ -1,17 +0,0 @@
# from ..utils import hex_to_rgb
class MTB_ColorInput:
RETURN_TYPES = ("COLOR",)
FUNCTION = "color"
CATEGORY = "mtb/color"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {"color": ("MTB_COLOR", {"default": "#ffffff"})},
}
def color(self, color):
return (color,)
__nodes__ = [MTB_ColorInput]
+1 -1
View File
@@ -34,7 +34,7 @@ class MTB_ImageRemoveBackgroundRembg:
),
"bgcolor": (
"COLOR",
{"default": "#000000","widgetType": "MTB_COLOR"},
{"default": "#000000"},
),
},
}
+1 -1
View File
@@ -145,7 +145,7 @@ class MTB_ModelPatchSeamless:
tilingX,
tilingY,
):
hacked_model = model.clone()
hacked_model = copy.deepcopy(model)
self.apply_circular(
hacked_model.model, startStep, stopStep, tilingX, tilingY
)
+1 -4
View File
@@ -43,10 +43,7 @@ class MTB_TransformImage:
["edge", "constant", "reflect", "symmetric"],
{"default": "edge"},
),
"constant_color": (
"COLOR",
{"default": "#000000", "widgetType": "MTB_COLOR"},
),
"constant_color": ("COLOR", {"default": "#000000"}),
},
"optional": {
"filter_type": (
+1 -1
View File
@@ -27,7 +27,7 @@ class MTB_LoadVitMatteModel:
def execute(self, *, kind: str, autodownload: bool):
dest = models_dir / "vitmatte"
dest.mkdir(exist_ok=True)
name = "dis" if kind == "Distinctions-646" else "com"
name = "dist" if kind == "Distinctions-646" else "com"
file = hf_hub_download(
repo_id="melmass/pytorch-scripts",
+2 -2
View File
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
[project]
name = "comfy-mtb"
version = "0.5.4"
version = "0.3.0"
description = "Animation oriented nodes pack for ComfyUI."
license = { text = "MIT" }
readme = "README.md"
@@ -63,7 +63,7 @@ DisplayName = "comfy-mtb"
Icon = "https://avatars.githubusercontent.com/u/7041726?v=4"
[tool.bumpversion]
current_version = "0.5.1"
current_version = "0.3.0"
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
serialize = ["{major}.{minor}.{patch}"]
search = "{current_version}"
+18
View File
@@ -1,5 +1,6 @@
import contextlib
import functools
import importlib
import math
import operator
import os
@@ -461,6 +462,23 @@ def _run_command(shell_cmd, ignored_lines_start):
print("Command executed successfully!")
def import_install(package_name):
package_spec = reqs_map.get(package_name, package_name)
try:
importlib.import_module(package_name)
except Exception: # (ImportError, ModuleNotFoundError):
run_command(
[
Path(sys.executable).as_posix(),
"-m",
"pip",
"install",
package_spec,
]
)
importlib.import_module(package_name)
# endregion
+144 -35
View File
@@ -12,6 +12,9 @@
import { app } from '../../scripts/app.js'
import { api } from '../../scripts/api.js'
if (!window.MTB) {
window.MTB = {}
}
// #region base utils
// - crude uuid
@@ -25,19 +28,6 @@ export function makeUUID() {
return uuid
}
// - basic debounce decorator
export function debounce(func, delay) {
let timeout
let debounced = function (...args) {
clearTimeout(timeout)
timeout = setTimeout(() => func.apply(this, args), delay)
}
debounced.cancel = () => {
clearTimeout(timeout)
}
return debounced
}
//- local storage manager
export class LocalStorageManager {
constructor(namespace) {
@@ -208,7 +198,6 @@ export function hideWidgetForGood(node, widget, suffix = '') {
widget.origComputeSize = widget.computeSize
widget.origSerializeValue = widget.serializeValue
widget.computeSize = () => [0, -4] // -4 is due to the gap litegraph adds between widgets automatically
widget.hidden = true
widget.type = CONVERTED_TYPE + suffix
// widget.serializeValue = () => {
// // Prevent serializing the widget if we have no input linked
@@ -290,6 +279,10 @@ export const getNamedWidget = (node, ...names) => {
* @returns {{to:LGraphNode, from:LGraphNode, type:'error' | 'incoming' | 'outgoing'}}
*/
export const nodesFromLink = (node, link) => {
if (typeof link === 'number') {
console.log('Resolving link from id', link)
link = app.graph.links[link]
}
const fromNode = app.graph.getNodeById(link.origin_id)
const toNode = app.graph.getNodeById(link.target_id)
@@ -635,21 +628,21 @@ function getBrightness(rgbObj) {
export function calculateTotalChildrenHeight(parentElement) {
let totalHeight = 0
if (!parentElement || !parentElement.children) {
return 0
}
for (const child of parentElement.children) {
const style = window.getComputedStyle(child)
const height = Number.parseFloat(style.height)
const marginTop = Number.parseFloat(style.marginTop)
const marginBottom = Number.parseFloat(style.marginBottom)
// Get height as an integer (without 'px')
const height = Number.parseInt(style.height, 10)
// Get vertical margin as integers
const marginTop = Number.parseInt(style.marginTop, 10)
const marginBottom = Number.parseInt(style.marginBottom, 10)
// Sum up height and vertical margins
totalHeight += height + marginTop + marginBottom
}
return Math.ceil(totalHeight)
return totalHeight
}
export const loadScript = (
@@ -660,15 +653,13 @@ export const loadScript = (
return new Promise((resolve, reject) => {
try {
// Check if the script already exists
let scriptEle = document.querySelector(`script[src="${FILE_URL}"]`)
if (scriptEle) {
scriptEle.addEventListener('load', (_ev) => {
resolve({ status: true })
})
const existingScript = document.querySelector(`script[src="${FILE_URL}"]`)
if (existingScript) {
resolve({ status: true, message: 'Script already loaded' })
return
}
scriptEle = document.createElement('script')
const scriptEle = document.createElement('script')
scriptEle.type = type
scriptEle.async = async
scriptEle.src = FILE_URL
@@ -687,8 +678,6 @@ export const loadScript = (
document.body.appendChild(scriptEle)
} catch (error) {
reject(error)
} finally {
infoLogger(`Finally loaded script: ${FILE_URL}`)
}
})
}
@@ -802,10 +791,12 @@ function loadParser(shiki) {
export const ensureMarkdownParser = async (callback) => {
infoLogger('Ensuring md parser')
const use_shiki = app.extensionManager.setting.get(
'mtb.noteplus.use-shiki',
false,
)
let use_shiki = false
try {
use_shiki = await api.getSetting('mtb.Use Shiki')
} catch (e) {
console.warn('Option not available yet', e)
}
if (window.MTB?.mdParser) {
infoLogger('Markdown parser found')
@@ -830,7 +821,8 @@ export const ensureMarkdownParser = async (callback) => {
callbackQueue.push(callback)
}
await await parserPromise
await parserPromise
await parserPromise
return window.MTB.mdParser
}
@@ -1083,6 +1075,66 @@ export const addDocumentation = (
// #endregion
// #region canvas / drawing
// calculate convex hull (Graham)
export function getConvexHull(points) {
if (points.length < 3) return points
// find the bottommost point (and leftmost if tied)
let bottom = 0
for (let i = 1; i < points.length; i++) {
if (
points[i][1] < points[bottom][1] ||
(points[i][1] === points[bottom][1] && points[i][0] < points[bottom][0])
) {
bottom = i
}
}
// swap bottom point to first position
;[points[0], points[bottom]] = [points[bottom], points[0]]
// sort points by polar angle with respect to base point
const basePoint = points[0]
points.sort((a, b) => {
if (a === basePoint) return -1
if (b === basePoint) return 1
const angleA = Math.atan2(a[1] - basePoint[1], a[0] - basePoint[0])
const angleB = Math.atan2(b[1] - basePoint[1], b[0] - basePoint[0])
if (angleA < angleB) return -1
if (angleA > angleB) return 1
// if angles are equal, sort by distance
const distA = (a[0] - basePoint[0]) ** 2 + (a[1] - basePoint[1]) ** 2
const distB = (b[0] - basePoint[0]) ** 2 + (b[1] - basePoint[1]) ** 2
return distA - distB
})
// build convex hull
const stack = [points[0], points[1]]
for (let i = 2; i < points.length; i++) {
while (
stack.length > 1 &&
!isLeftTurn(stack[stack.length - 2], stack[stack.length - 1], points[i])
) {
stack.pop()
}
stack.push(points[i])
}
return stack
}
function isLeftTurn(p1, p2, p3) {
return (
(p2[0] - p1[0]) * (p3[1] - p1[1]) - (p2[1] - p1[1]) * (p3[0] - p1[0]) > 0
)
}
// #endregion
// #region node extensions
/**
@@ -1157,6 +1209,8 @@ export const runAction = async (name, ...args) => {
const res = await req.json()
return res.result
}
window.MTB.run = runAction
export const getServerInfo = async () => {
const res = await api.fetchApi('/mtb/server-info')
return await res.json()
@@ -1169,3 +1223,58 @@ export const setServerInfo = async (opts) => {
}
// #endregion
// #region Authoring API / graph utilities
export const getAPIInputs = () => {
const inputs = {}
let counter = 1
for (const node of getNodes(true)) {
const widgets = node.widgets
if (node.properties.mtb_api && node.properties.useAPI) {
if (node.properties.mtb_api.inputs) {
for (const currentName in node.properties.mtb_api.inputs) {
const current = node.properties.mtb_api.inputs[currentName]
if (current.enabled) {
const inputName = current.name || currentName
const widget = widgets.find((w) => w.name === currentName)
if (!widget) continue
if (!(inputName in inputs)) {
inputs[inputName] = {
...current,
id: counter,
name: inputName,
type: current.type,
node_id: node.id,
widgets: [],
}
}
inputs[inputName].widgets.push(widget)
counter = counter + 1
}
}
}
}
}
return inputs
}
export const getNodes = (skip_unused) => {
const nodes = []
for (const outerNode of app.graph.computeExecutionOrder(false)) {
const skipNode =
(outerNode.mode === 2 || outerNode.mode === 4) && skip_unused
const innerNodes =
!skipNode && outerNode.getInnerNodes
? outerNode.getInnerNodes()
: [outerNode]
for (const node of innerNodes) {
if ((node.mode === 2 || node.mode === 4) && skip_unused) {
continue
}
nodes.push(node)
}
}
return nodes
}
// #endregion
+25 -20
View File
@@ -28,7 +28,7 @@ function createDebugSection(title) {
margin: '8px 0',
padding: '8px',
borderRadius: '4px',
backgroundColor: 'rgba(0,0,0,0.2)',
backgroundColor: 'rgba(0,0,0,0.2)'
})
const header = mtb_ui.makeElement('h3', {
@@ -37,7 +37,7 @@ function createDebugSection(title) {
borderBottom: '1px solid rgba(255,255,255,0.1)',
fontSize: '14px',
fontWeight: 'bold',
color: '#9f9',
color: '#9f9'
})
header.textContent = title
section.appendChild(header)
@@ -47,21 +47,21 @@ function createDebugSection(title) {
function createDebugContent(content, type) {
const wrapper = mtb_ui.makeElement('div', {
margin: '4px 0',
margin: '4px 0'
})
if (type === 'text') {
const text = mtb_ui.makeElement('p', {
margin: '2px 0',
fontFamily: 'monospace',
whiteSpace: 'pre-wrap',
whiteSpace: 'pre-wrap'
})
text.innerHTML = content
wrapper.appendChild(text)
} else if (type === 'image') {
const img = mtb_ui.makeElement('img', {
width: '100%',
borderRadius: '2px',
borderRadius: '2px'
})
img.src = content
wrapper.appendChild(img)
@@ -148,18 +148,18 @@ app.registerExtension({
const uiData = data.ui || data
if (uiData.items) {
uiData.items.forEach((item) => {
const inputName = item.input
if (!inputData[inputName]) {
inputData[inputName] = { text: [], b64_images: [] }
}
if (item.text) {
inputData[inputName].text.push(...item.text)
}
if (item.b64_images) {
inputData[inputName].b64_images.push(...item.b64_images)
}
})
uiData.items.forEach(item => {
const inputName = item.input
if (!inputData[inputName]) {
inputData[inputName] = { text: [], b64_images: [] }
}
if (item.text) {
inputData[inputName].text.push(...item.text)
}
if (item.b64_images) {
inputData[inputName].b64_images.push(...item.b64_images)
}
})
}
let widgetI = 1
@@ -171,18 +171,23 @@ app.registerExtension({
const section = createDebugSection(inputName)
if (content.text.length > 0) {
content.text.forEach((text) => {
content.text.forEach(text => {
section.appendChild(createDebugContent(text, 'text'))
})
}
if (content.b64_images.length > 0) {
content.b64_images.forEach((img) => {
content.b64_images.forEach(img => {
section.appendChild(createDebugContent(img, 'image'))
})
}
this.addDOMWidget(`debug_section_${widgetI}`, 'CUSTOM', section, {})
this.addDOMWidget(
`debug_section_${widgetI}`,
'CUSTOM',
section,
{}
)
widgetI++
}
+296 -296
View File
@@ -13,40 +13,40 @@ import { api } from '../../scripts/api.js'
import { app } from '../../scripts/app.js'
import { LocalStorageManager } from './comfy_shared.js'
const styles = {
lighbox: {
position: 'fixed',
top: 0,
left: 0,
width: '100vw',
height: '100vh',
background: 'rgba(0,0,0,0.5)',
display: 'none',
justifyContent: 'center',
alignItems: 'center',
zIndex: 999,
},
lightboxBtn: (extra) => ({
position: 'absolute',
top: '50%',
background: 'none',
border: 'none',
color: '#fff',
zIndex: 1000,
fontSize: '30px',
cursor: 'pointer',
pointerEvents: 'auto',
...extra,
}),
img_list: {
minHeight: '30px',
maxHeight: '300px',
width: '100vw',
position: 'absolute',
bottom: 0,
zIndex: 10,
background: '#333',
overflow: 'auto',
},
lighbox: {
position: 'fixed',
top: 0,
left: 0,
width: '100vw',
height: '100vh',
background: 'rgba(0,0,0,0.5)',
display: 'none',
justifyContent: 'center',
alignItems: 'center',
zIndex: 999,
},
lightboxBtn: (extra) => ({
position: 'absolute',
top: '50%',
background: 'none',
border: 'none',
color: '#fff',
zIndex: 1000,
fontSize: '30px',
cursor: 'pointer',
pointerEvents: 'auto',
...extra,
}),
img_list: {
minHeight: '30px',
maxHeight: '300px',
width: '100vw',
position: 'absolute',
bottom: 0,
zIndex: 10,
background: '#333',
overflow: 'auto',
},
}
let currentImageIndex = 0
@@ -58,299 +58,299 @@ const storage = new LocalStorageManager('mtb')
let activated = storage.get('image_feed', false)
app.registerExtension({
name: 'mtb.ImageFeed',
setup: () => {
app.ui.settings.addSetting({
id: 'mtb.Main.image-feed-enabled',
category: ['mtb', ' Main', 'image-feed-enabled'],
name: 'Enable Image Feed',
type: 'boolean',
defaultValue: false,
attrs: {
style: {
fontFamily: 'monospace',
},
},
async onChange(value) {
storage.set('image_feed', value)
activated = value
},
})
},
init: async () => {
if (!activated) {
return
}
const pythongossFeed = app.extensions.find(
(e) => e.name === 'pysssss.ImageFeed',
)
if (pythongossFeed) {
console.warn(
"[mtb] - Aborting the loading of mtb's imageFeed in favor of pysssss.ImageFeed",
)
activated = false // just in case other methods are added later on
return
}
// - HTML & CSS
//- lightbox
const lightboxContainer = document.createElement('div')
Object.assign(lightboxContainer.style, styles.lighbox)
name: 'mtb.ImageFeed',
setup: () => {
app.ui.settings.addSetting({
id: 'mtb.Main.image-feed-enabled',
category: ['mtb', 'Main', 'image-feed-enabled'],
name: 'Enable Image Feed',
type: 'boolean',
defaultValue: false,
attrs: {
style: {
fontFamily: 'monospace',
},
},
async onChange(value) {
storage.set('image_feed', value)
activated = value
},
})
},
init: async () => {
if (!activated) {
return
}
const pythongossFeed = app.extensions.find(
(e) => e.name === 'pysssss.ImageFeed',
)
if (pythongossFeed) {
console.warn(
"[mtb] - Aborting the loading of mtb's imageFeed in favor of pysssss.ImageFeed",
)
activated = false // just in case other methods are added later on
return
}
// - HTML & CSS
//- lightbox
const lightboxContainer = document.createElement('div')
Object.assign(lightboxContainer.style, styles.lighbox)
const lightboxImage = document.createElement('img')
Object.assign(lightboxImage.style, {
maxHeight: '100%',
maxWidth: '100%',
borderRadius: '5px',
})
const lightboxImage = document.createElement('img')
Object.assign(lightboxImage.style, {
maxHeight: '100%',
maxWidth: '100%',
borderRadius: '5px',
})
// previous and next buttons
const lightboxPrevBtn = document.createElement('button')
const lightboxNextBtn = document.createElement('button')
// previous and next buttons
const lightboxPrevBtn = document.createElement('button')
const lightboxNextBtn = document.createElement('button')
lightboxPrevBtn.textContent = '❮'
lightboxNextBtn.textContent = '❯'
lightboxPrevBtn.textContent = '❮'
lightboxNextBtn.textContent = '❯'
Object.assign(lightboxPrevBtn.style, styles.lightboxBtn({ left: '0%' }))
Object.assign(lightboxNextBtn.style, styles.lightboxBtn({ right: '0%' }))
Object.assign(lightboxPrevBtn.style, styles.lightboxBtn({ left: '0%' }))
Object.assign(lightboxNextBtn.style, styles.lightboxBtn({ right: '0%' }))
// close button
const lightboxCloseBtn = document.createElement('button')
Object.assign(
lightboxCloseBtn.style,
styles.lightboxBtn({ right: '0', top: '0' }),
)
lightboxCloseBtn.textContent = '❌'
// close button
const lightboxCloseBtn = document.createElement('button')
Object.assign(
lightboxCloseBtn.style,
styles.lightboxBtn({ right: '0', top: '0' }),
)
lightboxCloseBtn.textContent = '❌'
const lightboxButtons = document.createElement('div')
Object.assign(lightboxButtons.style, {
position: 'absolute',
top: '0%',
right: '0%',
// transform: "translate(50%, -50%)",
height: '100%',
width: '100%',
background: 'none',
border: 'none',
color: '#fff',
fontSize: '30px',
cursor: 'pointer',
pointerEvents: 'none',
})
const lightboxButtons = document.createElement('div')
Object.assign(lightboxButtons.style, {
position: 'absolute',
top: '0%',
right: '0%',
// transform: "translate(50%, -50%)",
height: '100%',
width: '100%',
background: 'none',
border: 'none',
color: '#fff',
fontSize: '30px',
cursor: 'pointer',
pointerEvents: 'none',
})
lightboxButtons.append(lightboxPrevBtn, lightboxNextBtn, lightboxCloseBtn)
lightboxContainer.append(lightboxButtons, lightboxImage)
lightboxButtons.append(lightboxPrevBtn, lightboxNextBtn, lightboxCloseBtn)
lightboxContainer.append(lightboxButtons, lightboxImage)
//- image list
const imageListContainer = document.createElement('div')
Object.assign(imageListContainer.style, styles.img_list)
//- image list
const imageListContainer = document.createElement('div')
Object.assign(imageListContainer.style, styles.img_list)
const createImgListBtn = (text, style) => {
const btn = document.createElement('button')
btn.type = 'button'
btn.textContent = text
Object.assign(btn.style, {
...style,
border: 'none',
color: '#fff',
background: 'none',
height: '20px',
cursor: 'pointer',
position: 'absolute',
top: '5px',
fontSize: '12px',
lineHeight: '12px',
})
imageListContainer.append(btn)
return btn
}
const showBtn = document.createElement('button')
const closeBtn = createImgListBtn('❌', {
width: '20px',
textIndent: '-4px',
right: '5px',
})
const loadButton = createImgListBtn('Load Session History', {
right: '90px',
})
const clearButton = createImgListBtn('Clear', {
right: '30px',
})
const createImgListBtn = (text, style) => {
const btn = document.createElement('button')
btn.type = 'button'
btn.textContent = text
Object.assign(btn.style, {
...style,
border: 'none',
color: '#fff',
background: 'none',
height: '20px',
cursor: 'pointer',
position: 'absolute',
top: '5px',
fontSize: '12px',
lineHeight: '12px',
})
imageListContainer.append(btn)
return btn
}
const showBtn = document.createElement('button')
const closeBtn = createImgListBtn('❌', {
width: '20px',
textIndent: '-4px',
right: '5px',
})
const loadButton = createImgListBtn('Load Session History', {
right: '90px',
})
const clearButton = createImgListBtn('Clear', {
right: '30px',
})
//- tools popup button
showBtn.classList.add('comfy-settings-btn')
Object.assign(showBtn.style, {
right: '16px',
cursor: 'pointer',
display: 'none',
})
//- tools popup button
showBtn.classList.add('comfy-settings-btn')
Object.assign(showBtn.style, {
right: '16px',
cursor: 'pointer',
display: 'none',
})
//- append to DOM
document.body.append(imageListContainer)
//- append to DOM
document.body.append(imageListContainer)
showBtn.textContent = '🖼'
showBtn.onclick = () => {
imageListContainer.style.display = 'block'
showBtn.style.display = 'none'
}
document.querySelector('.comfy-settings-btn').after(showBtn)
document.querySelector('.comfy-settings-btn').after(lightboxContainer)
showBtn.textContent = '🖼'
showBtn.onclick = () => {
imageListContainer.style.display = 'block'
showBtn.style.display = 'none'
}
document.querySelector('.comfy-settings-btn').after(showBtn)
document.querySelector('.comfy-settings-btn').after(lightboxContainer)
// for (const { output } of history) {
// if (output?.images) {
// for (const src of output.images) {
// const img = document.createElement("img");
// const but = document.createElement("button");
// for (const { output } of history) {
// if (output?.images) {
// for (const src of output.images) {
// const img = document.createElement("img");
// const but = document.createElement("button");
//- callbacks
closeBtn.onclick = () => {
imageListContainer.style.display = 'none'
showBtn.style.display = 'unset'
}
//- callbacks
closeBtn.onclick = () => {
imageListContainer.style.display = 'none'
showBtn.style.display = 'unset'
}
clearButton.onclick = () => {
imageListContainer.replaceChildren(closeBtn, clearButton, loadButton)
}
clearButton.onclick = () => {
imageListContainer.replaceChildren(closeBtn, clearButton, loadButton)
}
lightboxNextBtn.onclick = () => {
currentImageIndex = (currentImageIndex + 1) % imageUrls.length
const imageUrl = imageUrls[currentImageIndex]
lightboxImage.src = imageUrl
}
lightboxNextBtn.onclick = () => {
currentImageIndex = (currentImageIndex + 1) % imageUrls.length
const imageUrl = imageUrls[currentImageIndex]
lightboxImage.src = imageUrl
}
// Modify the lightboxPrevBtn onclick callback
lightboxPrevBtn.onclick = () => {
currentImageIndex =
(currentImageIndex - 1 + imageUrls.length) % imageUrls.length
const imageUrl = imageUrls[currentImageIndex]
lightboxImage.src = imageUrl
}
// Modify the lightboxPrevBtn onclick callback
lightboxPrevBtn.onclick = () => {
currentImageIndex =
(currentImageIndex - 1 + imageUrls.length) % imageUrls.length
const imageUrl = imageUrls[currentImageIndex]
lightboxImage.src = imageUrl
}
lightboxCloseBtn.onclick = () => {
lightboxContainer.style.display = 'none'
}
lightboxImage.onclick = lightboxNextBtn.onclick
/**
* This is the function that creates the image buttons for the image list
* They are wrapped in a button so that they can be clicked and open
* the image in the lightbox.
* @param {*} src
*/
const createImageBtn = (src) => {
console.debug(`making image ${src.filename}`)
const img = document.createElement('img')
const but = document.createElement('button')
lightboxCloseBtn.onclick = () => {
lightboxContainer.style.display = 'none'
}
lightboxImage.onclick = lightboxNextBtn.onclick
/**
* This is the function that creates the image buttons for the image list
* They are wrapped in a button so that they can be clicked and open
* the image in the lightbox.
* @param {*} src
*/
const createImageBtn = (src) => {
console.debug(`making image ${src.filename}`)
const img = document.createElement('img')
const but = document.createElement('button')
Object.assign(but.style, {
height: '120px',
width: '120px',
border: 'none',
padding: 0,
margin: 0,
})
Object.assign(img.style, {
width: '100%',
height: '100%',
objectFit: 'cover',
})
Object.assign(but.style, {
height: '120px',
width: '120px',
border: 'none',
padding: 0,
margin: 0,
})
Object.assign(img.style, {
width: '100%',
height: '100%',
objectFit: 'cover',
})
img.src = `/view?filename=${encodeURIComponent(src.filename)}&type=${
src.type
}&subfolder=${encodeURIComponent(src.subfolder)}`
img.src = `/view?filename=${encodeURIComponent(src.filename)}&type=${
src.type
}&subfolder=${encodeURIComponent(src.subfolder)}`
imageUrls.push(img.src)
imageUrls.push(img.src)
console.debug(img.src)
console.debug(img.src)
img.onload = () => {
but.style.width = `${120 * (img.naturalWidth / img.naturalHeight)}px`
}
img.onload = () => {
but.style.width = `${120 * (img.naturalWidth / img.naturalHeight)}px`
}
but.onclick = () => {
lightboxContainer.style.display = 'flex'
// add the same image to the lightbox
lightboxImage.src = img.src
// lighboxContainer.replaceChildren(lightboxButtons, img);
}
but.onclick = () => {
lightboxContainer.style.display = 'flex'
// add the same image to the lightbox
lightboxImage.src = img.src
// lighboxContainer.replaceChildren(lightboxButtons, img);
}
// add right click menu
but.addEventListener('contextmenu', (e) => {
e.preventDefault()
// add right click menu
but.addEventListener('contextmenu', (e) => {
e.preventDefault()
if (image_menu) {
image_menu.remove()
}
if (image_menu) {
image_menu.remove()
}
image_menu = document.createElement('div')
Object.assign(image_menu.style, {
position: 'absolute',
top: `${e.clientY}px`,
left: `${e.clientX}px`,
background: '#333',
color: '#fff',
padding: '5px',
borderRadius: '5px',
zIndex: 999,
})
const load_img = document.createElement('button')
load_img.textContent = 'Load'
load_img.onclick = () => {
app.handleFile(img.src)
}
image_menu = document.createElement('div')
Object.assign(image_menu.style, {
position: 'absolute',
top: `${e.clientY}px`,
left: `${e.clientX}px`,
background: '#333',
color: '#fff',
padding: '5px',
borderRadius: '5px',
zIndex: 999,
})
const load_img = document.createElement('button')
load_img.textContent = 'Load'
load_img.onclick = () => {
app.handleFile(img.src)
}
image_menu.appendChild(load_img)
document.body.appendChild(image_menu)
})
image_menu.appendChild(load_img)
document.body.appendChild(image_menu)
})
but.append(img)
imageListContainer.prepend(but)
}
but.append(img)
imageListContainer.prepend(but)
}
loadButton.onclick = async () => {
const all_history = await api.getHistory()
for (const history of all_history.History) {
if (history.outputs) {
for (const key of Object.keys(history.outputs)) {
console.debug(key)
if (history.outputs[key].images) {
for (const im of history.outputs[key].images) {
console.debug(im)
createImageBtn(im)
}
}
}
// for (const src of outputs.outputs.images) {
// console.debug(src)
// makeImage(`${src.subfolder}/${src.filename}`)
// }
}
}
}
loadButton.onclick = async () => {
const all_history = await api.getHistory()
for (const history of all_history.History) {
if (history.outputs) {
for (const key of Object.keys(history.outputs)) {
console.debug(key)
if (history.outputs[key].images) {
for (const im of history.outputs[key].images) {
console.debug(im)
createImageBtn(im)
}
}
}
// for (const src of outputs.outputs.images) {
// console.debug(src)
// makeImage(`${src.subfolder}/${src.filename}`)
// }
}
}
}
///////-------
///////-------
// const all_history = await api.getHistory()
// for (const history of all_history.History) {
// if (history.outputs) {
// for (const key of Object.keys(history.outputs)) {
// for (const im of history.outputs[key].images) {
// makeImage(im)
// }
// }
// // for (const src of outputs.outputs.images) {
// // console.debug(src)
// // makeImage(`${src.subfolder}/${src.filename}`)
// // }
// }
// }
// const all_history = await api.getHistory()
// for (const history of all_history.History) {
// if (history.outputs) {
// for (const key of Object.keys(history.outputs)) {
// for (const im of history.outputs[key].images) {
// makeImage(im)
// }
// }
// // for (const src of outputs.outputs.images) {
// // console.debug(src)
// // makeImage(`${src.subfolder}/${src.filename}`)
// // }
// }
// }
//- Hook into the API
api.addEventListener('executed', ({ detail }) => {
if (detail?.output?.images) {
for (const src of detail.output.images) {
console.debug(`Adding ${src} to image feed`)
createImageBtn(src)
}
}
})
},
//- Hook into the API
api.addEventListener('executed', ({ detail }) => {
if (detail?.output?.images) {
for (const src of detail.output.images) {
console.debug(`Adding ${src} to image feed`)
createImageBtn(src)
}
}
})
},
})
+505 -304
View File
@@ -2,8 +2,8 @@
import { app } from '../../scripts/app.js'
import { api } from '../../scripts/api.js'
import { infoLogger, successLogger, errorLogger } from './comfy_shared.js'
import * as mtb_ui from './mtb_ui.js'
import * as shared from './comfy_shared.js'
import {
@@ -13,38 +13,140 @@ import {
makeSelect,
makeSlider,
renderSidebar,
ContextMenu,
} from './mtb_ui.js'
let currentAbortController = null
/** cursor/offset of where we are at */
const offset = 0
// These are "global" variables mostly meant to sync user settings.
/** width of the images in the grid */
let currentWidth = 200
let saltUrls =
app.extensionManager.setting.get('mtb.io-sidebar.salt_urls') || false
let targetWidth =
app.extensionManager.setting.get('mtb.io-sidebar.img-size') || 512
let currentMode = 'input'
let subfolder = ''
let currentSort = 'None'
const IMAGE_NODES = ['LoadImage', 'VHS_LoadImagePath']
let clientOnce = false
/** reference to the dom element receiving the images */
let imgGrid = undefined
/** currently loaded image (as object urls) */
let loaded_images = undefined
/**
* stores the user's full local path to input/output directory
* This is then used to feed VHS Load Image (from path)
*/
let userDirectories = undefined
// const IMAGE_NODES = ['LoadImage', 'VHS_LoadImagePath']
const VIDEO_NODES = ['VHS_LoadVideo']
const PROCESSED_PROMPT_IDS = new Set()
let contextMenu = undefined
function debounce(func, wait) {
let timeout
return function executedFunction(...args) {
const later = () => {
infoLogger('Debouncing method')
clearTimeout(timeout)
func(...args)
}
clearTimeout(timeout)
timeout = setTimeout(later, wait)
}
}
const debouncedGetUrls = async (ms = 250) => {
if (loaded_images === undefined) {
return await getUrls(subfolder)
}
debounce(async (subfolder) => {
const urls = await getUrls(subfolder)
infoLogger('Loaded URLs (debounced): ', urls)
if (urls) {
loaded_images = await getImgsFromUrls(urls, imgGrid)
infoLogger('Loaded Images (debounced): ', loaded_images)
}
}, ms)
return loaded_images
}
/** Callback on clicking an image in the grid */
const updateImage = (node, image) => {
if (IMAGE_NODES.includes(node.type)) {
const w = node.widgets?.find((w) => w.name === 'image')
if (w) {
w.value = image
w.callback()
switch (node.type) {
case 'LoadImage': {
if (subfolder && subfolder !== '') {
app.extensionManager.toast.add({
severity: 'warn',
summary: 'Subfolder not supported',
detail: "The LoadImage node doesn't support subfolders",
life: 5000,
})
return
}
if (currentMode === 'output') {
app.extensionManager.toast.add({
severity: 'warn',
summary: 'Outputs not supported',
detail:
"The LoadImage node doesn't support loading outputs, use VHS Load Image Path and I'll resolve the full path.",
life: 5000,
})
return
}
// if (IMAGE_NODES.includes(node.type)) {
const w = node.widgets?.find((w) => w.name === 'image')
if (w) {
w.value = image
w.callback()
}
//}
break
}
} else if (VIDEO_NODES.includes(node.type)) {
const w = node.widgets?.find((w) => w.name === 'video')
if (w) {
node.updateParameters({ filename: image }, true)
case 'VHS_LoadImagePath': {
let value = image
if (!userDirectories?.output) {
app.extensionManager.toast.add({
severity: 'warn',
summary: 'User output directory not resolved',
detail: "We couldn't resolve the image full path.",
life: 5000,
})
return
}
if (subfolder && subfolder !== '') {
value = `${subfolder}/${image}`
}
value = `${userDirectories.output}/${value}`
const w = node.widgets?.find((w) => w.name === 'image')
if (w) {
console.log(w)
w.value = value
// TODO: VHS needs explicity value passsed here
w.callback(value)
}
break
}
case VIDEO_NODES.includes(node.type): {
const w = node.widgets?.find((w) => w.name === 'video')
if (w) {
node.updateParameters({ filename: image }, true)
}
break
}
default: {
console.warn('No method to update', node.type)
}
} else {
console.warn('No method to update', node.type)
}
}
@@ -53,19 +155,15 @@ const updateImage = (node, image) => {
* @param {ResultItem} resultItem
* @returns {string} - The request URL.
*/
const resultItemToQuery = (resultItem) => {
const res = [
const resultItemToQuery = (resultItem) =>
[
`/mtb/view?filename=${resultItem.filename}`,
`width=512`,
`type=${resultItem.type}`,
`subfolder=${resultItem.subfolder}`,
'preview=',
]
if (targetWidth > 0) {
res.splice(1, 0, `width=${targetWidth}`)
}
`preview=`,
].join('&')
return res.join('&')
}
/**
* Retrieves the unique prompt ID from a history task item.
* @param {HistoryTaskItem} historyTaskItem
@@ -91,7 +189,7 @@ const getNewOutputUrls = (mostRecentTask) => {
const imageOutputs = Object.values(nodeOutputs.images)
imageOutputs.forEach(
(resultItem) =>
(urls[resultItem.filename] = resultItemToQuery(resultItem)),
(urls[resultItem.filename] = resultItemToQuery(resultItem))
)
}
// Can process `animated` and `audio` outputs here.
@@ -120,94 +218,265 @@ const updateOutputsGrid = async () => {
}
const getImgsFromUrls = (urls, target, options = { prepend: false }) => {
const imgs = []
if (urls === undefined) {
return imgs
if (currentAbortController) {
currentAbortController.abort()
}
const elem = currentMode === 'video' ? 'video' : 'img'
infoLogger('getting images from urls', urls)
for (const [key, url] of Object.entries(urls)) {
const a = makeElement(elem)
a.src = url
a.width = currentWidth
if (currentMode === 'input') {
a.onclick = (_e) => {
if (subfolder !== '') {
app.extensionManager.toast.add({
severity: 'warn',
summary: 'Subfolder not supported',
detail: "The LoadImage node doesn't support subfolders",
life: 5000,
})
return
}
const selected = app.canvas.selected_nodes
if (selected && Object.keys(selected).length === 0) {
app.extensionManager.toast.add({
severity: 'warn',
summary: 'No node selected!',
detail:
'For now the only action when clicking images in the sidebar is to set the image on all selected LoadImage nodes.',
life: 5000,
})
return
}
currentAbortController = new AbortController()
const { signal } = currentAbortController
const imgs = []
if (!urls) return imgs
for (const [_id, node] of Object.entries(app.canvas.selected_nodes)) {
updateImage(node, key)
}
}
} else if (currentMode === 'output') {
a.onclick = (_e) => {
// window.MTB?.notify?.("Output import isn't supported yet...", 5000)
if (subfolder !== '') {
app.extensionManager.toast.add({
severity: 'warn',
summary: 'Subfolder not supported',
detail: "The LoadImage node doesn't support subfolders",
life: 5000,
})
return
}
const loadingIndicator = document.createElement('div')
loadingIndicator.className = 'mtb-loading-indicator'
if (target) target.appendChild(loadingIndicator)
app.extensionManager.toast.add({
severity: 'warn',
summary: 'Outputs not supported',
detail:
'For now only inputs can be clicked to load the image on the active LoadImage node.',
life: 5000,
const totalImages = Object.keys(urls).length
let loadedCount = 0
const updateLoadingStatus = () => {
loadingIndicator.textContent = `Loaded ${loadedCount} of ${totalImages} images`
}
updateLoadingStatus()
try {
const loadImage = async (key, url) => {
try {
const response = await fetch(url, { signal })
if (!response.ok) {
console.warn(`Failed to fetch ${key}: ${response.status}`)
return null
}
// throw new Error(`HTTP error! status: ${response.status}`)
const blob = await response.blob()
const imgUrl = URL.createObjectURL(blob)
const elem = makeElement(currentMode === 'video' ? 'video' : 'img')
elem.src = imgUrl
elem.width = currentWidth
// cleanup
elem.onload = () => URL.revokeObjectURL(imgUrl)
elem.onerror = () => URL.revokeObjectURL(imgUrl)
// Add click handler for input mode
// if (currentMode === 'input') {
// elem.onclick = (_e) => {
// Your existing click handler code
// }
// }
// Add context menu
elem.addEventListener('contextmenu', (e) => {
e.preventDefault()
const contextMenuItems = [
{
label: 'Add Node with Image',
icon: '🖼',
action: () => {
const node = app.graph.createNode('LoadImage')
updateImage(node, key)
},
},
{
label: 'Load Workflow from Image',
icon: '📋',
action: async () => {
try {
const response = await fetch(url)
const data = await response.blob()
// Assuming you have a function to extract workflow from image metadata
const workflow = await extractWorkflowFromImage(data)
if (workflow) {
app.loadGraphData(workflow)
}
} catch (error) {
app.extensionManager.toast.add({
severity: 'error',
summary: 'Error',
detail: 'Failed to load workflow from image',
life: 3000,
})
}
},
},
{
label: 'View Full Image',
icon: '🔍',
action: () => {
window.open(url, '_blank')
},
},
]
contextMenu.show(e.pageX, e.pageY, contextMenuItems, {
elem,
key,
url,
})
})
}
} 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
elem.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: 'Please select a node first.',
life: 5000,
})
return
}
for (const [_id, node] of Object.entries(selected)) {
updateImage(node, key)
}
}
for (const [_id, node] of Object.entries(app.canvas.selected_nodes)) {
updateImage(node, key)
loadedCount++
updateLoadingStatus()
return elem
} catch (error) {
if (error.name === 'AbortError') {
console.log('Fetch aborted')
return null
}
console.error('Error loading image:', error)
return null
}
}
imgs.push(a)
}
if (target !== undefined) {
if (options.prepend) target.prepend(...imgs)
const BATCH_SIZE = 20
for (let i = 0; i < Object.entries(urls).length; i += BATCH_SIZE) {
const batch = Object.entries(urls).slice(i, i + BATCH_SIZE)
const loadedImages = await Promise.all(
batch.map(([key, url]) => loadImage(key, url)),
)
const validImages = loadedImages.filter((img) => img !== null)
imgs.push(...validImages)
if (target) {
target.append(...validImages)
}
}
return imgs
// return
// const elem = currentMode === 'video' ? 'video' : 'img'
for (const [key, url] of Object.entries(urls)) {
const a = makeElement(elem)
a.src = url
a.width = currentWidth
const selected = app.canvas.selected_nodes
if (currentMode === 'input') {
a.onclick = (_e) => {
// if (subfolder !== '') {
// app.extensionManager.toast.add({
// severity: 'warn',
// summary: 'Subfolder not supported',
// detail: "The LoadImage node doesn't support subfolders",
// life: 5000,
// })
// return
// }
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 images in the sidebar is to set the image on all selected LoadImage nodes.',
life: 5000,
})
return
}
for (const [_id, node] of Object.entries(app.canvas.selected_nodes)) {
updateImage(node, key)
}
}
} else if (currentMode === 'output') {
a.onclick = (_e) => {
if (selected && Object.keys(selected).length === 0) {
return
}
for (const [_id, node] of Object.entries(app.canvas.selected_nodes)) {
updateImage(node, key)
}
// window.MTB?.notify?.("Output import isn't supported yet...", 5000)
// if (subfolder !== '') {
// app.extensionManager.toast.add({
// severity: 'warn',
// summary: 'Subfolder not supported',
// detail: "The LoadImage node doesn't support subfolders",
// life: 5000,
// })
// return
// }
//
// app.extensionManager.toast.add({
// severity: 'warn',
// summary: 'Outputs not supported',
// detail:
// 'For now only inputs can be clicked to load the image on the active LoadImage node.',
// life: 5000,
// })
}
} else {
a.autoplay = true
a.muted = true
a.loop = true
a.onclick = (_e) => {
const selected = app.canvas.selected_nodes
if (selected && Object.keys(selected).length === 0) {
app.extensionManager.toast.add({
severity: 'warn',
summary: 'No node selected!',
detail:
"For now the only action when clicking videos in the sidebar is to set the video on all selected 'Load Video (Upload)' nodes.",
life: 5000,
})
return
}
for (const [_id, node] of Object.entries(app.canvas.selected_nodes)) {
updateImage(node, key)
}
}
}
imgs.push(a)
}
if (target !== undefined) {
if (options.prepend) target.prepend(...imgs)
else target.append(...imgs)
}
return imgs
} finally {
// Keep loading indicator visible for a moment after completion
setTimeout(() => {
if (target && loadingIndicator.parentNode === target) {
loadingIndicator.remove()
}
}, 2000)
}
}
// Helper function to extract workflow from image metadata
async function extractWorkflowFromImage(blob) {
// Implementation depends on how the workflow data is stored in the image
// This is just a placeholder
try {
// You might need to use ExifReader or similar library to extract metadata
return null
} catch (error) {
console.error('Failed to extract workflow:', error)
return null
}
return imgs
}
const getModes = async () => {
@@ -216,11 +485,11 @@ const getModes = async () => {
}
const getUrls = async (subfolder) => {
const count = (await api.getSetting('mtb.io-sidebar.count')) || 1000
console.log('Sidebar count', count)
console.debug('Sidebar count', count)
if (currentMode === 'video') {
const output = await shared.runAction(
'getUserVideos',
targetWidth,
256,
count,
offset,
currentSort,
@@ -230,17 +499,108 @@ const getUrls = async (subfolder) => {
const output = await shared.runAction(
'getUserImages',
currentMode,
targetWidth,
count,
offset,
currentSort,
false,
subfolder,
saltUrls,
)
return output || {}
}
const build_ui = async (el) => {
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}`)
if (!userDirectories) {
userDirectories = {
input: allModes.input_root,
output: allModes.output_root,
}
infoLogger('User directories', userDirectories)
}
// const urls = await getUrls()
// const urls = await debouncedGetUrls(subfolder)
const cont = makeElement('div.mtb_sidebar')
contextMenu = new ContextMenu(cont)
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)
debouncedGetUrls(subfolder)
// if (urls) {
// loaded_images = 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)
// const urls = debouncedGetUrls(subfolder)
// const urls = await getUrls(subfolder)
debouncedGetUrls(subfolder)
// if (urls) {
// loaded_images = getImgsFromUrls(urls, imgGrid)
// }
}
})
const sizeSlider = makeSlider(64, 1024, currentWidth, 1)
imgTools.appendChild(orderSelect)
imgTools.appendChild(sizeSlider)
loaded_images = getImgsFromUrls(urls, imgGrid)
// infoLogger({ loaded_images })
sizeSlider.addEventListener('input', (e) => {
currentWidth = e.target.value
for (const img of loaded_images) {
img.style.width = `${e.target.value}px`
}
})
handle = renderSidebar(el, cont, [selector, imgGrid, imgTools])
}
//NOTE: do not load if using the old ui
if (window?.__COMFYUI_FRONTEND_VERSION__) {
// NOTE: removed this for now since I'm not actually exposing anything a client
@@ -249,110 +609,55 @@ if (window?.__COMFYUI_FRONTEND_VERSION__) {
const sidebar_extension = {
name: 'mtb.io-sidebar',
settings: [
{
// init: async () => {
// try {
// const res = await api.fetchApi('/mtb/server-info')
// const msg = await res.json()
// exposed = msg.exposed
// } catch (e) {
// console.error('Error:', e)
// }
// },
init: () => {
let handle
// const version = window?.__COMFYUI_FRONTEND_VERSION__
// console.log(`%c ${version}`, 'background: orange; color: white;')
ensureMTBStyles()
app.ui.settings.addSetting({
id: 'mtb.io-sidebar.count',
category: ['mtb', 'Input & Output Sidebar', 'count'],
name: 'Number of images to fetch',
type: 'number',
defaultValue: 1000,
tooltip:
"This setting affects the input/output sidebar to determine how many images to fetch per pagination (pagination is not yet supported so for now it's the static total)",
},
{
id: 'mtb.io-sidebar.salt_urls',
category: ['mtb', 'Input & Output Sidebar', 'salt_urls'],
name: 'Salt URLs',
type: 'boolean',
defaultValue: false,
onChange: (n, o) => {
saltUrls = n
attrs: {
style: {
// fontFamily: 'monospace',
},
},
tooltip:
'Adds a random query parameter to every urls to always invalidate caching.',
},
{
})
app.ui.settings.addSetting({
id: 'mtb.io-sidebar.img-size',
category: ['mtb', 'Input & Output Sidebar', 'img-size'],
name: 'Resize width of shown images',
name: 'Resolution of the images',
type: 'number',
defaultValue: 512,
type: (name, setter, value, attrs) => {
targetWidth = value
const container = mtb_ui.makeElement('div', {
display: 'flex',
alignItems: 'center',
gap: '8px',
})
console.log({ name, setter, value, attrs })
const baseId = name.replace(/[^a-zA-Z0-9]/g, '-').toLowerCase()
const checkboxId = `${baseId}-checkbox`
const numberInputId = `${baseId}-number`
const isCheckedInitially = value !== -1
// TODO: better way to get defaultValue?
const defaultValue = 512
const initialNumberValue = isCheckedInitially ? value : defaultValue
console.log('recreate')
const checkbox = mtb_ui.makeElement(
// harder to match styles (.p-toggleswitch-input)
// since it uses a div synced to the input...
'input',
{},
container,
)
checkbox.type = 'checkbox'
checkbox.id = checkboxId
checkbox.checked = isCheckedInitially
const numberInput = mtb_ui.makeElement(
'input.p-inputtext',
{},
container,
)
numberInput.type = 'number'
numberInput.id = numberInputId
numberInput.value = initialNumberValue
numberInput.disabled = !isCheckedInitially
numberInput.min = 128
checkbox.addEventListener('change', () => {
let valToSet = -1
if (checkbox.checked) {
numberInput.disabled = false
valToSet = Number.parseInt(numberInput.value, 10)
if (Number.isNaN(valToSet) || valToSet < numberInput.min) {
valToSet = defaultValue
numberInput.value = valToSet
}
} else {
numberInput.disabled = true
}
setter(valToSet)
})
numberInput.addEventListener('input', () => {
if (checkbox.checked) {
const numValue = Number.parseInt(numberInput.value, 10)
if (!Number.isNaN(numValue) && numberInput.value !== '') {
setter(numValue)
}
}
})
return container
tooltip: "It's recommended to keep it at 512px",
attrs: {
style: {
// fontFamily: 'monospace',
},
},
tooltip:
"If browsing large folders it's recommended to use this to avoid overflow/crash of the webpage. Image will get resized to this target width on the server before being sent to the client.",
},
{
})
app.ui.settings.addSetting({
id: 'mtb.io-sidebar.sort',
category: ['mtb', 'Input & Output Sidebar', 'sort'],
name: 'Default sort mode',
@@ -372,39 +677,7 @@ if (window?.__COMFYUI_FRONTEND_VERSION__) {
'Name',
'Name-Reverse',
],
},
{
id: 'mtb.io-sidebar.notice',
category: ['mtb', 'Input & Output Sidebar', 'sort'],
name: ' ',
type: (name, setter, value, attrs) => {
const container = mtb_ui.makeElement('div')
const notice =
'## Important\nIf you make **any** edits here you need to toggle off and back on the sidebar for it to take effect.'
if (window.MTB?.mdParser) {
MTB.mdParser.parse(notice).then((e) => {
container.innerHTML = e
})
} else {
shared.ensureMarkdownParser((p) => {
p.parse(notice).then((e) => {
container.innerHTML = e
})
})
}
return container
},
},
],
init: () => {
let handle
const version = window?.__COMFYUI_FRONTEND_VERSION__
console.log(`%c ${version}`, 'background: orange; color: white;')
ensureMTBStyles()
})
app.extensionManager.registerSidebarTab({
id: 'mtb-inputs-outputs',
@@ -420,81 +693,9 @@ if (window?.__COMFYUI_FRONTEND_VERSION__) {
handle = undefined
}
if (el.parentNode) {
el.parentNode.style.overflowY = 'clip'
if (!loaded_images) {
await build_ui(el)
}
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])
app.api.addEventListener('status', async () => {
if (currentMode !== 'output') return
updateOutputsGrid()
+28
View File
@@ -0,0 +1,28 @@
// NOTE: this will be the LT part of mtb API system
// I need to properly publish the source and fix a few things before
// import { app } from '../../scripts/app.js'
// // import { api } from '../../scripts/api.js'
//
// import * as shared from './comfy_shared.js'
// import { createOutliner } from './dist/mtb_inspector.js'
//
// if (window?.__COMFYUI_FRONTEND_VERSION__) {
// const version = window?.__COMFYUI_FRONTEND_VERSION__
// console.log(`%c ${version}`, 'background: orange; color: white;')
//
// const panel = app.extensionManager.registerSidebarTab({
// id: 'mtb-nodes',
// icon: 'pi pi-bolt',
// title: 'MTB',
// tooltip: 'MTB: API outliner',
// type: 'custom',
// // this is run everytime the tab's diplay is toggled on.
// render: (el) => {
// const outliner = createOutliner(el)
// const inputs = shared.getAPIInputs()
// console.log('INPUTS', inputs)
// outliner.$$set({ inputs })
// },
// })
// }
+86 -4
View File
@@ -174,16 +174,101 @@ export const ensureMTBStyles = () => {
.mtb_slider[type="range"]:active::-webkit-slider-thumb {
background-color: ${S.accent};
}
`
const contextMenus = `
.mtb_context_menu {
position: fixed;
background: var(--comfy-input-bg);
border: 1px solid var(--border-color);
border-radius: 4px;
padding: 4px 0;
min-width: 150px;
z-index: 1000;
box-shadow: 0 2px 5px rgba(0,0,0,0.2);
}
.mtb-context-menu-item {
padding: 6px 12px;
cursor: pointer;
display: flex;
align-items: center;
gap: 8px;
}
.mtb-context-menu-item:hover {
background: var(--comfy-input-hover);
}
.mtb-loading-indicator {
position: sticky;
bottom: 0;
left: 0;
right: 0;
background: var(--comfy-input-bg);
padding: 8px;
text-align: center;
border-top: 1px solid var(--border-color);
z-index: 100;
}
`
addNamedStyleSheet(
'mtb_ui',
`
${common}
${inputs}
${contextMenus}
`,
)
}
export class ContextMenu {
constructor(parent) {
this.menu = makeElement('div.mtb_context_menu', { display: 'none' })
const body = parent || document.body
body.appendChild(this.menu)
document.addEventListener('click', (e) => {
if (!this.menu.contains(e.target)) {
this.hide()
}
})
}
show(x, y, items, context) {
this.menu.innerHTML = ''
for (const item of items) {
const menuItem = makeElement('div.mtb-context-menu-item')
if (item.icon) {
const icon = makeElement(`i.${item.icon}`)
menuItem.appendChild(icon)
}
menuItem.appendChild(document.createTextNode(item.label))
menuItem.onclick = () => {
item.action(context)
this.hide()
}
this.menu.appendChild(menuItem)
}
this.menu.style.display = 'block'
const rect = this.menu.getBoundingClientRect()
const viewportWidth = window.innerWidth
const viewportHeight = window.innerHeight
x = Math.min(x, viewportWidth - rect.width)
y = Math.min(y, viewportHeight - rect.height)
this.menu.style.left = `${x}px`
this.menu.style.top = `${y}px`
}
hide() {
this.menu.style.display = 'none'
}
}
/**
* Wrap an element with a div
*
@@ -203,7 +288,7 @@ export const wrapElement = (element, style = {}) => {
* @param {Object} [style] - CSS styles to apply to the element.
* @returns {HTMLElement} - The created DOM element.
*/
export const makeElement = (kind, style, parent) => {
export const makeElement = (kind, style) => {
let [real_kind, className] = kind.split('.')
let id
@@ -224,9 +309,6 @@ export const makeElement = (kind, style, parent) => {
if (id) {
el.id = id
}
if (parent) {
parent.appendChild(el)
}
return el
}
+9 -49
View File
@@ -21,7 +21,7 @@ import { infoLogger } from './comfy_shared.js'
import { NumberInputWidget } from './numberInput.js'
// NOTE: new widget types registered by MTB Widgets
const newTypes = [/*'BOOL'*/ 'COLOR','MTB_COLOR', 'BBOX']
const newTypes = [/*'BOOL'*/ 'COLOR', 'BBOX']
const deprecated_nodes = {
// 'Animation Builder':
@@ -694,7 +694,7 @@ const mtb_widgets = {
app.ui.settings.addSetting({
id: 'mtb.Main.debug-enabled',
category: ['mtb', ' Main', 'debug-enabled'],
category: ['mtb', 'Main', 'debug-enabled'],
name: 'Enable Debug (py and js)',
type: 'boolean',
defaultValue: false,
@@ -739,6 +739,7 @@ const mtb_widgets = {
// },
COLOR: (node, inputName, inputData, _app) => {
console.debug('Registering color')
return {
widget: node.addCustomWidget(
MtbWidgets.COLOR(inputName, inputData[1]?.default || '#ff0000'),
@@ -747,16 +748,6 @@ const mtb_widgets = {
minHeight: 30,
}
},
MTB_COLOR: (node, inputName, inputData, _app) => {
return {
widget: node.addCustomWidget(
MtbWidgets.COLOR(inputName, inputData[1]?.default || '#ff0000'),
),
minWidth: 150,
minHeight: 30,
}
},
// BBOX: (node, inputName, inputData, app) => {
// console.debug("Registering bbox")
// return {
@@ -1021,15 +1012,12 @@ const mtb_widgets = {
)
loop_preview.value = 'Iteration: Idle'
let cancelQueue = false
const onReset = () => {
raw_iteration.value = 0
raw_loop.value = 0
value_preview.value = 'Idle'
loop_preview.value = 'Iteration: Idle'
cancelQueue = false
app.canvas.setDirty(true)
}
@@ -1038,43 +1026,15 @@ const mtb_widgets = {
this.addWidget('button', 'Reset', 'reset', onReset)
// run button
const chunkSize = 10
this.addWidget('button', 'Queue', 'queue', async () => {
onReset()
const totalPrompts = total_frames.value * loop_count.value
this.addWidget('button', 'Queue', 'queue', () => {
onReset() // this could maybe be a setting or checkbox
app.queuePrompt(0, total_frames.value * loop_count.value)
window.MTB?.notify?.(
`Starting a queue of ${totalPrompts} frames in chunks of ${chunkSize}...`,
`Started a queue of ${total_frames.value} frames (for ${
loop_count.value
} loop, so ${total_frames.value * loop_count.value})`,
5000,
)
for (let i = 0; i < totalPrompts; i += chunkSize) {
console.log({ cancelQueue })
if (cancelQueue) {
window.MTB?.notify?.(
`Queueing cancelled after ${i} frames.`,
3000,
)
break
}
const currentChunkSize = Math.min(chunkSize, totalPrompts - i)
await app.queuePrompt(0, currentChunkSize)
}
if (!cancelQueue) {
window.MTB?.notify?.(
`Finished queuing ${totalPrompts} frames.`,
5000,
)
}
})
this.addWidget('button', 'Cancel', 'cancel', () => {
cancelQueue = true
window.MTB?.notify?.(
'Cancellation requested. Waiting for current chunk to finish...',
3000,
)
})
this.onRemoved = () => {
+52 -64
View File
@@ -1,13 +1,10 @@
// web/note_plus.constants.js
export const DEFAULT_CSS = `/** here you can write css**/
h1 {
color: whitesmoke;
}`
export const DEFAULT_CSS = ''
export const DEFAULT_HTML = `<p style='color:red;font-family:monospace'>
Note+
</p>`
export const DEFAULT_MD = '# 📝 Note+'
export const DEFAULT_MD = '## Note+'
export const DEFAULT_MODE = 'markdown'
export const DEFAULT_THEME = 'one_dark'
@@ -58,57 +55,58 @@ We also support github callout:
`
export const THEMES = [
'ambiance',
'chaos',
'chrome',
'cloud9_day',
'cloud9_night',
'cloud9_night_low_color',
'cloud_editor',
'cloud_editor_dark',
'clouds',
'clouds_midnight',
'cobalt',
'crimson_editor',
'dawn',
'dracula',
'dreamweaver',
'eclipse',
'github',
'github_dark',
'gob',
'gruvbox',
'gruvbox_dark_hard',
'gruvbox_light_hard',
'idle_fingers',
'iplastic',
'katzenmilch',
'kr_theme',
'kuroir',
'merbivore',
'merbivore_soft',
'mono_industrial',
'monokai',
'nord_dark',
'one_dark',
'pastel_on_dark',
'solarized_dark',
'solarized_light',
'sqlserver',
'terminal',
'textmate',
'tomorrow',
'tomorrow_night',
'tomorrow_night_blue',
'tomorrow_night_bright',
'tomorrow_night_eighties',
'twilight',
'vibrant_ink',
'vscode',
'ambiance',
'chaos',
'chrome',
'cloud9_day',
'cloud9_night',
'cloud9_night_low_color',
'cloud_editor',
'cloud_editor_dark',
'clouds',
'clouds_midnight',
'cobalt',
'crimson_editor',
'dawn',
'dracula',
'dreamweaver',
'eclipse',
'github',
'github_dark',
'gob',
'gruvbox',
'gruvbox_dark_hard',
'gruvbox_light_hard',
'idle_fingers',
'iplastic',
'katzenmilch',
'kr_theme',
'kuroir',
'merbivore',
'merbivore_soft',
'mono_industrial',
'monokai',
'nord_dark',
'one_dark',
'pastel_on_dark',
'solarized_dark',
'solarized_light',
'sqlserver',
'terminal',
'textmate',
'tomorrow',
'tomorrow_night',
'tomorrow_night_blue',
'tomorrow_night_bright',
'tomorrow_night_eighties',
'twilight',
'vibrant_ink',
'vscode',
]
export const CSS_RESET = `
* {
font-family: monospace;
line-height: 1.25em;
}
.shiki{
@@ -118,8 +116,6 @@ export const CSS_RESET = `
.markdown-callout-title {
.octicon{
fill:white;
width:29px;
height:29px;
}
/* background: var(--current-color); */
color: var(--current-color);
@@ -128,8 +124,6 @@ export const CSS_RESET = `
/* border-start-start-radius: var(--radius); */
padding: 0.5em;
padding-inline-start: 1em;
display: flex;
align-items: center;
}
.markdown-callout-content {
padding: 1em;
@@ -142,12 +136,7 @@ export const CSS_RESET = `
border-left: 3px solid var(--current-color);
margin-bottom: 1em;
margin-top: 1em;
}
.markdown-callout p:nth-child(2) {
padding:1em;
}
.markdown-callout-tip {
--text-color: whitesmoke;
@@ -175,9 +164,8 @@ export const CSS_RESET = `
flex-direction:column;
align-items: flex-start;
width:95%;
/*margin-left: 20px;*/
/*margin-top:20px;*/
margin-left: 20px;
margin-top:20px;
/*background-color: rgba(255,0,0,0.5)!important;*/
}
+334 -377
View File
File diff suppressed because it is too large Load Diff
+3 -12
View File
@@ -41,16 +41,7 @@ const toastStyle = `
transition-duration: ${transition_time}ms;
`
function notify(message, timeout = 3000, old_mode = false) {
if (!old_mode) {
app.extensionManager.toast.add({
severity: 'info',
summary: 'MTB',
detail: message,
life: timeout,
})
return
}
function notify(message, timeout = 3000) {
log('Creating toast')
const container = document.getElementById('mtb-notify-container')
const toast = document.createElement('div')
@@ -68,7 +59,7 @@ function notify(message, timeout = 3000, old_mode = false) {
log('Transition out')
const totalHeight = Array.from(container.children).reduce(
(acc, child) => acc + child.offsetHeight + 10, // Add spacing of 10px between toasts
0,
0
)
container.style.height = `${totalHeight}px`
@@ -92,7 +83,7 @@ function notify(message, timeout = 3000, old_mode = false) {
// Update container's height to fit new toast
const totalHeight = Array.from(container.children).reduce(
(acc, child) => acc + child.offsetHeight + 10, // Add spacing of 10px between toasts
0,
0
)
container.style.height = `${totalHeight}px`