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 2417 additions and 4525 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 name: 📦 Publish to Comfy registry
on: on:
workflow_dispatch: workflow_dispatch:
push:
tags:
- '*'
permissions: permissions:
issues: write issues: write
@@ -18,5 +21,4 @@ jobs:
- name: 📦 Publish Custom Node - name: 📦 Publish Custom Node
uses: Comfy-Org/publish-node-action@v1 uses: Comfy-Org/publish-node-action@v1
with: with:
skip_checkout: 'true'
personal_access_token: ${{ secrets.COMFY_REGISTRY_TOKEN }} personal_access_token: ${{ secrets.COMFY_REGISTRY_TOKEN }}
-5
View File
@@ -1,16 +1,11 @@
__pycache__ __pycache__
*.py[cod] *.py[cod]
*.onnx *.onnx
wheels/ wheels/
node_modules/ node_modules/
compose.yaml compose.yaml
comfy_mtb.wsb comfy_mtb.wsb
Dockerfile Dockerfile
.DS_Store
node.zip
# I store the gh-pages worktrees (src & build) there # I store the gh-pages worktrees (src & build) there
.worktrees .worktrees
comfy.lock
+167 -77
View File
@@ -3,13 +3,15 @@
# File: __init__.py # File: __init__.py
# Project: comfy_mtb # Project: comfy_mtb
# Author: Mel Massadian # Author: Mel Massadian
# Copyright (c) 2023-2025 Mel Massadian # Copyright (c) 2023 Mel Massadian
# #
### ###
__version__ = "0.6.0" __version__ = "0.3.0"
import os import os
from collections import OrderedDict
from typing import Any
from aiohttp.web_request import Request from aiohttp.web_request import Request
@@ -34,8 +36,6 @@ from aiohttp import web
IN_COMFY = False IN_COMFY = False
PromptServer = None
try: try:
from server import PromptServer from server import PromptServer
@@ -77,7 +77,7 @@ def extract_nodes_from_source(filename: Path):
) )
break break
except SyntaxError: except SyntaxError:
log.error(f"Failed to parse ast from: {filename}") log.error("Failed to parse")
return nodes return nodes
@@ -242,37 +242,14 @@ if failed:
# - ENDPOINT # - ENDPOINT
# TODO: move that away and simplify existing endpoints if IN_COMFY and hasattr(PromptServer, "instance"):
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
img_cache = None img_cache = None
prompt_cache = None prompt_cache = None
import asyncio
import os
from io import BytesIO
from PIL import Image
from .repl import setup_custom_web_routes
setup_custom_web_routes(PromptServer.instance.app)
with contextlib.suppress(ImportError): with contextlib.suppress(ImportError):
from cachetools import TTLCache from cachetools import TTLCache
img_cache = TTLCache(maxsize=100, ttl=5) # 1 min TTL # img_cache = TTLCache(maxsize=100, ttl=5) # 1 min TTL
prompt_cache = TTLCache(maxsize=100, ttl=5) # 1 min TTL prompt_cache = TTLCache(maxsize=100, ttl=5) # 1 min TTL
node_dependency_mapping = get_node_dependencies() node_dependency_mapping = get_node_dependencies()
@@ -385,24 +362,134 @@ def register_routes():
# Return JSON for other requests # Return JSON for other requests
return web.json_response({"message": "Welcome to MTB!"}) 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): def get_cached_image(file_path: str, preview_params=None, channel=None):
cache_key = (file_path, preview_params, channel) cache_key = (file_path, preview_params, channel)
if img_cache and (cache_key in img_cache): try:
return img_cache[cache_key]
with Image.open(file_path) as img:
info = img.info
if preview_params:
img = process_preview(img, preview_params)
if channel:
img = process_channel(img, channel)
if prompt_cache:
prompt_cache[cache_key] = info
if img_cache: if img_cache:
img_cache[cache_key] = img.getvalue() cached_value = img_cache.get(cache_key)
return img_cache[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): def process_preview(img: Image.Image, preview_params):
image_format, quality, width = preview_params image_format, quality, width = preview_params
@@ -455,41 +542,44 @@ def register_routes():
# to load workflows in the sidebar # to load workflows in the sidebar
@PromptServer.instance.routes.get("/mtb/view") @PromptServer.instance.routes.get("/mtb/view")
async def view_image(request: Request): async def view_image(request: Request):
import folder_paths try:
import folder_paths
filename = request.rel_url.query.get("filename") await rate_limiter.acquire()
if not filename:
return web.Response(status=404)
filename, output_dir = folder_paths.annotated_filepath(filename) filename = request.rel_url.query.get("filename")
if filename[0] == "/" or ".." in filename: if not filename:
return web.Response(status=400) return web.Response(status=404)
if output_dir is None: filename, output_dir = folder_paths.annotated_filepath(filename)
rtype = request.rel_url.query.get("type", "output") if filename[0] == "/" or ".." in filename:
output_dir = folder_paths.get_directory_by_type(rtype) return web.Response(status=400)
if output_dir is None: if output_dir is None:
return web.Response(status=400) rtype = request.rel_url.query.get("type", "output")
output_dir = folder_paths.get_directory_by_type(rtype)
if "subfolder" in request.rel_url.query: if output_dir is None:
full_output_dir = os.path.join( return web.Response(status=400)
output_dir, request.rel_url.query["subfolder"]
) if "subfolder" in request.rel_url.query:
if ( full_output_dir = os.path.join(
os.path.commonpath( output_dir, request.rel_url.query["subfolder"]
(os.path.abspath(full_output_dir), output_dir)
) )
!= output_dir if (
): os.path.commonpath(
return web.Response(status=403) (os.path.abspath(full_output_dir), output_dir)
output_dir = full_output_dir )
!= output_dir
):
return web.Response(status=403)
output_dir = full_output_dir
filename = os.path.basename(filename) filename = os.path.basename(filename)
file = os.path.join(output_dir, filename) file = os.path.join(output_dir, filename)
if not os.path.isfile(file): if not os.path.isfile(file):
return web.Response(status=404) return web.Response(status=404)
ret_workflow = request.rel_url.query.get("workflow") ret_workflow = request.rel_url.query.get("workflow")
@@ -527,9 +617,13 @@ def register_routes():
width = request.rel_url.query.get("width") width = request.rel_url.query.get("width")
preview_info = (image_format, quality, 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") @PromptServer.instance.routes.get("/mtb/server-info")
async def get_debug(request: Request): async def get_debug(request: Request):
@@ -589,10 +683,6 @@ def register_routes():
return await endpoint.do_action(request) return await endpoint.do_action(request)
if IN_COMFY and hasattr(PromptServer, "instance"):
register_routes()
# - WAS Dictionary # - WAS Dictionary
MANIFEST = { MANIFEST = {
"name": "MTB Nodes", # The title that will be displayed on Node Class menu,. and Node Class view "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", "$schema": "https://biomejs.dev/schemas/1.6.1/schema.json",
"assist": { "actions": { "source": { "organizeImports": "on" } } }, "organizeImports": {
"enabled": true
},
"linter": { "linter": {
"enabled": true, "enabled": true,
"rules": { "rules": {
"recommended": true, "recommended": true,
"suspicious": { "suspicious": {
"noConsole": { "level": "warn", "options": { "allow": ["log"] } } "noConsoleLog": "warn"
}, },
"style": { "style": {
"noParameterAssign": "off", "noParameterAssign": "off",
"noShoutyConstants": "warn", "noShoutyConstants": "warn",
"useNamingConvention": "off", "useNamingConvention": "off"
"useAsConstAssertion": "error",
"useDefaultParameterLast": "error",
"useEnumInitializers": "error",
"useSelfClosingElements": "error",
"useSingleVarDeclarator": "error",
"noUnusedTemplateLiteral": "error",
"useNumberNamespace": "error",
"noInferrableTypes": "error",
"noUselessElse": "error"
} }
} }
}, },
+12 -10
View File
@@ -15,6 +15,7 @@ from .utils import (
backup_file, backup_file,
build_glob_patterns, build_glob_patterns,
glob_multiple, glob_multiple,
import_install,
reqs_map, reqs_map,
run_command, run_command,
styles_dir, styles_dir,
@@ -23,6 +24,7 @@ from .utils import (
endlog = mklog("mtb endpoint") endlog = mklog("mtb endpoint")
# - ACTIONS # - ACTIONS
import_install("requirements")
def ACTIONS_installDependency(dependency_names: list[str] | None = None): 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()] 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()] 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( def ACTIONS_getUserVideos(
@@ -110,15 +117,11 @@ def ACTIONS_getUserVideos(
def ACTIONS_getUserImages( def ACTIONS_getUserImages(
mode: Literal["input", "output"], mode: Literal["input", "output"],
target_width: int | str | None = None,
count=1000, count=1000,
offset=0, offset=0,
sort: str | None = None, sort: str | None = None,
include_subfolders: bool = False, include_subfolders: bool = False,
subfolder: str | None = None, subfolder=None,
# IIRC I copied this from Comfy base
# just keeping it until I properly checked implications
salt_urls=False,
): ):
# enabled = "MTB_EXPOSE" in os.environ # enabled = "MTB_EXPOSE" in os.environ
# if not enabled: # if not enabled:
@@ -126,12 +129,11 @@ def ACTIONS_getUserImages(
imgs = {} imgs = {}
count = count or 1000 count = count or 1000
target_width = int(target_width) if target_width else None
input_dir = Path(folder_paths.get_input_directory()) input_dir = Path(folder_paths.get_input_directory())
output_dir = Path(folder_paths.get_output_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: if subfolder:
entry_dir = entry_dir / subfolder entry_dir = entry_dir / subfolder
@@ -160,9 +162,9 @@ def ACTIONS_getUserImages(
imgs = { imgs = {
img.name: ( 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"{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) for i, img in enumerate(entries)
if offset <= i < offset + count if offset <= i < offset + count
+104 -211
View File
@@ -1,175 +1,85 @@
# NOTE: This file is only use for development you can ignore it # NOTE: This file is only use for development you can ignore it
use log.nu use private/log.nu
use nssm.nu *
use nutils.nu [ make-id upsert-all fwd-slash backup-file ]
use os.nu [ link ]
# --- utilities --- def get_root [--clean] {
def get_root [ --clean] { if $clean {
if $clean { $env.COMFY_CLEAN_ROOT
$env.COMFY.ROOTS.clean } else {
} else { $env.COMFY_ROOT
$env.COMFY.ROOTS.main }
}
} }
def --env path-add [pth] { export def "comfy build-web" [] {
$env.PATH = ($env.PATH | append ($pth | path expand)) cd $env.COMFY_MTB
cd web_source
npm run build
cp dist/*.js ../web/dist
}
export def "comfy dev-web" [] {
cd $env.COMFY_MTB
cd web_source
npm run dev
}
export def "daily run" [] {
let res = (comfy update --rebase)
comfy update --clean
comfy update_extensions
daily commit $res.from $res.to
} }
def short-date [] { def short-date [] {
format date "%Y-%m-%d" format date "%Y-%m-%d"
} }
export def spawn-for [timeout: duration task: closure] {
let input = $in
let parent_id = job id
let task_id = job spawn {
$input | do $task | job send --tag (job id) $parent_id
}
try {
job recv --tag $task_id --timeout $timeout
} catch {
job kill $task_id
error make {
msg: "Task timed out."
label: {
text: "timed out"
span: (metadata $task).span
}
}
}
}
# --- exports --
export def "comfy profile" [timeout = 60sec] {
let to_match = "To see the GUI go to"
pyinstrument -r html main.py ...($env.COMFY.ARGS)
| tee -e {
each {
let stde = $in
print -ne $stde
if $to_match in $stde {
print $"(ansi gb)Profiling Done!(ansi reset)"
let process = (ps -l | where name =~ python | where command =~ pyinstrument | last)
kill -f $process.pid
}
}
}
| complete
| get stdout
| save $"profiled_(date now | format date '%s').html"
}
export def "comfy profile-plus" [] {
let timestamp = (date now | format date "%s")
let log_name = $"cprofile_run_($timestamp)"
let profiled = (python -m cProfile main.py --port 3000 --preview-method auto | tee -e { print -ne } | complete)
let out = (
$profiled.stdout
| lines
# skip summary
| skip 4
| str join "\n"
)
# save result
$out | save $"raw_($log_name).txt"
# process
$out
| from ssv
| upsert-all { into float } tottime percall cumtime
| save $"($log_name).nuon"
}
export def restart-server [] {
nssm restart -c comfy
}
# build the web components of mtb
export def "comfy build-web" [] {
cd $env.COMFY.ROOTS.mtb
if ("./web/dist" | path exists) {
rm -rt ./web/dist
}
cd web_source
^$env.NPM_BINARY run build
cp -r dist ../web/dist
}
# start the dev server for web components
export def "comfy dev-web" [] {
cd $env.COMFY.ROOTS.mtb
cd web_source
^$env.NPM_BINARY run dev
}
# daily check / update
export def "daily run" [] {
let res = (comfy update --rebase)
comfy update --clean
comfy update_extensions
daily commit $res.from_commit $res.to_commit
}
# was daily run today? # was daily run today?
export def "daily was-run" [] { export def "daily was-run" [] {
let daily = ($env.COMFY.ROOTS.mtb | path join daily.nuon) let daily = ($env.COMFY_MTB | path join daily.nuon)
if ($daily | path exists) { if ($daily | path exists) {
let last = (open $daily | sort-by date | get date | last | short-date) let last = (open $daily | sort-by date | get date | last | short-date)
let today = (date now | short-date) let today = (date now | short-date)
return ($last == $today) return ($last == $today)
} }
return false return false
} }
export def "daily commit" [from_commit: string to_commit: string] { export def "daily commit" [from:string, to:string] {
let daily = ($env.COMFY.ROOTS.mtb | path join daily.nuon) let daily = ($env.COMFY_MTB | path join daily.nuon)
let commit = [{date: (date now) from_commit: $from_commit to_commit: $to_commit}] let commit = [{date: (date now) from:$from to:$to}]
let dailies = ( let dailies = (if ($daily | path exists) {
if ($daily | path exists) { open $daily | append $commit
open $daily | append $commit } else {
} else {
$commit $commit
} })
)
$dailies | save -f $daily $dailies | save -f $daily
log success "Commited daily check" log success "Commited daily check"
} }
# start the comfy server # start the comfy server
export def "comfy start" [ export def "comfy start" [--clean,--old-ui, --listen, --skip-daily(-s)] {
--clean if (not (daily was-run)) and not $skip_daily {
--old-ui log info "Running daily checks"
--listen daily run
--skip-daily (-s) }
] { let root = get_root --clean=($clean)
if not (daily was-run) and not $skip_daily { cd $root
log info "Running daily checks"
daily run
}
let root = (get_root --clean=$clean)
cd $root
log info "Running Server" log info "Running Server"
MTB_DEBUG=true python main.py --port 3000 ...(if $old_ui { ["--front-end-version" "Comfy-Org/ComfyUI_legacy_frontend@latest"] } else { [--front-end-version Comfy-Org/ComfyUI_frontend@latest] }) --preview-method auto ...(if $listen { ["--listen"] } else { [] }) MTB_DEBUG=true python main.py --port 3000 ...(if $old_ui { ["--front-end-version", "Comfy-Org/ComfyUI_legacy_frontend@latest"]} else {[ --front-end-version Comfy-Org/ComfyUI_frontend@latest]}) --preview-method auto ...(if $listen {["--listen"]} else {[]})
} }
# update comfy itself and merge master in current branch # update comfy itself and merge master in current branch
export def "comfy update" [ export def "comfy update" [
--clean # comfy clean instance --clean # ??
--rebase # Rebase instead of merge --rebase # Rebase instead of merge
] { ] {
let root = get_root --clean=$clean let root = get_root --clean=$clean
@@ -184,28 +94,21 @@ export def "comfy update" [
log info "Backing up and removing models symlinks" log info "Backing up and removing models symlinks"
# preparing root for pull # preparing root for pull
let pyproject = if not $clean { if not $clean {
log info "Backing up the pyproject.toml..."
let proj = (backup-file --root pyproject.toml)
log info "Restoring the original pyproject"
git checkout pyproject.toml git checkout pyproject.toml
cd $models cd $models
# find and store all symlinks # find and store all symlinks
log info "Checking for links in models..." let links = (ls -la |
let links = ( where not ($it.target | is-empty) |
ls -la | where not ($it.target | is-empty) | select name target | sort-by name select name target |
) sort-by name)
log info $"Found links: ($links)"
if not ($links | is-empty) { if not ($links | is-empty) {
log info "Backing up the symlinks..."
backup-file --root links.nuon
$links | save -f links.nuon $links | save -f links.nuon
# remove them # remove them
open links.nuon | each {|p| rm $p.name } open links.nuon | each {|p| rm $p.name }
} }
$proj
} else { } else {
# just remove symlinks # just remove symlinks
rm $models rm $models
@@ -238,6 +141,7 @@ export def "comfy update" [
if $rebase { if $rebase {
log info "Rebasing changes" log info "Rebasing changes"
git rebase master git rebase master
} else { } else {
log info "Merging changes" log info "Merging changes"
git merge master git merge master
@@ -248,10 +152,9 @@ export def "comfy update" [
if not $clean { if not $clean {
rm pyproject.toml rm pyproject.toml
log info "Using our own pyproject..." cp pyproject-mel.toml pyproject.toml
cp $pyproject pyproject.toml
cd $models cd $models
log info "Relinking models..."
# resymlink them # resymlink them
open links.nuon | each {|p| link -a $p.target $p.name } open links.nuon | each {|p| link -a $p.target $p.name }
} else { } else {
@@ -264,82 +167,72 @@ export def "comfy update" [
log success $"Update successful \(($commit_count) new commits\)" log success $"Update successful \(($commit_count) new commits\)"
return {from_commit: $current_commit to_commit: $new_commit} return {from:$current_commit to:$new_commit}
} }
export def "comfy toggle_extensions" [ export def "comfy toggle_extensions" [--clean] {
--clean let root = get_root --clean=($clean)
] { cd $root
let root = get_root --clean=$clean cd custom_nodes
cd $root let exts = (ls | where type in ["dir","symlink"] | get name)
cd custom_nodes let choices = ($exts | input list -m "choose extension to toggle")
let exts = (ls | where type in ["dir" "symlink"] | get name) if ($choices | is-empty) {
let choices = ($exts | input list -m "choose extension to toggle") return
if ($choices | is-empty) { }
return
}
log info "Choices" $choices log info "Choices" $choices
let filtered = $choices | wrap name | upsert enabled {|p| not ($p.name | str ends-with ".disabled") } let filtered = $choices | wrap name | upsert enabled {|p| not ($p.name | str ends-with ".disabled")}
log info "Filtered" $filtered log info "Filtered" $filtered
$filtered | each {|f| $filtered | each {|f|
let new_name = ($f.name | str replace ".disabled" "") let new_name = ($f.name | str replace ".disabled" "")
let new_name = if $f.enabled { let new_name = if $f.enabled {
$"($new_name).disabled" $"($new_name).disabled"
} else { } else {
$new_name $new_name
}
log info $"Moving ($f.name) to ($new_name)"
mv $f.name $new_name
} }
log info $"Moving ($f.name) to ($new_name)"
mv $f.name $new_name
}
} }
# git pull all extensions # git pull all extensions
export def "comfy update_extensions" [ --clean] { export def "comfy update_extensions" [--clean] {
let root = get_root --clean=$clean let root = get_root --clean=($clean)
cd $root cd $root
cd custom_nodes cd custom_nodes
git multipull . -s -q git multipull . -s -q
} }
# manual set version of mtb def --env path-add [pth] {
export def "comfy-mtb set-version" [version: string] { $env.PATH = ($env.PATH | append ($pth | path expand))
# let pyproject = open pyproject.toml
# let current_version = $pyproject.project.version
# $pyproject | upsert project.version $version | save -f pyproject.toml
# taplo format pyproject.toml
sd "(__version__ = )\"(.*)\"" $"${1}\"($version)\"" __init__.py
sd "(version = )(.*)" $"${1}\"($version)\"" pyproject.toml
# log info $"⬆️ Bump version: ($current_version) → ($version)"
} }
# -- env
export-env { export-env {
$env.PYTHONUTF8 = 1 $env.PYTHONUTF8 = 1
$env.COMFY = { $env.COMFY_MTB = ("." | path expand)
base_url : "https://mel-pc.tail3c8eb.ts.net" # $env.CUDA_ROOT = 'C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.1\'
ARGS: [--port 3000 --preview-method auto]
ROOTS: {
mtb: ("." | path expand | fwd-slash)
main: ("../.." | path expand | fwd-slash)
clean: ($env.COMFY_ROOT | path dirname | path join ComfyClean | fwd-slash)
}
}
$env.NPM_BINARY = "bun"
$env.CUDA_HOME = $env.CUDA_ROOT $env.CUDA_HOME = $env.CUDA_ROOT
#
path-add 'C:/Portable/TensorRT-8.6.0.12/lib'
#
if $nu.os-info.family == 'windows' {
path-add "G:/BIN/TensorRT-10.7.0.23/lib"
path-add "G:/BIN/cudnn-windows-x86_64-9.6.0.74_cuda12-archive/bin"
}
#
path-add ($env.CUDA_ROOT | path join bin)
overlay use "../../.venv/Scripts/activate.nu" $env.COMFY_ROOT = ("../.." | path expand)
$env.COMFY_CLEAN_ROOT = ($env.COMFY_ROOT | path dirname | path join ComfyClean)
path-add 'C:/Portable/TensorRT-8.6.0.12/lib'
if $nu.os-info.family == 'windows' {
path-add 'G:\BIN\TensorRT-10.7.0.23\lib'
path-add 'G:\BIN\cudnn-windows-x86_64-9.6.0.74_cuda12-archive\bin'
}
path-add ($env.CUDA_ROOT | path join bin)
overlay use ../../.venv/Scripts/activate.nu
} }
-3
View File
@@ -1,3 +0,0 @@
{
"use_repl": false
}
+1
View File
@@ -43,6 +43,7 @@ pip_map = {
"tb-nightly": "tensorboard", "tb-nightly": "tensorboard",
"protobuf": "google.protobuf", "protobuf": "google.protobuf",
"qrcode[pil]": "qrcode", "qrcode[pil]": "qrcode",
"requirements-parser": "requirements",
# Add more mappings as needed # Add more mappings as needed
} }
+13 -14
View File
@@ -1,16 +1,20 @@
from typing import TYPE_CHECKING, Any, TypedDict from typing import Any, TypedDict
import torch import torch
import torchaudio import torchaudio
from comfy.model_management import get_torch_device from comfy.model_management import get_torch_device
from huggingface_hub import snapshot_download from huggingface_hub import snapshot_download
from transformers import (
WhisperForConditionalGeneration,
WhisperProcessor,
)
if TYPE_CHECKING: # from transformers import (
from transformers import ( # AutoFeatureExtractor,
WhisperForConditionalGeneration, # WhisperForConditionalGeneration,
WhisperProcessor, # WhisperModel,
) # WhisperProcessor,
# )
from ..log import log from ..log import log
from ..utils import get_model_path from ..utils import get_model_path
@@ -97,8 +101,8 @@ class MtbAudio:
class WhisperPipeline(TypedDict): class WhisperPipeline(TypedDict):
"""Whisper model pipeline.""" """Whisper model pipeline."""
processor: "WhisperProcessor" processor: WhisperProcessor
model: "WhisperForConditionalGeneration" model: WhisperForConditionalGeneration
class MTB_LoadWhisper: class MTB_LoadWhisper:
@@ -144,11 +148,6 @@ class MTB_LoadWhisper:
def load(self, model_size="tiny", download_missing=False): def load(self, model_size="tiny", download_missing=False):
"""Load Whisper model and processor.""" """Load Whisper model and processor."""
from transformers import (
WhisperForConditionalGeneration,
WhisperProcessor,
)
whisper_dir = get_model_path("whisper") whisper_dir = get_model_path("whisper")
tag = f"whisper-{model_size}" tag = f"whisper-{model_size}"
model_dir = whisper_dir / tag model_dir = whisper_dir / tag
-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 torch
import torchvision.transforms.functional as TF from PIL import Image, ImageDraw, ImageFilter
from ..log import log from ..log import log
from ..utils import np2tensor, pil2tensor, tensor2np, tensor2pil
class BoundingBox(NamedTuple):
"""The bounding box tuple."""
x: int
y: int
width: int
height: int
class MTB_Bbox: class MTB_Bbox:
"""A literal bounding box.""" """The bounding box (BBOX) custom type used by other nodes"""
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
@@ -46,14 +37,12 @@ class MTB_Bbox:
FUNCTION = "do_crop" FUNCTION = "do_crop"
CATEGORY = "mtb/crop" CATEGORY = "mtb/crop"
def do_crop( def do_crop(self, x: int, y: int, width: int, height: int): # bbox
self, x: int, y: int, width: int, height: int return ((x, y, width, height),)
) -> tuple[BoundingBox]: # bbox
return (BoundingBox(x, y, width, height),)
class MTB_SplitBbox: class MTB_SplitBbox:
"""Split the components of a bbox.""" """Split the components of a bbox"""
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
@@ -66,8 +55,8 @@ class MTB_SplitBbox:
RETURN_TYPES = ("INT", "INT", "INT", "INT") RETURN_TYPES = ("INT", "INT", "INT", "INT")
RETURN_NAMES = ("x", "y", "width", "height") RETURN_NAMES = ("x", "y", "width", "height")
def split_bbox(self, bbox: BoundingBox) -> BoundingBox: def split_bbox(self, bbox):
return bbox return (bbox[0], bbox[1], bbox[2], bbox[3])
class MTB_UpscaleBboxBy: class MTB_UpscaleBboxBy:
@@ -85,23 +74,26 @@ class MTB_UpscaleBboxBy:
FUNCTION = "upscale" 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 x, y, width, height = bbox
center_x = x + width / 2 center_x = x + width // 2
center_y = y + height / 2 center_y = y + height // 2
new_width = int(width * scale) new_width = int(width * scale)
new_height = int(height * scale) new_height = int(height * scale)
new_x = int(center_x - new_width / 2) new_x = center_x - new_width // 2
new_y = int(center_y - new_height / 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: class MTB_BboxFromMask:
"""From a mask extract the bounding box.""" """From a mask extract the bounding box"""
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
@@ -111,7 +103,7 @@ class MTB_BboxFromMask:
"invert": ("BOOLEAN", {"default": False}), "invert": ("BOOLEAN", {"default": False}),
}, },
"optional": { "optional": {
"image": ("IMAGE", {"tooltip": "Optional image"}), "image": ("IMAGE",),
}, },
} }
@@ -127,44 +119,52 @@ class MTB_BboxFromMask:
CATEGORY = "mtb/crop" CATEGORY = "mtb/crop"
def extract_bounding_box( def extract_bounding_box(
self, self, mask: torch.Tensor, invert: bool, image=None
mask: torch.Tensor, ):
*, # if image != None:
invert: bool = False, # if mask.size(0) != image.size(0):
image: torch.Tensor | None = None, # if mask.size(0) != 1:
) -> tuple[BoundingBox, torch.Tensor | None]: # log.error(
mask = 1 - mask if invert else mask # 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})"
non_zero_indices = torch.nonzero(mask) # )
if non_zero_indices.numel() == 0: # raise Exception(
log.warning( # 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})"
"BboxFromMask: Mask is empty. Returning a (0,0,0,0) bbox." # )
)
return (BoundingBox(0, 0, 0, 0), image)
min_coords = torch.min(non_zero_indices, dim=0).values # we invert it
max_coords = torch.max(non_zero_indices, dim=0).values _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() non_zero_indices = np.nonzero(alpha_channel)
max_y, max_x = max_coords[1].item(), max_coords[2].item()
width = max_x - min_x + 1 min_x, max_x = np.min(non_zero_indices[1]), np.max(non_zero_indices[1])
height = max_y - min_y + 1 min_y, max_y = np.min(non_zero_indices[0]), np.max(non_zero_indices[0])
bounding_box = BoundingBox( # Create a bounding box tuple
int(min_x), int(min_y), int(width), int(height) 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: 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 The BBOX input takes precedence over the tuple input
""" """
@@ -204,38 +204,35 @@ class MTB_Crop:
def do_crop( def do_crop(
self, self,
image: torch.Tensor, image: torch.Tensor,
*, mask=None,
mask: torch.Tensor | None = None, x=0,
x: int = 0, y=0,
y: int = 0, width=256,
width: int = 256, height=256,
height: int = 256, bbox=None,
bbox: BoundingBox | None = None,
): ):
image = image.numpy()
if mask is not None:
mask = mask.numpy()
if bbox is not None: if bbox is not None:
x, y, width, height = bbox 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_image = image[:, y : y + height, x : x + width, :]
cropped_mask = ( cropped_mask = None
mask[:, y : y + height, x : x + width] if mask is not None:
if mask is not None cropped_mask = (
else None mask[:, y : y + height, x : x + width]
) if mask is not None
crop_data = BoundingBox(x, y, width, height) else None
)
crop_data = (x, y, width, height)
return ( return (
cropped_image, torch.from_numpy(cropped_image),
cropped_mask if cropped_mask is not None else None, torch.from_numpy(cropped_mask)
if cropped_mask is not None
else None,
crop_data, crop_data,
) )
@@ -249,33 +246,35 @@ class MTB_Crop:
# return (x_left, y_top, x_right, y_bottom) # 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: if not target_size:
return bbox return bbox
new_bbox = BoundingBox( new_bbox = (
bbox.x, bbox[0],
bbox.y, bbox[1],
min(target_size[0] - bbox.x, bbox.width), min(target_size[0] - bbox[0], bbox[2]),
min(target_size[1] - bbox.y, bbox.height), min(target_size[1] - bbox[1], bbox[3]),
) )
if new_bbox != bbox: 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 return new_bbox
def bbox_to_region( def bbox_to_region(bbox, target_size=None):
bbox: BoundingBox, target_size: tuple[int, int] | None = None
):
bbox = bbox_check(bbox, target_size) bbox = bbox_check(bbox, target_size)
# to region # 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: 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 @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
@@ -292,113 +291,91 @@ class MTB_Uncrop:
} }
RETURN_TYPES = ("IMAGE",) RETURN_TYPES = ("IMAGE",)
FUNCTION = "do_uncrop" FUNCTION = "do_crop"
CATEGORY = "mtb/crop" CATEGORY = "mtb/crop"
def do_uncrop( def do_crop(self, image, crop_image, bbox, border_blending):
self, def inset_border(image, border_width=20, border_color=(0)):
image: torch.Tensor, width, height = image.size
crop_image: torch.Tensor, bordered_image = Image.new(
bbox: BoundingBox, image.mode, (width, height), border_color
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."
) )
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: crop_img = crop.convert("RGB")
image = image.repeat(len(crop_image), 1, 1, 1)
batch_size, bg_h, bg_w, _ = image.shape log.debug(f"Crop image size: {crop_img.size}")
_, fg_h, fg_w, _ = crop_image.shape log.debug(f"Image size: {img.size}")
x, y, width, height = bbox
if (width, height) != (fg_w, fg_h): if border_blending > 1.0:
log.warning( border_blending = 1.0
f"Uncrop: crop_image size {(fg_w, fg_h)} " elif border_blending < 0.0:
"differs from bbox {(width, height)}. Resizing to fit bbox." 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) blend.putalpha(mask)
resized_crop = torch.nn.functional.interpolate( img = Image.alpha_composite(img.convert("RGBA"), blend)
resized_crop, out_images.append(img.convert("RGB"))
size=(height, width),
mode="bicubic",
align_corners=False,
)
resized_crop = resized_crop.permute(0, 2, 3, 1)
pbar.update(1) return (pil2tensor(out_images),)
# 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),)
class MTB_BBoxForceDimensions: class MTB_BBoxForceDimensions:
"""
Resize a BBOX to new dimensions while keeping its center.
Optionally constrains the BBOX to stay within image boundaries.
"""
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
return { return {
@@ -406,7 +383,6 @@ class MTB_BBoxForceDimensions:
"bbox": ("BBOX",), "bbox": ("BBOX",),
"width": ("INT", {"default": 512, "min": 1, "max": 8192}), "width": ("INT", {"default": 512, "min": 1, "max": 8192}),
"height": ("INT", {"default": 512, "min": 1, "max": 8192}), "height": ("INT", {"default": 512, "min": 1, "max": 8192}),
"constrain_to_image": ("BOOLEAN", {"default": True}),
}, },
"optional": { "optional": {
"image": ("IMAGE",), "image": ("IMAGE",),
@@ -419,12 +395,10 @@ class MTB_BBoxForceDimensions:
def force_dimensions( def force_dimensions(
self, self,
*,
bbox: tuple[int, int, int, int], bbox: tuple[int, int, int, int],
width: int, width: int,
height: int, height: int,
constrain_to_image: bool = True, image: torch.Tensor = None,
image: torch.Tensor | None = None,
) -> tuple[tuple[int, int, int, int]]: ) -> tuple[tuple[int, int, int, int]]:
x, y, curr_width, curr_height = bbox x, y, curr_width, curr_height = bbox
@@ -434,14 +408,27 @@ class MTB_BBoxForceDimensions:
new_x = center_x - width // 2 new_x = center_x - width // 2
new_y = center_y - height // 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] img_height, img_width = image.shape[1:3]
new_x = max(0, min(new_x, img_width - width)) x_overflow = max(0, new_x + width - img_width) + min(0, new_x)
new_y = max(0, min(new_y, img_height - height)) y_overflow = max(0, new_y + height - img_height) + min(0, new_y)
width = min(width, img_width) if width > img_width or height > img_height:
height = min(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__ = [ __nodes__ = [
+179 -605
View File
@@ -1,70 +1,33 @@
import base64 import base64
import io import io
import textwrap import json
from collections.abc import Callable from pathlib import Path
from functools import wraps
from typing import Any, Literal, Protocol, TypedDict, runtime_checkable
import folder_paths
import torch import torch
from rich import inspect
from rich.console import Console
from ..log import log from ..log import log
from ..utils import LazyProxyTensor, get_torch_tensor_info, tensor2pil from ..utils import tensor2pil
try:
import matplotlib.pyplot as plt
import numpy as np
plt.style.use("dark_background")
MATPLOTLIB_AVAILABLE = True
except ImportError:
MATPLOTLIB_AVAILABLE = False
# region Decorator def get_detailed_type_info(obj):
def metadata(**meta_kwargs: Any) -> Callable[[Any], Any]: type_info = []
"""Add metadata to method (`__meta__` dict)."""
def decorator(func: Callable[[Any], Any]) -> Callable[[Any], Any]:
@wraps(func)
def wrapper(*args, **kwargs):
return func(*args, **kwargs)
wrapper.__meta__ = meta_kwargs
return wrapper
return decorator
# endregion
class UIResult(TypedDict):
kind: Literal["text", "b64_images"]
data: str
def indent_results(results: list[UIResult], by: str = " "):
for res in results:
if res["kind"] == "text":
log.debug(f"Indenting: {res['data']}")
res["data"] = textwrap.indent(res["data"], by)
return results
ProcessorResult = list[UIResult]
def _get_detailed_type_info(obj) -> str:
type_info: list[str] = []
type_name = type(obj).__name__ type_name = type(obj).__name__
type_info.append(f"Type: {type_name}") type_info.append(f"Type: {type_name}")
if isinstance(obj, torch.Tensor): if isinstance(obj, torch.Tensor):
return get_torch_tensor_info(obj) type_info.extend(
[
elif isinstance(obj, list | tuple): f"Shape: {obj.shape}",
f"Dtype: {obj.dtype}",
f"Device: {obj.device}",
f"Requires grad: {obj.requires_grad}",
f"Stride: {obj.stride()}",
f"Contiguous: {obj.is_contiguous()}",
]
)
elif isinstance(obj, (list, tuple)):
type_info.extend( type_info.extend(
[ [
f"Length: {len(obj)}", f"Length: {len(obj)}",
@@ -84,184 +47,122 @@ def _get_detailed_type_info(obj) -> str:
attributes = [attr for attr in dir(obj) if not attr.startswith("_")] attributes = [attr for attr in dir(obj) if not attr.startswith("_")]
type_info.append(f"Attributes: {attributes}") type_info.append(f"Attributes: {attributes}")
return "\n".join(type_info) return type_info
def _apply_rich_results(processed, mode="none", title=""):
processing_text = False
acc = ""
reshaped: list[UIResult] = []
for i in range(len(processed)):
if processed[i]["kind"] == "text":
if not processing_text:
processing_text = True
acc += processed[i]["data"] + "\n"
if len(processed) == (i + 1):
reshaped.append(
UIResult(
kind="text", data=_apply_rich(acc, mode, title=title)
)
)
else:
if processing_text:
processing_text = False
reshaped.append(
UIResult(
kind="text", data=_apply_rich(acc, mode, title=title)
)
)
acc = ""
reshaped.append(processed[i])
return reshaped
# for item in processed:
# region processors # region processors
def _apply_rich( def process_tensor(tensor: torch.Tensor, as_type=False):
formatted: str | list[str], rich_mode: str | None = None, *, title="" log.debug(f"Tensor: {tensor.shape}")
) -> str:
if rich_mode is None: if as_type:
return ( return {
formatted if isinstance(formatted, str) else "\n".join(formatted) "text": [f"Tensor of shape {tensor.shape} of type {tensor.dtype}"]
}
is_mask = len(tensor.shape) == 3
if is_mask:
tensor = tensor.unsqueeze(-1).repeat(1, 1, 1, 3)
image = tensor2pil(tensor)
b64_imgs = []
for im in image:
if is_mask:
im = im.convert("L")
buffered = io.BytesIO()
im.save(buffered, format="PNG")
b64_imgs.append(
"data:image/png;base64,"
+ base64.b64encode(buffered.getvalue()).decode("utf-8")
) )
from rich.console import Console return {"b64_images": b64_imgs}
console = Console(record=True)
if isinstance(formatted, list): def process_list(anything, as_type=False):
for line in formatted: text = []
console.print(line) if not anything:
return {"text": []}
if as_type:
type_info = get_detailed_type_info(anything)
type_info.extend(get_detailed_type_info(anything[0]))
return {"text": type_info}
first_element = anything[0]
if (
isinstance(first_element, list)
and first_element
and isinstance(first_element[0], torch.Tensor)
):
text.append(
"List of List of Tensors: "
f"{first_element[0].shape} (x{len(anything)})"
)
elif isinstance(first_element, torch.Tensor):
text.append(
f"List of Tensors: {first_element.shape} (x{len(anything)})"
)
else: else:
console.print(formatted) text.append(f"Array ({len(anything)}): {anything}")
CSV_CODE_FORMAT = """ return {"text": text}
<svg class="rich-terminal" viewBox="0 0 {width} {height}" xmlns="http://www.w3.org/2000/svg">
<!-- Generated with Rich https://www.textualize.io -->
<style>
@font-face {{
font-family: "Fira Code";
src: local("FiraCode-Regular"),
url("https://cdnjs.cloudflare.com/ajax/libs/firacode/6.2.0/woff2/FiraCode-Regular.woff2") format("woff2"),
url("https://cdnjs.cloudflare.com/ajax/libs/firacode/6.2.0/woff/FiraCode-Regular.woff") format("woff");
font-style: normal;
font-weight: 400;
}}
@font-face {{
font-family: "Fira Code";
src: local("FiraCode-Bold"),
url("https://cdnjs.cloudflare.com/ajax/libs/firacode/6.2.0/woff2/FiraCode-Bold.woff2") format("woff2"),
url("https://cdnjs.cloudflare.com/ajax/libs/firacode/6.2.0/woff/FiraCode-Bold.woff") format("woff");
font-style: bold;
font-weight: 700;
}}
.{unique_id}-matrix {{ def process_dict(anything, as_type=False):
font-family: Fira Code, monospace; text = []
font-size: {char_height}px; if as_type:
line-height: {line_height}px; return {"text": get_detailed_type_info(anything)}
font-variant-east-asian: full-width;
}}
.{unique_id}-title {{ if "samples" in anything:
font-size: 18px; is_empty = (
font-weight: bold; "(empty)" if torch.count_nonzero(anything["samples"]) == 0 else ""
font-family: arial; )
}} text.append(f"Latent Samples: {anything['samples'].shape} {is_empty}")
{styles} elif "waveform" in anything:
</style> is_empty = (
"(empty) " if torch.count_nonzero(anything["samples"]) == 0 else ""
<defs>
<clipPath id="{unique_id}-clip-terminal">
<rect x="0" y="0" width="{terminal_width}" height="{terminal_height}" />
</clipPath>
{lines}
</defs>
{chrome}
<g clip-path="url(#{unique_id}-clip-terminal)">
{backgrounds}
<g class="{unique_id}-matrix">
{matrix}
</g>
</g>
</svg>
"""
if rich_mode == "svg-window":
return console.export_svg(title=title, code_format=CSV_CODE_FORMAT)
elif rich_mode == "svg":
return console.export_svg(
title=title,
code_format=CSV_CODE_FORMAT.replace("{chrome}", ""),
) )
elif rich_mode == "html": text.append(
CONSOLE_HTML_FORMAT = textwrap.dedent(""" f"Audio Samples: {anything['waveform'].shape}{is_empty} | sample rate {anything['sample_rate']}"
<div style="color:{foreground};">
<code style="font-family:inherit">{code}</code>
</div>
""").strip()
import rich.terminal_theme
return console.export_html(
inline_styles=True,
code_format=CONSOLE_HTML_FORMAT,
theme=rich.terminal_theme.MONOKAI,
) )
log.error(f"Unknown rich mode: {rich_mode}") else:
return formatted if isinstance(formatted, str) else "\n".join(formatted) log.debug(f"Unhandled dict: {anything.keys()}")
text.append(json.dumps(anything, indent=2))
return {"text": text}
def process_bool(anything, as_type=False):
return {"text": ["True" if anything else "False"]}
def process_text(anything, as_type=False):
if as_type:
return {"text": get_detailed_type_info(anything)}
return {"text": [str(anything)]}
# endregion # endregion
# region conditions
# those are pretty dumb there is now probably a better way..
def is_condition(item):
return (
isinstance(item, list)
and all(isinstance(i, list) for i in item)
and isinstance(item[0][0], torch.Tensor)
)
# endregion
RICH_MODE = Literal["none", "html", "svg", "svg-window"]
@runtime_checkable
class Processor(Protocol):
"""Generic protocol for processor functions."""
def __call__(
self, item: Any, *, as_type: bool = False, deep: bool = False
) -> ProcessorResult: ...
class MTB_Debug: class MTB_Debug:
"""A debug node.""" """Experimental node to debug any Comfy values.
support for more types and widgets is planned.
"""
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
return { return {
"required": {"output_to_console": ("BOOLEAN", {"default": False})}, "required": {"output_to_console": ("BOOLEAN", {"default": False})},
"optional": { "optional": {"as_detailed_types": ("BOOLEAN", {"default": False})},
"as_detailed_types": ("BOOLEAN", {"default": False}),
"deep_inspect": ("BOOLEAN", {"default": False}),
"rich_mode": (
("none", "html", "svg", "svg-window"),
{"default": "none"},
),
},
} }
RETURN_TYPES = () RETURN_TYPES = ()
@@ -269,426 +170,99 @@ class MTB_Debug:
CATEGORY = "mtb/debug" CATEGORY = "mtb/debug"
OUTPUT_NODE = True OUTPUT_NODE = True
_processors: dict[type, Processor]
def __init__(self):
self._condition_processors = {is_condition: self._process_condition}
self._class_name_processors = {
"CLIP": self._process_clip,
"VAE": self._process_vae,
}
self._processors = {
torch.nn.Module: self._process_module,
torch.Tensor: self._process_tensor,
LazyProxyTensor: self._process_repr,
list: self._process_container,
tuple: self._process_container,
dict: self._process_dict,
bool: self._process_bool,
str: self._process_primitive,
int: self._process_primitive,
float: self._process_primitive,
type(None): self._process_primitive,
}
# - Dispatchers ------------------------------------------------------------
def _dispatch_processor(
self, item: Any, *, as_type=False, deep=False
) -> ProcessorResult:
"""Find and calls the appropriate processor for the given item."""
# first conditions
for c, process in self._condition_processors.items():
if c(item):
return process(item, as_type=as_type, deep=deep)
# named class
class_name = type(item).__name__
if class_name in self._class_name_processors:
return self._class_name_processors[class_name](
item, as_type=as_type, deep=deep
)
# type based or unknown
processor = self._processors.get(type(item), self._process_unknown)
res = processor(item, as_type=as_type, deep=deep)
return res
def do_debug( def do_debug(
self, self, output_to_console: bool, as_detailed_types: bool, **kwargs
**kwargs,
): ):
output = {"ui": {"items": []}} output = {"ui": {"items": []}}
settings = {k: kwargs.pop(k) for k in self.INPUT_TYPES()["optional"]} if output_to_console:
output_to_console = kwargs.pop("output_to_console") for k, v in kwargs.items():
as_type = settings.get("as_detailed_types", False) log.info(f"{k}: {v}")
deep = settings.get("deep_inspect", False)
rich_mode = settings.get("rich_mode", "none")
for input_name, item in kwargs.items(): for input_name, anything in kwargs.items():
processed = self._dispatch_processor( processor = processors.get(type(anything), process_text)
item, as_type=as_type, deep=deep
)
if processed is None:
continue
if rich_mode != "none": processed = processor(anything, as_detailed_types)
title = f"{input_name} ({type(item).__name__})"
processed = _apply_rich_results(processed, rich_mode, title)
if output_to_console: item = {
log.info(f"- Input '{input_name}':") "input": input_name,
for p in processed: **processed,
if p["kind"] == "text": }
log.info(f" {p['data']}") output["ui"]["items"].append(item)
if p["kind"] == "b64_image":
log.info(f" (contains {len(p['data'])} images)")
output["ui"]["items"].append(
{"input": input_name, "items": processed}
)
return output return output
def _process_unknown(
self, item: Any, *, as_type=False, deep=False
) -> ProcessorResult:
console = Console(
record=True,
width=120,
)
console.print(f"Generic {type(item).__name__}", emoji=True) class MTB_SaveTensors:
if as_type: """Save torch tensors (image, mask or latent) to disk.
inspect(item, console=console, all=deep, methods=deep, docs=deep)
else:
console.print(item, emoji=True)
text_output = console.export_text(clear=True) useful to debug things outside comfy.
"""
return [UIResult(kind="text", data=text_output.strip())] def __init__(self):
self.output_dir = folder_paths.get_output_directory()
self.type = "mtb/debug"
def _process_repr( @classmethod
self, item: Any, as_type=False, deep=False def INPUT_TYPES(cls):
) -> ProcessorResult: return {
return [{"kind": "text", "data": item.__repr__()}] "required": {
"filename_prefix": ("STRING", {"default": "ComfyPickle"}),
},
"optional": {
"image": ("IMAGE",),
"mask": ("MASK",),
"latent": ("LATENT",),
},
}
def _process_primitive( FUNCTION = "save"
self, item: Any, *, as_type=False, deep=False OUTPUT_NODE = True
) -> ProcessorResult: RETURN_TYPES = ()
if as_type: CATEGORY = "mtb/debug"
return self._process_unknown(item, as_type=as_type, deep=deep)
return [UIResult(kind="text", data=str(item))] def save(
self,
filename_prefix,
image: torch.Tensor | None = None,
mask: torch.Tensor | None = None,
latent: torch.Tensor | None = None,
):
(
full_output_folder,
filename,
counter,
subfolder,
filename_prefix,
) = folder_paths.get_save_image_path(filename_prefix, self.output_dir)
full_output_folder = Path(full_output_folder)
if image is not None:
image_file = f"{filename}_image_{counter:05}.pt"
torch.save(image, full_output_folder / image_file)
# np.save(full_output_folder/ image_file, image.cpu().numpy())
def _process_bool( if mask is not None:
self, item: bool, *, as_type=False, deep=False mask_file = f"{filename}_mask_{counter:05}.pt"
) -> ProcessorResult: # noqa: FBT001 torch.save(mask, full_output_folder / mask_file)
return [{"kind": "text", "data": "True" if item else "False"}] # np.save(full_output_folder/ mask_file, mask.cpu().numpy())
def _process_clip( if latent is not None:
self, item: Any, *, as_type=False, deep=False # for latent we must use pickle
) -> ProcessorResult: latent_file = f"{filename}_latent_{counter:05}.pt"
try: torch.save(latent, full_output_folder / latent_file)
clip_model = getattr(item, "cond_stage_model", None) # pickle.dump(latent, open(full_output_folder/ latent_file, "wb"))
tokenizer = getattr(item, "tokenizer", None)
text = [UIResult(kind="text", data="CLIP")] # np.save(full_output_folder / latent_file,
if clip_model: # latent[""].cpu().numpy())
text.append(UIResult(kind="text", data="CLIP Model:"))
model_summary = self._process_module(
clip_model, as_type=as_type
)
if model_summary:
text.extend(indent_results(model_summary, " "))
else:
text.append(
UIResult(
kind="text",
data="[error] failed to get informations about clip model",
)
)
if tokenizer: return f"{filename_prefix}_{counter:05}"
text.append(UIResult(kind="text", data="Tokenizer:"))
vocab_size = getattr(tokenizer, "vocab_size", "N/A")
text.append(
UIResult(
kind="text",
data=f" Class: {type(tokenizer).__name__}\n Vocab Size: {vocab_size}",
)
)
return text
except Exception as e:
log.error(f"Failed to process CLIP object: {e}")
return self._process_unknown(item, as_type=as_type, deep=deep)
def _process_condition(
self, item: Any, *, as_type=False, deep=False
) -> ProcessorResult:
count = len(item)
result = [UIResult(kind="text", data=f"Conditions: {count}")]
for cond in item:
result.extend(self._preview_conditioning_tensor(cond[0]))
return result
def _process_vae(
self, item: Any, *, as_type=False, deep=False
) -> ProcessorResult:
try:
vae_model = getattr(
item, "first_stage_model", getattr(item, "vae", item)
)
text = [
UIResult(kind="text", data="VAE"),
UIResult(kind="text", data="Internal Model:"),
]
model_summary = self._process_module(
vae_model, as_type=as_type, deep=deep
)
text.extend(indent_results(model_summary, " "))
return text
except Exception as e:
log.error(f"Failed to process VAE object: {e}")
return self._process_unknown(item, as_type=as_type, deep=deep)
def _process_module(
self, item: torch.nn.Module, *, as_type=False, deep=False
) -> ProcessorResult:
if as_type and deep:
return self._process_unknown(item, as_type=as_type, deep=deep)
total_params = sum(p.numel() for p in item.parameters())
trainable_params = sum(
p.numel() for p in item.parameters() if p.requires_grad
)
try:
device = next(item.parameters()).device
except StopIteration:
device = "cpu (no parameters)"
train_percent = (
f"{trainable_params / total_params:.2%}"
if total_params > 0
else "0.00%"
)
text = [
f"Model: {type(item).__name__} on {device}",
textwrap.dedent(f"""
- Parameters: {total_params:,}
- Trainable: {trainable_params:,} ({train_percent})
""").strip(),
]
return [{"kind": "text", "data": d} for d in text]
def _process_tensor(
self, item: torch.Tensor, *, as_type=False, deep=False
) -> ProcessorResult:
is_latent = item.ndim == 4 and item.shape[1] == 4
is_image = (
not is_latent and item.ndim == 4 and item.shape[3] in [1, 3, 4]
)
is_conditioning = item.ndim == 3 and item.shape[2] in [
768,
1024,
1152,
1280,
2048,
4096,
]
is_mask = (item.ndim == 2) or (item.ndim == 3 and not is_conditioning)
if as_type:
type_name = "Unknown Tensor"
if is_latent:
type_name = "Latent Tensor"
elif is_image:
type_name = "Image Tensor"
elif is_conditioning:
type_name = "CLIP Conditioning Tensor"
elif is_mask:
type_name = "Mask Tensor"
return [
{
"kind": "text",
"data": get_torch_tensor_info(item, name=type_name),
}
]
if is_image or is_mask:
return self._render_image_tensor(item)
if is_latent:
return self._preview_latent_tensor(item)
if is_conditioning:
return self._preview_conditioning_tensor(item)
return self._process_unknown(item, as_type=as_type, deep=deep)
def _visualize_tensor_heatmap(
self, tensor_2d: torch.Tensor, title: str
) -> str | None:
if not MATPLOTLIB_AVAILABLE:
log.warning("Matplotlib not found. Skipping tensor visualization.")
return None
if tensor_2d.ndim != 2:
log.warning(
f"Cannot visualize tensor with {tensor_2d.ndim} dimensions. Requires 2."
)
return None
fig, ax = plt.subplots(figsize=(6, 4), dpi=100)
im = ax.imshow(tensor_2d.cpu().numpy(), cmap="viridis", aspect="auto")
fig.colorbar(im, ax=ax)
ax.set_title(title)
fig.tight_layout()
buf = io.BytesIO()
fig.savefig(buf, format="png", bbox_inches="tight", pad_inches=0.1)
plt.close(fig)
buf.seek(0)
return "data:image/png;base64," + base64.b64encode(buf.read()).decode(
"utf-8"
)
def _render_image_tensor(self, item: torch.Tensor) -> ProcessorResult:
is_mask = (item.ndim == 2) or (item.ndim == 3 and item.shape[-1] != 3)
img_tensor = (
item.unsqueeze(0) if item.ndim == 3 and not is_mask else item
)
img_tensor = item.unsqueeze(0) if item.ndim == 2 else img_tensor
images = tensor2pil(img_tensor)
b64_imgs = []
for im in images:
if is_mask:
im = im.convert("L")
buffered = io.BytesIO()
im.save(buffered, format="PNG")
b64_imgs.append(
"data:image/png;base64,"
+ base64.b64encode(buffered.getvalue()).decode("utf-8")
)
return [UIResult(kind="b64_images", data=b64_imgs)]
def _preview_latent_tensor(self, item: torch.Tensor) -> ProcessorResult:
is_empty = "(empty)" if torch.count_nonzero(item) == 0 else ""
stats = [
f"Min: {item.min():.4f}",
f"Max: {item.max():.4f}",
f"Mean: {item.mean():.4f}",
]
text = [
get_torch_tensor_info(item, name="Latent Tensor"),
is_empty,
] + stats
result = [UIResult(kind="text", data=t) for t in text]
vis_tensor = item[0].mean(dim=0)
heatmap_b64 = self._visualize_tensor_heatmap(
vis_tensor, "Latent Energy (Channel Mean)"
)
if heatmap_b64:
result.append(UIResult(kind="b64_images", data=[heatmap_b64]))
return result
def _preview_conditioning_tensor(
self, item: torch.Tensor
) -> ProcessorResult:
_batch, tokens, embed_dim = item.shape
text = [
get_torch_tensor_info(item, name="CLIP Conditioning Tensor"),
f"Token Count: {tokens}",
f"Embedding Dim: {embed_dim}",
]
result = [UIResult(kind="text", data=d) for d in text]
heatmap_b64 = self._visualize_tensor_heatmap(
item[0], "Token Embeddings (approx)"
)
if heatmap_b64:
result.append(UIResult(kind="b64_images", data=[heatmap_b64]))
return result
def _process_container(
self, item: list | tuple, *, as_type=False, deep=False
) -> ProcessorResult:
if not item:
return [UIResult(kind="text", data=f"Empty {type(item).__name__}")]
container_type = type(item).__name__
element_type = type(item[0]).__name__
all_match = all(type(i) is type(item[0]) for i in item)
result = [
UIResult(
kind="text",
data=f"{container_type} of {len(item)} x {element_type}",
),
UIResult(kind="text", data=f"(mixed types: {not all_match})"),
]
if not as_type or (as_type and deep):
for i, sub_item in enumerate(item):
res = self._dispatch_processor(
sub_item, as_type=as_type, deep=deep
)
if res:
text = res[0].get("data", "Unknown")
res[0]["data"] = f"[{i}]: {text}"
result.extend(res)
return result
first_item_result = self._dispatch_processor(
item[0], as_type=as_type, deep=deep
)
if not first_item_result:
return result
return (
result
+ [UIResult(kind="text", data="Preview of first element:")]
+ indent_results(first_item_result, " - ")
)
def _process_dict(
self, item: dict, *, as_type=False, deep=False
) -> ProcessorResult:
if "pooled_output" in item and isinstance(
item["pooled_output"], torch.Tensor
):
return self._dispatch_processor(
item["pooled_output"], as_type=as_type, deep=deep
)
if "samples" in item and isinstance(item.get("samples"), torch.Tensor):
return self._dispatch_processor(
item["samples"], as_type=as_type, deep=deep
)
if "waveform" in item and isinstance(
item.get("waveform"), torch.Tensor
):
waveform = item["waveform"]
is_empty = "(empty) " if torch.count_nonzero(waveform) == 0 else ""
text = textwrap.dedent(f"""
Audio Waveform: {waveform.shape}{is_empty}
Sample Rate: {item.get("sample_rate", "N/A")}
""").strip()
return [{"kind": "text", "data": text}]
log.debug(
f"Processing generic dict with rich inspector: {item.keys()}"
)
return self._process_unknown(item, as_type=as_type, deep=deep)
__nodes__ = [MTB_Debug] processors = {
torch.Tensor: process_tensor,
list: process_list,
dict: process_dict,
bool: process_bool,
}
__nodes__ = [MTB_Debug, MTB_SaveTensors]
-70
View File
@@ -1,70 +0,0 @@
import folder_paths
import torch
class MTB_SaveTensors:
"""Save torch tensors (image, mask or latent) to disk.
useful to debug things outside comfy.
"""
def __init__(self):
self.output_dir = folder_paths.get_output_directory()
self.type = "mtb/debug"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"filename_prefix": ("STRING", {"default": "ComfyPickle"}),
},
"optional": {
"image": ("IMAGE",),
"mask": ("MASK",),
"latent": ("LATENT",),
},
}
FUNCTION = "save"
OUTPUT_NODE = True
RETURN_TYPES = ()
CATEGORY = "mtb/debug"
def save(
self,
filename_prefix,
image: torch.Tensor | None = None,
mask: torch.Tensor | None = None,
latent: torch.Tensor | None = None,
):
(
full_output_folder,
filename,
counter,
subfolder,
filename_prefix,
) = folder_paths.get_save_image_path(filename_prefix, self.output_dir)
full_output_folder = Path(full_output_folder)
if image is not None:
image_file = f"{filename}_image_{counter:05}.pt"
torch.save(image, full_output_folder / image_file)
# np.save(full_output_folder/ image_file, image.cpu().numpy())
if mask is not None:
mask_file = f"{filename}_mask_{counter:05}.pt"
torch.save(mask, full_output_folder / mask_file)
# np.save(full_output_folder/ mask_file, mask.cpu().numpy())
if latent is not None:
# for latent we must use pickle
latent_file = f"{filename}_latent_{counter:05}.pt"
torch.save(latent, full_output_folder / latent_file)
# pickle.dump(latent, open(full_output_folder/ latent_file, "wb"))
# np.save(full_output_folder / latent_file,
# latent[""].cpu().numpy())
return f"{filename_prefix}_{counter:05}"
__nodes__ = [MTB_SaveTensors]
+4 -7
View File
@@ -4,8 +4,12 @@ import sys
from pathlib import Path from pathlib import Path
import comfy.model_management as model_management import comfy.model_management as model_management
import cv2
import insightface
import numpy as np import numpy as np
import onnxruntime
import torch import torch
from insightface.model_zoo.inswapper import INSwapper
from PIL import Image from PIL import Image
from ..errors import ModelNotFound from ..errors import ModelNotFound
@@ -39,8 +43,6 @@ class MTB_LoadFaceAnalysisModel:
DEPRECATED = True DEPRECATED = True
def load_model(self, faceswap_model: str): def load_model(self, faceswap_model: str):
import insightface
if faceswap_model == "antelopev2": if faceswap_model == "antelopev2":
download_antelopev2() download_antelopev2()
@@ -79,9 +81,6 @@ class MTB_LoadFaceSwapModel:
DEPRECATED = True DEPRECATED = True
def load_model(self, faceswap_model: str): def load_model(self, faceswap_model: str):
import onnxruntime
from insightface.model_zoo.inswapper import INSwapper
model_path = get_model_path("insightface", faceswap_model) model_path = get_model_path("insightface", faceswap_model)
if not model_path or not model_path.exists(): if not model_path or not model_path.exists():
raise ModelNotFound(f"{faceswap_model} ({model_path})") raise ModelNotFound(f"{faceswap_model} ({model_path})")
@@ -213,8 +212,6 @@ def swap_face(
face_swapper_model, face_swapper_model,
faces_index: set[int] | None = None, faces_index: set[int] | None = None,
) -> Image.Image: ) -> Image.Image:
import cv2
if faces_index is None: if faces_index is None:
faces_index = {0} faces_index = {0}
log.debug(f"Swapping faces: {faces_index}") log.debug(f"Swapping faces: {faces_index}")
+9 -23
View File
@@ -7,13 +7,6 @@ from PIL import Image, ImageDraw, ImageFont
from ..log import log from ..log import log
from ..utils import comfy_dir, font_path, pil2tensor from ..utils import comfy_dir, font_path, pil2tensor
# try:
# from cairosvg import svg2png
# HAS_CAIRO = True
# except ImportError:
# HAS_CAIRO = False
# class MtbExamples: # class MtbExamples:
# """MTB Example Images""" # """MTB Example Images"""
@@ -306,7 +299,7 @@ by default it fallsback to a default font.
def text_to_image( def text_to_image(
self, self,
text: str | list[str], text: str,
font, font,
wrap, wrap,
trim, trim,
@@ -348,9 +341,11 @@ by default it fallsback to a default font.
color = (255, 255, 255, 255) color = (255, 255, 255, 255)
background = (0, 0, 0, 255) background = (0, 0, 0, 255)
def render_text(text_to_render: str, alpha=None) -> Image.Image: def render_text(text_to_render, alpha=None):
if trim: if trim:
text_to_render = text_to_render.strip() text_to_render = (
text_to_render.encode("ascii", "ignore").decode().strip()
)
if wrap: if wrap:
wrap_width = (((width / 100) * h_coverage) / font_size) * 2 wrap_width = (((width / 100) * h_coverage) / font_size) * 2
lines = textwrap.wrap(text_to_render, width=wrap_width) lines = textwrap.wrap(text_to_render, width=wrap_width)
@@ -423,9 +418,7 @@ by default it fallsback to a default font.
active_chunks.append((chunk["text"], alpha)) active_chunks.append((chunk["text"], alpha))
for chunk_text, alpha in active_chunks: for chunk_text, alpha in active_chunks:
chunk_img = render_text( chunk_img = render_text(chunk_text, alpha)
chunk_text.encode("ascii", "ignore").decode(), alpha
)
frame = Image.alpha_composite(frame, chunk_img) frame = Image.alpha_composite(frame, chunk_img)
frames.append(frame) frames.append(frame)
@@ -433,16 +426,9 @@ by default it fallsback to a default font.
frame_tensors = [pil2tensor(frame) for frame in frames] frame_tensors = [pil2tensor(frame) for frame in frames]
return (torch.cat(frame_tensors, dim=0),) return (torch.cat(frame_tensors, dim=0),)
else: else:
results = [] text_img = render_text(text)
if not isinstance(text, list): result = Image.alpha_composite(base_img, text_img)
text = [text] return (pil2tensor(result),)
for t in text:
text_img = render_text(t)
result = Image.alpha_composite(base_img, text_img)
results.append(result)
return (pil2tensor(results),)
__nodes__ = [ __nodes__ = [
+7 -168
View File
@@ -4,22 +4,18 @@ import re
import urllib.parse import urllib.parse
import urllib.request import urllib.request
from math import pi from math import pi
from typing import Any
import comfy.model_management as mm import comfy.model_management as model_management
import comfy.utils import comfy.utils
import numpy as np import numpy as np
import torch import torch
from comfy.comfy_types.node_typing import IO as CIO
from PIL import Image from PIL import Image
from ..log import log from ..log import log
from ..utils import ( from ..utils import (
EASINGS, EASINGS,
LazyProxyTensor,
apply_easing, apply_easing,
get_server_info, get_server_info,
get_torch_tensor_info,
numpy_NFOV, numpy_NFOV,
pil2tensor, pil2tensor,
tensor2np, tensor2np,
@@ -136,64 +132,12 @@ class MTB_ApplyTextTemplate:
CATEGORY = "mtb/utils" CATEGORY = "mtb/utils"
FUNCTION = "execute" FUNCTION = "execute"
def execute(self, *, template: str, **kwargs) -> tuple[str | list[str]]: def execute(self, *, template: str, **kwargs):
keys = list(kwargs.keys()) res = f"{template}"
values = list(kwargs.values()) for k, v in kwargs.items():
res = res.replace(f"{{{k}}}", f"{v}")
has_list = any(isinstance(v, list) for v in values) return (res,)
target_length = -1
if has_list:
first_list = next(x for x in values if isinstance(x, list))
# all_list = all(isinstance(x, list) for x in kwargs.values())
# if not all_list:
# raise ValueError(
# "Text template supports either str or list[str] but not a mix of the two (yet?)"
# )
target_length = len(first_list)
same_length = all(
len(v) == target_length for v in values if isinstance(v, list)
)
if not same_length:
raise ValueError(
"Text template received multiple list[str] but their size is varying, they should match..."
)
if has_list:
results = []
# do a padded loop, not the most efficient but easy
# to handle for now
for it in range(target_length):
res = f"{template}"
for k, v in kwargs.items():
if isinstance(v, list):
res = self.apply_res(res, k, v[it])
else:
res = self.apply_res(res, k, v)
results.append(res)
return (results,)
else:
res = f"{template}"
for k, v in kwargs.items():
res = self.apply_res(res, k, v)
return (res,)
def apply_res(self, res, key, value):
if isinstance(value, float):
value = f"{value:.3f}"
elif isinstance(value, torch.Tensor):
value = get_torch_tensor_info(value)
else:
log.debug(
f"Falling back to default string conversion for {key} of type {type(value).__name__}"
)
return res.replace(f"{{{key}}}", f"{value}")
class MTB_MatchDimensions: class MTB_MatchDimensions:
@@ -397,7 +341,7 @@ class MTB_AutoPanEquilateral:
frames.append(frame) frames.append(frame)
mm.throw_exception_if_processing_interrupted() model_management.throw_exception_if_processing_interrupted()
pbar.update(1) pbar.update(1)
return (pil2tensor(frames),) return (pil2tensor(frames),)
@@ -923,108 +867,6 @@ class MTB_TensorOps:
return (result,) 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,)
class MTB_ProxyTensor:
"""Wraps an input tensor into a LazyProxyTensor.
builds upon an idea by @AustinMroz
"""
NODE_NAME = "ProxyTensor"
NODE_DISPLAY_NAME = "Proxy Tensor"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"tensor": ("IMAGE",),
"target_dtype": (
["float32", "float16", "bfloat16"],
{"default": "float32"},
),
"target_device": (
["keep", "cpu", "gpu"],
{
"default": "keep",
"tooltip": "CAUTION: This isn't compatible with most nodes for now",
},
),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("proxy_tensor",)
FUNCTION = "execute"
CATEGORY = "mtb/utils"
def execute(
self,
tensor: torch.Tensor,
target_dtype: str = "float32",
target_device: str = "keep",
):
torch_dtype: torch.dtype = getattr(torch, target_dtype)
if target_device == "gpu":
torch_device = mm.get_torch_device()
elif target_device == "cpu":
torch_device = torch.device("cpu")
else:
torch_device = tensor.device
proxy = LazyProxyTensor(tensor, torch_dtype, torch_device)
log.info(f"Created Proxy Tensor: \n{proxy}")
return (proxy,)
__nodes__ = [ __nodes__ = [
MTB_StringReplace, MTB_StringReplace,
MTB_FitNumber, MTB_FitNumber,
@@ -1040,7 +882,4 @@ __nodes__ = [
MTB_FloatToFloats, MTB_FloatToFloats,
MTB_FloatsToInts, MTB_FloatsToInts,
MTB_TensorOps, MTB_TensorOps,
MTB_BooleanNot,
MTB_GetItem,
MTB_ProxyTensor,
] ]
+45 -101
View File
@@ -3,12 +3,11 @@ import json
import math import math
import os import os
import comfy.utils import comfy.model_management as model_management
import folder_paths import folder_paths
import numpy as np import numpy as np
import torch import torch
import torch.nn.functional as F import torch.nn.functional as F
from comfy import model_management
from PIL import Image, ImageOps from PIL import Image, ImageOps
from PIL.PngImagePlugin import PngInfo from PIL.PngImagePlugin import PngInfo
from skimage.filters import gaussian from skimage.filters import gaussian
@@ -75,10 +74,7 @@ class MTB_ExtractCoordinatesFromImage:
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
return { return {
"required": { "required": {
"threshold": ( "threshold": ("FLOAT",),
"FLOAT",
{"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01},
),
"max_points": ("INT", {"default": 50, "min": 0}), "max_points": ("INT", {"default": 50, "min": 0}),
}, },
"optional": {"image": ("IMAGE",), "mask": ("MASK",)}, "optional": {"image": ("IMAGE",), "mask": ("MASK",)},
@@ -91,124 +87,72 @@ class MTB_ExtractCoordinatesFromImage:
image: torch.Tensor | None = None, image: torch.Tensor | None = None,
mask: torch.Tensor | None = None, mask: torch.Tensor | None = None,
) -> tuple[list[list[tuple[int, int]]], torch.Tensor]: ) -> 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: if image is not None:
batch_count, height, width, _channel_count = image.shape batch_count, height, width, channel_count = image.shape
input_device = image.device imgs = image
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."
)
else: else:
if mask.ndim == 2: if mask is None:
mask = mask.unsqueeze(0) raise ValueError("Must provide either image or mask")
if mask.ndim != 3:
raise ValueError(
f"Mask has unexpected ndim: {mask.ndim} when image is not provided. Expected 2 or 3."
)
batch_count, height, width = mask.shape 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]]] = [] all_points: list[list[tuple[int, int]]] = []
debug_images = torch.zeros( debug_images = torch.zeros(
(batch_count, height, width, 3), (batch_count, height, width, 3),
dtype=torch.uint8, dtype=torch.uint8,
device=input_device, device=imgs.device,
) )
points_tensor = torch.tensor( for i, img in enumerate(imgs):
[255, 255, 255], dtype=torch.uint8, device=input_device if channel_count == 1:
) alpha_channel = img if len(img.shape) == 2 else img[:, :, 0]
elif channel_count == 2:
for i in range(batch_count): alpha_channel = img[:, :, 1]
value_threshold: torch.Tensor elif channel_count == 4:
if image is not None: alpha_channel = img[:, :, 3]
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}"
)
else: else:
mask_slice = mask[i] # get intensity
value_threshold = mask_slice alpha_channel = img[:, :, :3].max(dim=2)[0]
condition = value_threshold > threshold points = (alpha_channel > threshold).nonzero(as_tuple=False)
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_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: points = [(int(y.item()), int(x.item())) for x, y in points]
# shuffle and pick max_points randomly all_points.append(points)
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
)
current_points = [ for x, y in points:
(int(p[1].item()), int(p[0].item())) for p in points_yx self._draw_circle(debug_images[i], (x, y), 5)
]
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,
)
return (all_points, debug_images) return (all_points, debug_images)
@staticmethod @staticmethod
def _draw_circle( def _draw_circle(
image: torch.Tensor, image: torch.Tensor, center: tuple[int, int], radius: int
center: tuple[int, int],
radius: int,
color_tensor: torch.Tensor,
): ):
"""Draw a 5px circle on the image.""" """Draw a 5px circle on the image."""
x0, y0 = center x0, y0 = center
h, w, _ = image.shape for x in range(-radius, radius + 1):
min_x_bbox = max(0, x0 - radius) for y in range(-radius, radius + 1):
max_x_bbox = min(w - 1, x0 + radius) in_radius = x**2 + y**2 <= radius**2
min_y_bbox = max(0, y0 - radius) in_bounds = (
max_y_bbox = min(h - 1, y0 + radius) 0 <= x0 + x < image.shape[1]
and 0 <= y0 + y < image.shape[0]
for py in range(min_y_bbox, max_y_bbox + 1): )
for px in range(min_x_bbox, max_x_bbox + 1): if in_radius and in_bounds:
if (px - x0) ** 2 + (py - y0) ** 2 <= radius**2: image[y0 + y, x0 + x] = torch.tensor(
image[py, px] = color_tensor [255, 255, 255],
dtype=torch.uint8,
device=image.device,
)
class MTB_ColorCorrectGPU: class MTB_ColorCorrectGPU:
+2 -9
View File
@@ -21,11 +21,7 @@ class MTB_StackImages:
"match_method": ( "match_method": (
["error", "smallest", "largest"], ["error", "smallest", "largest"],
{"default": "error"}, {"default": "error"},
), )
"output_rgb": (
"BOOLEAN",
{"default": True, "tooltip": "Output RGB instead of RGBA"},
),
}, },
} }
@@ -33,7 +29,7 @@ class MTB_StackImages:
FUNCTION = "stack" FUNCTION = "stack"
CATEGORY = "mtb/image utils" 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: if not kwargs:
raise ValueError("At least one tensor must be provided.") raise ValueError("At least one tensor must be provided.")
@@ -102,9 +98,6 @@ class MTB_StackImages:
stacked_tensor = torch.cat(normalized_tensors, dim=dim) stacked_tensor = torch.cat(normalized_tensors, dim=dim)
if output_rgb:
stacked_tensor = stacked_tensor[:, :, :, :3]
return (stacked_tensor,) return (stacked_tensor,)
def normalize_to_rgba(self, tensor): def normalize_to_rgba(self, tensor):
+38 -5
View File
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
[project] [project]
name = "comfy-mtb" name = "comfy-mtb"
version = "0.6.0" version = "0.3.0"
description = "Animation oriented nodes pack for ComfyUI." description = "Animation oriented nodes pack for ComfyUI."
license = { text = "MIT" } license = { text = "MIT" }
readme = "README.md" readme = "README.md"
@@ -62,6 +62,39 @@ PublisherId = "mel"
DisplayName = "comfy-mtb" DisplayName = "comfy-mtb"
Icon = "https://avatars.githubusercontent.com/u/7041726?v=4" Icon = "https://avatars.githubusercontent.com/u/7041726?v=4"
[tool.bumpversion]
current_version = "0.3.0"
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
serialize = ["{major}.{minor}.{patch}"]
search = "{current_version}"
replace = "{new_version}"
regex = false
ignore_missing_version = false
ignore_missing_files = false
tag = true
sign_tags = true
tag_name = "v{new_version}"
tag_message = "⬆️ Bump version: {current_version} → {new_version}"
allow_dirty = true
commit = true
message = "⬆️ Bump version: {current_version} → {new_version}"
commit_args = ""
[[tool.bumpversion.files]]
filename = "__init__.py"
search = "__version__ = \"{current_version}\""
replace = "__version__ = \"{new_version}\""
[[tool.bumpversion.files]]
filename = "pyproject.toml"
search = "version = \"{current_version}\""
replace = "version = \"{new_version}\""
# [[tool.bumpversion.files]]
# filename = "your_package/__init__.py"
# search = "__version__ = '{current_version}'"
# replace = "__version__ = '{new_version}'"
# INFO: All those remaining keys are meant for local dev # INFO: All those remaining keys are meant for local dev
[tool.pyright] [tool.pyright]
include = ["."] include = ["."]
@@ -78,18 +111,18 @@ stubPath = "src/stubs"
reportMissingImports = true reportMissingImports = true
reportMissingTypeStubs = false reportMissingTypeStubs = false
reportExplicitAny = false
typeCheckingMode = "basic" typeCheckingMode = "basic"
pythonVersion = "3.11"
pythonPlatform = "Windows"
pythonVersion = "3.10"
pythonPlatform = "Windows"
[tool.pytest.ini_options] [tool.pytest.ini_options]
log_level = "DEBUG" log_level = "DEBUG"
log_cli = true log_cli = true
markers = [ markers = [
"wip: tests that aren't fully finished yet", "wip: tests that aren't fully finished yet",
'''heavy: marks tests as heavy (deselect with '-m "not heavy"')''', "heavy: marks tests as heavy (deselect with '-m \"not heavy\"')",
] ]
filterwarnings = ["ignore::UserWarning", 'ignore::DeprecationWarning'] filterwarnings = ["ignore::UserWarning", 'ignore::DeprecationWarning']
+18
View File
@@ -0,0 +1,18 @@
{
"exclude": [
"**/node_modules",
"**/__pycache__",
],
"ignore": [
"extern"
],
"defineConstant": {
"DEBUG": true
},
"venvPath": "../../../.venv/",
"reportMissingImports": true,
"reportMissingTypeStubs": false,
"pythonVersion": "3.10",
"pythonPlatform": "All",
"reportOptionalMemberAccess": "none"
}
-637
View File
@@ -1,637 +0,0 @@
import base64
import code
import io
import re
import sys
from contextlib import redirect_stderr, redirect_stdout
# import matplotlib.pyplot as plt
import numpy as np
import torch
from aiohttp import web
from PIL import Image
from rich.console import Console
from rich.traceback import Traceback
from .log import log
try:
import pyflakes.api
import pyflakes.reporter
_HAS_LINT = True
except ImportError:
print(
"ComfyREPL: pyflakes not found. Linting will be disabled. Install with 'pip install pyflakes'."
)
_HAS_LINT = False
# --- Linting Library ---
# try:
# import ruff
# import ruff.lint
# import ruff.lint.linter
# import ruff.settings
#
# _HAS_LINT = True
# except ImportError:
# print(
# "ComfyREPL: ruff not found. Linting will be disabled. Install with 'pip install ruff'."
# )
# _HAS_LINT = False
# --- Audio/Video Libraries ---
try:
import scipy.io.wavfile
_HAS_SCIPY = True
except ImportError:
print(
"ComfyREPL: SciPy not found. Audio display will be disabled. Install with 'pip install scipy'."
)
_HAS_SCIPY = False
try:
import imageio
import imageio.plugins.ffmpeg # Ensure ffmpeg plugin is available
_HAS_IMAGEIO = True
except ImportError:
print(
"ComfyREPL: Imageio or imageio-ffmpeg not found. Video display will be disabled. Install with 'pip install imageio imageio-ffmpeg'."
)
_HAS_IMAGEIO = False
# --- Audio Display ---
class AudioDisplay:
def __init__(self, samples, sample_rate):
if not _HAS_SCIPY:
raise ImportError("Audio display requires scipy and numpy.")
if not isinstance(samples, (np.ndarray, torch.Tensor)):
raise TypeError(
"Audio samples must be a numpy array or torch tensor."
)
if isinstance(samples, torch.Tensor):
samples = samples.detach().cpu().numpy()
# Ensure samples are in a format scipy.io.wavfile can handle (e.g., int16, float32)
if samples.dtype == np.float64:
samples = samples.astype(np.float32)
elif samples.dtype == np.int64:
# Or scale to int32 if range requires
samples = samples.astype(np.int16)
self.samples = samples
self.sample_rate = sample_rate
def _to_wav_base64(self):
buffer = io.BytesIO()
try:
scipy.io.wavfile.write(buffer, self.sample_rate, self.samples)
audio_base64 = base64.b64encode(buffer.getvalue()).decode("utf-8")
return audio_base64
except Exception as e:
return f"<div style='color: red;'>Error encoding audio: {e}</div>"
def _repr_html_(self):
base64_data = self._to_wav_base64()
if base64_data.startswith("<div"):
return base64_data
return f'<audio controls src="data:audio/wav;base64,{base64_data}" style="margin: 5px 0;"/>'
def render_audio(samples, sample_rate):
"""
Render audio samples as an HTML audio player.
Args:
samples (np.ndarray or torch.Tensor): Audio samples.
sample_rate (int): Sample rate in Hz.
Returns
-------
AudioDisplay: An object that will render as an HTML audio player.
"""
return AudioDisplay(samples, sample_rate)
# --- Display Classes ---
class VideoDisplay:
def __init__(self, frames, fps=24, options=None):
if not _HAS_IMAGEIO: # numpy/PIL/torch needed for frames
raise ImportError(
"Video display requires imageio, imageio-ffmpeg, and image libraries (numpy, Pillow, torch)."
)
self.frames = []
for frame in frames:
if isinstance(frame, Image.Image):
self.frames.append(np.array(frame))
elif isinstance(frame, np.ndarray):
# Ensure HWC and uint8
if frame.ndim == 3 and frame.shape[0] in [1, 3, 4]: # CHW
frame = np.transpose(frame, (1, 2, 0))
if frame.dtype != np.uint8:
frame = (
(frame * 255).astype(np.uint8)
if frame.max() <= 1.0
else frame.astype(np.uint8)
)
self.frames.append(frame)
elif isinstance(frame, torch.Tensor):
np_frame = frame.detach().cpu().numpy()
if np_frame.ndim == 3 and np_frame.shape[0] in [
1,
3,
4,
]: # CHW
np_frame = np.transpose(np_frame, (1, 2, 0))
if np_frame.dtype != np.uint8:
np_frame = (
(np_frame * 255).astype(np.uint8)
if np_frame.max() <= 1.0
else np_frame.astype(np.uint8)
)
self.frames.append(np_frame)
else:
raise TypeError(
f"Unsupported frame type: {type(frame)}. Must be PIL.Image, numpy.ndarray, or torch.Tensor."
)
self.fps = fps
self.options = options if options is not None else {}
def _to_mp4_base64(self):
buffer = io.BytesIO()
try:
# Use imageio to write frames to an in-memory MP4 file
imageio.mimwrite(
buffer,
self.frames,
format="mp4",
fps=self.fps,
codec="libx264",
quality=8,
) # quality 1-10
video_base64 = base64.b64encode(buffer.getvalue()).decode("utf-8")
return video_base64
except Exception as e:
return f"<div style='color: red;'>Error encoding video: {e}</div>"
def _repr_html_(self):
base64_data = self._to_mp4_base64()
if base64_data.startswith("<div"): # Check if it's an error message
return base64_data
# Build HTML options string
option_str = ""
for key, value in self.options.items():
if isinstance(value, bool) and value:
option_str += f" {key}"
elif isinstance(value, str):
option_str += f' {key}="{value}"'
else:
option_str += f' {key}="{value}"' # Fallback for numbers etc.
return f'<video controls src="data:video/mp4;base64,{base64_data}" style="max-width: 100%; height: auto; border: 1px solid #555; margin: 5px 0;"{option_str}/>'
def render_video(batch_tensor_or_array_of_pil_images, fps=24, options=None):
"""
Render video frames as an HTML video player.
Args:
batch_tensor_or_array_of_pil_images (list of PIL.Image, np.ndarray, or torch.Tensor):
A list of frames, or a single batch tensor/array (B, H, W, C) or (B, C, H, W).
fps (int): Frames per second.
options (dict): Dictionary of HTML <video> tag attributes (e.g., {"loop": True, "autoplay": True}).
Returns
-------
VideoDisplay: An object that will render as an HTML video player.
"""
frames_list = []
if isinstance(
batch_tensor_or_array_of_pil_images, (np.ndarray, torch.Tensor)
):
# Assume it's a batch tensor/array
for i in range(batch_tensor_or_array_of_pil_images.shape[0]):
frames_list.append(batch_tensor_or_array_of_pil_images[i])
elif isinstance(batch_tensor_or_array_of_pil_images, list):
frames_list = batch_tensor_or_array_of_pil_images
else:
raise TypeError(
"Input for render_video must be a list of frames or a batch tensor/array."
)
return VideoDisplay(frames_list, fps, options)
class ComfyREPLBackend:
def __init__(self):
self.repl_consoles: dict[str, code.InteractiveConsole] = {}
# self.repl_console = None
self.image_outputs = []
self.audio_outputs = []
self.video_outputs = []
self._original_displayhook = sys.displayhook
# self._init_repl_console()
@staticmethod
def _init_repl_console():
"""Define the globals that will be available in the REPL session."""
repl_globals = {"__builtins__": __builtins__}
# repl_globals["plt"] = plt
repl_globals["np"] = np
repl_globals["Image"] = Image
repl_globals["torch"] = torch
repl_globals["repl_display"] = _repl_display_image
if _HAS_SCIPY:
repl_globals["render_audio"] = render_audio
if _HAS_IMAGEIO:
repl_globals["render_video"] = render_video
return code.InteractiveConsole(locals=repl_globals)
def _custom_displayhook(self, value):
"""
Displayhook that capture and process image, audio, video objects.
For other objects, it fallsback to the original displayhook.
"""
if value is None:
return
# Attempt to handle as an image
if (
isinstance(value, (Image.Image, np.ndarray, torch.Tensor))
# or (
# hasattr(value, "figure")
# and isinstance(value.figure, plt.Figure)
# )
# or isinstance(value, plt.Figure)
):
img_html = _repl_display_image(value)
self.image_outputs.append(img_html)
return
# Attempt to handle as audio
elif isinstance(value, AudioDisplay):
audio_html = value._repr_html_()
self.audio_outputs.append(audio_html)
return
# Attempt to handle as video
elif isinstance(value, VideoDisplay):
video_html = value._repr_html_()
self.video_outputs.append(video_html)
return
else:
# If not a special media type, let the original displayhook handle it.
self._original_displayhook(value)
def _console_to_html(
self, stream: io.StringIO | Traceback, width: int = 120
) -> str:
if isinstance(stream, io.StringIO):
captured_text_output = stream.getvalue()
else:
captured_text_output = Traceback
html_console = Console(
file=io.StringIO(), record=True, force_terminal=True, width=width
)
html_console.print(captured_text_output)
return html_console.export_html(inline_styles=True)
def get_console(self, node_name: str):
console = self.repl_consoles.get(node_name)
if console:
return console
console = self._init_repl_console()
self.repl_consoles[node_name] = console
return self.repl_consoles[node_name]
def execute_code(self, node_name: str, code: str):
# Clear outputs from previous execution
self.image_outputs = []
self.audio_outputs = []
self.video_outputs = []
output_html = ""
error_message = None
repl_console = self.get_console(node_name)
string_io = io.StringIO()
# Temporarily patch sys.displayhook
sys.displayhook = self._custom_displayhook
try:
with redirect_stdout(string_io), redirect_stderr(string_io):
for line in code.splitlines():
repl_console.push(line)
full_rich_html = self._console_to_html(string_io)
match = re.search(
r"<body.*?>(.*?)</body>", full_rich_html, re.DOTALL
)
if match:
output_html = match.group(1)
else:
output_html = full_rich_html
# Append any captured media HTML *after* the rich text output
for img_html in self.image_outputs:
output_html += img_html
for audio_html in self.audio_outputs:
output_html += audio_html
for video_html in self.video_outputs:
output_html += video_html
except Exception as e:
exc_type, exc_value, exc_traceback = sys.exc_info()
rich_traceback = Traceback.from_exception(
exc_type,
exc_value,
exc_traceback,
show_locals=True,
suppress=[__file__],
)
# error_console = Console(
# file=io.StringIO(), record=True, force_terminal=True, width=120
# )
# error_console.print(rich_traceback)
# full_error_html = error_console.export_html(inline_styles=True)
full_error_html = self._console_to_html(rich_traceback)
match = re.search(
r"<body.*?>(.*?)</body>", full_error_html, re.DOTALL
)
output_html = match.group(1) if match else full_error_html
error_message = str(e)
finally:
sys.displayhook = (
self._original_displayhook
) # Always restore original displayhook
return {"output_html": output_html, "error": error_message}
def lint_code(self, node_name: str, code: str):
diagnostics = []
if not _HAS_LINT:
diagnostics.append(
{
"row": 0,
"column": 0,
"text": "Pyflakes not installed. Linting disabled. Install with 'uv add pyflakes'.",
"type": "warning",
}
)
return web.json_response({"diagnostics": diagnostics})
# Use a custom reporter to capture messages
class PyflakesReporter(pyflakes.reporter.Reporter):
def __init__(self):
self.messages = []
# Suppress stdout/stderr from pyflakes itself
self._stdout = io.StringIO()
self._stderr = io.StringIO()
super().__init__(self._stdout, self._stderr)
def flake(self, message):
# Ace editor expects 0-indexed row, pyflakes gives 1-indexed lineno
self.messages.append(
{
"row": message.lineno - 1,
"column": message.col,
"text": str(message),
"type": "warning", # pyflakes usually gives warnings
}
)
def unexpectedError(self, filename, msg):
self.messages.append(
{
"row": 0,
"column": 0,
"text": f"Pyflakes internal error: {msg}",
"type": "error",
}
)
def syntaxError(self, filename, msg, lineno, offset, text):
log.info(f"Received {text} to syntax error")
self.messages.append(
{
"row": lineno - 1, # Ace is 0-indexed
"column": offset,
"text": f"Syntax Error: {msg}",
"type": "error",
}
)
reporter = PyflakesReporter()
pyflakes.api.check(code, node_name, reporter)
return {"diagnostics": reporter.messages}
def lint_code_ruff(self, code: str):
diagnostics = []
if not _HAS_LINT:
diagnostics.append(
{
"row": 0,
"column": 0,
"text": "Ruff not installed. Linting disabled. Install with 'pip install ruff'.",
"type": "warning",
}
)
return {"diagnostics": diagnostics}
# Define the builtins/globals that Ruff should recognize
# These are the names we inject into the REPL's scope
repl_builtins = [
"repl_display",
"render_audio",
"render_video",
# "plt",
"np",
"Image",
"torch",
]
try:
# Lint the code using Ruff's programmatic API
result = ruff.lint.linter.lint_stdin(
code.encode("utf-8"),
path="<stdin>",
builtins=repl_builtins,
)
for diagnostic in result.diagnostics:
diag_type = "warning" # Default
# Ruff's error codes: F (Pyflakes), E (Pycodestyle), W (Pycodestyle warning), I (isort), N (naming), etc.
# F821: Undefined name (often an error)
if (
diagnostic.kind.code.startswith("E")
or diagnostic.kind.code == "F821"
):
diag_type = "error"
elif diagnostic.kind.code.startswith("W"):
diag_type = "warning"
diagnostics.append(
{
"row": diagnostic.location.row - 1, # Ace is 0-indexed
"column": diagnostic.location.column
- 1, # Ace is 0-indexed
"text": diagnostic.message,
"type": diag_type,
}
)
except Exception as e:
diagnostics.append(
{
"row": 0,
"column": 0,
"text": f"Ruff internal error: {e}",
"type": "error",
}
)
return {"diagnostics": diagnostics}
def _repl_display_image(img_data):
"""
Internal function to convert image data (PIL, numpy, torch, matplotlib) to base64 HTML.
"""
pil_img = None
# fig = None
if isinstance(img_data, Image.Image):
pil_img = img_data
elif isinstance(img_data, np.ndarray):
# Handle different numpy array shapes (HWC, CHW)
if img_data.ndim == 3:
if img_data.shape[0] in [1, 3, 4]: # Likely CHW
if img_data.shape[0] == 1: # Grayscale
img_data = img_data.squeeze(0)
else: # Color
img_data = np.transpose(img_data, (1, 2, 0)) # CHW to HWC
# Ensure it's uint8 for PIL, assuming float [0,1] or int [0,255]
if img_data.dtype != np.uint8:
img_data = (
(img_data * 255).astype(np.uint8)
if img_data.max() <= 1.0
else img_data.astype(np.uint8)
)
pil_img = Image.fromarray(img_data)
elif isinstance(img_data, torch.Tensor):
# Move to CPU, convert to numpy
np_img = img_data.detach().cpu().numpy()
# Handle different tensor shapes (CHW, HWC)
if np_img.ndim == 3:
if np_img.shape[0] in [1, 3, 4]: # Likely CHW
if np_img.shape[0] == 1: # Grayscale
np_img = np_img.squeeze(0)
else: # Color
np_img = np.transpose(np_img, (1, 2, 0)) # CHW to HWC
# Ensure it's uint8 for PIL, assuming float [0,1] or int [0,255]
if np_img.dtype != np.uint8:
np_img = (
(np_img * 255).astype(np.uint8)
if np_img.max() <= 1.0
else np_img.astype(np.uint8)
)
pil_img = Image.fromarray(np_img)
# elif hasattr(img_data, "figure") and isinstance(
# img_data.figure, plt.Figure
# ):
# # If it's a matplotlib Axes object, get its figure
# fig = img_data.figure
# elif isinstance(img_data, plt.Figure):
# fig = img_data
else:
return f"<div style='color: red;'>Unsupported image type for display: {type(img_data)}</div>"
buffer = io.BytesIO()
try:
if pil_img:
pil_img.save(buffer, format="PNG")
# elif fig:
# fig.savefig(
# buffer, format="PNG", bbox_inches="tight", pad_inches=0.1
# )
# plt.close(
# fig
# ) # Close the figure to prevent it from showing up in other contexts
else:
return (
"<div style='color: red;'>Could not process image data.</div>"
)
except Exception as e:
return f"<div style='color: red;'>Error saving image: {e}</div>"
img_base64 = base64.b64encode(buffer.getvalue()).decode("utf-8")
return f'<img src="data:image/png;base64,{img_base64}" style="max-width: 100%; height: auto; border: 1px solid #555; margin: 5px 0;"/>'
# Instantiate the backend class globally
_comfy_repl_backend = ComfyREPLBackend()
# Update aiohttp handlers to use the backend instance
async def repl_execute_code_handler(request):
data = await request.json()
name = data.get("name")
if name is None: # we send an error
return web.Response(
status=417, reason="Expectation Failed", text="Missing name key"
)
code = data.get("code", "")
result = _comfy_repl_backend.execute_code(name, code)
return web.json_response(result)
async def repl_lint_code_handler(request):
data = await request.json()
name = data.get("name")
if name is None: # we send an error
return web.Response(
status=417, reason="Expectation Failed", text="Missing name key"
)
# raise web.HTTPExpectationFailed(
# reason="Missing name key (reason)", text="Missing name key (text)"
# )
code = data.get("code", "")
result = _comfy_repl_backend.lint_code(name, code)
return web.json_response(result)
def setup_custom_web_routes(app: web.Application):
"""
Function to register our custom web routes with the ComfyUI server.
"""
log.info("ComfyREPL: Registering /mtb/execute route...")
app.router.add_post("/mtb/execute", repl_execute_code_handler)
app.router.add_post("/mtb/lint", repl_lint_code_handler)
# You can add more routes here if needed, e.g., for clearing state.
+23 -196
View File
@@ -1,5 +1,6 @@
import contextlib import contextlib
import functools import functools
import importlib
import math import math
import operator import operator
import os import os
@@ -8,14 +9,11 @@ import shutil
import socket import socket
import subprocess import subprocess
import sys import sys
import textwrap
import uuid import uuid
import warnings
from collections.abc import Callable, Sequence from collections.abc import Callable, Sequence
from enum import Enum from enum import Enum
from functools import reduce from functools import reduce
from pathlib import Path from pathlib import Path
from types import EllipsisType
from typing import TypeVar from typing import TypeVar
from urllib.parse import urlparse from urllib.parse import urlparse
@@ -464,6 +462,25 @@ def _run_command(shell_cmd, ignored_lines_start):
print("Command executed successfully!") 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 # endregion
@@ -527,196 +544,6 @@ PIL_FILTER_MAP = {
# region TENSOR Utilities # region TENSOR Utilities
class LazyProxyTensor:
"""Memory-efficient proxy that wrap a tensor but presents itself as a different dtype (e.g., float32).
It mimics a torch.Tensor's read-only attributes and methods. Data conversion
and normalization happen lazily on access (e.g., via slicing), avoiding
the high memory cost of a full conversion.
Supported source dtypes:
- torch.uint8 (normalized from [0, 255])
- torch.uint16 (normalized from [0, 65535])
- All float types (passed through, assumed to be in [0, 1] range)"
"""
_source_tensor: torch.Tensor
_target_dtype: torch.dtype
_target_element_size: int
_scale_divisor: float
_warned_inefficient_access: bool
def __init__(
self, source_tensor, target_dtype=torch.float32, target_device=None
):
if not isinstance(source_tensor, torch.Tensor):
raise ValueError("Input must be a torch.Tensor.")
self._source_tensor = source_tensor
self._target_dtype = target_dtype
self._target_device = (
target_device
if target_device is not None
else source_tensor.device
)
# Determine the normalization divisor based on source dtype
# fmt: off
if source_tensor.dtype == torch.uint8: self._scale_divisor = 255.0
elif source_tensor.dtype == torch.uint16: self._scale_divisor = 65535.0
elif torch.is_floating_point(source_tensor): self._scale_divisor = 1.0
else: raise ValueError(f"Unsupported source dtype for LazyProxyTensor: {source_tensor.dtype}")
# fmt: on
self._target_element_size = torch.empty(
(), dtype=self._target_dtype
).element_size()
self._warned_inefficient_access = False
def is_contiguous(self, *args, **kwargs):
return self._source_tensor.is_contiguous(*args, **kwargs)
def stride(self, *args, **kwargs):
return self._source_tensor.stride(*args, **kwargs)
@property
def shape(self):
return self._source_tensor.shape
@property
def requires_grad(self):
return False
def nelement(self):
"""Return the total number of elements in the (pretend) tensor."""
return self._source_tensor.nelement()
def element_size(self):
"""Return the size in bytes of an individual (pretend) float element."""
return self._target_element_size
@property
def dtype(self):
return self._target_dtype
@property
def device(self):
return self._target_device
def __len__(self):
return self._source_tensor.shape[0]
def __getitem__(self, key):
if (
self._source_tensor.device != self._target_device
and not self._warned_inefficient_access
):
warnings.warn(
"Inefficient access pattern detected for LazyProxyTensor. "
"You are slicing a device-proxied tensor, which causes slow, "
"repeated data transfers. For performance, use the .iter_chunks() method."
)
self._warned_inefficient_access = True
subset = self._source_tensor[key]
return (
subset.to(self._target_device).to(self._target_dtype)
/ self._scale_divisor
)
# def __iter__(self):
# for i in range(len(self)):
# yield self[i]
def iter_chunks(self, chunk_size=16):
for i in range(0, len(self), chunk_size):
chunk = self._source_tensor[i : i + chunk_size]
yield (
chunk.to(self._target_device, non_blocking=True).to(
self._target_dtype
)
/ self._scale_divisor
)
def squeeze(self, dim: str | EllipsisType | None = None):
squeezed = self._source_tensor.squeeze(dim)
return LazyProxyTensor(squeezed, self._target_dtype)
def unsqueeze(self, dim: int = 0):
unsqueezed = self._source_tensor.unsqueeze(dim)
return LazyProxyTensor(unsqueezed, self._target_dtype)
def repeat(self, *sizes):
repeated = self._source_tensor.repeat(*sizes)
return LazyProxyTensor(repeated, self._target_dtype)
def _format_mem_size(self, mem_bytes):
if mem_bytes > 1e9:
return f"{mem_bytes / 1e9:.2f} GB"
if mem_bytes > 1e6:
return f"{mem_bytes / 1e6:.2f} MB"
if mem_bytes > 1e3:
return f"{mem_bytes / 1e3:.2f} KB"
return f"{mem_bytes} B"
def __repr__(self):
actual_info = get_torch_tensor_info(self._source_tensor, name="Source")
target_info = get_torch_tensor_info(self, name="Target")
info = f"""
{target_info}
{actual_info}
"""
return textwrap.dedent(info).strip()
def get_torch_tensor_info(
tensor: torch.Tensor | LazyProxyTensor | np.ndarray,
*,
name: str | None = None,
):
mem_str = "N/A"
is_tensor = isinstance(tensor, torch.Tensor | LazyProxyTensor)
if is_tensor:
mem_bytes = tensor.element_size() * tensor.nelement()
else:
mem_bytes = tensor.itemsize * tensor.size
if mem_bytes > 1e9:
mem_str = f"{mem_bytes / 1e9:.2f} GB"
elif mem_bytes > 1e6:
mem_str = f"{mem_bytes / 1e6:.2f} MB"
elif mem_bytes > 1e3:
mem_str = f"{mem_bytes / 1e3:.2f} KB"
else:
mem_str = f"{mem_bytes} B"
device = "N/A"
grad = "False"
type_name = name or "Tensor" if is_tensor else "Numpy Array"
if is_tensor:
device = tensor.device
grad = str(tensor.requires_grad)
text = f"""
{type_name}
shape: {tensor.shape}
dtype: {str(tensor.dtype).replace("torch.", "")}
device: {device}
requires grad: {grad}
memory: {mem_str}
"""
return textwrap.dedent(text).strip()
def to_numpy(image: torch.Tensor) -> npt.NDArray[np.uint8]: def to_numpy(image: torch.Tensor) -> npt.NDArray[np.uint8]:
"""Converts a tensor to a ndarray with proper scaling and type conversion.""" """Converts a tensor to a ndarray with proper scaling and type conversion."""
np_array = np.clip(255.0 * image.cpu().numpy(), 0, 255).astype(np.uint8) np_array = np.clip(255.0 * image.cpu().numpy(), 0, 255).astype(np.uint8)
@@ -727,12 +554,12 @@ def handle_batch(
tensor: torch.Tensor, tensor: torch.Tensor,
func: Callable[[torch.Tensor], Image.Image | npt.NDArray[np.uint8]], func: Callable[[torch.Tensor], Image.Image | npt.NDArray[np.uint8]],
) -> list[Image.Image] | list[npt.NDArray[np.uint8]]: ) -> list[Image.Image] | list[npt.NDArray[np.uint8]]:
"""Handle batch processing for a given tensor and conversion function.""" """Handles batch processing for a given tensor and conversion function."""
return [func(tensor[i]) for i in range(tensor.shape[0])] return [func(tensor[i]) for i in range(tensor.shape[0])]
def tensor2pil(tensor: torch.Tensor) -> list[Image.Image]: def tensor2pil(tensor: torch.Tensor) -> list[Image.Image]:
"""Convert a batch of tensors to a list of PIL Images.""" """Converts a batch of tensors to a list of PIL Images."""
def single_tensor2pil(t: torch.Tensor) -> Image.Image: def single_tensor2pil(t: torch.Tensor) -> Image.Image:
np_array = to_numpy(t) np_array = to_numpy(t)
@@ -749,7 +576,7 @@ def tensor2pil(tensor: torch.Tensor) -> list[Image.Image]:
def pil2tensor(images: Image.Image | list[Image.Image]) -> torch.Tensor: def pil2tensor(images: Image.Image | list[Image.Image]) -> torch.Tensor:
"""Convert a PIL Image or a list of PIL Images to a tensor.""" """Converts a PIL Image or a list of PIL Images to a tensor."""
def single_pil2tensor(image: Image.Image) -> torch.Tensor: def single_pil2tensor(image: Image.Image) -> torch.Tensor:
np_image = np.array(image).astype(np.float32) / 255.0 np_image = np.array(image).astype(np.float32) / 255.0
+157 -234
View File
@@ -12,52 +12,10 @@
import { app } from '../../scripts/app.js' import { app } from '../../scripts/app.js'
import { api } from '../../scripts/api.js' import { api } from '../../scripts/api.js'
// #region base utils if (!window.MTB) {
window.MTB = {}
/**
* Computes the convex hull of a set of points using the Monotone Chain algorithm.
*
* @param {Array<Array<number>>} points An array of points, where each point is an array of two numbers [x, y].
* @returns {Array<Array<number>>} The points forming the convex hull, in counter-clockwise order.
*/
export const getConvexHull = (points) => {
if (points.length <= 3) {
return points
}
points.sort((a, b) => a[0] - b[0] || a[1] - b[1])
const lower = []
for (const p of points) {
while (
lower.length >= 2 &&
cross_product(lower[lower.length - 2], lower[lower.length - 1], p) <= 0
) {
lower.pop()
}
lower.push(p)
}
const upper = []
for (let i = points.length - 1; i >= 0; i--) {
const p = points[i]
while (
upper.length >= 2 &&
cross_product(upper[upper.length - 2], upper[upper.length - 1], p) <= 0
) {
upper.pop()
}
upper.push(p)
}
function cross_product(o, a, b) {
return (a[0] - o[0]) * (b[1] - o[1]) - (a[1] - o[1]) * (b[0] - o[0])
}
return lower
.slice(0, lower.length - 1)
.concat(upper.slice(0, upper.length - 1))
} }
// #region base utils
// - crude uuid // - crude uuid
export function makeUUID() { export function makeUUID() {
@@ -70,19 +28,6 @@ export function makeUUID() {
return uuid 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 //- local storage manager
export class LocalStorageManager { export class LocalStorageManager {
constructor(namespace) { constructor(namespace) {
@@ -253,7 +198,6 @@ export function hideWidgetForGood(node, widget, suffix = '') {
widget.origComputeSize = widget.computeSize widget.origComputeSize = widget.computeSize
widget.origSerializeValue = widget.serializeValue widget.origSerializeValue = widget.serializeValue
widget.computeSize = () => [0, -4] // -4 is due to the gap litegraph adds between widgets automatically widget.computeSize = () => [0, -4] // -4 is due to the gap litegraph adds between widgets automatically
widget.hidden = true
widget.type = CONVERTED_TYPE + suffix widget.type = CONVERTED_TYPE + suffix
// widget.serializeValue = () => { // widget.serializeValue = () => {
// // Prevent serializing the widget if we have no input linked // // Prevent serializing the widget if we have no input linked
@@ -336,9 +280,9 @@ export const getNamedWidget = (node, ...names) => {
*/ */
export const nodesFromLink = (node, link) => { export const nodesFromLink = (node, link) => {
if (typeof link === 'number') { if (typeof link === 'number') {
link = app.graph.getLink(link) console.log('Resolving link from id', link)
link = app.graph.links[link]
} }
const fromNode = app.graph.getNodeById(link.origin_id) const fromNode = app.graph.getNodeById(link.origin_id)
const toNode = app.graph.getNodeById(link.target_id) const toNode = app.graph.getNodeById(link.target_id)
@@ -429,54 +373,12 @@ export function getWidgetType(config) {
// #endregion // #endregion
// function to test if input is a dynamic one
const isDynamicInput = (input) => {
infoLogger('Checking if input dynamic', { input })
// return input.name.startsWith(connectionPrefix)
return input._isDynamic === true
}
// Add a dynamic input, update node properties and slot colors!
const addDynamicInput = (node, name, kind) => {
const input = node.addInput(name, kind)
input._isDynamic = true
update_dynamic_properties(node)
set_slot_colors(node, ['cyan', undefined], isDynamicInput)
return input
}
const set_slot_colors = (node, colors, condition) => {
if (!condition) {
condition = (_s) => true
}
for (const slot of node.slots) {
infoLogger('Candidate', { slot, accepted: condition(slot) })
if (condition(slot)) {
slot.color_off = colors[0]
slot.color_on = colors[1]
}
}
}
const update_dynamic_properties = (node) => {
const dyn = []
for (const input of node.inputs) {
if (isDynamicInput(input)) {
dyn.push(input.name)
}
}
node.setProperty('dynamic_connections', dyn)
}
// #region dynamic connections // #region dynamic connections
/** /**
* @param {NodeType} nodeType The nodetype to attach the documentation to * @param {NodeType} nodeType The nodetype to attach the documentation to
* @param {str} prefix A prefix added to each dynamic inputs * @param {str} prefix A prefix added to each dynamic inputs
* @param {str | [str]} inputType The datatype(s) of those dynamic inputs * @param {str | [str]} inputType The datatype(s) of those dynamic inputs
* @param {{separator?:string,rename_menu?:'label'|'name', start_index?:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} [opts] Extra options * @param {{separator?:string, start_index?:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} [opts] Extra options
* @returns * @returns
*/ */
export const setupDynamicConnections = ( export const setupDynamicConnections = (
@@ -490,115 +392,20 @@ export const setupDynamicConnections = (
Object.getOwnPropertyDescriptors(nodeType).title.value, Object.getOwnPropertyDescriptors(nodeType).title.value,
) )
/** @type {{separator:string,rename_menu?:"label"|"name" start_index:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} */ /** @type {{separator:string, start_index:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} */
const options = Object.assign( const options = Object.assign(
{ {
separator: '_', separator: '_',
start_index: 1, start_index: 1,
rename_menu: 'label',
}, },
opts || {}, opts || {},
) )
const is_valid_name = (node, val) => {
return true
}
nodeType.prototype.getSlotMenuOptions = (slot) => {
if (!slot.input) {
return
}
infoLogger('Slot Menu', { slot })
return [
{
content: `Rename Input (${options.rename_menu})`,
callback: () => {
const dialog = app.canvas.createDialog(
"<span class='name'>Name</span><input autofocus type='text'/><button>OK</button>",
{},
)
const dialogInput = dialog.querySelector('input')
if (dialogInput) {
if (options.rename_menu === 'label') {
dialogInput.value = slot.input.label || slot.input.name || ''
} else if (options.rename_menu === 'name') {
dialogInput.value = slot.input.name || ''
}
}
const inner = () => {
// TODO: check if name exists or other guards
const val = dialogInput.value
if (!is_valid_name(slot.node, val)) {
dialog.close()
return
}
app.graph.beforeChange()
if (options.rename_menu === 'label') {
slot.input.label = val
} else if (options.rename_menu === 'name') {
slot.input.name = val
slot.input.label = val
}
app.graph.afterChange()
dialog.close()
}
dialog.querySelector('button').addEventListener('click', inner)
dialogInput.addEventListener('keydown', (e) => {
dialog.is_modified = true
if (e.keyCode === 27) {
dialog.close()
} else if (e.keyCode === 13) {
inner()
} else if (e.keyCode !== 13 && e.target?.localName !== 'textarea') {
return
}
e.preventDefault()
e.stopPropagation()
})
dialogInput.focus()
},
},
]
}
const onConfigure = nodeType.prototype.onConfigure
nodeType.prototype.onConfigure = function (data) {
const r = onConfigure ? onConfigure.apply(this, data) : undefined
// Set or restore serialized properties, lt seems to auto serialize/deserialize to/from string
if (!('dynamic_connections' in this.properties)) {
// this.addProperty('dynamic_connections', [], 'string')
this.setProperty('dynamic_connections', [])
} else {
const dynamic_connections = this.properties.dynamic_connections
if (typeof dynamic_connections !== 'object') {
return r
}
for (const name of dynamic_connections) {
infoLogger(`Would dynamize: ${name}`)
const input = this.inputs.find((i) => i.name === name)
if (input) {
infoLogger('Input found', { input })
input._isDynamic = true
}
}
}
// set color
set_slot_colors(this, ['cyan', undefined], isDynamicInput)
return r
}
const onNodeCreated = nodeType.prototype.onNodeCreated const onNodeCreated = nodeType.prototype.onNodeCreated
const inputList = typeof inputType === 'object' const inputList = typeof inputType === 'object'
nodeType.prototype.onNodeCreated = function () { nodeType.prototype.onNodeCreated = function () {
const r = onNodeCreated ? onNodeCreated.apply(this, []) : undefined const r = onNodeCreated ? onNodeCreated.apply(this, []) : undefined
this.addInput(
const input = addDynamicInput(
this,
`${prefix}${options.separator}${options.start_index}`, `${prefix}${options.separator}${options.start_index}`,
inputList ? '*' : inputType, inputList ? '*' : inputType,
) )
@@ -670,7 +477,10 @@ export const dynamic_connection = (
opts || {}, opts || {},
) )
if (node.inputs.length > 0 && !isDynamicInput(node.inputs[index])) { // function to test if input is a dynamic one
const isDynamicInput = (inputName) => inputName.startsWith(connectionPrefix)
if (node.inputs.length > 0 && !isDynamicInput(node.inputs[index].name)) {
return return
} }
@@ -680,7 +490,6 @@ export const dynamic_connection = (
const nameArray = options.nameArray || [] const nameArray = options.nameArray || []
const clean_inputs = () => { const clean_inputs = () => {
if (node.id < 0) return // being duplicated
if (node.inputs.length === 0) return if (node.inputs.length === 0) return
let w_count = node.widgets?.length || 0 let w_count = node.widgets?.length || 0
@@ -690,7 +499,7 @@ export const dynamic_connection = (
const to_remove = [] const to_remove = []
for (let n = 1; n < node.inputs.length; n++) { for (let n = 1; n < node.inputs.length; n++) {
const element = node.inputs[n] const element = node.inputs[n]
if (!element.link && isDynamicInput(element)) { if (!element.link && isDynamicInput(element.name)) {
if (node.widgets) { if (node.widgets) {
const w = node.widgets.find((w) => w.name === element.name) const w = node.widgets.find((w) => w.name === element.name)
if (w) { if (w) {
@@ -704,12 +513,9 @@ export const dynamic_connection = (
} }
for (let i = 0; i < to_remove.length; i++) { for (let i = 0; i < to_remove.length; i++) {
const id = to_remove[i] const id = to_remove[i]
try {
node.removeInput(id) node.removeInput(id)
i_count -= 1 i_count -= 1
} catch (err) {
errorLogger('Cannot remove input', err)
}
} }
node.inputs.length = i_count node.inputs.length = i_count
@@ -723,7 +529,7 @@ export const dynamic_connection = (
for (let i = 0; i < node.inputs.length; i++) { for (let i = 0; i < node.inputs.length; i++) {
let name = '' let name = ''
// rename only prefixed inputs // rename only prefixed inputs
if (node.inputs[i].name.startsWith(connectionPrefix)) { if (isDynamicInput(node.inputs[i].name)) {
// prefixed => rename and increase index // prefixed => rename and increase index
name = `${connectionPrefix}${prefixed_idx}` name = `${connectionPrefix}${prefixed_idx}`
prefixed_idx += 1 prefixed_idx += 1
@@ -777,8 +583,9 @@ export const dynamic_connection = (
if (node.inputs.length === 0) return if (node.inputs.length === 0) return
// add an extra input // add an extra input
if (node.inputs[node.inputs.length - 1].link !== null) { if (node.inputs[node.inputs.length - 1].link !== null) {
// count only the prefixed inputs
const nextIndex = node.inputs.reduce( const nextIndex = node.inputs.reduce(
(acc, cur) => (isDynamicInput(cur) ? ++acc : acc), (acc, cur) => (isDynamicInput(cur.name) ? ++acc : acc),
0, 0,
) )
@@ -788,7 +595,7 @@ export const dynamic_connection = (
: `${connectionPrefix}${nextIndex + options.start_index}` : `${connectionPrefix}${nextIndex + options.start_index}`
infoLogger(`Adding input ${nextIndex + 1} (${name})`) infoLogger(`Adding input ${nextIndex + 1} (${name})`)
addDynamicInput(node, name, conType) node.addInput(name, conType)
} }
} }
} }
@@ -821,21 +628,21 @@ function getBrightness(rgbObj) {
export function calculateTotalChildrenHeight(parentElement) { export function calculateTotalChildrenHeight(parentElement) {
let totalHeight = 0 let totalHeight = 0
if (!parentElement || !parentElement.children) {
return 0
}
for (const child of parentElement.children) { for (const child of parentElement.children) {
const style = window.getComputedStyle(child) const style = window.getComputedStyle(child)
const height = Number.parseFloat(style.height) // Get height as an integer (without 'px')
const marginTop = Number.parseFloat(style.marginTop) const height = Number.parseInt(style.height, 10)
const marginBottom = Number.parseFloat(style.marginBottom)
// 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 totalHeight += height + marginTop + marginBottom
} }
return Math.ceil(totalHeight) return totalHeight
} }
export const loadScript = ( export const loadScript = (
@@ -846,15 +653,13 @@ export const loadScript = (
return new Promise((resolve, reject) => { return new Promise((resolve, reject) => {
try { try {
// Check if the script already exists // Check if the script already exists
let scriptEle = document.querySelector(`script[src="${FILE_URL}"]`) const existingScript = document.querySelector(`script[src="${FILE_URL}"]`)
if (scriptEle) { if (existingScript) {
scriptEle.addEventListener('load', (_ev) => { resolve({ status: true, message: 'Script already loaded' })
resolve({ status: true })
})
return return
} }
scriptEle = document.createElement('script') const scriptEle = document.createElement('script')
scriptEle.type = type scriptEle.type = type
scriptEle.async = async scriptEle.async = async
scriptEle.src = FILE_URL scriptEle.src = FILE_URL
@@ -873,8 +678,6 @@ export const loadScript = (
document.body.appendChild(scriptEle) document.body.appendChild(scriptEle)
} catch (error) { } catch (error) {
reject(error) reject(error)
} finally {
infoLogger(`Finally loaded script: ${FILE_URL}`)
} }
}) })
} }
@@ -988,10 +791,12 @@ function loadParser(shiki) {
export const ensureMarkdownParser = async (callback) => { export const ensureMarkdownParser = async (callback) => {
infoLogger('Ensuring md parser') infoLogger('Ensuring md parser')
const use_shiki = app.extensionManager.setting.get( let use_shiki = false
'mtb.noteplus.use-shiki', try {
false, use_shiki = await api.getSetting('mtb.Use Shiki')
) } catch (e) {
console.warn('Option not available yet', e)
}
if (window.MTB?.mdParser) { if (window.MTB?.mdParser) {
infoLogger('Markdown parser found') infoLogger('Markdown parser found')
@@ -1016,7 +821,8 @@ export const ensureMarkdownParser = async (callback) => {
callbackQueue.push(callback) callbackQueue.push(callback)
} }
await await parserPromise await parserPromise
await parserPromise
return window.MTB.mdParser return window.MTB.mdParser
} }
@@ -1269,6 +1075,66 @@ export const addDocumentation = (
// #endregion // #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 // #region node extensions
/** /**
@@ -1343,6 +1209,8 @@ export const runAction = async (name, ...args) => {
const res = await req.json() const res = await req.json()
return res.result return res.result
} }
window.MTB.run = runAction
export const getServerInfo = async () => { export const getServerInfo = async () => {
const res = await api.fetchApi('/mtb/server-info') const res = await api.fetchApi('/mtb/server-info')
return await res.json() return await res.json()
@@ -1355,3 +1223,58 @@ export const setServerInfo = async (opts) => {
} }
// #endregion // #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
+99 -75
View File
@@ -11,12 +11,7 @@
/// <reference path="../types/typedefs.js" /> /// <reference path="../types/typedefs.js" />
import { app } from '../../scripts/app.js' import { app } from '../../scripts/app.js'
import * as shared from './comfy_shared.js'
import {
setupDynamicConnections,
cleanupNode,
infoLogger,
} from './comfy_shared.js'
import * as mtb_ui from './mtb_ui.js' import * as mtb_ui from './mtb_ui.js'
function escapeHtml(unsafe) { function escapeHtml(unsafe) {
@@ -33,7 +28,7 @@ function createDebugSection(title) {
margin: '8px 0', margin: '8px 0',
padding: '8px', padding: '8px',
borderRadius: '4px', borderRadius: '4px',
backgroundColor: 'rgba(0,0,0,0.2)', backgroundColor: 'rgba(0,0,0,0.2)'
}) })
const header = mtb_ui.makeElement('h3', { const header = mtb_ui.makeElement('h3', {
@@ -42,7 +37,7 @@ function createDebugSection(title) {
borderBottom: '1px solid rgba(255,255,255,0.1)', borderBottom: '1px solid rgba(255,255,255,0.1)',
fontSize: '14px', fontSize: '14px',
fontWeight: 'bold', fontWeight: 'bold',
color: '#9f9', color: '#9f9'
}) })
header.textContent = title header.textContent = title
section.appendChild(header) section.appendChild(header)
@@ -50,25 +45,25 @@ function createDebugSection(title) {
return section return section
} }
function createDebugContent(item) { function createDebugContent(content, type) {
const wrapper = mtb_ui.makeElement('div', { const wrapper = mtb_ui.makeElement('div', {
margin: '4px 0', margin: '4px 0'
}) })
if (item.kind === 'text') { if (type === 'text') {
const text = mtb_ui.makeElement('div', { const text = mtb_ui.makeElement('p', {
margin: '2px 0', margin: '2px 0',
fontFamily: 'monospace', fontFamily: 'monospace',
whiteSpace: 'pre-wrap', whiteSpace: 'pre-wrap'
}) })
text.innerHTML = item.data text.innerHTML = content
wrapper.appendChild(text) wrapper.appendChild(text)
} else if (item.kind === 'b64_images') { } else if (type === 'image') {
const img = mtb_ui.makeElement('img', { const img = mtb_ui.makeElement('img', {
width: '100%', width: '100%',
borderRadius: '2px', borderRadius: '2px'
}) })
img.src = item.data img.src = content
wrapper.appendChild(img) wrapper.appendChild(img)
} }
@@ -85,85 +80,115 @@ app.registerExtension({
*/ */
async beforeRegisterNodeDef(nodeType, nodeData, app) { async beforeRegisterNodeDef(nodeType, nodeData, app) {
if (nodeData.name === 'Debug (mtb)') { if (nodeData.name === 'Debug (mtb)') {
const clear_widgets = (target) => { const onNodeCreated = nodeType.prototype.onNodeCreated
if (target.widgets) { nodeType.prototype.onNodeCreated = function (...args) {
let tgt_len = target.widgets.length this.options = {}
for (let i = 0; i < target.widgets.length; i++) { const r = onNodeCreated ? onNodeCreated.apply(this, args) : undefined
if ( this.addInput('anything_1', '*')
![ return r
'output_to_console',
'deep_inspect',
'as_detailed_types',
'rich_mode',
].includes(target.widgets[i].name)
) {
target.widgets[i].onRemove?.()
target.widgets[i].onRemoved?.()
tgt_len -= 1
}
}
target.widgets.length = tgt_len
}
} }
const original_getExtraMenuOptions = const onConnectionsChange = nodeType.prototype.onConnectionsChange
nodeType.prototype.getExtraMenuOptions /**
nodeType.prototype.getExtraMenuOptions = function (_, options) { * @param {OnConnectionsChangeParams} args
original_getExtraMenuOptions?.apply(this, arguments) */
options.push({ nodeType.prototype.onConnectionsChange = function (...args) {
content: '🐛 Clear Outputs', const [_type, index, connected, link_info, ioSlot] = args
callback: async () => { const r = onConnectionsChange
clear_widgets(this) ? onConnectionsChange.apply(this, args)
}, : undefined
// TODO: remove all widgets on disconnect once computed
shared.dynamic_connection(this, index, connected, 'anything_', '*', {
link: link_info,
ioSlot: ioSlot,
}) })
}
setupDynamicConnections(nodeType, 'var', '*') //- infer type
if (link_info) {
// const fromNode = this.graph._nodes.find(
// (otherNode) => otherNode.id === link_info.origin_id,
// )
// const fromNode = app.graph.getNodeById(link_info.origin_id)
const { from } = shared.nodesFromLink(this, link_info)
if (!from || this.inputs.length === 0) return
const type = from.outputs[link_info.origin_slot].type
this.inputs[index].type = type
// this.inputs[index].label = type.toLowerCase()
}
//- restore dynamic input
if (!connected) {
this.inputs[index].type = '*'
this.inputs[index].label = `anything_${index + 1}`
}
return r
}
const onExecuted = nodeType.prototype.onExecuted const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (...args) { nodeType.prototype.onExecuted = function (...args) {
onExecuted?.apply(this, args) onExecuted?.apply(this, args)
const [data, ..._rest] = args const [data, ..._rest] = args
clear_widgets(this) if (this.widgets) {
let tgt_len = this.widgets.length
for (let i = 0; i < this.widgets.length; i++) {
if (
this.widgets[i].name !== 'output_to_console' &&
this.widgets[i].name !== 'as_detailed_types'
) {
this.widgets[i].onRemove?.()
this.widgets[i].onRemoved?.()
tgt_len -= 1
}
}
this.widgets.length = tgt_len
}
const inputData = {} const inputData = {}
const uiData = data.ui || data const uiData = data.ui || data
const name_to_label = this.inputs.reduce((acc, input) => {
acc[input.name] = input.label || input.name
return acc
}, {})
if (uiData.items) { if (uiData.items) {
uiData.items.forEach((item) => { uiData.items.forEach(item => {
const inputName = item.input const inputName = item.input
inputData[inputName] = item.items 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)
}
})
} }
const mainDebugContainer = mtb_ui.makeElement('div', {
width: '100%', let widgetI = 1
})
let hasContent = false
for (const [inputName, content] of Object.entries(inputData)) { for (const [inputName, content] of Object.entries(inputData)) {
if (!content || content?.length === 0) { if (content.text.length === 0 && content.b64_images.length === 0) {
continue continue
} }
hasContent = true
const section = createDebugSection(name_to_label[inputName]) const section = createDebugSection(inputName)
for (const item of content) { if (content.text.length > 0) {
section.appendChild(createDebugContent(item)) content.text.forEach(text => {
section.appendChild(createDebugContent(text, 'text'))
})
} }
mainDebugContainer.appendChild(section)
} if (content.b64_images.length > 0) {
if (hasContent) { content.b64_images.forEach(img => {
this.addDOMWidget('debug_output', 'CUSTOM', mainDebugContainer, { section.appendChild(createDebugContent(img, 'image'))
hideOnZoom: false, })
}) }
this.addDOMWidget(
`debug_section_${widgetI}`,
'CUSTOM',
section,
{}
)
widgetI++
} }
this.onRemoved = function () { this.onRemoved = function () {
@@ -174,9 +199,8 @@ app.registerExtension({
widget.onRemoved?.() widget.onRemoved?.()
widget.onRemove?.() widget.onRemove?.()
} }
cleanupNode(this) shared.cleanupNode(this)
} }
this.setDirtyCanvas(true, true)
} }
} }
}, },
+296 -296
View File
@@ -13,40 +13,40 @@ import { api } from '../../scripts/api.js'
import { app } from '../../scripts/app.js' import { app } from '../../scripts/app.js'
import { LocalStorageManager } from './comfy_shared.js' import { LocalStorageManager } from './comfy_shared.js'
const styles = { const styles = {
lighbox: { lighbox: {
position: 'fixed', position: 'fixed',
top: 0, top: 0,
left: 0, left: 0,
width: '100vw', width: '100vw',
height: '100vh', height: '100vh',
background: 'rgba(0,0,0,0.5)', background: 'rgba(0,0,0,0.5)',
display: 'none', display: 'none',
justifyContent: 'center', justifyContent: 'center',
alignItems: 'center', alignItems: 'center',
zIndex: 999, zIndex: 999,
}, },
lightboxBtn: (extra) => ({ lightboxBtn: (extra) => ({
position: 'absolute', position: 'absolute',
top: '50%', top: '50%',
background: 'none', background: 'none',
border: 'none', border: 'none',
color: '#fff', color: '#fff',
zIndex: 1000, zIndex: 1000,
fontSize: '30px', fontSize: '30px',
cursor: 'pointer', cursor: 'pointer',
pointerEvents: 'auto', pointerEvents: 'auto',
...extra, ...extra,
}), }),
img_list: { img_list: {
minHeight: '30px', minHeight: '30px',
maxHeight: '300px', maxHeight: '300px',
width: '100vw', width: '100vw',
position: 'absolute', position: 'absolute',
bottom: 0, bottom: 0,
zIndex: 10, zIndex: 10,
background: '#333', background: '#333',
overflow: 'auto', overflow: 'auto',
}, },
} }
let currentImageIndex = 0 let currentImageIndex = 0
@@ -58,299 +58,299 @@ const storage = new LocalStorageManager('mtb')
let activated = storage.get('image_feed', false) let activated = storage.get('image_feed', false)
app.registerExtension({ app.registerExtension({
name: 'mtb.ImageFeed', name: 'mtb.ImageFeed',
setup: () => { setup: () => {
app.ui.settings.addSetting({ app.ui.settings.addSetting({
id: 'mtb.Main.image-feed-enabled', id: 'mtb.Main.image-feed-enabled',
category: ['mtb', ' Main', 'image-feed-enabled'], category: ['mtb', 'Main', 'image-feed-enabled'],
name: 'Enable Image Feed', name: 'Enable Image Feed',
type: 'boolean', type: 'boolean',
defaultValue: false, defaultValue: false,
attrs: { attrs: {
style: { style: {
fontFamily: 'monospace', fontFamily: 'monospace',
}, },
}, },
async onChange(value) { async onChange(value) {
storage.set('image_feed', value) storage.set('image_feed', value)
activated = value activated = value
}, },
}) })
}, },
init: async () => { init: async () => {
if (!activated) { if (!activated) {
return return
} }
const pythongossFeed = app.extensions.find( const pythongossFeed = app.extensions.find(
(e) => e.name === 'pysssss.ImageFeed', (e) => e.name === 'pysssss.ImageFeed',
) )
if (pythongossFeed) { if (pythongossFeed) {
console.warn( console.warn(
"[mtb] - Aborting the loading of mtb's imageFeed in favor of pysssss.ImageFeed", "[mtb] - Aborting the loading of mtb's imageFeed in favor of pysssss.ImageFeed",
) )
activated = false // just in case other methods are added later on activated = false // just in case other methods are added later on
return return
} }
// - HTML & CSS // - HTML & CSS
//- lightbox //- lightbox
const lightboxContainer = document.createElement('div') const lightboxContainer = document.createElement('div')
Object.assign(lightboxContainer.style, styles.lighbox) Object.assign(lightboxContainer.style, styles.lighbox)
const lightboxImage = document.createElement('img') const lightboxImage = document.createElement('img')
Object.assign(lightboxImage.style, { Object.assign(lightboxImage.style, {
maxHeight: '100%', maxHeight: '100%',
maxWidth: '100%', maxWidth: '100%',
borderRadius: '5px', borderRadius: '5px',
}) })
// previous and next buttons // previous and next buttons
const lightboxPrevBtn = document.createElement('button') const lightboxPrevBtn = document.createElement('button')
const lightboxNextBtn = document.createElement('button') const lightboxNextBtn = document.createElement('button')
lightboxPrevBtn.textContent = '❮' lightboxPrevBtn.textContent = '❮'
lightboxNextBtn.textContent = '❯' lightboxNextBtn.textContent = '❯'
Object.assign(lightboxPrevBtn.style, styles.lightboxBtn({ left: '0%' })) Object.assign(lightboxPrevBtn.style, styles.lightboxBtn({ left: '0%' }))
Object.assign(lightboxNextBtn.style, styles.lightboxBtn({ right: '0%' })) Object.assign(lightboxNextBtn.style, styles.lightboxBtn({ right: '0%' }))
// close button // close button
const lightboxCloseBtn = document.createElement('button') const lightboxCloseBtn = document.createElement('button')
Object.assign( Object.assign(
lightboxCloseBtn.style, lightboxCloseBtn.style,
styles.lightboxBtn({ right: '0', top: '0' }), styles.lightboxBtn({ right: '0', top: '0' }),
) )
lightboxCloseBtn.textContent = '❌' lightboxCloseBtn.textContent = '❌'
const lightboxButtons = document.createElement('div') const lightboxButtons = document.createElement('div')
Object.assign(lightboxButtons.style, { Object.assign(lightboxButtons.style, {
position: 'absolute', position: 'absolute',
top: '0%', top: '0%',
right: '0%', right: '0%',
// transform: "translate(50%, -50%)", // transform: "translate(50%, -50%)",
height: '100%', height: '100%',
width: '100%', width: '100%',
background: 'none', background: 'none',
border: 'none', border: 'none',
color: '#fff', color: '#fff',
fontSize: '30px', fontSize: '30px',
cursor: 'pointer', cursor: 'pointer',
pointerEvents: 'none', pointerEvents: 'none',
}) })
lightboxButtons.append(lightboxPrevBtn, lightboxNextBtn, lightboxCloseBtn) lightboxButtons.append(lightboxPrevBtn, lightboxNextBtn, lightboxCloseBtn)
lightboxContainer.append(lightboxButtons, lightboxImage) lightboxContainer.append(lightboxButtons, lightboxImage)
//- image list //- image list
const imageListContainer = document.createElement('div') const imageListContainer = document.createElement('div')
Object.assign(imageListContainer.style, styles.img_list) Object.assign(imageListContainer.style, styles.img_list)
const createImgListBtn = (text, style) => { const createImgListBtn = (text, style) => {
const btn = document.createElement('button') const btn = document.createElement('button')
btn.type = 'button' btn.type = 'button'
btn.textContent = text btn.textContent = text
Object.assign(btn.style, { Object.assign(btn.style, {
...style, ...style,
border: 'none', border: 'none',
color: '#fff', color: '#fff',
background: 'none', background: 'none',
height: '20px', height: '20px',
cursor: 'pointer', cursor: 'pointer',
position: 'absolute', position: 'absolute',
top: '5px', top: '5px',
fontSize: '12px', fontSize: '12px',
lineHeight: '12px', lineHeight: '12px',
}) })
imageListContainer.append(btn) imageListContainer.append(btn)
return btn return btn
} }
const showBtn = document.createElement('button') const showBtn = document.createElement('button')
const closeBtn = createImgListBtn('❌', { const closeBtn = createImgListBtn('❌', {
width: '20px', width: '20px',
textIndent: '-4px', textIndent: '-4px',
right: '5px', right: '5px',
}) })
const loadButton = createImgListBtn('Load Session History', { const loadButton = createImgListBtn('Load Session History', {
right: '90px', right: '90px',
}) })
const clearButton = createImgListBtn('Clear', { const clearButton = createImgListBtn('Clear', {
right: '30px', right: '30px',
}) })
//- tools popup button //- tools popup button
showBtn.classList.add('comfy-settings-btn') showBtn.classList.add('comfy-settings-btn')
Object.assign(showBtn.style, { Object.assign(showBtn.style, {
right: '16px', right: '16px',
cursor: 'pointer', cursor: 'pointer',
display: 'none', display: 'none',
}) })
//- append to DOM //- append to DOM
document.body.append(imageListContainer) document.body.append(imageListContainer)
showBtn.textContent = '🖼' showBtn.textContent = '🖼'
showBtn.onclick = () => { showBtn.onclick = () => {
imageListContainer.style.display = 'block' imageListContainer.style.display = 'block'
showBtn.style.display = 'none' showBtn.style.display = 'none'
} }
document.querySelector('.comfy-settings-btn').after(showBtn) document.querySelector('.comfy-settings-btn').after(showBtn)
document.querySelector('.comfy-settings-btn').after(lightboxContainer) document.querySelector('.comfy-settings-btn').after(lightboxContainer)
// for (const { output } of history) { // for (const { output } of history) {
// if (output?.images) { // if (output?.images) {
// for (const src of output.images) { // for (const src of output.images) {
// const img = document.createElement("img"); // const img = document.createElement("img");
// const but = document.createElement("button"); // const but = document.createElement("button");
//- callbacks //- callbacks
closeBtn.onclick = () => { closeBtn.onclick = () => {
imageListContainer.style.display = 'none' imageListContainer.style.display = 'none'
showBtn.style.display = 'unset' showBtn.style.display = 'unset'
} }
clearButton.onclick = () => { clearButton.onclick = () => {
imageListContainer.replaceChildren(closeBtn, clearButton, loadButton) imageListContainer.replaceChildren(closeBtn, clearButton, loadButton)
} }
lightboxNextBtn.onclick = () => { lightboxNextBtn.onclick = () => {
currentImageIndex = (currentImageIndex + 1) % imageUrls.length currentImageIndex = (currentImageIndex + 1) % imageUrls.length
const imageUrl = imageUrls[currentImageIndex] const imageUrl = imageUrls[currentImageIndex]
lightboxImage.src = imageUrl lightboxImage.src = imageUrl
} }
// Modify the lightboxPrevBtn onclick callback // Modify the lightboxPrevBtn onclick callback
lightboxPrevBtn.onclick = () => { lightboxPrevBtn.onclick = () => {
currentImageIndex = currentImageIndex =
(currentImageIndex - 1 + imageUrls.length) % imageUrls.length (currentImageIndex - 1 + imageUrls.length) % imageUrls.length
const imageUrl = imageUrls[currentImageIndex] const imageUrl = imageUrls[currentImageIndex]
lightboxImage.src = imageUrl lightboxImage.src = imageUrl
} }
lightboxCloseBtn.onclick = () => { lightboxCloseBtn.onclick = () => {
lightboxContainer.style.display = 'none' lightboxContainer.style.display = 'none'
} }
lightboxImage.onclick = lightboxNextBtn.onclick lightboxImage.onclick = lightboxNextBtn.onclick
/** /**
* This is the function that creates the image buttons for the image list * 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 * They are wrapped in a button so that they can be clicked and open
* the image in the lightbox. * the image in the lightbox.
* @param {*} src * @param {*} src
*/ */
const createImageBtn = (src) => { const createImageBtn = (src) => {
console.debug(`making image ${src.filename}`) console.debug(`making image ${src.filename}`)
const img = document.createElement('img') const img = document.createElement('img')
const but = document.createElement('button') const but = document.createElement('button')
Object.assign(but.style, { Object.assign(but.style, {
height: '120px', height: '120px',
width: '120px', width: '120px',
border: 'none', border: 'none',
padding: 0, padding: 0,
margin: 0, margin: 0,
}) })
Object.assign(img.style, { Object.assign(img.style, {
width: '100%', width: '100%',
height: '100%', height: '100%',
objectFit: 'cover', objectFit: 'cover',
}) })
img.src = `/view?filename=${encodeURIComponent(src.filename)}&type=${ img.src = `/view?filename=${encodeURIComponent(src.filename)}&type=${
src.type src.type
}&subfolder=${encodeURIComponent(src.subfolder)}` }&subfolder=${encodeURIComponent(src.subfolder)}`
imageUrls.push(img.src) imageUrls.push(img.src)
console.debug(img.src) console.debug(img.src)
img.onload = () => { img.onload = () => {
but.style.width = `${120 * (img.naturalWidth / img.naturalHeight)}px` but.style.width = `${120 * (img.naturalWidth / img.naturalHeight)}px`
} }
but.onclick = () => { but.onclick = () => {
lightboxContainer.style.display = 'flex' lightboxContainer.style.display = 'flex'
// add the same image to the lightbox // add the same image to the lightbox
lightboxImage.src = img.src lightboxImage.src = img.src
// lighboxContainer.replaceChildren(lightboxButtons, img); // lighboxContainer.replaceChildren(lightboxButtons, img);
} }
// add right click menu // add right click menu
but.addEventListener('contextmenu', (e) => { but.addEventListener('contextmenu', (e) => {
e.preventDefault() e.preventDefault()
if (image_menu) { if (image_menu) {
image_menu.remove() image_menu.remove()
} }
image_menu = document.createElement('div') image_menu = document.createElement('div')
Object.assign(image_menu.style, { Object.assign(image_menu.style, {
position: 'absolute', position: 'absolute',
top: `${e.clientY}px`, top: `${e.clientY}px`,
left: `${e.clientX}px`, left: `${e.clientX}px`,
background: '#333', background: '#333',
color: '#fff', color: '#fff',
padding: '5px', padding: '5px',
borderRadius: '5px', borderRadius: '5px',
zIndex: 999, zIndex: 999,
}) })
const load_img = document.createElement('button') const load_img = document.createElement('button')
load_img.textContent = 'Load' load_img.textContent = 'Load'
load_img.onclick = () => { load_img.onclick = () => {
app.handleFile(img.src) app.handleFile(img.src)
} }
image_menu.appendChild(load_img) image_menu.appendChild(load_img)
document.body.appendChild(image_menu) document.body.appendChild(image_menu)
}) })
but.append(img) but.append(img)
imageListContainer.prepend(but) imageListContainer.prepend(but)
} }
loadButton.onclick = async () => { loadButton.onclick = async () => {
const all_history = await api.getHistory() const all_history = await api.getHistory()
for (const history of all_history.History) { for (const history of all_history.History) {
if (history.outputs) { if (history.outputs) {
for (const key of Object.keys(history.outputs)) { for (const key of Object.keys(history.outputs)) {
console.debug(key) console.debug(key)
if (history.outputs[key].images) { if (history.outputs[key].images) {
for (const im of history.outputs[key].images) { for (const im of history.outputs[key].images) {
console.debug(im) console.debug(im)
createImageBtn(im) createImageBtn(im)
} }
} }
} }
// for (const src of outputs.outputs.images) { // for (const src of outputs.outputs.images) {
// console.debug(src) // console.debug(src)
// makeImage(`${src.subfolder}/${src.filename}`) // makeImage(`${src.subfolder}/${src.filename}`)
// } // }
} }
} }
} }
///////------- ///////-------
// const all_history = await api.getHistory() // const all_history = await api.getHistory()
// for (const history of all_history.History) { // for (const history of all_history.History) {
// if (history.outputs) { // if (history.outputs) {
// for (const key of Object.keys(history.outputs)) { // for (const key of Object.keys(history.outputs)) {
// for (const im of history.outputs[key].images) { // for (const im of history.outputs[key].images) {
// makeImage(im) // makeImage(im)
// } // }
// } // }
// // for (const src of outputs.outputs.images) { // // for (const src of outputs.outputs.images) {
// // console.debug(src) // // console.debug(src)
// // makeImage(`${src.subfolder}/${src.filename}`) // // makeImage(`${src.subfolder}/${src.filename}`)
// // } // // }
// } // }
// } // }
//- Hook into the API //- Hook into the API
api.addEventListener('executed', ({ detail }) => { api.addEventListener('executed', ({ detail }) => {
if (detail?.output?.images) { if (detail?.output?.images) {
for (const src of detail.output.images) { for (const src of detail.output.images) {
console.debug(`Adding ${src} to image feed`) console.debug(`Adding ${src} to image feed`)
createImageBtn(src) createImageBtn(src)
} }
} }
}) })
}, },
}) })
+505 -304
View File
@@ -2,8 +2,8 @@
import { app } from '../../scripts/app.js' import { app } from '../../scripts/app.js'
import { api } from '../../scripts/api.js' import { api } from '../../scripts/api.js'
import { infoLogger, successLogger, errorLogger } from './comfy_shared.js'
import * as mtb_ui from './mtb_ui.js'
import * as shared from './comfy_shared.js' import * as shared from './comfy_shared.js'
import { import {
@@ -13,38 +13,140 @@ import {
makeSelect, makeSelect,
makeSlider, makeSlider,
renderSidebar, renderSidebar,
ContextMenu,
} from './mtb_ui.js' } from './mtb_ui.js'
let currentAbortController = null
/** cursor/offset of where we are at */
const offset = 0 const offset = 0
// These are "global" variables mostly meant to sync user settings. /** width of the images in the grid */
let currentWidth = 200 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 currentMode = 'input'
let subfolder = '' let subfolder = ''
let currentSort = 'None' 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 VIDEO_NODES = ['VHS_LoadVideo']
const PROCESSED_PROMPT_IDS = new Set() 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) => { const updateImage = (node, image) => {
if (IMAGE_NODES.includes(node.type)) { switch (node.type) {
const w = node.widgets?.find((w) => w.name === 'image') case 'LoadImage': {
if (w) { if (subfolder && subfolder !== '') {
w.value = image app.extensionManager.toast.add({
w.callback() 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)) { case 'VHS_LoadImagePath': {
const w = node.widgets?.find((w) => w.name === 'video') let value = image
if (w) {
node.updateParameters({ filename: image }, true) 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 * @param {ResultItem} resultItem
* @returns {string} - The request URL. * @returns {string} - The request URL.
*/ */
const resultItemToQuery = (resultItem) => { const resultItemToQuery = (resultItem) =>
const res = [ [
`/mtb/view?filename=${resultItem.filename}`, `/mtb/view?filename=${resultItem.filename}`,
`width=512`,
`type=${resultItem.type}`, `type=${resultItem.type}`,
`subfolder=${resultItem.subfolder}`, `subfolder=${resultItem.subfolder}`,
'preview=', `preview=`,
] ].join('&')
if (targetWidth > 0) {
res.splice(1, 0, `width=${targetWidth}`)
}
return res.join('&')
}
/** /**
* Retrieves the unique prompt ID from a history task item. * Retrieves the unique prompt ID from a history task item.
* @param {HistoryTaskItem} historyTaskItem * @param {HistoryTaskItem} historyTaskItem
@@ -91,7 +189,7 @@ const getNewOutputUrls = (mostRecentTask) => {
const imageOutputs = Object.values(nodeOutputs.images) const imageOutputs = Object.values(nodeOutputs.images)
imageOutputs.forEach( imageOutputs.forEach(
(resultItem) => (resultItem) =>
(urls[resultItem.filename] = resultItemToQuery(resultItem)), (urls[resultItem.filename] = resultItemToQuery(resultItem))
) )
} }
// Can process `animated` and `audio` outputs here. // Can process `animated` and `audio` outputs here.
@@ -120,94 +218,265 @@ const updateOutputsGrid = async () => {
} }
const getImgsFromUrls = (urls, target, options = { prepend: false }) => { const getImgsFromUrls = (urls, target, options = { prepend: false }) => {
const imgs = [] if (currentAbortController) {
if (urls === undefined) { currentAbortController.abort()
return imgs
} }
const elem = currentMode === 'video' ? 'video' : 'img' infoLogger('getting images from urls', urls)
for (const [key, url] of Object.entries(urls)) { currentAbortController = new AbortController()
const a = makeElement(elem) const { signal } = currentAbortController
a.src = url const imgs = []
a.width = currentWidth if (!urls) return imgs
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
}
for (const [_id, node] of Object.entries(app.canvas.selected_nodes)) { const loadingIndicator = document.createElement('div')
updateImage(node, key) loadingIndicator.className = 'mtb-loading-indicator'
} if (target) target.appendChild(loadingIndicator)
}
} else if (currentMode === 'output') {
a.onclick = (_e) => {
// window.MTB?.notify?.("Output import isn't supported yet...", 5000)
if (subfolder !== '') {
app.extensionManager.toast.add({
severity: 'warn',
summary: 'Subfolder not supported',
detail: "The LoadImage node doesn't support subfolders",
life: 5000,
})
return
}
app.extensionManager.toast.add({ const totalImages = Object.keys(urls).length
severity: 'warn', let loadedCount = 0
summary: 'Outputs not supported', const updateLoadingStatus = () => {
detail: loadingIndicator.textContent = `Loaded ${loadedCount} of ${totalImages} images`
'For now only inputs can be clicked to load the image on the active LoadImage node.', }
life: 5000, 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 elem.onclick = (_e) => {
a.loop = true const selected = app.canvas.selected_nodes
a.onclick = (_e) => { if (!selected || Object.keys(selected).length === 0) {
const selected = app.canvas.selected_nodes app.extensionManager.toast.add({
if (selected && Object.keys(selected).length === 0) { severity: 'warn',
app.extensionManager.toast.add({ summary: 'No node selected!',
severity: 'warn', detail: 'Please select a node first.',
summary: 'No node selected!', life: 5000,
detail: })
"For now the only action when clicking videos in the sidebar is to set the video on all selected 'Load Video (Upload)' nodes.", return
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)) { loadedCount++
updateImage(node, key) 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) const BATCH_SIZE = 20
} for (let i = 0; i < Object.entries(urls).length; i += BATCH_SIZE) {
if (target !== undefined) { const batch = Object.entries(urls).slice(i, i + BATCH_SIZE)
if (options.prepend) target.prepend(...imgs) 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) 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 () => { const getModes = async () => {
@@ -216,11 +485,11 @@ const getModes = async () => {
} }
const getUrls = async (subfolder) => { const getUrls = async (subfolder) => {
const count = (await api.getSetting('mtb.io-sidebar.count')) || 1000 const count = (await api.getSetting('mtb.io-sidebar.count')) || 1000
console.log('Sidebar count', count) console.debug('Sidebar count', count)
if (currentMode === 'video') { if (currentMode === 'video') {
const output = await shared.runAction( const output = await shared.runAction(
'getUserVideos', 'getUserVideos',
targetWidth, 256,
count, count,
offset, offset,
currentSort, currentSort,
@@ -230,17 +499,108 @@ const getUrls = async (subfolder) => {
const output = await shared.runAction( const output = await shared.runAction(
'getUserImages', 'getUserImages',
currentMode, currentMode,
targetWidth,
count, count,
offset, offset,
currentSort, currentSort,
false, false,
subfolder, subfolder,
saltUrls,
) )
return output || {} 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 //NOTE: do not load if using the old ui
if (window?.__COMFYUI_FRONTEND_VERSION__) { if (window?.__COMFYUI_FRONTEND_VERSION__) {
// NOTE: removed this for now since I'm not actually exposing anything a client // 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 = { const sidebar_extension = {
name: 'mtb.io-sidebar', 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', id: 'mtb.io-sidebar.count',
category: ['mtb', 'Input & Output Sidebar', 'count'], category: ['mtb', 'Input & Output Sidebar', 'count'],
name: 'Number of images to fetch', name: 'Number of images to fetch',
type: 'number', type: 'number',
defaultValue: 1000, defaultValue: 1000,
tooltip: 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)", "This setting affects the input/output sidebar to determine how many images to fetch per pagination (pagination is not yet supported so for now it's the static total)",
}, attrs: {
{ style: {
id: 'mtb.io-sidebar.salt_urls', // fontFamily: 'monospace',
category: ['mtb', 'Input & Output Sidebar', 'salt_urls'], },
name: 'Salt URLs',
type: 'boolean',
defaultValue: false,
onChange: (n, o) => {
saltUrls = n
}, },
tooltip: })
'Adds a random query parameter to every urls to always invalidate caching.',
}, app.ui.settings.addSetting({
{
id: 'mtb.io-sidebar.img-size', id: 'mtb.io-sidebar.img-size',
category: ['mtb', 'Input & Output 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, 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 }) tooltip: "It's recommended to keep it at 512px",
attrs: {
const baseId = name.replace(/[^a-zA-Z0-9]/g, '-').toLowerCase() style: {
const checkboxId = `${baseId}-checkbox` // fontFamily: 'monospace',
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: app.ui.settings.addSetting({
"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.",
},
{
id: 'mtb.io-sidebar.sort', id: 'mtb.io-sidebar.sort',
category: ['mtb', 'Input & Output Sidebar', 'sort'], category: ['mtb', 'Input & Output Sidebar', 'sort'],
name: 'Default sort mode', name: 'Default sort mode',
@@ -372,39 +677,7 @@ if (window?.__COMFYUI_FRONTEND_VERSION__) {
'Name', 'Name',
'Name-Reverse', '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({ app.extensionManager.registerSidebarTab({
id: 'mtb-inputs-outputs', id: 'mtb-inputs-outputs',
@@ -420,81 +693,9 @@ if (window?.__COMFYUI_FRONTEND_VERSION__) {
handle = undefined handle = undefined
} }
if (el.parentNode) { if (!loaded_images) {
el.parentNode.style.overflowY = 'clip' 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 () => { app.api.addEventListener('status', async () => {
if (currentMode !== 'output') return if (currentMode !== 'output') return
updateOutputsGrid() updateOutputsGrid()
-529
View File
@@ -1,529 +0,0 @@
/** Python REPL for the frontend (uses rich)*/
import { app } from '../../scripts/app.js'
import * as shared from './comfy_shared.js'
import * as mtb_ui from './mtb_ui.js'
class ComfyREPL extends LiteGraph.LGraphNode {
constructor() {
super()
this.shape = LiteGraph.BOX_SHAPE
this.isVirtualNode = true
this.category = 'mtb/repl'
this.title = '🐍 REPL (mtb)'
this.uuid = shared.makeUUID()
this.size = [600, 400]
// Create a container for our custom widgets
this.widget = this.addDOMWidget('HTML', 'html', this.createREPLWidget())
this.loadAceEditor()
// Store input and output for persistence
this.properties = {
inputCode: '',
outputHistory: '',
}
this.outputArea.innerHTML = this.properties.outputHistory
this.outputArea.scrollTop = this.outputArea.scrollHeight
// Debounced linting function
this.debouncedLint = shared.debounce(this.lintCode.bind(this), 500)
// Resizing state variables
this.isResizing = false
this.initialMouseY = 0
this.initialInputHeight = 0
this.initialOutputHeight = 0
}
loadAceEditor() {
if (window.MTB?.ace_loaded) {
return
}
let NEED_PATCH = false
if (window.ace) {
shared.infoLogger(
'A global ace was found in scope, to avoid issues with it we will patch it',
)
NEED_PATCH = true
// window._backupAce = window.ace
// window.ace = null
}
shared
.loadScript('/mtb_async/ace/ace.js')
.then((m) => {
shared.infoLogger('ACE was loaded', m)
// window.MTB_ACE = window.ace
window.MTB.ace_loaded = true
this.initAceEditor()
this.aceEditor.setValue(this.properties.inputCode, -1)
})
.catch((e) => {
shared.errorLogger(e)
})
.finally(() => {
if (NEED_PATCH) {
console.log('Patching back window object')
window.ace = window._backupAce
}
})
}
initAceEditor() {
if (!window.MTB.ace_loaded) {
console.error('ACE editor not loaded. Cannot set up editors.')
return
}
if (!this.inputDiv) {
console.error('Input div not found for Ace editor initialization.')
return
}
this.aceEditor = ace.edit(this.inputDiv)
this.aceEditor.setTheme('ace/theme/monokai') //"ace/theme/dracula", "ace/theme/github"
this.aceEditor.session.setMode('ace/mode/python')
this.aceEditor.setOptions({
enableBasicAutocompletion: true,
enableLiveAutocompletion: true,
enableSnippets: true,
fontSize: '14px',
fontFamily: 'monospace',
showPrintMargin: false,
wrap: true,
tabSize: 4,
useSoftTabs: true,
highlightActiveLine: true,
highlightSelectedWord: true,
cursorStyle: 'ace', // "ace" | "slim" | "smooth" | "wide"
behavioursEnabled: true,
displayIndentGuides: true,
fixedWidthGutter: true,
scrollPastEnd: 0.5,
})
// Custom keybinding for Ctrl+Enter
this.aceEditor.commands.addCommand({
name: 'runCode',
bindKey: { win: 'Ctrl-Enter', mac: 'Command-Enter' },
exec: () => this.executeCode(),
})
// Listen for changes to trigger linting
let lintDisabled = false
this.aceEditor.session.on('change', () => {
if (!lintDisabled) {
this.debouncedLint()
}
})
this.outputArea.scrollTop = this.outputArea.scrollHeight
}
addOutput(html) {
this.outputArea.innerHTML += html
this.properties.outputHistory += html
this.outputArea.scrollTop = this.outputArea.scrollHeight
}
createREPLWidget() {
const container = mtb_ui.makeElement('div', {
display: 'flex',
flexDirection: 'column',
width: '100%',
height: '100%',
boxSizing: 'border-box',
padding: '5px',
})
this.inputDiv = mtb_ui.makeElement(
'div',
{
width: 'calc(100% - 10px)',
height: '100px',
backgroundColor: '#333',
color: '#eee',
border: '1px solid #555',
borderRadius: '4px',
marginBottom: '5px',
boxSizing: 'border-box',
overflow: 'hidden',
},
container,
)
// Resizable Handle
this.handleDiv = mtb_ui.makeElement(
'div',
{
width: '100%',
height: '5px',
backgroundColor: '#666',
cursor: 'ns-resize',
marginBottom: '5px',
borderRadius: '2px',
},
container,
)
this.handleDiv.addEventListener('mousedown', this.startResizing.bind(this))
// Run Button
this.runButton = mtb_ui.makeElement(
'button',
{
width: '100%',
padding: '8px',
backgroundColor: '#555',
color: '#fff',
border: 'none',
borderRadius: '4px',
cursor: 'pointer',
marginBottom: '5px',
fontSize: '14px',
},
container,
)
this.runButton.textContent = 'Run Code (Ctrl+Enter)'
this.runButton.onclick = () => this.executeCode()
// Clear Button
this.clearButton = mtb_ui.makeElement(
'button',
{
width: '100%',
padding: '8px',
backgroundColor: '#555',
color: '#fff',
border: 'none',
borderRadius: '4px',
cursor: 'pointer',
marginBottom: '5px',
fontSize: '14px',
},
container,
)
this.clearButton.textContent = 'Clear Output'
this.clearButton.onclick = () => {
this.outputArea.innerHTML = ''
this.properties.outputHistory = ''
}
// Output Area
this.outputArea = mtb_ui.makeElement(
'div',
{
flexGrow: '1',
width: 'calc(100% - 10px)',
backgroundColor: '#222',
color: '#ddd',
border: '1px solid #555',
borderRadius: '4px',
padding: '5px',
fontFamily: 'monospace',
fontSize: '14px',
overflowY: 'auto',
whiteSpace: 'pre-wrap',
boxSizing: 'border-box',
},
container,
)
return container
}
// --- Resizing Logic ---
startResizing(e) {
if (!this.inputDiv) {
shared.infoLogger("The input div isn't ready", this)
shared.errorLogger("The input div isn't ready")
return
}
this.isResizing = true
this.initialMouseY = e.clientY
this.initialInputHeight = this.inputDiv.offsetHeight
this.initialOutputHeight = this.outputArea.offsetHeight
document.addEventListener('mousemove', this.doResize.bind(this))
document.addEventListener('mouseup', this.stopResizing.bind(this))
document.body.style.cursor = 'ns-resize' // Change cursor globally
}
doResize(e) {
if (!this.isResizing) return
const deltaY = e.clientY - this.initialMouseY
let new_input_height = this.initialInputHeight + deltaY
let new_output_height = this.initialOutputHeight - deltaY
const minInputHeight = 50 // Minimum height for Ace editor
const minOutputHeight = 50 // Minimum height for output area
// Clamp heights to minimums
if (new_input_height < minInputHeight) {
new_input_height = minInputHeight
new_output_height =
this.initialInputHeight + this.initialOutputHeight - minInputHeight
}
if (new_output_height < minOutputHeight) {
new_output_height = minOutputHeight
new_input_height =
this.initialInputHeight + this.initialOutputHeight - minOutputHeight
}
this.inputDiv.style.height = `${new_input_height}px`
this.outputArea.style.height = `${new_output_height}px`
// Update the stored ratio for persistence
const totalDynamicHeight =
this.inputDiv.offsetHeight + this.outputArea.offsetHeight
if (totalDynamicHeight > 0) {
this.properties.inputHeightRatio = new_input_height / totalDynamicHeight
}
this.aceEditor.resize() // Important for Ace to redraw
}
stopResizing() {
this.isResizing = false
document.removeEventListener('mousemove', this.doResize)
document.removeEventListener('mouseup', this.stopResizing)
document.body.style.cursor = '' // Restore default cursor
}
// --- End Resizing Logic ---
async executeCode() {
const code = this.aceEditor.getValue()
if (!code.trim()) {
return
}
const inputPrompt = `<div style="color:#888; margin-top: 10px;">>>> ${code}</div>`
this.addOutput(inputPrompt)
try {
const response = await fetch('/mtb/execute', {
method: 'POST',
headers: {
'Content-Type': 'application/json',
},
body: JSON.stringify({ code: code, name: this.uuid }),
})
if (!response.ok) {
throw new Error(`HTTP error! status: ${response.status}`)
}
const result = await response.json()
console.debug('Received from backend', result)
const outputHtml = result.output_html || ''
const error = result.error
if (error) {
this.addOutput(
`<div style="color: #f00; font-weight: bold;">Error:</div>${outputHtml}`,
)
} else {
this.addOutput(outputHtml)
}
} catch (e) {
const errorMessage = `<div style="color: #f00;">Frontend Error: ${e.message}</div>`
this.addOutput(errorMessage)
console.error('ComfyREPL Frontend Error:', e)
} finally {
// Not clearing
// this.inputArea.value = '' // Clear input after execution
// this.properties.inputCode = '' // Clear persisted input
}
}
async lintCode() {
if (!this.aceEditor) {
return
}
const code = this.aceEditor.getValue()
if (!code.trim()) {
this.aceEditor.session.setAnnotations([]) // Clear annotations if empty
return
}
try {
const response = await fetch('/mtb/lint', {
// New linting endpoint
method: 'POST',
headers: {
'Content-Type': 'application/json',
},
body: JSON.stringify({ code: code, name: this.uuid }),
})
if (!response.ok) {
console.log(response)
throw new Error(
`HTTP error! status: ${response.status} ${response.statusText}`,
)
}
const result = await response.json()
// result.diagnostics should be an array of {row, column, text, type}
this.aceEditor.session.setAnnotations(result.diagnostics)
} catch (e) {
console.error('ComfyREPL Linting Error:', e)
this.aceEditor.session.setAnnotations([
{
row: 0,
column: 0,
text: `Linting failed: ${e.message}`,
type: 'error',
},
])
}
}
// Restore properties when loading a graph
onConfigure() {
if (this.properties.inputCode && this.aceEditor) {
this.aceEditor.setValue(this.properties.inputCode, -1)
}
// if (this.properties.inputCode) {
// this.inputArea.value = this.properties.inputCode
// }
if (this.properties.outputHistory) {
this.outputArea.innerHTML = this.properties.outputHistory
this.outputArea.scrollTop = this.outputArea.scrollHeight
}
if (this.properties.uuid) {
this.uuid = this.properties.uuid
}
this.debouncedLint()
this.onResize(this.size)
}
// Save properties when saving a graph
onSerialize(o) {
if (this.aceEditor) {
o.properties.inputCode = this.aceEditor.getValue() //this.inputArea.value
}
o.properties.outputHistory = this.outputArea.innerHTML
o.properties.uuid = this.uuid
o.properties.inputHeightRatio = this.properties.inputHeightRatio
}
onRemoved() {
// Clean up DOM elements when node is removed
if (this.widget?.element?.parentNode) {
this.widget.element.parentNode.removeChild(this.widget.element)
}
// Destroy Ace editor instance to prevent memory leaks
if (this.aceEditor) {
this.aceEditor.destroy()
this.aceEditor.container.remove() // Remove the Ace container div from DOM
}
// Clean up global event listeners if node is removed while resizing
document.removeEventListener('mousemove', this.doResize)
document.removeEventListener('mouseup', this.stopResizing)
document.body.style.cursor = ''
}
// LiteGraph method to handle node resizing
onResize(size) {
// Call parent method if it exists (important for LiteGraph's internal sizing)
if (super.onResize) {
super.onResize(size)
}
// Adjust container size
const container = this.widget.element
container.style.width = `${size[0] - 10}px` // Account for padding
container.style.height = `${size[1] - 10}px`
// Adjust input and output area widths
this.inputDiv.style.width = 'calc(100% - 10px)'
this.outputArea.style.width = 'calc(100% - 10px)'
//
// const old = () => {
// // Calculate remaining height for output area
// // Ace editor manages its own height within this.inputDiv, so we use offsetHeight
// const inputHeight = this.inputDiv.offsetHeight
// const runButtonHeight = this.runButton.offsetHeight
// const clearButtonHeight = this.clearButton.offsetHeight
// const totalFixedHeight =
// inputHeight + runButtonHeight + clearButtonHeight + 15 // 15 for margins/padding
//
// const remainingHeight = size[1] - 10 - totalFixedHeight
// this.outputArea.style.height = `${Math.max(50, remainingHeight)}px` // Min height 50px
// }
// Calculate dynamic heights
const containerHeight = size[1] - 10
const handleHeight = this.handleDiv.offsetHeight
const buttonHeights =
this.runButton.offsetHeight + this.clearButton.offsetHeight + 15 // Sum of button heights + margins
const dynamicContentHeight = containerHeight - buttonHeights - handleHeight
const minInputHeight = 50
const minOutputHeight = 50
let inputHeight = Math.max(
minInputHeight,
dynamicContentHeight * (this.properties.inputHeightRatio || 1.0),
)
let outputHeight = Math.max(
minOutputHeight,
dynamicContentHeight - inputHeight,
)
//
// // Re-distribute if one hits its minimum
// if (
// inputHeight === minInputHeight &&
// dynamicContentHeight - minInputHeight > minOutputHeight
// ) {
// outputHeight = dynamicContentHeight - minInputHeight
// } else if (
// outputHeight === minOutputHeight &&
// dynamicContentHeight - minOutputHeight > minInputHeight
// ) {
// inputHeight = dynamicContentHeight - minOutputHeight
// }
//
// // Final check to ensure total height matches available dynamic space
// const currentTotal = inputHeight + outputHeight
// if (currentTotal !== dynamicContentHeight) {
// // Adjust one of them if there's a small discrepancy due to rounding
// if (inputHeight > minInputHeight) {
// inputHeight += dynamicContentHeight - currentTotal
// } else if (outputHeight > minOutputHeight) {
// outputHeight += dynamicContentHeight - currentTotal
// }
// }
this.inputDiv.style.height = `${inputHeight}px`
this.outputArea.style.height = `${outputHeight}px`
// Update the ratio based on the actual heights set
if (dynamicContentHeight > 0) {
this.properties.inputHeightRatio = inputHeight / dynamicContentHeight
}
// Inform Ace editor about the resize so it can redraw its content
if (this.aceEditor) {
this.aceEditor.resize()
}
}
}
const repl = {
name: 'mtb.repl',
registerCustomNodes() {
LiteGraph.registerNodeType('Python REPL', ComfyREPL)
},
}
app.registerExtension(repl)
+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 { .mtb_slider[type="range"]:active::-webkit-slider-thumb {
background-color: ${S.accent}; 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( addNamedStyleSheet(
'mtb_ui', 'mtb_ui',
` `
${common} ${common}
${inputs} ${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 * 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. * @param {Object} [style] - CSS styles to apply to the element.
* @returns {HTMLElement} - The created DOM 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 [real_kind, className] = kind.split('.')
let id let id
@@ -224,9 +309,6 @@ export const makeElement = (kind, style, parent) => {
if (id) { if (id) {
el.id = id el.id = id
} }
if (parent) {
parent.appendChild(el)
}
return el return el
} }
+45 -84
View File
@@ -694,7 +694,7 @@ const mtb_widgets = {
app.ui.settings.addSetting({ app.ui.settings.addSetting({
id: 'mtb.Main.debug-enabled', id: 'mtb.Main.debug-enabled',
category: ['mtb', ' Main', 'debug-enabled'], category: ['mtb', 'Main', 'debug-enabled'],
name: 'Enable Debug (py and js)', name: 'Enable Debug (py and js)',
type: 'boolean', type: 'boolean',
defaultValue: false, defaultValue: false,
@@ -1012,15 +1012,12 @@ const mtb_widgets = {
) )
loop_preview.value = 'Iteration: Idle' loop_preview.value = 'Iteration: Idle'
let cancelQueue = false
const onReset = () => { const onReset = () => {
raw_iteration.value = 0 raw_iteration.value = 0
raw_loop.value = 0 raw_loop.value = 0
value_preview.value = 'Idle' value_preview.value = 'Idle'
loop_preview.value = 'Iteration: Idle' loop_preview.value = 'Iteration: Idle'
cancelQueue = false
app.canvas.setDirty(true) app.canvas.setDirty(true)
} }
@@ -1029,42 +1026,15 @@ const mtb_widgets = {
this.addWidget('button', 'Reset', 'reset', onReset) this.addWidget('button', 'Reset', 'reset', onReset)
// run button // run button
const chunkSize = 10 this.addWidget('button', 'Queue', 'queue', () => {
this.addWidget('button', 'Queue', 'queue', async () => { onReset() // this could maybe be a setting or checkbox
onReset() app.queuePrompt(0, total_frames.value * loop_count.value)
const totalPrompts = total_frames.value * loop_count.value
window.MTB?.notify?.( 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, 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 = () => { this.onRemoved = () => {
@@ -1196,9 +1166,7 @@ const mtb_widgets = {
//NOTE: dynamic nodes //NOTE: dynamic nodes
case 'Apply Text Template (mtb)': { case 'Apply Text Template (mtb)': {
shared.setupDynamicConnections(nodeType, 'var', '*', { shared.setupDynamicConnections(nodeType, 'var', '*')
rename_menu: 'name',
})
break break
} }
case 'Save Data Bundle (mtb)': { case 'Save Data Bundle (mtb)': {
@@ -1330,54 +1298,47 @@ const mtb_widgets = {
const related = new Set([this.id]) const related = new Set([this.id])
const visited = new Set() const visited = new Set()
if (this.outputs[0].links) { if (this.outputs[0].links) {
for (const linkId of this.outputs[0].links) { const initLink = this.outputs[0].links[0]
const { to: loopEnd } = shared.nodesFromLink(this, linkId) const { to: loopEnd } = shared.nodesFromLink(this, initLink)
const canReachEnd = (node, visited = new Set()) => { const canReachEnd = (node, visited = new Set()) => {
if (node === loopEnd) return true if (node === loopEnd) return true
if (visited.has(node.id)) return false if (visited.has(node.id)) return false
visited.add(node.id) visited.add(node.id)
for (const output of node.outputs || []) { for (const output of node.outputs || []) {
if (!output.links) continue if (!output.links) continue
for (const linkId of output.links) { for (const linkId of output.links) {
const { to: nextNode } = shared.nodesFromLink( const { to: nextNode } = shared.nodesFromLink(node, linkId)
node, if (!nextNode) continue
linkId, if (canReachEnd(nextNode, visited)) {
) return true
if (!nextNode) continue
if (canReachEnd(nextNode, visited)) {
return true
}
}
}
return false
}
const traverseNodes = (node) => {
if (visited.has(node.id)) return
visited.add(node.id)
// can reach the end
if (node !== this && node !== loopEnd && !canReachEnd(node)) {
return
}
related.add(node.id)
for (const output of node.outputs || []) {
if (!output.links) continue
for (const linkId of output.links) {
const { to: nextNode } = shared.nodesFromLink(
node,
linkId,
)
if (!nextNode) continue
traverseNodes(nextNode)
} }
} }
} }
return false
traverseNodes(this)
} }
const traverseNodes = (node) => {
if (visited.has(node.id)) return
visited.add(node.id)
// can reach the end
if (node !== this && node !== loopEnd && !canReachEnd(node)) {
return
}
related.add(node.id)
for (const output of node.outputs || []) {
if (!output.links) continue
for (const linkId of output.links) {
const { to: nextNode } = shared.nodesFromLink(node, linkId)
if (!nextNode) continue
traverseNodes(nextNode)
}
}
}
traverseNodes(this)
} }
this.related_to_flow = Array.from(related) this.related_to_flow = Array.from(related)
this.computed_flow = true this.computed_flow = true
+52 -64
View File
@@ -1,13 +1,10 @@
// web/note_plus.constants.js // web/note_plus.constants.js
export const DEFAULT_CSS = `/** here you can write css**/ export const DEFAULT_CSS = ''
h1 {
color: whitesmoke;
}`
export const DEFAULT_HTML = `<p style='color:red;font-family:monospace'> export const DEFAULT_HTML = `<p style='color:red;font-family:monospace'>
Note+ Note+
</p>` </p>`
export const DEFAULT_MD = '# 📝 Note+' export const DEFAULT_MD = '## Note+'
export const DEFAULT_MODE = 'markdown' export const DEFAULT_MODE = 'markdown'
export const DEFAULT_THEME = 'one_dark' export const DEFAULT_THEME = 'one_dark'
@@ -58,57 +55,58 @@ We also support github callout:
` `
export const THEMES = [ export const THEMES = [
'ambiance', 'ambiance',
'chaos', 'chaos',
'chrome', 'chrome',
'cloud9_day', 'cloud9_day',
'cloud9_night', 'cloud9_night',
'cloud9_night_low_color', 'cloud9_night_low_color',
'cloud_editor', 'cloud_editor',
'cloud_editor_dark', 'cloud_editor_dark',
'clouds', 'clouds',
'clouds_midnight', 'clouds_midnight',
'cobalt', 'cobalt',
'crimson_editor', 'crimson_editor',
'dawn', 'dawn',
'dracula', 'dracula',
'dreamweaver', 'dreamweaver',
'eclipse', 'eclipse',
'github', 'github',
'github_dark', 'github_dark',
'gob', 'gob',
'gruvbox', 'gruvbox',
'gruvbox_dark_hard', 'gruvbox_dark_hard',
'gruvbox_light_hard', 'gruvbox_light_hard',
'idle_fingers', 'idle_fingers',
'iplastic', 'iplastic',
'katzenmilch', 'katzenmilch',
'kr_theme', 'kr_theme',
'kuroir', 'kuroir',
'merbivore', 'merbivore',
'merbivore_soft', 'merbivore_soft',
'mono_industrial', 'mono_industrial',
'monokai', 'monokai',
'nord_dark', 'nord_dark',
'one_dark', 'one_dark',
'pastel_on_dark', 'pastel_on_dark',
'solarized_dark', 'solarized_dark',
'solarized_light', 'solarized_light',
'sqlserver', 'sqlserver',
'terminal', 'terminal',
'textmate', 'textmate',
'tomorrow', 'tomorrow',
'tomorrow_night', 'tomorrow_night',
'tomorrow_night_blue', 'tomorrow_night_blue',
'tomorrow_night_bright', 'tomorrow_night_bright',
'tomorrow_night_eighties', 'tomorrow_night_eighties',
'twilight', 'twilight',
'vibrant_ink', 'vibrant_ink',
'vscode', 'vscode',
] ]
export const CSS_RESET = ` export const CSS_RESET = `
* { * {
font-family: monospace;
line-height: 1.25em; line-height: 1.25em;
} }
.shiki{ .shiki{
@@ -118,8 +116,6 @@ export const CSS_RESET = `
.markdown-callout-title { .markdown-callout-title {
.octicon{ .octicon{
fill:white; fill:white;
width:29px;
height:29px;
} }
/* background: var(--current-color); */ /* background: var(--current-color); */
color: var(--current-color); color: var(--current-color);
@@ -128,8 +124,6 @@ export const CSS_RESET = `
/* border-start-start-radius: var(--radius); */ /* border-start-start-radius: var(--radius); */
padding: 0.5em; padding: 0.5em;
padding-inline-start: 1em; padding-inline-start: 1em;
display: flex;
align-items: center;
} }
.markdown-callout-content { .markdown-callout-content {
padding: 1em; padding: 1em;
@@ -142,12 +136,7 @@ export const CSS_RESET = `
border-left: 3px solid var(--current-color); border-left: 3px solid var(--current-color);
margin-bottom: 1em; margin-bottom: 1em;
margin-top: 1em; margin-top: 1em;
} }
.markdown-callout p:nth-child(2) {
padding:1em;
}
.markdown-callout-tip { .markdown-callout-tip {
--text-color: whitesmoke; --text-color: whitesmoke;
@@ -175,9 +164,8 @@ export const CSS_RESET = `
flex-direction:column; flex-direction:column;
align-items: flex-start; align-items: flex-start;
width:95%; width:95%;
/*margin-left: 20px;*/ margin-left: 20px;
/*margin-top:20px;*/ margin-top:20px;
/*background-color: rgba(255,0,0,0.5)!important;*/ /*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; transition-duration: ${transition_time}ms;
` `
function notify(message, timeout = 3000, old_mode = false) { function notify(message, timeout = 3000) {
if (!old_mode) {
app.extensionManager.toast.add({
severity: 'info',
summary: 'MTB',
detail: message,
life: timeout,
})
return
}
log('Creating toast') log('Creating toast')
const container = document.getElementById('mtb-notify-container') const container = document.getElementById('mtb-notify-container')
const toast = document.createElement('div') const toast = document.createElement('div')
@@ -68,7 +59,7 @@ function notify(message, timeout = 3000, old_mode = false) {
log('Transition out') log('Transition out')
const totalHeight = Array.from(container.children).reduce( const totalHeight = Array.from(container.children).reduce(
(acc, child) => acc + child.offsetHeight + 10, // Add spacing of 10px between toasts (acc, child) => acc + child.offsetHeight + 10, // Add spacing of 10px between toasts
0, 0
) )
container.style.height = `${totalHeight}px` 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 // Update container's height to fit new toast
const totalHeight = Array.from(container.children).reduce( const totalHeight = Array.from(container.children).reduce(
(acc, child) => acc + child.offsetHeight + 10, // Add spacing of 10px between toasts (acc, child) => acc + child.offsetHeight + 10, // Add spacing of 10px between toasts
0, 0
) )
container.style.height = `${totalHeight}px` container.style.height = `${totalHeight}px`