Compare commits
50
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
821a031bfc | ||
|
|
3d12bd29a8 | ||
|
|
3494a4767e | ||
|
|
41d444ae70 | ||
|
|
41d79d4677 | ||
|
|
54f4963a7b | ||
|
|
d585b16ee7 | ||
|
|
499cd218aa | ||
|
|
fd00e12724 | ||
|
|
e37e4648e5 | ||
|
|
06685a5418 | ||
|
|
40145ddf40 | ||
|
|
c1d74c0d69 | ||
|
|
eca2ea5da9 | ||
|
|
c6d1f73cfc | ||
|
|
4ea7c0b67f | ||
|
|
8316914f02 | ||
|
|
5c0e020c73 | ||
|
|
d00722e9ea | ||
|
|
0106c13250 | ||
|
|
55226058d4 | ||
|
|
50e0f7b357 | ||
|
|
71f601094a | ||
|
|
ea750b5e8b | ||
|
|
ff2e99f73e | ||
|
|
efc6855073 | ||
|
|
0853b7fb6a | ||
|
|
10aa493dd8 | ||
|
|
f038d76748 | ||
|
|
940a781f29 | ||
|
|
c7248344cc | ||
|
|
fab33a40a2 | ||
|
|
8f83e8d4d7 | ||
|
|
6c59d5c32d | ||
|
|
e98f3f626f | ||
|
|
7e89e96e9d | ||
|
|
177b6eeef3 | ||
|
|
502a583409 | ||
|
|
a7966355c1 | ||
|
|
321abea51a | ||
|
|
63be3f26fd | ||
|
|
c4f40e299f | ||
|
|
b541670a5b | ||
|
|
4574c6451c | ||
|
|
7fb27804e1 | ||
|
|
9a7e022df1 | ||
|
|
2c483fd1d2 | ||
|
|
0967d439f5 | ||
|
|
319c02d658 | ||
|
|
265cb953ec |
@@ -0,0 +1,7 @@
|
||||
**/GFPGAN/inputs/**
|
||||
**/GFPGAN/tests/**
|
||||
**/frame_interpolation/photos/*
|
||||
moment.gif
|
||||
node.zip
|
||||
.DS_Store
|
||||
|
||||
@@ -1,9 +1,6 @@
|
||||
name: 📦 Publish to Comfy registry
|
||||
on:
|
||||
workflow_dispatch:
|
||||
push:
|
||||
tags:
|
||||
- '*'
|
||||
|
||||
permissions:
|
||||
issues: write
|
||||
@@ -21,4 +18,5 @@ jobs:
|
||||
- name: 📦 Publish Custom Node
|
||||
uses: Comfy-Org/publish-node-action@v1
|
||||
with:
|
||||
skip_checkout: 'true'
|
||||
personal_access_token: ${{ secrets.COMFY_REGISTRY_TOKEN }}
|
||||
|
||||
@@ -1,11 +1,16 @@
|
||||
__pycache__
|
||||
*.py[cod]
|
||||
*.onnx
|
||||
|
||||
wheels/
|
||||
node_modules/
|
||||
compose.yaml
|
||||
comfy_mtb.wsb
|
||||
Dockerfile
|
||||
|
||||
.DS_Store
|
||||
node.zip
|
||||
|
||||
# I store the gh-pages worktrees (src & build) there
|
||||
.worktrees
|
||||
comfy.lock
|
||||
|
||||
+33
-11
@@ -3,11 +3,11 @@
|
||||
# File: __init__.py
|
||||
# Project: comfy_mtb
|
||||
# Author: Mel Massadian
|
||||
# Copyright (c) 2023 Mel Massadian
|
||||
# Copyright (c) 2023-2025 Mel Massadian
|
||||
#
|
||||
###
|
||||
|
||||
__version__ = "0.3.0"
|
||||
__version__ = "0.6.0"
|
||||
|
||||
import os
|
||||
|
||||
@@ -34,6 +34,8 @@ from aiohttp import web
|
||||
|
||||
IN_COMFY = False
|
||||
|
||||
PromptServer = None
|
||||
|
||||
try:
|
||||
from server import PromptServer
|
||||
|
||||
@@ -75,7 +77,7 @@ def extract_nodes_from_source(filename: Path):
|
||||
)
|
||||
break
|
||||
except SyntaxError:
|
||||
log.error("Failed to parse")
|
||||
log.error(f"Failed to parse ast from: {filename}")
|
||||
return nodes
|
||||
|
||||
|
||||
@@ -240,10 +242,33 @@ if failed:
|
||||
# - ENDPOINT
|
||||
|
||||
|
||||
if IN_COMFY and hasattr(PromptServer, "instance"):
|
||||
# TODO: move that away and simplify existing endpoints
|
||||
|
||||
|
||||
def register_routes():
|
||||
if not PromptServer:
|
||||
log.error("No prompt server, are you inside comfy?")
|
||||
|
||||
if PromptServer.instance.app.frozen:
|
||||
log.warning(
|
||||
"The router is frozen and cannot be further edited."
|
||||
"If you are hot reloading mtb this is expected."
|
||||
)
|
||||
return
|
||||
|
||||
img_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):
|
||||
from cachetools import TTLCache
|
||||
|
||||
@@ -360,13 +385,6 @@ if IN_COMFY and hasattr(PromptServer, "instance"):
|
||||
# Return JSON for other requests
|
||||
return web.json_response({"message": "Welcome to MTB!"})
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
from io import BytesIO
|
||||
|
||||
from aiohttp import web
|
||||
from PIL import Image
|
||||
|
||||
def get_cached_image(file_path: str, preview_params=None, channel=None):
|
||||
cache_key = (file_path, preview_params, channel)
|
||||
if img_cache and (cache_key in img_cache):
|
||||
@@ -571,6 +589,10 @@ if IN_COMFY and hasattr(PromptServer, "instance"):
|
||||
return await endpoint.do_action(request)
|
||||
|
||||
|
||||
if IN_COMFY and hasattr(PromptServer, "instance"):
|
||||
register_routes()
|
||||
|
||||
|
||||
# - WAS Dictionary
|
||||
MANIFEST = {
|
||||
"name": "MTB Nodes", # The title that will be displayed on Node Class menu,. and Node Class view
|
||||
|
||||
+13
-6
@@ -1,19 +1,26 @@
|
||||
{
|
||||
"$schema": "https://biomejs.dev/schemas/1.6.1/schema.json",
|
||||
"organizeImports": {
|
||||
"enabled": true
|
||||
},
|
||||
"$schema": "https://biomejs.dev/schemas/2.0.5/schema.json",
|
||||
"assist": { "actions": { "source": { "organizeImports": "on" } } },
|
||||
"linter": {
|
||||
"enabled": true,
|
||||
"rules": {
|
||||
"recommended": true,
|
||||
"suspicious": {
|
||||
"noConsoleLog": "warn"
|
||||
"noConsole": { "level": "warn", "options": { "allow": ["log"] } }
|
||||
},
|
||||
"style": {
|
||||
"noParameterAssign": "off",
|
||||
"noShoutyConstants": "warn",
|
||||
"useNamingConvention": "off"
|
||||
"useNamingConvention": "off",
|
||||
"useAsConstAssertion": "error",
|
||||
"useDefaultParameterLast": "error",
|
||||
"useEnumInitializers": "error",
|
||||
"useSelfClosingElements": "error",
|
||||
"useSingleVarDeclarator": "error",
|
||||
"noUnusedTemplateLiteral": "error",
|
||||
"useNumberNamespace": "error",
|
||||
"noInferrableTypes": "error",
|
||||
"noUselessElse": "error"
|
||||
}
|
||||
}
|
||||
},
|
||||
|
||||
+9
-6
@@ -15,7 +15,6 @@ from .utils import (
|
||||
backup_file,
|
||||
build_glob_patterns,
|
||||
glob_multiple,
|
||||
import_install,
|
||||
reqs_map,
|
||||
run_command,
|
||||
styles_dir,
|
||||
@@ -24,7 +23,6 @@ from .utils import (
|
||||
endlog = mklog("mtb endpoint")
|
||||
|
||||
# - ACTIONS
|
||||
import_install("requirements")
|
||||
|
||||
|
||||
def ACTIONS_installDependency(dependency_names: list[str] | None = None):
|
||||
@@ -112,11 +110,15 @@ def ACTIONS_getUserVideos(
|
||||
|
||||
def ACTIONS_getUserImages(
|
||||
mode: Literal["input", "output"],
|
||||
target_width: int | str | None = None,
|
||||
count=1000,
|
||||
offset=0,
|
||||
sort: str | None = None,
|
||||
include_subfolders: bool = False,
|
||||
subfolder=None,
|
||||
subfolder: str | None = 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
|
||||
# if not enabled:
|
||||
@@ -124,11 +126,12 @@ def ACTIONS_getUserImages(
|
||||
|
||||
imgs = {}
|
||||
count = count or 1000
|
||||
target_width = int(target_width) if target_width else None
|
||||
|
||||
input_dir = Path(folder_paths.get_input_directory())
|
||||
output_dir = Path(folder_paths.get_output_directory())
|
||||
|
||||
entry_dir = input_dir if mode == "input" else output_dir
|
||||
entry_dir: Path = input_dir if mode == "input" else output_dir
|
||||
if subfolder:
|
||||
entry_dir = entry_dir / subfolder
|
||||
|
||||
@@ -157,9 +160,9 @@ def ACTIONS_getUserImages(
|
||||
|
||||
imgs = {
|
||||
img.name: (
|
||||
f"/mtb/view?filename={img.name}&width=512&type={mode}&subfolder={subfolder or ''}"
|
||||
f"/mtb/view?filename={img.name}{f'&width={target_width}' if target_width and target_width > 0 else ''}&type={mode}&subfolder={subfolder or ''}"
|
||||
f"{img.parent.relative_to(entry_dir) if include_subfolders else ''}"
|
||||
f"&preview=&rand={secrets.randbelow(424242)}"
|
||||
f"&preview={f'&rand={secrets.randbelow(424242)}' if salt_urls else ''}"
|
||||
)
|
||||
for i, img in enumerate(entries)
|
||||
if offset <= i < offset + count
|
||||
|
||||
@@ -1,85 +1,175 @@
|
||||
# NOTE: This file is only use for development you can ignore it
|
||||
|
||||
use private/log.nu
|
||||
use log.nu
|
||||
use nssm.nu *
|
||||
use nutils.nu [ make-id upsert-all fwd-slash backup-file ]
|
||||
use os.nu [ link ]
|
||||
|
||||
def get_root [--clean] {
|
||||
if $clean {
|
||||
$env.COMFY_CLEAN_ROOT
|
||||
} else {
|
||||
$env.COMFY_ROOT
|
||||
}
|
||||
# --- utilities ---
|
||||
def get_root [ --clean] {
|
||||
if $clean {
|
||||
$env.COMFY.ROOTS.clean
|
||||
} else {
|
||||
$env.COMFY.ROOTS.main
|
||||
}
|
||||
}
|
||||
|
||||
export def "comfy build-web" [] {
|
||||
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 --env path-add [pth] {
|
||||
$env.PATH = ($env.PATH | append ($pth | path expand))
|
||||
}
|
||||
|
||||
def short-date [] {
|
||||
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?
|
||||
export def "daily was-run" [] {
|
||||
|
||||
let daily = ($env.COMFY_MTB | path join daily.nuon)
|
||||
let daily = ($env.COMFY.ROOTS.mtb | path join daily.nuon)
|
||||
|
||||
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)
|
||||
return ($last == $today)
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
export def "daily commit" [from:string, to:string] {
|
||||
let daily = ($env.COMFY_MTB | path join daily.nuon)
|
||||
let commit = [{date: (date now) from:$from to:$to}]
|
||||
export def "daily commit" [from_commit: string to_commit: string] {
|
||||
let daily = ($env.COMFY.ROOTS.mtb | path join daily.nuon)
|
||||
let commit = [{date: (date now) from_commit: $from_commit to_commit: $to_commit}]
|
||||
|
||||
let dailies = (if ($daily | path exists) {
|
||||
open $daily | append $commit
|
||||
} else {
|
||||
let dailies = (
|
||||
if ($daily | path exists) {
|
||||
open $daily | append $commit
|
||||
} else {
|
||||
$commit
|
||||
})
|
||||
}
|
||||
)
|
||||
|
||||
$dailies | save -f $daily
|
||||
log success "Commited daily check"
|
||||
}
|
||||
|
||||
# start the comfy server
|
||||
export def "comfy start" [--clean,--old-ui, --listen, --skip-daily(-s)] {
|
||||
if (not (daily was-run)) and not $skip_daily {
|
||||
log info "Running daily checks"
|
||||
daily run
|
||||
}
|
||||
let root = get_root --clean=($clean)
|
||||
cd $root
|
||||
export def "comfy start" [
|
||||
--clean
|
||||
--old-ui
|
||||
--listen
|
||||
--skip-daily (-s)
|
||||
] {
|
||||
if not (daily was-run) and not $skip_daily {
|
||||
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
|
||||
export def "comfy update" [
|
||||
--clean # ??
|
||||
--rebase # Rebase instead of merge
|
||||
--clean # comfy clean instance
|
||||
--rebase # Rebase instead of merge
|
||||
] {
|
||||
let root = get_root --clean=$clean
|
||||
|
||||
@@ -94,21 +184,28 @@ export def "comfy update" [
|
||||
log info "Backing up and removing models symlinks"
|
||||
|
||||
# preparing root for pull
|
||||
if not $clean {
|
||||
let pyproject = 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
|
||||
cd $models
|
||||
# find and store all symlinks
|
||||
let links = (ls -la |
|
||||
where not ($it.target | is-empty) |
|
||||
select name target |
|
||||
sort-by name)
|
||||
|
||||
log info "Checking for links in models..."
|
||||
let links = (
|
||||
ls -la | where not ($it.target | is-empty) | select name target | sort-by name
|
||||
)
|
||||
log info $"Found links: ($links)"
|
||||
|
||||
if not ($links | is-empty) {
|
||||
log info "Backing up the symlinks..."
|
||||
backup-file --root links.nuon
|
||||
$links | save -f links.nuon
|
||||
# remove them
|
||||
open links.nuon | each {|p| rm $p.name }
|
||||
}
|
||||
$proj
|
||||
} else {
|
||||
# just remove symlinks
|
||||
rm $models
|
||||
@@ -141,7 +238,6 @@ export def "comfy update" [
|
||||
if $rebase {
|
||||
log info "Rebasing changes"
|
||||
git rebase master
|
||||
|
||||
} else {
|
||||
log info "Merging changes"
|
||||
git merge master
|
||||
@@ -152,9 +248,10 @@ export def "comfy update" [
|
||||
|
||||
if not $clean {
|
||||
rm pyproject.toml
|
||||
cp pyproject-mel.toml pyproject.toml
|
||||
log info "Using our own pyproject..."
|
||||
cp $pyproject pyproject.toml
|
||||
cd $models
|
||||
|
||||
log info "Relinking models..."
|
||||
# resymlink them
|
||||
open links.nuon | each {|p| link -a $p.target $p.name }
|
||||
} else {
|
||||
@@ -167,72 +264,82 @@ export def "comfy update" [
|
||||
|
||||
log success $"Update successful \(($commit_count) new commits\)"
|
||||
|
||||
return {from:$current_commit to:$new_commit}
|
||||
|
||||
|
||||
return {from_commit: $current_commit to_commit: $new_commit}
|
||||
}
|
||||
|
||||
export def "comfy toggle_extensions" [--clean] {
|
||||
let root = get_root --clean=($clean)
|
||||
cd $root
|
||||
cd custom_nodes
|
||||
let exts = (ls | where type in ["dir","symlink"] | get name)
|
||||
let choices = ($exts | input list -m "choose extension to toggle")
|
||||
if ($choices | is-empty) {
|
||||
return
|
||||
}
|
||||
export def "comfy toggle_extensions" [
|
||||
--clean
|
||||
] {
|
||||
let root = get_root --clean=$clean
|
||||
cd $root
|
||||
cd custom_nodes
|
||||
let exts = (ls | where type in ["dir" "symlink"] | get name)
|
||||
let choices = ($exts | input list -m "choose extension to toggle")
|
||||
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
|
||||
$filtered | each {|f|
|
||||
let new_name = ($f.name | str replace ".disabled" "")
|
||||
log info "Filtered" $filtered
|
||||
$filtered | each {|f|
|
||||
let new_name = ($f.name | str replace ".disabled" "")
|
||||
|
||||
let new_name = if $f.enabled {
|
||||
$"($new_name).disabled"
|
||||
} else {
|
||||
$new_name
|
||||
}
|
||||
log info $"Moving ($f.name) to ($new_name)"
|
||||
mv $f.name $new_name
|
||||
let new_name = if $f.enabled {
|
||||
$"($new_name).disabled"
|
||||
} else {
|
||||
$new_name
|
||||
}
|
||||
log info $"Moving ($f.name) to ($new_name)"
|
||||
mv $f.name $new_name
|
||||
}
|
||||
}
|
||||
|
||||
# git pull all extensions
|
||||
export def "comfy update_extensions" [--clean] {
|
||||
let root = get_root --clean=($clean)
|
||||
cd $root
|
||||
cd custom_nodes
|
||||
git multipull . -s -q
|
||||
export def "comfy update_extensions" [ --clean] {
|
||||
let root = get_root --clean=$clean
|
||||
cd $root
|
||||
cd custom_nodes
|
||||
git multipull . -s -q
|
||||
}
|
||||
|
||||
def --env path-add [pth] {
|
||||
$env.PATH = ($env.PATH | append ($pth | path expand))
|
||||
|
||||
# manual set version of mtb
|
||||
export def "comfy-mtb set-version" [version: string] {
|
||||
# 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 {
|
||||
$env.PYTHONUTF8 = 1
|
||||
$env.COMFY_MTB = ("." | path expand)
|
||||
# $env.CUDA_ROOT = 'C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.1\'
|
||||
|
||||
$env.CUDA_HOME = $env.CUDA_ROOT
|
||||
|
||||
$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'
|
||||
$env.COMFY = {
|
||||
base_url : "https://mel-pc.tail3c8eb.ts.net"
|
||||
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
|
||||
#
|
||||
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
|
||||
|
||||
overlay use "../../.venv/Scripts/activate.nu"
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
{
|
||||
"use_repl": false
|
||||
}
|
||||
@@ -43,7 +43,6 @@ pip_map = {
|
||||
"tb-nightly": "tensorboard",
|
||||
"protobuf": "google.protobuf",
|
||||
"qrcode[pil]": "qrcode",
|
||||
"requirements-parser": "requirements",
|
||||
# Add more mappings as needed
|
||||
}
|
||||
|
||||
|
||||
+14
-13
@@ -1,20 +1,16 @@
|
||||
from typing import Any, TypedDict
|
||||
from typing import TYPE_CHECKING, Any, TypedDict
|
||||
|
||||
import torch
|
||||
import torchaudio
|
||||
from comfy.model_management import get_torch_device
|
||||
from huggingface_hub import snapshot_download
|
||||
from transformers import (
|
||||
WhisperForConditionalGeneration,
|
||||
WhisperProcessor,
|
||||
)
|
||||
|
||||
# from transformers import (
|
||||
# AutoFeatureExtractor,
|
||||
# WhisperForConditionalGeneration,
|
||||
# WhisperModel,
|
||||
# WhisperProcessor,
|
||||
# )
|
||||
if TYPE_CHECKING:
|
||||
from transformers import (
|
||||
WhisperForConditionalGeneration,
|
||||
WhisperProcessor,
|
||||
)
|
||||
|
||||
from ..log import log
|
||||
from ..utils import get_model_path
|
||||
|
||||
@@ -101,8 +97,8 @@ class MtbAudio:
|
||||
class WhisperPipeline(TypedDict):
|
||||
"""Whisper model pipeline."""
|
||||
|
||||
processor: WhisperProcessor
|
||||
model: WhisperForConditionalGeneration
|
||||
processor: "WhisperProcessor"
|
||||
model: "WhisperForConditionalGeneration"
|
||||
|
||||
|
||||
class MTB_LoadWhisper:
|
||||
@@ -148,6 +144,11 @@ class MTB_LoadWhisper:
|
||||
|
||||
def load(self, model_size="tiny", download_missing=False):
|
||||
"""Load Whisper model and processor."""
|
||||
from transformers import (
|
||||
WhisperForConditionalGeneration,
|
||||
WhisperProcessor,
|
||||
)
|
||||
|
||||
whisper_dir = get_model_path("whisper")
|
||||
tag = f"whisper-{model_size}"
|
||||
model_dir = whisper_dir / tag
|
||||
|
||||
+190
@@ -0,0 +1,190 @@
|
||||
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]
|
||||
+194
-181
@@ -1,13 +1,22 @@
|
||||
import numpy as np
|
||||
from typing import NamedTuple
|
||||
|
||||
import torch
|
||||
from PIL import Image, ImageDraw, ImageFilter
|
||||
import torchvision.transforms.functional as TF
|
||||
|
||||
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:
|
||||
"""The bounding box (BBOX) custom type used by other nodes"""
|
||||
"""A literal bounding box."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
@@ -37,12 +46,14 @@ class MTB_Bbox:
|
||||
FUNCTION = "do_crop"
|
||||
CATEGORY = "mtb/crop"
|
||||
|
||||
def do_crop(self, x: int, y: int, width: int, height: int): # bbox
|
||||
return ((x, y, width, height),)
|
||||
def do_crop(
|
||||
self, x: int, y: int, width: int, height: int
|
||||
) -> tuple[BoundingBox]: # bbox
|
||||
return (BoundingBox(x, y, width, height),)
|
||||
|
||||
|
||||
class MTB_SplitBbox:
|
||||
"""Split the components of a bbox"""
|
||||
"""Split the components of a bbox."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
@@ -55,8 +66,8 @@ class MTB_SplitBbox:
|
||||
RETURN_TYPES = ("INT", "INT", "INT", "INT")
|
||||
RETURN_NAMES = ("x", "y", "width", "height")
|
||||
|
||||
def split_bbox(self, bbox):
|
||||
return (bbox[0], bbox[1], bbox[2], bbox[3])
|
||||
def split_bbox(self, bbox: BoundingBox) -> BoundingBox:
|
||||
return bbox
|
||||
|
||||
|
||||
class MTB_UpscaleBboxBy:
|
||||
@@ -74,26 +85,23 @@ class MTB_UpscaleBboxBy:
|
||||
|
||||
FUNCTION = "upscale"
|
||||
|
||||
def upscale(
|
||||
self, bbox: tuple[int, int, int, int], scale: float
|
||||
) -> tuple[tuple[int, int, int, int]]:
|
||||
def upscale(self, bbox: BoundingBox, scale: float) -> tuple[BoundingBox]:
|
||||
x, y, width, height = bbox
|
||||
|
||||
center_x = x + width // 2
|
||||
center_y = y + height // 2
|
||||
center_x = x + width / 2
|
||||
center_y = y + height / 2
|
||||
|
||||
new_width = int(width * scale)
|
||||
new_height = int(height * scale)
|
||||
|
||||
new_x = center_x - new_width // 2
|
||||
new_y = center_y - new_height // 2
|
||||
new_x = int(center_x - new_width / 2)
|
||||
new_y = int(center_y - new_height / 2)
|
||||
|
||||
scaled = (new_x, new_y, new_width, new_height)
|
||||
return (scaled,)
|
||||
return (BoundingBox(new_x, new_y, new_width, new_height),)
|
||||
|
||||
|
||||
class MTB_BboxFromMask:
|
||||
"""From a mask extract the bounding box"""
|
||||
"""From a mask extract the bounding box."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
@@ -103,7 +111,7 @@ class MTB_BboxFromMask:
|
||||
"invert": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"optional": {
|
||||
"image": ("IMAGE",),
|
||||
"image": ("IMAGE", {"tooltip": "Optional image"}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -119,52 +127,44 @@ class MTB_BboxFromMask:
|
||||
CATEGORY = "mtb/crop"
|
||||
|
||||
def extract_bounding_box(
|
||||
self, mask: torch.Tensor, invert: bool, image=None
|
||||
):
|
||||
# if image != None:
|
||||
# if mask.size(0) != image.size(0):
|
||||
# if mask.size(0) != 1:
|
||||
# log.error(
|
||||
# f"Batch count mismatch for mask and image, it can either be 1 mask for X images, or X masks for X images (mask: {mask.shape} | image: {image.shape})"
|
||||
# )
|
||||
self,
|
||||
mask: torch.Tensor,
|
||||
*,
|
||||
invert: bool = False,
|
||||
image: torch.Tensor | None = None,
|
||||
) -> tuple[BoundingBox, torch.Tensor | None]:
|
||||
mask = 1 - mask if invert else mask
|
||||
non_zero_indices = torch.nonzero(mask)
|
||||
|
||||
# raise Exception(
|
||||
# f"Batch count mismatch for mask and image, it can either be 1 mask for X images, or X masks for X images (mask: {mask.shape} | image: {image.shape})"
|
||||
# )
|
||||
if non_zero_indices.numel() == 0:
|
||||
log.warning(
|
||||
"BboxFromMask: Mask is empty. Returning a (0,0,0,0) bbox."
|
||||
)
|
||||
return (BoundingBox(0, 0, 0, 0), image)
|
||||
|
||||
# we invert it
|
||||
_mask = tensor2pil(1.0 - mask)[0] if invert else tensor2pil(mask)[0]
|
||||
alpha_channel = np.array(_mask)
|
||||
min_coords = torch.min(non_zero_indices, dim=0).values
|
||||
max_coords = torch.max(non_zero_indices, dim=0).values
|
||||
|
||||
non_zero_indices = np.nonzero(alpha_channel)
|
||||
min_y, min_x = min_coords[1].item(), min_coords[2].item()
|
||||
max_y, max_x = max_coords[1].item(), max_coords[2].item()
|
||||
|
||||
min_x, max_x = np.min(non_zero_indices[1]), np.max(non_zero_indices[1])
|
||||
min_y, max_y = np.min(non_zero_indices[0]), np.max(non_zero_indices[0])
|
||||
width = max_x - min_x + 1
|
||||
height = max_y - min_y + 1
|
||||
|
||||
# Create a bounding box tuple
|
||||
if image != None:
|
||||
# Convert the image to a NumPy array
|
||||
imgs = tensor2np(image)
|
||||
out = []
|
||||
for img in imgs:
|
||||
# Crop the image from the bounding box
|
||||
img = img[min_y:max_y, min_x:max_x, :]
|
||||
log.debug(f"Cropped image to shape {img.shape}")
|
||||
out.append(img)
|
||||
|
||||
image = np2tensor(out)
|
||||
log.debug(f"Cropped images shape: {image.shape}")
|
||||
bounding_box = (min_x, min_y, max_x - min_x, max_y - min_y)
|
||||
return (
|
||||
bounding_box,
|
||||
image,
|
||||
bounding_box = BoundingBox(
|
||||
int(min_x), int(min_y), int(width), int(height)
|
||||
)
|
||||
|
||||
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:
|
||||
"""Crops an image and an optional mask to a given bounding box
|
||||
"""Crop an image and an optional mask to a given bounding box.
|
||||
|
||||
The bounding box can be given as a tuple of (x, y, width, height) or as a BBOX type
|
||||
The BBOX input takes precedence over the tuple input
|
||||
"""
|
||||
|
||||
@@ -204,35 +204,38 @@ class MTB_Crop:
|
||||
def do_crop(
|
||||
self,
|
||||
image: torch.Tensor,
|
||||
mask=None,
|
||||
x=0,
|
||||
y=0,
|
||||
width=256,
|
||||
height=256,
|
||||
bbox=None,
|
||||
*,
|
||||
mask: torch.Tensor | None = None,
|
||||
x: int = 0,
|
||||
y: int = 0,
|
||||
width: int = 256,
|
||||
height: int = 256,
|
||||
bbox: BoundingBox | None = None,
|
||||
):
|
||||
image = image.numpy()
|
||||
if mask is not None:
|
||||
mask = mask.numpy()
|
||||
|
||||
if bbox is not None:
|
||||
x, y, width, height = bbox
|
||||
|
||||
cropped_image = image[:, y : y + height, x : x + width, :]
|
||||
cropped_mask = None
|
||||
if mask is not None:
|
||||
cropped_mask = (
|
||||
mask[:, y : y + height, x : x + width]
|
||||
if mask is not None
|
||||
else None
|
||||
if width <= 0 or height <= 0:
|
||||
log.error(
|
||||
"Crop dimensions must be positive. Check the BBOX or widget inputs."
|
||||
)
|
||||
crop_data = (x, y, width, height)
|
||||
return (
|
||||
torch.zeros_like(image),
|
||||
torch.zeros_like(mask) if mask is not None else None,
|
||||
(x, y, width, height),
|
||||
)
|
||||
|
||||
cropped_image = image[:, y : y + height, x : x + width, :]
|
||||
cropped_mask = (
|
||||
mask[:, y : y + height, x : x + width]
|
||||
if mask is not None
|
||||
else None
|
||||
)
|
||||
crop_data = BoundingBox(x, y, width, height)
|
||||
|
||||
return (
|
||||
torch.from_numpy(cropped_image),
|
||||
torch.from_numpy(cropped_mask)
|
||||
if cropped_mask is not None
|
||||
else None,
|
||||
cropped_image,
|
||||
cropped_mask if cropped_mask is not None else None,
|
||||
crop_data,
|
||||
)
|
||||
|
||||
@@ -246,35 +249,33 @@ class MTB_Crop:
|
||||
# return (x_left, y_top, x_right, y_bottom)
|
||||
|
||||
|
||||
def bbox_check(bbox, target_size=None):
|
||||
def bbox_check(bbox: BoundingBox, target_size: tuple[int, int] | None = None):
|
||||
if not target_size:
|
||||
return bbox
|
||||
|
||||
new_bbox = (
|
||||
bbox[0],
|
||||
bbox[1],
|
||||
min(target_size[0] - bbox[0], bbox[2]),
|
||||
min(target_size[1] - bbox[1], bbox[3]),
|
||||
new_bbox = BoundingBox(
|
||||
bbox.x,
|
||||
bbox.y,
|
||||
min(target_size[0] - bbox.x, bbox.width),
|
||||
min(target_size[1] - bbox.y, bbox.height),
|
||||
)
|
||||
if new_bbox != bbox:
|
||||
log.warn(f"BBox too big, constrained to {new_bbox}")
|
||||
log.warning(f"BBox too big, constrained to {new_bbox}")
|
||||
|
||||
return new_bbox
|
||||
|
||||
|
||||
def bbox_to_region(bbox, target_size=None):
|
||||
def bbox_to_region(
|
||||
bbox: BoundingBox, target_size: tuple[int, int] | None = None
|
||||
):
|
||||
bbox = bbox_check(bbox, target_size)
|
||||
|
||||
# to region
|
||||
return (bbox[0], bbox[1], bbox[0] + bbox[2], bbox[1] + bbox[3])
|
||||
return (bbox.x, bbox.y, bbox.x + bbox.width, bbox.y + bbox.height)
|
||||
|
||||
|
||||
class MTB_Uncrop:
|
||||
"""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
|
||||
"""
|
||||
"""Uncrop an image to a given bounding box."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
@@ -291,91 +292,113 @@ class MTB_Uncrop:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "do_crop"
|
||||
|
||||
FUNCTION = "do_uncrop"
|
||||
CATEGORY = "mtb/crop"
|
||||
|
||||
def do_crop(self, image, crop_image, bbox, border_blending):
|
||||
def inset_border(image, border_width=20, border_color=(0)):
|
||||
width, height = image.size
|
||||
bordered_image = Image.new(
|
||||
image.mode, (width, height), border_color
|
||||
def do_uncrop(
|
||||
self,
|
||||
image: torch.Tensor,
|
||||
crop_image: torch.Tensor,
|
||||
bbox: BoundingBox,
|
||||
border_blending: float = 0.25,
|
||||
):
|
||||
if len(image) > 1 and len(image) != len(crop_image):
|
||||
raise ValueError(
|
||||
"Uncrop: Batch size of background 'image' must be 1 or match the 'crop_image' batch size."
|
||||
)
|
||||
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
|
||||
import comfy.utils
|
||||
|
||||
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"
|
||||
)
|
||||
pbar = comfy.utils.ProgressBar(4)
|
||||
|
||||
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]
|
||||
device = image.device
|
||||
|
||||
# uncrop the image based on the bounding box
|
||||
bb_x, bb_y, bb_width, bb_height = bbox
|
||||
log.debug(f"Working on device: {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
|
||||
crop_image = crop_image.to(device)
|
||||
|
||||
crop_img = crop.convert("RGB")
|
||||
if len(image) == 1 and len(crop_image) > 1:
|
||||
image = image.repeat(len(crop_image), 1, 1, 1)
|
||||
|
||||
log.debug(f"Crop image size: {crop_img.size}")
|
||||
log.debug(f"Image size: {img.size}")
|
||||
batch_size, bg_h, bg_w, _ = image.shape
|
||||
_, fg_h, fg_w, _ = crop_image.shape
|
||||
x, y, width, height = bbox
|
||||
|
||||
if border_blending > 1.0:
|
||||
border_blending = 1.0
|
||||
elif border_blending < 0.0:
|
||||
border_blending = 0.0
|
||||
|
||||
blend_ratio = (max(crop_img.size) / 2) * float(border_blending)
|
||||
|
||||
blend = img.convert("RGBA")
|
||||
mask = Image.new("L", img.size, 0)
|
||||
|
||||
mask_block = Image.new("L", (bb_width, bb_height), 255)
|
||||
mask_block = inset_border(mask_block, int(blend_ratio / 2), (0))
|
||||
|
||||
mask.paste(mask_block, paste_region)
|
||||
log.debug(f"Blend size: {blend.size} | kind {blend.mode}")
|
||||
log.debug(
|
||||
f"Crop image size: {crop_img.size} | kind {crop_img.mode}"
|
||||
)
|
||||
log.debug(f"BBox: {paste_region}")
|
||||
blend.paste(crop_img, paste_region)
|
||||
|
||||
mask = mask.filter(ImageFilter.BoxBlur(radius=blend_ratio / 4))
|
||||
mask = mask.filter(
|
||||
ImageFilter.GaussianBlur(radius=blend_ratio / 4)
|
||||
if (width, height) != (fg_w, fg_h):
|
||||
log.warning(
|
||||
f"Uncrop: crop_image size {(fg_w, fg_h)} "
|
||||
"differs from bbox {(width, height)}. Resizing to fit bbox."
|
||||
)
|
||||
|
||||
blend.putalpha(mask)
|
||||
img = Image.alpha_composite(img.convert("RGBA"), blend)
|
||||
out_images.append(img.convert("RGB"))
|
||||
resized_crop = crop_image.permute(0, 3, 1, 2)
|
||||
resized_crop = torch.nn.functional.interpolate(
|
||||
resized_crop,
|
||||
size=(height, width),
|
||||
mode="bicubic",
|
||||
align_corners=False,
|
||||
)
|
||||
resized_crop = resized_crop.permute(0, 2, 3, 1)
|
||||
|
||||
return (pil2tensor(out_images),)
|
||||
pbar.update(1)
|
||||
# paste coords
|
||||
paste_x1 = max(x, 0)
|
||||
paste_y1 = max(y, 0)
|
||||
paste_x2 = min(x + width, bg_w)
|
||||
paste_y2 = min(y + height, bg_h)
|
||||
|
||||
# region from crop (bound)
|
||||
crop_x1 = max(0, -x)
|
||||
crop_y1 = max(0, -y)
|
||||
crop_x2 = crop_x1 + (paste_x2 - paste_x1)
|
||||
crop_y2 = crop_y1 + (paste_y2 - paste_y1)
|
||||
|
||||
if paste_x1 >= paste_x2 or paste_y1 >= paste_y2:
|
||||
log.warning(
|
||||
"Uncrop: BBOX is entirely outside the image boundaries. Returning original image."
|
||||
)
|
||||
return (image,)
|
||||
|
||||
pbar.update(1)
|
||||
source_slice = resized_crop[:, crop_y1:crop_y2, crop_x1:crop_x2, :]
|
||||
|
||||
final_image = image.clone()
|
||||
final_image[:, paste_y1:paste_y2, paste_x1:paste_x2, :] = source_slice
|
||||
|
||||
pbar.update(1)
|
||||
|
||||
blend_radius = int(max(width, height) * border_blending * 0.5)
|
||||
if blend_radius > 0:
|
||||
_device = device
|
||||
if torch.cuda.is_available():
|
||||
_device = torch.device("cuda")
|
||||
|
||||
log.debug("Processing blending")
|
||||
alpha_mask = torch.zeros((batch_size, bg_h, bg_w), device=_device)
|
||||
alpha_mask[:, paste_y1:paste_y2, paste_x1:paste_x2] = 1.0
|
||||
|
||||
kernel_size = 2 * blend_radius + 1
|
||||
|
||||
log.debug("Gaussian blur...")
|
||||
alpha_mask = TF.gaussian_blur(
|
||||
alpha_mask.unsqueeze(1), kernel_size=[kernel_size, kernel_size]
|
||||
).squeeze(1)
|
||||
alpha_mask = alpha_mask.unsqueeze(-1)
|
||||
|
||||
log.debug("Applying blending")
|
||||
final_image = final_image.to(_device) * alpha_mask + image.to(
|
||||
_device
|
||||
) * (1.0 - alpha_mask)
|
||||
|
||||
pbar.update(1)
|
||||
return (final_image.to(device),)
|
||||
|
||||
|
||||
class MTB_BBoxForceDimensions:
|
||||
"""
|
||||
Resize a BBOX to new dimensions while keeping its center.
|
||||
|
||||
Optionally constrains the BBOX to stay within image boundaries.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
@@ -383,6 +406,7 @@ class MTB_BBoxForceDimensions:
|
||||
"bbox": ("BBOX",),
|
||||
"width": ("INT", {"default": 512, "min": 1, "max": 8192}),
|
||||
"height": ("INT", {"default": 512, "min": 1, "max": 8192}),
|
||||
"constrain_to_image": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
"optional": {
|
||||
"image": ("IMAGE",),
|
||||
@@ -395,10 +419,12 @@ class MTB_BBoxForceDimensions:
|
||||
|
||||
def force_dimensions(
|
||||
self,
|
||||
*,
|
||||
bbox: tuple[int, int, int, int],
|
||||
width: int,
|
||||
height: int,
|
||||
image: torch.Tensor = None,
|
||||
constrain_to_image: bool = True,
|
||||
image: torch.Tensor | None = None,
|
||||
) -> tuple[tuple[int, int, int, int]]:
|
||||
x, y, curr_width, curr_height = bbox
|
||||
|
||||
@@ -408,27 +434,14 @@ class MTB_BBoxForceDimensions:
|
||||
new_x = center_x - width // 2
|
||||
new_y = center_y - height // 2
|
||||
|
||||
if image is not None:
|
||||
if constrain_to_image and image is not None:
|
||||
img_height, img_width = image.shape[1:3]
|
||||
x_overflow = max(0, new_x + width - img_width) + min(0, new_x)
|
||||
y_overflow = max(0, new_y + height - img_height) + min(0, new_y)
|
||||
if width > img_width or height > img_height:
|
||||
x_exceed = width - img_width if width > img_width else 0
|
||||
y_exceed = height - img_height if height > img_height else 0
|
||||
raise ValueError(
|
||||
f"Target bbox dimensions ({width}x{height}) exceed image bounds ({img_width}x{img_height}) "
|
||||
f"by {x_exceed}px horizontally and {y_exceed}px vertically"
|
||||
)
|
||||
new_x = max(0, min(new_x, img_width - width))
|
||||
new_y = max(0, min(new_y, img_height - height))
|
||||
width = min(width, img_width)
|
||||
height = min(height, img_height)
|
||||
|
||||
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),)
|
||||
return ((new_x, new_y, width, height),)
|
||||
|
||||
|
||||
__nodes__ = [
|
||||
|
||||
+605
-179
@@ -1,33 +1,70 @@
|
||||
import base64
|
||||
import io
|
||||
import json
|
||||
from pathlib import Path
|
||||
import textwrap
|
||||
from collections.abc import Callable
|
||||
from functools import wraps
|
||||
from typing import Any, Literal, Protocol, TypedDict, runtime_checkable
|
||||
|
||||
import folder_paths
|
||||
import torch
|
||||
from rich import inspect
|
||||
from rich.console import Console
|
||||
|
||||
from ..log import log
|
||||
from ..utils import tensor2pil
|
||||
from ..utils import LazyProxyTensor, get_torch_tensor_info, tensor2pil
|
||||
|
||||
try:
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
|
||||
plt.style.use("dark_background")
|
||||
MATPLOTLIB_AVAILABLE = True
|
||||
except ImportError:
|
||||
MATPLOTLIB_AVAILABLE = False
|
||||
|
||||
|
||||
def get_detailed_type_info(obj):
|
||||
type_info = []
|
||||
# region Decorator
|
||||
def metadata(**meta_kwargs: Any) -> Callable[[Any], Any]:
|
||||
"""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_info.append(f"Type: {type_name}")
|
||||
|
||||
if isinstance(obj, torch.Tensor):
|
||||
type_info.extend(
|
||||
[
|
||||
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)):
|
||||
return get_torch_tensor_info(obj)
|
||||
|
||||
elif isinstance(obj, list | tuple):
|
||||
type_info.extend(
|
||||
[
|
||||
f"Length: {len(obj)}",
|
||||
@@ -47,122 +84,184 @@ def get_detailed_type_info(obj):
|
||||
attributes = [attr for attr in dir(obj) if not attr.startswith("_")]
|
||||
type_info.append(f"Attributes: {attributes}")
|
||||
|
||||
return type_info
|
||||
return "\n".join(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
|
||||
def process_tensor(tensor: torch.Tensor, as_type=False):
|
||||
log.debug(f"Tensor: {tensor.shape}")
|
||||
|
||||
if as_type:
|
||||
return {
|
||||
"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")
|
||||
def _apply_rich(
|
||||
formatted: str | list[str], rich_mode: str | None = None, *, title=""
|
||||
) -> str:
|
||||
if rich_mode is None:
|
||||
return (
|
||||
formatted if isinstance(formatted, str) else "\n".join(formatted)
|
||||
)
|
||||
|
||||
return {"b64_images": b64_imgs}
|
||||
from rich.console import Console
|
||||
|
||||
console = Console(record=True)
|
||||
|
||||
def process_list(anything, as_type=False):
|
||||
text = []
|
||||
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)})"
|
||||
)
|
||||
if isinstance(formatted, list):
|
||||
for line in formatted:
|
||||
console.print(line)
|
||||
else:
|
||||
text.append(f"Array ({len(anything)}): {anything}")
|
||||
console.print(formatted)
|
||||
|
||||
return {"text": text}
|
||||
CSV_CODE_FORMAT = """
|
||||
<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;
|
||||
}}
|
||||
|
||||
def process_dict(anything, as_type=False):
|
||||
text = []
|
||||
if as_type:
|
||||
return {"text": get_detailed_type_info(anything)}
|
||||
.{unique_id}-matrix {{
|
||||
font-family: Fira Code, monospace;
|
||||
font-size: {char_height}px;
|
||||
line-height: {line_height}px;
|
||||
font-variant-east-asian: full-width;
|
||||
}}
|
||||
|
||||
if "samples" in anything:
|
||||
is_empty = (
|
||||
"(empty)" if torch.count_nonzero(anything["samples"]) == 0 else ""
|
||||
)
|
||||
text.append(f"Latent Samples: {anything['samples'].shape} {is_empty}")
|
||||
.{unique_id}-title {{
|
||||
font-size: 18px;
|
||||
font-weight: bold;
|
||||
font-family: arial;
|
||||
}}
|
||||
|
||||
elif "waveform" in anything:
|
||||
is_empty = (
|
||||
"(empty) " if torch.count_nonzero(anything["samples"]) == 0 else ""
|
||||
{styles}
|
||||
</style>
|
||||
|
||||
<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}", ""),
|
||||
)
|
||||
|
||||
text.append(
|
||||
f"Audio Samples: {anything['waveform'].shape}{is_empty} | sample rate {anything['sample_rate']}"
|
||||
elif rich_mode == "html":
|
||||
CONSOLE_HTML_FORMAT = textwrap.dedent("""
|
||||
<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,
|
||||
)
|
||||
|
||||
else:
|
||||
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)]}
|
||||
log.error(f"Unknown rich mode: {rich_mode}")
|
||||
return formatted if isinstance(formatted, str) else "\n".join(formatted)
|
||||
|
||||
|
||||
# endregion
|
||||
|
||||
|
||||
class MTB_Debug:
|
||||
"""Experimental node to debug any Comfy values.
|
||||
# region conditions
|
||||
|
||||
support for more types and widgets is planned.
|
||||
"""
|
||||
|
||||
# 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:
|
||||
"""A debug node."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {"output_to_console": ("BOOLEAN", {"default": False})},
|
||||
"optional": {"as_detailed_types": ("BOOLEAN", {"default": False})},
|
||||
"optional": {
|
||||
"as_detailed_types": ("BOOLEAN", {"default": False}),
|
||||
"deep_inspect": ("BOOLEAN", {"default": False}),
|
||||
"rich_mode": (
|
||||
("none", "html", "svg", "svg-window"),
|
||||
{"default": "none"},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ()
|
||||
@@ -170,99 +269,426 @@ class MTB_Debug:
|
||||
CATEGORY = "mtb/debug"
|
||||
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(
|
||||
self, output_to_console: bool, as_detailed_types: bool, **kwargs
|
||||
self,
|
||||
**kwargs,
|
||||
):
|
||||
output = {"ui": {"items": []}}
|
||||
|
||||
if output_to_console:
|
||||
for k, v in kwargs.items():
|
||||
log.info(f"{k}: {v}")
|
||||
settings = {k: kwargs.pop(k) for k in self.INPUT_TYPES()["optional"]}
|
||||
output_to_console = kwargs.pop("output_to_console")
|
||||
as_type = settings.get("as_detailed_types", False)
|
||||
deep = settings.get("deep_inspect", False)
|
||||
rich_mode = settings.get("rich_mode", "none")
|
||||
|
||||
for input_name, anything in kwargs.items():
|
||||
processor = processors.get(type(anything), process_text)
|
||||
for input_name, item in kwargs.items():
|
||||
processed = self._dispatch_processor(
|
||||
item, as_type=as_type, deep=deep
|
||||
)
|
||||
if processed is None:
|
||||
continue
|
||||
|
||||
processed = processor(anything, as_detailed_types)
|
||||
if rich_mode != "none":
|
||||
title = f"{input_name} ({type(item).__name__})"
|
||||
processed = _apply_rich_results(processed, rich_mode, title)
|
||||
|
||||
item = {
|
||||
"input": input_name,
|
||||
**processed,
|
||||
}
|
||||
output["ui"]["items"].append(item)
|
||||
if output_to_console:
|
||||
log.info(f"- Input '{input_name}':")
|
||||
for p in processed:
|
||||
if p["kind"] == "text":
|
||||
log.info(f" {p['data']}")
|
||||
if p["kind"] == "b64_image":
|
||||
log.info(f" (contains {len(p['data'])} images)")
|
||||
|
||||
output["ui"]["items"].append(
|
||||
{"input": input_name, "items": processed}
|
||||
)
|
||||
return output
|
||||
|
||||
def _process_unknown(
|
||||
self, item: Any, *, as_type=False, deep=False
|
||||
) -> ProcessorResult:
|
||||
console = Console(
|
||||
record=True,
|
||||
width=120,
|
||||
)
|
||||
|
||||
class MTB_SaveTensors:
|
||||
"""Save torch tensors (image, mask or latent) to disk.
|
||||
console.print(f"Generic {type(item).__name__}", emoji=True)
|
||||
if as_type:
|
||||
inspect(item, console=console, all=deep, methods=deep, docs=deep)
|
||||
else:
|
||||
console.print(item, emoji=True)
|
||||
|
||||
useful to debug things outside comfy.
|
||||
"""
|
||||
text_output = console.export_text(clear=True)
|
||||
|
||||
def __init__(self):
|
||||
self.output_dir = folder_paths.get_output_directory()
|
||||
self.type = "mtb/debug"
|
||||
return [UIResult(kind="text", data=text_output.strip())]
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"filename_prefix": ("STRING", {"default": "ComfyPickle"}),
|
||||
},
|
||||
"optional": {
|
||||
"image": ("IMAGE",),
|
||||
"mask": ("MASK",),
|
||||
"latent": ("LATENT",),
|
||||
},
|
||||
}
|
||||
def _process_repr(
|
||||
self, item: Any, as_type=False, deep=False
|
||||
) -> ProcessorResult:
|
||||
return [{"kind": "text", "data": item.__repr__()}]
|
||||
|
||||
FUNCTION = "save"
|
||||
OUTPUT_NODE = True
|
||||
RETURN_TYPES = ()
|
||||
CATEGORY = "mtb/debug"
|
||||
def _process_primitive(
|
||||
self, item: Any, *, as_type=False, deep=False
|
||||
) -> ProcessorResult:
|
||||
if as_type:
|
||||
return self._process_unknown(item, as_type=as_type, deep=deep)
|
||||
|
||||
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())
|
||||
return [UIResult(kind="text", data=str(item))]
|
||||
|
||||
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())
|
||||
def _process_bool(
|
||||
self, item: bool, *, as_type=False, deep=False
|
||||
) -> ProcessorResult: # noqa: FBT001
|
||||
return [{"kind": "text", "data": "True" if item else "False"}]
|
||||
|
||||
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"))
|
||||
def _process_clip(
|
||||
self, item: Any, *, as_type=False, deep=False
|
||||
) -> ProcessorResult:
|
||||
try:
|
||||
clip_model = getattr(item, "cond_stage_model", None)
|
||||
tokenizer = getattr(item, "tokenizer", None)
|
||||
|
||||
# np.save(full_output_folder / latent_file,
|
||||
# latent[""].cpu().numpy())
|
||||
text = [UIResult(kind="text", data="CLIP")]
|
||||
if clip_model:
|
||||
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",
|
||||
)
|
||||
)
|
||||
|
||||
return f"{filename_prefix}_{counter:05}"
|
||||
if tokenizer:
|
||||
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)
|
||||
|
||||
|
||||
processors = {
|
||||
torch.Tensor: process_tensor,
|
||||
list: process_list,
|
||||
dict: process_dict,
|
||||
bool: process_bool,
|
||||
}
|
||||
|
||||
__nodes__ = [MTB_Debug, MTB_SaveTensors]
|
||||
__nodes__ = [MTB_Debug]
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
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]
|
||||
+7
-4
@@ -4,12 +4,8 @@ import sys
|
||||
from pathlib import Path
|
||||
|
||||
import comfy.model_management as model_management
|
||||
import cv2
|
||||
import insightface
|
||||
import numpy as np
|
||||
import onnxruntime
|
||||
import torch
|
||||
from insightface.model_zoo.inswapper import INSwapper
|
||||
from PIL import Image
|
||||
|
||||
from ..errors import ModelNotFound
|
||||
@@ -43,6 +39,8 @@ class MTB_LoadFaceAnalysisModel:
|
||||
DEPRECATED = True
|
||||
|
||||
def load_model(self, faceswap_model: str):
|
||||
import insightface
|
||||
|
||||
if faceswap_model == "antelopev2":
|
||||
download_antelopev2()
|
||||
|
||||
@@ -81,6 +79,9 @@ class MTB_LoadFaceSwapModel:
|
||||
DEPRECATED = True
|
||||
|
||||
def load_model(self, faceswap_model: str):
|
||||
import onnxruntime
|
||||
from insightface.model_zoo.inswapper import INSwapper
|
||||
|
||||
model_path = get_model_path("insightface", faceswap_model)
|
||||
if not model_path or not model_path.exists():
|
||||
raise ModelNotFound(f"{faceswap_model} ({model_path})")
|
||||
@@ -212,6 +213,8 @@ def swap_face(
|
||||
face_swapper_model,
|
||||
faces_index: set[int] | None = None,
|
||||
) -> Image.Image:
|
||||
import cv2
|
||||
|
||||
if faces_index is None:
|
||||
faces_index = {0}
|
||||
log.debug(f"Swapping faces: {faces_index}")
|
||||
|
||||
+23
-9
@@ -7,6 +7,13 @@ from PIL import Image, ImageDraw, ImageFont
|
||||
from ..log import log
|
||||
from ..utils import comfy_dir, font_path, pil2tensor
|
||||
|
||||
# try:
|
||||
# from cairosvg import svg2png
|
||||
# HAS_CAIRO = True
|
||||
# except ImportError:
|
||||
# HAS_CAIRO = False
|
||||
|
||||
|
||||
# class MtbExamples:
|
||||
# """MTB Example Images"""
|
||||
|
||||
@@ -299,7 +306,7 @@ by default it fallsback to a default font.
|
||||
|
||||
def text_to_image(
|
||||
self,
|
||||
text: str,
|
||||
text: str | list[str],
|
||||
font,
|
||||
wrap,
|
||||
trim,
|
||||
@@ -341,11 +348,9 @@ by default it fallsback to a default font.
|
||||
color = (255, 255, 255, 255)
|
||||
background = (0, 0, 0, 255)
|
||||
|
||||
def render_text(text_to_render, alpha=None):
|
||||
def render_text(text_to_render: str, alpha=None) -> Image.Image:
|
||||
if trim:
|
||||
text_to_render = (
|
||||
text_to_render.encode("ascii", "ignore").decode().strip()
|
||||
)
|
||||
text_to_render = text_to_render.strip()
|
||||
if wrap:
|
||||
wrap_width = (((width / 100) * h_coverage) / font_size) * 2
|
||||
lines = textwrap.wrap(text_to_render, width=wrap_width)
|
||||
@@ -418,7 +423,9 @@ by default it fallsback to a default font.
|
||||
active_chunks.append((chunk["text"], alpha))
|
||||
|
||||
for chunk_text, alpha in active_chunks:
|
||||
chunk_img = render_text(chunk_text, alpha)
|
||||
chunk_img = render_text(
|
||||
chunk_text.encode("ascii", "ignore").decode(), alpha
|
||||
)
|
||||
frame = Image.alpha_composite(frame, chunk_img)
|
||||
|
||||
frames.append(frame)
|
||||
@@ -426,9 +433,16 @@ by default it fallsback to a default font.
|
||||
frame_tensors = [pil2tensor(frame) for frame in frames]
|
||||
return (torch.cat(frame_tensors, dim=0),)
|
||||
else:
|
||||
text_img = render_text(text)
|
||||
result = Image.alpha_composite(base_img, text_img)
|
||||
return (pil2tensor(result),)
|
||||
results = []
|
||||
if not isinstance(text, list):
|
||||
text = [text]
|
||||
|
||||
for t in text:
|
||||
text_img = render_text(t)
|
||||
result = Image.alpha_composite(base_img, text_img)
|
||||
results.append(result)
|
||||
|
||||
return (pil2tensor(results),)
|
||||
|
||||
|
||||
__nodes__ = [
|
||||
|
||||
+168
-7
@@ -4,18 +4,22 @@ import re
|
||||
import urllib.parse
|
||||
import urllib.request
|
||||
from math import pi
|
||||
from typing import Any
|
||||
|
||||
import comfy.model_management as model_management
|
||||
import comfy.model_management as mm
|
||||
import comfy.utils
|
||||
import numpy as np
|
||||
import torch
|
||||
from comfy.comfy_types.node_typing import IO as CIO
|
||||
from PIL import Image
|
||||
|
||||
from ..log import log
|
||||
from ..utils import (
|
||||
EASINGS,
|
||||
LazyProxyTensor,
|
||||
apply_easing,
|
||||
get_server_info,
|
||||
get_torch_tensor_info,
|
||||
numpy_NFOV,
|
||||
pil2tensor,
|
||||
tensor2np,
|
||||
@@ -132,12 +136,64 @@ class MTB_ApplyTextTemplate:
|
||||
CATEGORY = "mtb/utils"
|
||||
FUNCTION = "execute"
|
||||
|
||||
def execute(self, *, template: str, **kwargs):
|
||||
res = f"{template}"
|
||||
for k, v in kwargs.items():
|
||||
res = res.replace(f"{{{k}}}", f"{v}")
|
||||
def execute(self, *, template: str, **kwargs) -> tuple[str | list[str]]:
|
||||
keys = list(kwargs.keys())
|
||||
values = list(kwargs.values())
|
||||
|
||||
return (res,)
|
||||
has_list = any(isinstance(v, list) for v in values)
|
||||
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:
|
||||
@@ -341,7 +397,7 @@ class MTB_AutoPanEquilateral:
|
||||
|
||||
frames.append(frame)
|
||||
|
||||
model_management.throw_exception_if_processing_interrupted()
|
||||
mm.throw_exception_if_processing_interrupted()
|
||||
pbar.update(1)
|
||||
|
||||
return (pil2tensor(frames),)
|
||||
@@ -867,6 +923,108 @@ class MTB_TensorOps:
|
||||
return (result,)
|
||||
|
||||
|
||||
class MTB_GetItem:
|
||||
"""Generic index based getter for common types"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"container": (CIO.ANY,),
|
||||
"index": ("INT", {"default": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = (CIO.ANY,)
|
||||
RETURN_NAMES = ("item",)
|
||||
FUNCTION = "get_item"
|
||||
CATEGORY = "mtb/utils"
|
||||
|
||||
def get_item(self, container: Any, index: int):
|
||||
if "__getitem__" in dir(container):
|
||||
log.debug(f"Container is {type(container)}")
|
||||
res = container[index]
|
||||
if type(res) is torch.Tensor:
|
||||
res = res.unsqueeze(0)
|
||||
|
||||
return (res,)
|
||||
|
||||
|
||||
class MTB_BooleanNot:
|
||||
"""Inverts a boolean."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"bool_in": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("BOOLEAN",)
|
||||
RETURN_NAMES = ("inverted_bool",)
|
||||
FUNCTION = "invert"
|
||||
CATEGORY = "mtb/utils"
|
||||
|
||||
def invert(self, bool_in: bool):
|
||||
return (not bool_in,)
|
||||
|
||||
|
||||
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__ = [
|
||||
MTB_StringReplace,
|
||||
MTB_FitNumber,
|
||||
@@ -882,4 +1040,7 @@ __nodes__ = [
|
||||
MTB_FloatToFloats,
|
||||
MTB_FloatsToInts,
|
||||
MTB_TensorOps,
|
||||
MTB_BooleanNot,
|
||||
MTB_GetItem,
|
||||
MTB_ProxyTensor,
|
||||
]
|
||||
|
||||
+103
-47
@@ -3,11 +3,12 @@ import json
|
||||
import math
|
||||
import os
|
||||
|
||||
import comfy.model_management as model_management
|
||||
import comfy.utils
|
||||
import folder_paths
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from comfy import model_management
|
||||
from PIL import Image, ImageOps
|
||||
from PIL.PngImagePlugin import PngInfo
|
||||
from skimage.filters import gaussian
|
||||
@@ -74,7 +75,10 @@ class MTB_ExtractCoordinatesFromImage:
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"threshold": ("FLOAT",),
|
||||
"threshold": (
|
||||
"FLOAT",
|
||||
{"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01},
|
||||
),
|
||||
"max_points": ("INT", {"default": 50, "min": 0}),
|
||||
},
|
||||
"optional": {"image": ("IMAGE",), "mask": ("MASK",)},
|
||||
@@ -87,72 +91,124 @@ class MTB_ExtractCoordinatesFromImage:
|
||||
image: torch.Tensor | None = None,
|
||||
mask: torch.Tensor | None = None,
|
||||
) -> tuple[list[list[tuple[int, int]]], torch.Tensor]:
|
||||
if image is not None:
|
||||
batch_count, height, width, channel_count = image.shape
|
||||
imgs = image
|
||||
else:
|
||||
if mask is None:
|
||||
raise ValueError("Must provide either image or mask")
|
||||
batch_count, height, width = mask.shape
|
||||
channel_count = 1
|
||||
imgs = mask
|
||||
if image is None and mask is None:
|
||||
raise ValueError("Must provide either image or mask")
|
||||
|
||||
if channel_count not in [1, 2, 3, 4]:
|
||||
raise ValueError(f"Incorrect channel count: {channel_count}")
|
||||
if image is not None:
|
||||
batch_count, height, width, _channel_count = image.shape
|
||||
input_device = image.device
|
||||
if mask is not None:
|
||||
if mask.ndim == 2:
|
||||
mask = mask.unsqueeze(0)
|
||||
if mask.ndim != 3:
|
||||
raise ValueError(
|
||||
f"Mask has unexpected ndim: {mask.ndim}. Expected 2 or 3."
|
||||
)
|
||||
|
||||
b_mask, h_mask, w_mask = mask.shape
|
||||
if not (h_mask == height and w_mask == width):
|
||||
raise ValueError(
|
||||
f"Image dimensions ({height}x{width}) and mask dimensions ({h_mask}x{w_mask}) are spatially incompatible."
|
||||
)
|
||||
if b_mask == 1 and batch_count > 1:
|
||||
mask = mask.expand(batch_count, height, width)
|
||||
|
||||
elif b_mask != batch_count:
|
||||
raise ValueError(
|
||||
f"Image batch size ({batch_count}) and mask batch size ({b_mask}) are incompatible and mask cannot be broadcast."
|
||||
)
|
||||
else:
|
||||
if mask.ndim == 2:
|
||||
mask = mask.unsqueeze(0)
|
||||
|
||||
if mask.ndim != 3:
|
||||
raise ValueError(
|
||||
f"Mask has unexpected ndim: {mask.ndim} when image is not provided. Expected 2 or 3."
|
||||
)
|
||||
|
||||
batch_count, height, width = mask.shape
|
||||
input_device = mask.device
|
||||
|
||||
all_points: list[list[tuple[int, int]]] = []
|
||||
debug_images = torch.zeros(
|
||||
(batch_count, height, width, 3),
|
||||
dtype=torch.uint8,
|
||||
device=imgs.device,
|
||||
device=input_device,
|
||||
)
|
||||
|
||||
for i, img in enumerate(imgs):
|
||||
if channel_count == 1:
|
||||
alpha_channel = img if len(img.shape) == 2 else img[:, :, 0]
|
||||
elif channel_count == 2:
|
||||
alpha_channel = img[:, :, 1]
|
||||
elif channel_count == 4:
|
||||
alpha_channel = img[:, :, 3]
|
||||
points_tensor = torch.tensor(
|
||||
[255, 255, 255], dtype=torch.uint8, device=input_device
|
||||
)
|
||||
|
||||
for i in range(batch_count):
|
||||
value_threshold: torch.Tensor
|
||||
if image is not None:
|
||||
img_slice = image[i]
|
||||
img_channels = img_slice.shape[2]
|
||||
if img_channels == 1 or img_channels == 2:
|
||||
value_threshold = img_slice[:, :, 0]
|
||||
elif img_channels == 3 or img_channels == 4:
|
||||
value_threshold = img_slice[:, :, :3].max(dim=2)[0]
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported image channel count: {img_channels} for image at batch index {i}"
|
||||
)
|
||||
else:
|
||||
# get intensity
|
||||
alpha_channel = img[:, :, :3].max(dim=2)[0]
|
||||
mask_slice = mask[i]
|
||||
value_threshold = mask_slice
|
||||
|
||||
points = (alpha_channel > threshold).nonzero(as_tuple=False)
|
||||
condition = value_threshold > threshold
|
||||
if image is not None and mask is not None:
|
||||
mask_slice = mask[i]
|
||||
mask_active_condition = mask_slice > 0.0
|
||||
condition = condition & mask_active_condition
|
||||
|
||||
if len(points) > max_points:
|
||||
indices = torch.randperm(points.size(0), device=img.device)[
|
||||
:max_points
|
||||
]
|
||||
points = points[indices]
|
||||
points_yx = condition.nonzero(as_tuple=False)
|
||||
|
||||
points = [(int(y.item()), int(x.item())) for x, y in points]
|
||||
all_points.append(points)
|
||||
if points_yx.size(0) > max_points:
|
||||
# shuffle and pick max_points randomly
|
||||
indices = torch.randperm(
|
||||
points_yx.size(0), device=input_device
|
||||
)[:max_points]
|
||||
points_yx = points_yx[indices]
|
||||
elif max_points == 0:
|
||||
points_yx = torch.empty(
|
||||
(0, 2), dtype=torch.long, device=input_device
|
||||
)
|
||||
|
||||
for x, y in points:
|
||||
self._draw_circle(debug_images[i], (x, y), 5)
|
||||
current_points = [
|
||||
(int(p[1].item()), int(p[0].item())) for p in points_yx
|
||||
]
|
||||
all_points.append(current_points)
|
||||
for x_coord, y_coord in current_points:
|
||||
self._draw_circle(
|
||||
debug_images[i],
|
||||
(x_coord, y_coord),
|
||||
radius=5,
|
||||
color_tensor=points_tensor,
|
||||
)
|
||||
|
||||
return (all_points, debug_images)
|
||||
|
||||
@staticmethod
|
||||
def _draw_circle(
|
||||
image: torch.Tensor, center: tuple[int, int], radius: int
|
||||
image: torch.Tensor,
|
||||
center: tuple[int, int],
|
||||
radius: int,
|
||||
color_tensor: torch.Tensor,
|
||||
):
|
||||
"""Draw a 5px circle on the image."""
|
||||
x0, y0 = center
|
||||
for x in range(-radius, radius + 1):
|
||||
for y in range(-radius, radius + 1):
|
||||
in_radius = x**2 + y**2 <= radius**2
|
||||
in_bounds = (
|
||||
0 <= x0 + x < image.shape[1]
|
||||
and 0 <= y0 + y < image.shape[0]
|
||||
)
|
||||
if in_radius and in_bounds:
|
||||
image[y0 + y, x0 + x] = torch.tensor(
|
||||
[255, 255, 255],
|
||||
dtype=torch.uint8,
|
||||
device=image.device,
|
||||
)
|
||||
h, w, _ = image.shape
|
||||
min_x_bbox = max(0, x0 - radius)
|
||||
max_x_bbox = min(w - 1, x0 + radius)
|
||||
min_y_bbox = max(0, y0 - radius)
|
||||
max_y_bbox = min(h - 1, y0 + radius)
|
||||
|
||||
for py in range(min_y_bbox, max_y_bbox + 1):
|
||||
for px in range(min_x_bbox, max_x_bbox + 1):
|
||||
if (px - x0) ** 2 + (py - y0) ** 2 <= radius**2:
|
||||
image[py, px] = color_tensor
|
||||
|
||||
|
||||
class MTB_ColorCorrectGPU:
|
||||
|
||||
@@ -21,7 +21,11 @@ class MTB_StackImages:
|
||||
"match_method": (
|
||||
["error", "smallest", "largest"],
|
||||
{"default": "error"},
|
||||
)
|
||||
),
|
||||
"output_rgb": (
|
||||
"BOOLEAN",
|
||||
{"default": True, "tooltip": "Output RGB instead of RGBA"},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -29,7 +33,7 @@ class MTB_StackImages:
|
||||
FUNCTION = "stack"
|
||||
CATEGORY = "mtb/image utils"
|
||||
|
||||
def stack(self, vertical, match_method="error", **kwargs):
|
||||
def stack(self, vertical, match_method="error", output_rgb=True, **kwargs):
|
||||
if not kwargs:
|
||||
raise ValueError("At least one tensor must be provided.")
|
||||
|
||||
@@ -98,6 +102,9 @@ class MTB_StackImages:
|
||||
|
||||
stacked_tensor = torch.cat(normalized_tensors, dim=dim)
|
||||
|
||||
if output_rgb:
|
||||
stacked_tensor = stacked_tensor[:, :, :, :3]
|
||||
|
||||
return (stacked_tensor,)
|
||||
|
||||
def normalize_to_rgba(self, tensor):
|
||||
|
||||
+5
-38
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "comfy-mtb"
|
||||
version = "0.3.0"
|
||||
version = "0.6.0"
|
||||
description = "Animation oriented nodes pack for ComfyUI."
|
||||
license = { text = "MIT" }
|
||||
readme = "README.md"
|
||||
@@ -62,39 +62,6 @@ PublisherId = "mel"
|
||||
DisplayName = "comfy-mtb"
|
||||
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
|
||||
[tool.pyright]
|
||||
include = ["."]
|
||||
@@ -111,18 +78,18 @@ stubPath = "src/stubs"
|
||||
|
||||
reportMissingImports = true
|
||||
reportMissingTypeStubs = false
|
||||
reportExplicitAny = false
|
||||
typeCheckingMode = "basic"
|
||||
|
||||
pythonVersion = "3.10"
|
||||
pythonVersion = "3.11"
|
||||
pythonPlatform = "Windows"
|
||||
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
log_level = "DEBUG"
|
||||
log_cli = true
|
||||
markers = [
|
||||
"wip: tests that aren't fully finished yet",
|
||||
"heavy: marks tests as heavy (deselect with '-m \"not heavy\"')",
|
||||
|
||||
'''heavy: marks tests as heavy (deselect with '-m "not heavy"')''',
|
||||
]
|
||||
filterwarnings = ["ignore::UserWarning", 'ignore::DeprecationWarning']
|
||||
|
||||
|
||||
@@ -1,18 +0,0 @@
|
||||
{
|
||||
"exclude": [
|
||||
"**/node_modules",
|
||||
"**/__pycache__",
|
||||
],
|
||||
"ignore": [
|
||||
"extern"
|
||||
],
|
||||
"defineConstant": {
|
||||
"DEBUG": true
|
||||
},
|
||||
"venvPath": "../../../.venv/",
|
||||
"reportMissingImports": true,
|
||||
"reportMissingTypeStubs": false,
|
||||
"pythonVersion": "3.10",
|
||||
"pythonPlatform": "All",
|
||||
"reportOptionalMemberAccess": "none"
|
||||
}
|
||||
@@ -0,0 +1,637 @@
|
||||
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.
|
||||
@@ -1,6 +1,5 @@
|
||||
import contextlib
|
||||
import functools
|
||||
import importlib
|
||||
import math
|
||||
import operator
|
||||
import os
|
||||
@@ -9,11 +8,14 @@ import shutil
|
||||
import socket
|
||||
import subprocess
|
||||
import sys
|
||||
import textwrap
|
||||
import uuid
|
||||
import warnings
|
||||
from collections.abc import Callable, Sequence
|
||||
from enum import Enum
|
||||
from functools import reduce
|
||||
from pathlib import Path
|
||||
from types import EllipsisType
|
||||
from typing import TypeVar
|
||||
from urllib.parse import urlparse
|
||||
|
||||
@@ -462,25 +464,6 @@ def _run_command(shell_cmd, ignored_lines_start):
|
||||
print("Command executed successfully!")
|
||||
|
||||
|
||||
def import_install(package_name):
|
||||
package_spec = reqs_map.get(package_name, package_name)
|
||||
|
||||
try:
|
||||
importlib.import_module(package_name)
|
||||
|
||||
except Exception: # (ImportError, ModuleNotFoundError):
|
||||
run_command(
|
||||
[
|
||||
Path(sys.executable).as_posix(),
|
||||
"-m",
|
||||
"pip",
|
||||
"install",
|
||||
package_spec,
|
||||
]
|
||||
)
|
||||
importlib.import_module(package_name)
|
||||
|
||||
|
||||
# endregion
|
||||
|
||||
|
||||
@@ -544,6 +527,196 @@ PIL_FILTER_MAP = {
|
||||
|
||||
|
||||
# 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]:
|
||||
"""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)
|
||||
@@ -554,12 +727,12 @@ def handle_batch(
|
||||
tensor: torch.Tensor,
|
||||
func: Callable[[torch.Tensor], Image.Image | npt.NDArray[np.uint8]],
|
||||
) -> list[Image.Image] | list[npt.NDArray[np.uint8]]:
|
||||
"""Handles batch processing for a given tensor and conversion function."""
|
||||
"""Handle batch processing for a given tensor and conversion function."""
|
||||
return [func(tensor[i]) for i in range(tensor.shape[0])]
|
||||
|
||||
|
||||
def tensor2pil(tensor: torch.Tensor) -> list[Image.Image]:
|
||||
"""Converts a batch of tensors to a list of PIL Images."""
|
||||
"""Convert a batch of tensors to a list of PIL Images."""
|
||||
|
||||
def single_tensor2pil(t: torch.Tensor) -> Image.Image:
|
||||
np_array = to_numpy(t)
|
||||
@@ -576,7 +749,7 @@ def tensor2pil(tensor: torch.Tensor) -> list[Image.Image]:
|
||||
|
||||
|
||||
def pil2tensor(images: Image.Image | list[Image.Image]) -> torch.Tensor:
|
||||
"""Converts a PIL Image or a list of PIL Images to a tensor."""
|
||||
"""Convert a PIL Image or a list of PIL Images to a tensor."""
|
||||
|
||||
def single_pil2tensor(image: Image.Image) -> torch.Tensor:
|
||||
np_image = np.array(image).astype(np.float32) / 255.0
|
||||
|
||||
+236
-90
@@ -14,6 +14,51 @@ import { api } from '../../scripts/api.js'
|
||||
|
||||
// #region base utils
|
||||
|
||||
/**
|
||||
* 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))
|
||||
}
|
||||
|
||||
// - crude uuid
|
||||
export function makeUUID() {
|
||||
let dt = new Date().getTime()
|
||||
@@ -25,6 +70,19 @@ export function makeUUID() {
|
||||
return uuid
|
||||
}
|
||||
|
||||
// - basic debounce decorator
|
||||
export function debounce(func, delay) {
|
||||
let timeout
|
||||
let debounced = function (...args) {
|
||||
clearTimeout(timeout)
|
||||
timeout = setTimeout(() => func.apply(this, args), delay)
|
||||
}
|
||||
debounced.cancel = () => {
|
||||
clearTimeout(timeout)
|
||||
}
|
||||
return debounced
|
||||
}
|
||||
|
||||
//- local storage manager
|
||||
export class LocalStorageManager {
|
||||
constructor(namespace) {
|
||||
@@ -195,6 +253,7 @@ export function hideWidgetForGood(node, widget, suffix = '') {
|
||||
widget.origComputeSize = widget.computeSize
|
||||
widget.origSerializeValue = widget.serializeValue
|
||||
widget.computeSize = () => [0, -4] // -4 is due to the gap litegraph adds between widgets automatically
|
||||
widget.hidden = true
|
||||
widget.type = CONVERTED_TYPE + suffix
|
||||
// widget.serializeValue = () => {
|
||||
// // Prevent serializing the widget if we have no input linked
|
||||
@@ -276,6 +335,10 @@ export const getNamedWidget = (node, ...names) => {
|
||||
* @returns {{to:LGraphNode, from:LGraphNode, type:'error' | 'incoming' | 'outgoing'}}
|
||||
*/
|
||||
export const nodesFromLink = (node, link) => {
|
||||
if (typeof link === 'number') {
|
||||
link = app.graph.getLink(link)
|
||||
}
|
||||
|
||||
const fromNode = app.graph.getNodeById(link.origin_id)
|
||||
const toNode = app.graph.getNodeById(link.target_id)
|
||||
|
||||
@@ -366,12 +429,54 @@ export function getWidgetType(config) {
|
||||
|
||||
// #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
|
||||
/**
|
||||
* @param {NodeType} nodeType The nodetype to attach the documentation to
|
||||
* @param {str} prefix A prefix added to each dynamic inputs
|
||||
* @param {str | [str]} inputType The datatype(s) of those dynamic inputs
|
||||
* @param {{separator?:string, start_index?:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} [opts] Extra options
|
||||
* @param {{separator?:string,rename_menu?:'label'|'name', start_index?:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} [opts] Extra options
|
||||
* @returns
|
||||
*/
|
||||
export const setupDynamicConnections = (
|
||||
@@ -385,20 +490,115 @@ export const setupDynamicConnections = (
|
||||
Object.getOwnPropertyDescriptors(nodeType).title.value,
|
||||
)
|
||||
|
||||
/** @type {{separator:string, start_index:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} */
|
||||
/** @type {{separator:string,rename_menu?:"label"|"name" start_index:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} */
|
||||
const options = Object.assign(
|
||||
{
|
||||
separator: '_',
|
||||
start_index: 1,
|
||||
rename_menu: 'label',
|
||||
},
|
||||
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 inputList = typeof inputType === 'object'
|
||||
|
||||
nodeType.prototype.onNodeCreated = function () {
|
||||
const r = onNodeCreated ? onNodeCreated.apply(this, []) : undefined
|
||||
this.addInput(
|
||||
|
||||
const input = addDynamicInput(
|
||||
this,
|
||||
`${prefix}${options.separator}${options.start_index}`,
|
||||
inputList ? '*' : inputType,
|
||||
)
|
||||
@@ -470,10 +670,7 @@ export const dynamic_connection = (
|
||||
opts || {},
|
||||
)
|
||||
|
||||
// function to test if input is a dynamic one
|
||||
const isDynamicInput = (inputName) => inputName.startsWith(connectionPrefix)
|
||||
|
||||
if (node.inputs.length > 0 && !isDynamicInput(node.inputs[index].name)) {
|
||||
if (node.inputs.length > 0 && !isDynamicInput(node.inputs[index])) {
|
||||
return
|
||||
}
|
||||
|
||||
@@ -483,6 +680,7 @@ export const dynamic_connection = (
|
||||
const nameArray = options.nameArray || []
|
||||
|
||||
const clean_inputs = () => {
|
||||
if (node.id < 0) return // being duplicated
|
||||
if (node.inputs.length === 0) return
|
||||
|
||||
let w_count = node.widgets?.length || 0
|
||||
@@ -492,7 +690,7 @@ export const dynamic_connection = (
|
||||
const to_remove = []
|
||||
for (let n = 1; n < node.inputs.length; n++) {
|
||||
const element = node.inputs[n]
|
||||
if (!element.link && isDynamicInput(element.name)) {
|
||||
if (!element.link && isDynamicInput(element)) {
|
||||
if (node.widgets) {
|
||||
const w = node.widgets.find((w) => w.name === element.name)
|
||||
if (w) {
|
||||
@@ -506,9 +704,12 @@ export const dynamic_connection = (
|
||||
}
|
||||
for (let i = 0; i < to_remove.length; i++) {
|
||||
const id = to_remove[i]
|
||||
|
||||
node.removeInput(id)
|
||||
i_count -= 1
|
||||
try {
|
||||
node.removeInput(id)
|
||||
i_count -= 1
|
||||
} catch (err) {
|
||||
errorLogger('Cannot remove input', err)
|
||||
}
|
||||
}
|
||||
node.inputs.length = i_count
|
||||
|
||||
@@ -522,7 +723,7 @@ export const dynamic_connection = (
|
||||
for (let i = 0; i < node.inputs.length; i++) {
|
||||
let name = ''
|
||||
// rename only prefixed inputs
|
||||
if (isDynamicInput(node.inputs[i].name)) {
|
||||
if (node.inputs[i].name.startsWith(connectionPrefix)) {
|
||||
// prefixed => rename and increase index
|
||||
name = `${connectionPrefix}${prefixed_idx}`
|
||||
prefixed_idx += 1
|
||||
@@ -576,9 +777,8 @@ export const dynamic_connection = (
|
||||
if (node.inputs.length === 0) return
|
||||
// add an extra input
|
||||
if (node.inputs[node.inputs.length - 1].link !== null) {
|
||||
// count only the prefixed inputs
|
||||
const nextIndex = node.inputs.reduce(
|
||||
(acc, cur) => (isDynamicInput(cur.name) ? ++acc : acc),
|
||||
(acc, cur) => (isDynamicInput(cur) ? ++acc : acc),
|
||||
0,
|
||||
)
|
||||
|
||||
@@ -588,7 +788,7 @@ export const dynamic_connection = (
|
||||
: `${connectionPrefix}${nextIndex + options.start_index}`
|
||||
|
||||
infoLogger(`Adding input ${nextIndex + 1} (${name})`)
|
||||
node.addInput(name, conType)
|
||||
addDynamicInput(node, name, conType)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -621,21 +821,21 @@ function getBrightness(rgbObj) {
|
||||
export function calculateTotalChildrenHeight(parentElement) {
|
||||
let totalHeight = 0
|
||||
|
||||
if (!parentElement || !parentElement.children) {
|
||||
return 0
|
||||
}
|
||||
|
||||
for (const child of parentElement.children) {
|
||||
const style = window.getComputedStyle(child)
|
||||
|
||||
// Get height as an integer (without 'px')
|
||||
const height = Number.parseInt(style.height, 10)
|
||||
const height = Number.parseFloat(style.height)
|
||||
const marginTop = Number.parseFloat(style.marginTop)
|
||||
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
|
||||
}
|
||||
|
||||
return totalHeight
|
||||
return Math.ceil(totalHeight)
|
||||
}
|
||||
|
||||
export const loadScript = (
|
||||
@@ -646,13 +846,15 @@ export const loadScript = (
|
||||
return new Promise((resolve, reject) => {
|
||||
try {
|
||||
// Check if the script already exists
|
||||
const existingScript = document.querySelector(`script[src="${FILE_URL}"]`)
|
||||
if (existingScript) {
|
||||
resolve({ status: true, message: 'Script already loaded' })
|
||||
let scriptEle = document.querySelector(`script[src="${FILE_URL}"]`)
|
||||
if (scriptEle) {
|
||||
scriptEle.addEventListener('load', (_ev) => {
|
||||
resolve({ status: true })
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
const scriptEle = document.createElement('script')
|
||||
scriptEle = document.createElement('script')
|
||||
scriptEle.type = type
|
||||
scriptEle.async = async
|
||||
scriptEle.src = FILE_URL
|
||||
@@ -671,6 +873,8 @@ export const loadScript = (
|
||||
document.body.appendChild(scriptEle)
|
||||
} catch (error) {
|
||||
reject(error)
|
||||
} finally {
|
||||
infoLogger(`Finally loaded script: ${FILE_URL}`)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -784,12 +988,10 @@ function loadParser(shiki) {
|
||||
|
||||
export const ensureMarkdownParser = async (callback) => {
|
||||
infoLogger('Ensuring md parser')
|
||||
let use_shiki = false
|
||||
try {
|
||||
use_shiki = await api.getSetting('mtb.Use Shiki')
|
||||
} catch (e) {
|
||||
console.warn('Option not available yet', e)
|
||||
}
|
||||
const use_shiki = app.extensionManager.setting.get(
|
||||
'mtb.noteplus.use-shiki',
|
||||
false,
|
||||
)
|
||||
|
||||
if (window.MTB?.mdParser) {
|
||||
infoLogger('Markdown parser found')
|
||||
@@ -814,8 +1016,7 @@ export const ensureMarkdownParser = async (callback) => {
|
||||
callbackQueue.push(callback)
|
||||
}
|
||||
|
||||
await parserPromise
|
||||
await parserPromise
|
||||
await await parserPromise
|
||||
|
||||
return window.MTB.mdParser
|
||||
}
|
||||
@@ -1154,58 +1355,3 @@ export const setServerInfo = async (opts) => {
|
||||
}
|
||||
|
||||
// #endregion
|
||||
|
||||
// #region Authoring API / graph utilities
|
||||
export const getAPIInputs = () => {
|
||||
const inputs = {}
|
||||
let counter = 1
|
||||
for (const node of getNodes(true)) {
|
||||
const widgets = node.widgets
|
||||
|
||||
if (node.properties.mtb_api && node.properties.useAPI) {
|
||||
if (node.properties.mtb_api.inputs) {
|
||||
for (const currentName in node.properties.mtb_api.inputs) {
|
||||
const current = node.properties.mtb_api.inputs[currentName]
|
||||
if (current.enabled) {
|
||||
const inputName = current.name || currentName
|
||||
const widget = widgets.find((w) => w.name === currentName)
|
||||
if (!widget) continue
|
||||
if (!(inputName in inputs)) {
|
||||
inputs[inputName] = {
|
||||
...current,
|
||||
id: counter,
|
||||
name: inputName,
|
||||
type: current.type,
|
||||
node_id: node.id,
|
||||
widgets: [],
|
||||
}
|
||||
}
|
||||
inputs[inputName].widgets.push(widget)
|
||||
counter = counter + 1
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return inputs
|
||||
}
|
||||
|
||||
export const getNodes = (skip_unused) => {
|
||||
const nodes = []
|
||||
for (const outerNode of app.graph.computeExecutionOrder(false)) {
|
||||
const skipNode =
|
||||
(outerNode.mode === 2 || outerNode.mode === 4) && skip_unused
|
||||
const innerNodes =
|
||||
!skipNode && outerNode.getInnerNodes
|
||||
? outerNode.getInnerNodes()
|
||||
: [outerNode]
|
||||
for (const node of innerNodes) {
|
||||
if ((node.mode === 2 || node.mode === 4) && skip_unused) {
|
||||
continue
|
||||
}
|
||||
nodes.push(node)
|
||||
}
|
||||
}
|
||||
return nodes
|
||||
}
|
||||
// #endregion
|
||||
|
||||
+76
-100
@@ -11,7 +11,12 @@
|
||||
/// <reference path="../types/typedefs.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'
|
||||
|
||||
function escapeHtml(unsafe) {
|
||||
@@ -28,7 +33,7 @@ function createDebugSection(title) {
|
||||
margin: '8px 0',
|
||||
padding: '8px',
|
||||
borderRadius: '4px',
|
||||
backgroundColor: 'rgba(0,0,0,0.2)'
|
||||
backgroundColor: 'rgba(0,0,0,0.2)',
|
||||
})
|
||||
|
||||
const header = mtb_ui.makeElement('h3', {
|
||||
@@ -37,7 +42,7 @@ function createDebugSection(title) {
|
||||
borderBottom: '1px solid rgba(255,255,255,0.1)',
|
||||
fontSize: '14px',
|
||||
fontWeight: 'bold',
|
||||
color: '#9f9'
|
||||
color: '#9f9',
|
||||
})
|
||||
header.textContent = title
|
||||
section.appendChild(header)
|
||||
@@ -45,25 +50,25 @@ function createDebugSection(title) {
|
||||
return section
|
||||
}
|
||||
|
||||
function createDebugContent(content, type) {
|
||||
function createDebugContent(item) {
|
||||
const wrapper = mtb_ui.makeElement('div', {
|
||||
margin: '4px 0'
|
||||
margin: '4px 0',
|
||||
})
|
||||
|
||||
if (type === 'text') {
|
||||
const text = mtb_ui.makeElement('p', {
|
||||
if (item.kind === 'text') {
|
||||
const text = mtb_ui.makeElement('div', {
|
||||
margin: '2px 0',
|
||||
fontFamily: 'monospace',
|
||||
whiteSpace: 'pre-wrap'
|
||||
whiteSpace: 'pre-wrap',
|
||||
})
|
||||
text.innerHTML = content
|
||||
text.innerHTML = item.data
|
||||
wrapper.appendChild(text)
|
||||
} else if (type === 'image') {
|
||||
} else if (item.kind === 'b64_images') {
|
||||
const img = mtb_ui.makeElement('img', {
|
||||
width: '100%',
|
||||
borderRadius: '2px'
|
||||
borderRadius: '2px',
|
||||
})
|
||||
img.src = content
|
||||
img.src = item.data
|
||||
wrapper.appendChild(img)
|
||||
}
|
||||
|
||||
@@ -80,115 +85,85 @@ app.registerExtension({
|
||||
*/
|
||||
async beforeRegisterNodeDef(nodeType, nodeData, app) {
|
||||
if (nodeData.name === 'Debug (mtb)') {
|
||||
const onNodeCreated = nodeType.prototype.onNodeCreated
|
||||
nodeType.prototype.onNodeCreated = function (...args) {
|
||||
this.options = {}
|
||||
const r = onNodeCreated ? onNodeCreated.apply(this, args) : undefined
|
||||
this.addInput('anything_1', '*')
|
||||
return r
|
||||
const clear_widgets = (target) => {
|
||||
if (target.widgets) {
|
||||
let tgt_len = target.widgets.length
|
||||
for (let i = 0; i < target.widgets.length; i++) {
|
||||
if (
|
||||
![
|
||||
'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 onConnectionsChange = nodeType.prototype.onConnectionsChange
|
||||
/**
|
||||
* @param {OnConnectionsChangeParams} args
|
||||
*/
|
||||
nodeType.prototype.onConnectionsChange = function (...args) {
|
||||
const [_type, index, connected, link_info, ioSlot] = args
|
||||
const r = onConnectionsChange
|
||||
? 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,
|
||||
const original_getExtraMenuOptions =
|
||||
nodeType.prototype.getExtraMenuOptions
|
||||
nodeType.prototype.getExtraMenuOptions = function (_, options) {
|
||||
original_getExtraMenuOptions?.apply(this, arguments)
|
||||
options.push({
|
||||
content: '🐛 Clear Outputs',
|
||||
callback: async () => {
|
||||
clear_widgets(this)
|
||||
},
|
||||
})
|
||||
|
||||
//- 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
|
||||
}
|
||||
|
||||
setupDynamicConnections(nodeType, 'var', '*')
|
||||
|
||||
const onExecuted = nodeType.prototype.onExecuted
|
||||
nodeType.prototype.onExecuted = function (...args) {
|
||||
onExecuted?.apply(this, args)
|
||||
const [data, ..._rest] = args
|
||||
|
||||
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
|
||||
}
|
||||
clear_widgets(this)
|
||||
|
||||
const inputData = {}
|
||||
|
||||
const uiData = data.ui || data
|
||||
|
||||
if (uiData.items) {
|
||||
uiData.items.forEach(item => {
|
||||
const inputName = item.input
|
||||
if (!inputData[inputName]) {
|
||||
inputData[inputName] = { text: [], b64_images: [] }
|
||||
}
|
||||
if (item.text) {
|
||||
inputData[inputName].text.push(...item.text)
|
||||
}
|
||||
if (item.b64_images) {
|
||||
inputData[inputName].b64_images.push(...item.b64_images)
|
||||
}
|
||||
})
|
||||
}
|
||||
const name_to_label = this.inputs.reduce((acc, input) => {
|
||||
acc[input.name] = input.label || input.name
|
||||
return acc
|
||||
}, {})
|
||||
|
||||
let widgetI = 1
|
||||
|
||||
if (uiData.items) {
|
||||
uiData.items.forEach((item) => {
|
||||
const inputName = item.input
|
||||
inputData[inputName] = item.items
|
||||
})
|
||||
}
|
||||
const mainDebugContainer = mtb_ui.makeElement('div', {
|
||||
width: '100%',
|
||||
})
|
||||
let hasContent = false
|
||||
for (const [inputName, content] of Object.entries(inputData)) {
|
||||
if (content.text.length === 0 && content.b64_images.length === 0) {
|
||||
if (!content || content?.length === 0) {
|
||||
continue
|
||||
}
|
||||
hasContent = true
|
||||
|
||||
const section = createDebugSection(inputName)
|
||||
const section = createDebugSection(name_to_label[inputName])
|
||||
|
||||
if (content.text.length > 0) {
|
||||
content.text.forEach(text => {
|
||||
section.appendChild(createDebugContent(text, 'text'))
|
||||
})
|
||||
for (const item of content) {
|
||||
section.appendChild(createDebugContent(item))
|
||||
}
|
||||
|
||||
if (content.b64_images.length > 0) {
|
||||
content.b64_images.forEach(img => {
|
||||
section.appendChild(createDebugContent(img, 'image'))
|
||||
})
|
||||
}
|
||||
|
||||
this.addDOMWidget(
|
||||
`debug_section_${widgetI}`,
|
||||
'CUSTOM',
|
||||
section,
|
||||
{}
|
||||
)
|
||||
widgetI++
|
||||
mainDebugContainer.appendChild(section)
|
||||
}
|
||||
if (hasContent) {
|
||||
this.addDOMWidget('debug_output', 'CUSTOM', mainDebugContainer, {
|
||||
hideOnZoom: false,
|
||||
})
|
||||
}
|
||||
|
||||
this.onRemoved = function () {
|
||||
@@ -199,8 +174,9 @@ app.registerExtension({
|
||||
widget.onRemoved?.()
|
||||
widget.onRemove?.()
|
||||
}
|
||||
shared.cleanupNode(this)
|
||||
cleanupNode(this)
|
||||
}
|
||||
this.setDirtyCanvas(true, true)
|
||||
}
|
||||
}
|
||||
},
|
||||
|
||||
+296
-296
@@ -13,40 +13,40 @@ import { api } from '../../scripts/api.js'
|
||||
import { app } from '../../scripts/app.js'
|
||||
import { LocalStorageManager } from './comfy_shared.js'
|
||||
const styles = {
|
||||
lighbox: {
|
||||
position: 'fixed',
|
||||
top: 0,
|
||||
left: 0,
|
||||
width: '100vw',
|
||||
height: '100vh',
|
||||
background: 'rgba(0,0,0,0.5)',
|
||||
display: 'none',
|
||||
justifyContent: 'center',
|
||||
alignItems: 'center',
|
||||
zIndex: 999,
|
||||
},
|
||||
lightboxBtn: (extra) => ({
|
||||
position: 'absolute',
|
||||
top: '50%',
|
||||
background: 'none',
|
||||
border: 'none',
|
||||
color: '#fff',
|
||||
zIndex: 1000,
|
||||
fontSize: '30px',
|
||||
cursor: 'pointer',
|
||||
pointerEvents: 'auto',
|
||||
...extra,
|
||||
}),
|
||||
img_list: {
|
||||
minHeight: '30px',
|
||||
maxHeight: '300px',
|
||||
width: '100vw',
|
||||
position: 'absolute',
|
||||
bottom: 0,
|
||||
zIndex: 10,
|
||||
background: '#333',
|
||||
overflow: 'auto',
|
||||
},
|
||||
lighbox: {
|
||||
position: 'fixed',
|
||||
top: 0,
|
||||
left: 0,
|
||||
width: '100vw',
|
||||
height: '100vh',
|
||||
background: 'rgba(0,0,0,0.5)',
|
||||
display: 'none',
|
||||
justifyContent: 'center',
|
||||
alignItems: 'center',
|
||||
zIndex: 999,
|
||||
},
|
||||
lightboxBtn: (extra) => ({
|
||||
position: 'absolute',
|
||||
top: '50%',
|
||||
background: 'none',
|
||||
border: 'none',
|
||||
color: '#fff',
|
||||
zIndex: 1000,
|
||||
fontSize: '30px',
|
||||
cursor: 'pointer',
|
||||
pointerEvents: 'auto',
|
||||
...extra,
|
||||
}),
|
||||
img_list: {
|
||||
minHeight: '30px',
|
||||
maxHeight: '300px',
|
||||
width: '100vw',
|
||||
position: 'absolute',
|
||||
bottom: 0,
|
||||
zIndex: 10,
|
||||
background: '#333',
|
||||
overflow: 'auto',
|
||||
},
|
||||
}
|
||||
|
||||
let currentImageIndex = 0
|
||||
@@ -58,299 +58,299 @@ const storage = new LocalStorageManager('mtb')
|
||||
let activated = storage.get('image_feed', false)
|
||||
|
||||
app.registerExtension({
|
||||
name: 'mtb.ImageFeed',
|
||||
setup: () => {
|
||||
app.ui.settings.addSetting({
|
||||
id: 'mtb.Main.image-feed-enabled',
|
||||
category: ['mtb', 'Main', 'image-feed-enabled'],
|
||||
name: 'Enable Image Feed',
|
||||
type: 'boolean',
|
||||
defaultValue: false,
|
||||
attrs: {
|
||||
style: {
|
||||
fontFamily: 'monospace',
|
||||
},
|
||||
},
|
||||
async onChange(value) {
|
||||
storage.set('image_feed', value)
|
||||
activated = value
|
||||
},
|
||||
})
|
||||
},
|
||||
init: async () => {
|
||||
if (!activated) {
|
||||
return
|
||||
}
|
||||
const pythongossFeed = app.extensions.find(
|
||||
(e) => e.name === 'pysssss.ImageFeed',
|
||||
)
|
||||
if (pythongossFeed) {
|
||||
console.warn(
|
||||
"[mtb] - Aborting the loading of mtb's imageFeed in favor of pysssss.ImageFeed",
|
||||
)
|
||||
activated = false // just in case other methods are added later on
|
||||
return
|
||||
}
|
||||
// - HTML & CSS
|
||||
//- lightbox
|
||||
const lightboxContainer = document.createElement('div')
|
||||
Object.assign(lightboxContainer.style, styles.lighbox)
|
||||
name: 'mtb.ImageFeed',
|
||||
setup: () => {
|
||||
app.ui.settings.addSetting({
|
||||
id: 'mtb.Main.image-feed-enabled',
|
||||
category: ['mtb', ' Main', 'image-feed-enabled'],
|
||||
name: 'Enable Image Feed',
|
||||
type: 'boolean',
|
||||
defaultValue: false,
|
||||
attrs: {
|
||||
style: {
|
||||
fontFamily: 'monospace',
|
||||
},
|
||||
},
|
||||
async onChange(value) {
|
||||
storage.set('image_feed', value)
|
||||
activated = value
|
||||
},
|
||||
})
|
||||
},
|
||||
init: async () => {
|
||||
if (!activated) {
|
||||
return
|
||||
}
|
||||
const pythongossFeed = app.extensions.find(
|
||||
(e) => e.name === 'pysssss.ImageFeed',
|
||||
)
|
||||
if (pythongossFeed) {
|
||||
console.warn(
|
||||
"[mtb] - Aborting the loading of mtb's imageFeed in favor of pysssss.ImageFeed",
|
||||
)
|
||||
activated = false // just in case other methods are added later on
|
||||
return
|
||||
}
|
||||
// - HTML & CSS
|
||||
//- lightbox
|
||||
const lightboxContainer = document.createElement('div')
|
||||
Object.assign(lightboxContainer.style, styles.lighbox)
|
||||
|
||||
const lightboxImage = document.createElement('img')
|
||||
Object.assign(lightboxImage.style, {
|
||||
maxHeight: '100%',
|
||||
maxWidth: '100%',
|
||||
borderRadius: '5px',
|
||||
})
|
||||
const lightboxImage = document.createElement('img')
|
||||
Object.assign(lightboxImage.style, {
|
||||
maxHeight: '100%',
|
||||
maxWidth: '100%',
|
||||
borderRadius: '5px',
|
||||
})
|
||||
|
||||
// previous and next buttons
|
||||
const lightboxPrevBtn = document.createElement('button')
|
||||
const lightboxNextBtn = document.createElement('button')
|
||||
// previous and next buttons
|
||||
const lightboxPrevBtn = document.createElement('button')
|
||||
const lightboxNextBtn = document.createElement('button')
|
||||
|
||||
lightboxPrevBtn.textContent = '❮'
|
||||
lightboxNextBtn.textContent = '❯'
|
||||
lightboxPrevBtn.textContent = '❮'
|
||||
lightboxNextBtn.textContent = '❯'
|
||||
|
||||
Object.assign(lightboxPrevBtn.style, styles.lightboxBtn({ left: '0%' }))
|
||||
Object.assign(lightboxNextBtn.style, styles.lightboxBtn({ right: '0%' }))
|
||||
Object.assign(lightboxPrevBtn.style, styles.lightboxBtn({ left: '0%' }))
|
||||
Object.assign(lightboxNextBtn.style, styles.lightboxBtn({ right: '0%' }))
|
||||
|
||||
// close button
|
||||
const lightboxCloseBtn = document.createElement('button')
|
||||
Object.assign(
|
||||
lightboxCloseBtn.style,
|
||||
styles.lightboxBtn({ right: '0', top: '0' }),
|
||||
)
|
||||
lightboxCloseBtn.textContent = '❌'
|
||||
// close button
|
||||
const lightboxCloseBtn = document.createElement('button')
|
||||
Object.assign(
|
||||
lightboxCloseBtn.style,
|
||||
styles.lightboxBtn({ right: '0', top: '0' }),
|
||||
)
|
||||
lightboxCloseBtn.textContent = '❌'
|
||||
|
||||
const lightboxButtons = document.createElement('div')
|
||||
Object.assign(lightboxButtons.style, {
|
||||
position: 'absolute',
|
||||
top: '0%',
|
||||
right: '0%',
|
||||
// transform: "translate(50%, -50%)",
|
||||
height: '100%',
|
||||
width: '100%',
|
||||
background: 'none',
|
||||
border: 'none',
|
||||
color: '#fff',
|
||||
fontSize: '30px',
|
||||
cursor: 'pointer',
|
||||
pointerEvents: 'none',
|
||||
})
|
||||
const lightboxButtons = document.createElement('div')
|
||||
Object.assign(lightboxButtons.style, {
|
||||
position: 'absolute',
|
||||
top: '0%',
|
||||
right: '0%',
|
||||
// transform: "translate(50%, -50%)",
|
||||
height: '100%',
|
||||
width: '100%',
|
||||
background: 'none',
|
||||
border: 'none',
|
||||
color: '#fff',
|
||||
fontSize: '30px',
|
||||
cursor: 'pointer',
|
||||
pointerEvents: 'none',
|
||||
})
|
||||
|
||||
lightboxButtons.append(lightboxPrevBtn, lightboxNextBtn, lightboxCloseBtn)
|
||||
lightboxContainer.append(lightboxButtons, lightboxImage)
|
||||
lightboxButtons.append(lightboxPrevBtn, lightboxNextBtn, lightboxCloseBtn)
|
||||
lightboxContainer.append(lightboxButtons, lightboxImage)
|
||||
|
||||
//- image list
|
||||
const imageListContainer = document.createElement('div')
|
||||
Object.assign(imageListContainer.style, styles.img_list)
|
||||
//- image list
|
||||
const imageListContainer = document.createElement('div')
|
||||
Object.assign(imageListContainer.style, styles.img_list)
|
||||
|
||||
const createImgListBtn = (text, style) => {
|
||||
const btn = document.createElement('button')
|
||||
btn.type = 'button'
|
||||
btn.textContent = text
|
||||
Object.assign(btn.style, {
|
||||
...style,
|
||||
border: 'none',
|
||||
color: '#fff',
|
||||
background: 'none',
|
||||
height: '20px',
|
||||
cursor: 'pointer',
|
||||
position: 'absolute',
|
||||
top: '5px',
|
||||
fontSize: '12px',
|
||||
lineHeight: '12px',
|
||||
})
|
||||
imageListContainer.append(btn)
|
||||
return btn
|
||||
}
|
||||
const showBtn = document.createElement('button')
|
||||
const closeBtn = createImgListBtn('❌', {
|
||||
width: '20px',
|
||||
textIndent: '-4px',
|
||||
right: '5px',
|
||||
})
|
||||
const loadButton = createImgListBtn('Load Session History', {
|
||||
right: '90px',
|
||||
})
|
||||
const clearButton = createImgListBtn('Clear', {
|
||||
right: '30px',
|
||||
})
|
||||
const createImgListBtn = (text, style) => {
|
||||
const btn = document.createElement('button')
|
||||
btn.type = 'button'
|
||||
btn.textContent = text
|
||||
Object.assign(btn.style, {
|
||||
...style,
|
||||
border: 'none',
|
||||
color: '#fff',
|
||||
background: 'none',
|
||||
height: '20px',
|
||||
cursor: 'pointer',
|
||||
position: 'absolute',
|
||||
top: '5px',
|
||||
fontSize: '12px',
|
||||
lineHeight: '12px',
|
||||
})
|
||||
imageListContainer.append(btn)
|
||||
return btn
|
||||
}
|
||||
const showBtn = document.createElement('button')
|
||||
const closeBtn = createImgListBtn('❌', {
|
||||
width: '20px',
|
||||
textIndent: '-4px',
|
||||
right: '5px',
|
||||
})
|
||||
const loadButton = createImgListBtn('Load Session History', {
|
||||
right: '90px',
|
||||
})
|
||||
const clearButton = createImgListBtn('Clear', {
|
||||
right: '30px',
|
||||
})
|
||||
|
||||
//- tools popup button
|
||||
showBtn.classList.add('comfy-settings-btn')
|
||||
Object.assign(showBtn.style, {
|
||||
right: '16px',
|
||||
cursor: 'pointer',
|
||||
display: 'none',
|
||||
})
|
||||
//- tools popup button
|
||||
showBtn.classList.add('comfy-settings-btn')
|
||||
Object.assign(showBtn.style, {
|
||||
right: '16px',
|
||||
cursor: 'pointer',
|
||||
display: 'none',
|
||||
})
|
||||
|
||||
//- append to DOM
|
||||
document.body.append(imageListContainer)
|
||||
//- append to DOM
|
||||
document.body.append(imageListContainer)
|
||||
|
||||
showBtn.textContent = '🖼'
|
||||
showBtn.onclick = () => {
|
||||
imageListContainer.style.display = 'block'
|
||||
showBtn.style.display = 'none'
|
||||
}
|
||||
document.querySelector('.comfy-settings-btn').after(showBtn)
|
||||
document.querySelector('.comfy-settings-btn').after(lightboxContainer)
|
||||
showBtn.textContent = '🖼'
|
||||
showBtn.onclick = () => {
|
||||
imageListContainer.style.display = 'block'
|
||||
showBtn.style.display = 'none'
|
||||
}
|
||||
document.querySelector('.comfy-settings-btn').after(showBtn)
|
||||
document.querySelector('.comfy-settings-btn').after(lightboxContainer)
|
||||
|
||||
// for (const { output } of history) {
|
||||
// if (output?.images) {
|
||||
// for (const src of output.images) {
|
||||
// const img = document.createElement("img");
|
||||
// const but = document.createElement("button");
|
||||
// for (const { output } of history) {
|
||||
// if (output?.images) {
|
||||
// for (const src of output.images) {
|
||||
// const img = document.createElement("img");
|
||||
// const but = document.createElement("button");
|
||||
|
||||
//- callbacks
|
||||
closeBtn.onclick = () => {
|
||||
imageListContainer.style.display = 'none'
|
||||
showBtn.style.display = 'unset'
|
||||
}
|
||||
//- callbacks
|
||||
closeBtn.onclick = () => {
|
||||
imageListContainer.style.display = 'none'
|
||||
showBtn.style.display = 'unset'
|
||||
}
|
||||
|
||||
clearButton.onclick = () => {
|
||||
imageListContainer.replaceChildren(closeBtn, clearButton, loadButton)
|
||||
}
|
||||
clearButton.onclick = () => {
|
||||
imageListContainer.replaceChildren(closeBtn, clearButton, loadButton)
|
||||
}
|
||||
|
||||
lightboxNextBtn.onclick = () => {
|
||||
currentImageIndex = (currentImageIndex + 1) % imageUrls.length
|
||||
const imageUrl = imageUrls[currentImageIndex]
|
||||
lightboxImage.src = imageUrl
|
||||
}
|
||||
lightboxNextBtn.onclick = () => {
|
||||
currentImageIndex = (currentImageIndex + 1) % imageUrls.length
|
||||
const imageUrl = imageUrls[currentImageIndex]
|
||||
lightboxImage.src = imageUrl
|
||||
}
|
||||
|
||||
// Modify the lightboxPrevBtn onclick callback
|
||||
lightboxPrevBtn.onclick = () => {
|
||||
currentImageIndex =
|
||||
(currentImageIndex - 1 + imageUrls.length) % imageUrls.length
|
||||
const imageUrl = imageUrls[currentImageIndex]
|
||||
lightboxImage.src = imageUrl
|
||||
}
|
||||
// Modify the lightboxPrevBtn onclick callback
|
||||
lightboxPrevBtn.onclick = () => {
|
||||
currentImageIndex =
|
||||
(currentImageIndex - 1 + imageUrls.length) % imageUrls.length
|
||||
const imageUrl = imageUrls[currentImageIndex]
|
||||
lightboxImage.src = imageUrl
|
||||
}
|
||||
|
||||
lightboxCloseBtn.onclick = () => {
|
||||
lightboxContainer.style.display = 'none'
|
||||
}
|
||||
lightboxImage.onclick = lightboxNextBtn.onclick
|
||||
/**
|
||||
* This is the function that creates the image buttons for the image list
|
||||
* They are wrapped in a button so that they can be clicked and open
|
||||
* the image in the lightbox.
|
||||
* @param {*} src
|
||||
*/
|
||||
const createImageBtn = (src) => {
|
||||
console.debug(`making image ${src.filename}`)
|
||||
const img = document.createElement('img')
|
||||
const but = document.createElement('button')
|
||||
lightboxCloseBtn.onclick = () => {
|
||||
lightboxContainer.style.display = 'none'
|
||||
}
|
||||
lightboxImage.onclick = lightboxNextBtn.onclick
|
||||
/**
|
||||
* This is the function that creates the image buttons for the image list
|
||||
* They are wrapped in a button so that they can be clicked and open
|
||||
* the image in the lightbox.
|
||||
* @param {*} src
|
||||
*/
|
||||
const createImageBtn = (src) => {
|
||||
console.debug(`making image ${src.filename}`)
|
||||
const img = document.createElement('img')
|
||||
const but = document.createElement('button')
|
||||
|
||||
Object.assign(but.style, {
|
||||
height: '120px',
|
||||
width: '120px',
|
||||
border: 'none',
|
||||
padding: 0,
|
||||
margin: 0,
|
||||
})
|
||||
Object.assign(img.style, {
|
||||
width: '100%',
|
||||
height: '100%',
|
||||
objectFit: 'cover',
|
||||
})
|
||||
Object.assign(but.style, {
|
||||
height: '120px',
|
||||
width: '120px',
|
||||
border: 'none',
|
||||
padding: 0,
|
||||
margin: 0,
|
||||
})
|
||||
Object.assign(img.style, {
|
||||
width: '100%',
|
||||
height: '100%',
|
||||
objectFit: 'cover',
|
||||
})
|
||||
|
||||
img.src = `/view?filename=${encodeURIComponent(src.filename)}&type=${
|
||||
src.type
|
||||
}&subfolder=${encodeURIComponent(src.subfolder)}`
|
||||
img.src = `/view?filename=${encodeURIComponent(src.filename)}&type=${
|
||||
src.type
|
||||
}&subfolder=${encodeURIComponent(src.subfolder)}`
|
||||
|
||||
imageUrls.push(img.src)
|
||||
imageUrls.push(img.src)
|
||||
|
||||
console.debug(img.src)
|
||||
console.debug(img.src)
|
||||
|
||||
img.onload = () => {
|
||||
but.style.width = `${120 * (img.naturalWidth / img.naturalHeight)}px`
|
||||
}
|
||||
img.onload = () => {
|
||||
but.style.width = `${120 * (img.naturalWidth / img.naturalHeight)}px`
|
||||
}
|
||||
|
||||
but.onclick = () => {
|
||||
lightboxContainer.style.display = 'flex'
|
||||
// add the same image to the lightbox
|
||||
lightboxImage.src = img.src
|
||||
// lighboxContainer.replaceChildren(lightboxButtons, img);
|
||||
}
|
||||
but.onclick = () => {
|
||||
lightboxContainer.style.display = 'flex'
|
||||
// add the same image to the lightbox
|
||||
lightboxImage.src = img.src
|
||||
// lighboxContainer.replaceChildren(lightboxButtons, img);
|
||||
}
|
||||
|
||||
// add right click menu
|
||||
but.addEventListener('contextmenu', (e) => {
|
||||
e.preventDefault()
|
||||
// add right click menu
|
||||
but.addEventListener('contextmenu', (e) => {
|
||||
e.preventDefault()
|
||||
|
||||
if (image_menu) {
|
||||
image_menu.remove()
|
||||
}
|
||||
if (image_menu) {
|
||||
image_menu.remove()
|
||||
}
|
||||
|
||||
image_menu = document.createElement('div')
|
||||
Object.assign(image_menu.style, {
|
||||
position: 'absolute',
|
||||
top: `${e.clientY}px`,
|
||||
left: `${e.clientX}px`,
|
||||
background: '#333',
|
||||
color: '#fff',
|
||||
padding: '5px',
|
||||
borderRadius: '5px',
|
||||
zIndex: 999,
|
||||
})
|
||||
const load_img = document.createElement('button')
|
||||
load_img.textContent = 'Load'
|
||||
load_img.onclick = () => {
|
||||
app.handleFile(img.src)
|
||||
}
|
||||
image_menu = document.createElement('div')
|
||||
Object.assign(image_menu.style, {
|
||||
position: 'absolute',
|
||||
top: `${e.clientY}px`,
|
||||
left: `${e.clientX}px`,
|
||||
background: '#333',
|
||||
color: '#fff',
|
||||
padding: '5px',
|
||||
borderRadius: '5px',
|
||||
zIndex: 999,
|
||||
})
|
||||
const load_img = document.createElement('button')
|
||||
load_img.textContent = 'Load'
|
||||
load_img.onclick = () => {
|
||||
app.handleFile(img.src)
|
||||
}
|
||||
|
||||
image_menu.appendChild(load_img)
|
||||
document.body.appendChild(image_menu)
|
||||
})
|
||||
image_menu.appendChild(load_img)
|
||||
document.body.appendChild(image_menu)
|
||||
})
|
||||
|
||||
but.append(img)
|
||||
imageListContainer.prepend(but)
|
||||
}
|
||||
but.append(img)
|
||||
imageListContainer.prepend(but)
|
||||
}
|
||||
|
||||
loadButton.onclick = async () => {
|
||||
const all_history = await api.getHistory()
|
||||
for (const history of all_history.History) {
|
||||
if (history.outputs) {
|
||||
for (const key of Object.keys(history.outputs)) {
|
||||
console.debug(key)
|
||||
if (history.outputs[key].images) {
|
||||
for (const im of history.outputs[key].images) {
|
||||
console.debug(im)
|
||||
createImageBtn(im)
|
||||
}
|
||||
}
|
||||
}
|
||||
// for (const src of outputs.outputs.images) {
|
||||
// console.debug(src)
|
||||
// makeImage(`${src.subfolder}/${src.filename}`)
|
||||
// }
|
||||
}
|
||||
}
|
||||
}
|
||||
loadButton.onclick = async () => {
|
||||
const all_history = await api.getHistory()
|
||||
for (const history of all_history.History) {
|
||||
if (history.outputs) {
|
||||
for (const key of Object.keys(history.outputs)) {
|
||||
console.debug(key)
|
||||
if (history.outputs[key].images) {
|
||||
for (const im of history.outputs[key].images) {
|
||||
console.debug(im)
|
||||
createImageBtn(im)
|
||||
}
|
||||
}
|
||||
}
|
||||
// for (const src of outputs.outputs.images) {
|
||||
// console.debug(src)
|
||||
// makeImage(`${src.subfolder}/${src.filename}`)
|
||||
// }
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
///////-------
|
||||
///////-------
|
||||
|
||||
// const all_history = await api.getHistory()
|
||||
// for (const history of all_history.History) {
|
||||
// if (history.outputs) {
|
||||
// for (const key of Object.keys(history.outputs)) {
|
||||
// for (const im of history.outputs[key].images) {
|
||||
// makeImage(im)
|
||||
// }
|
||||
// }
|
||||
// // for (const src of outputs.outputs.images) {
|
||||
// // console.debug(src)
|
||||
// // makeImage(`${src.subfolder}/${src.filename}`)
|
||||
// // }
|
||||
// }
|
||||
// }
|
||||
// const all_history = await api.getHistory()
|
||||
// for (const history of all_history.History) {
|
||||
// if (history.outputs) {
|
||||
// for (const key of Object.keys(history.outputs)) {
|
||||
// for (const im of history.outputs[key].images) {
|
||||
// makeImage(im)
|
||||
// }
|
||||
// }
|
||||
// // for (const src of outputs.outputs.images) {
|
||||
// // console.debug(src)
|
||||
// // makeImage(`${src.subfolder}/${src.filename}`)
|
||||
// // }
|
||||
// }
|
||||
// }
|
||||
|
||||
//- Hook into the API
|
||||
api.addEventListener('executed', ({ detail }) => {
|
||||
if (detail?.output?.images) {
|
||||
for (const src of detail.output.images) {
|
||||
console.debug(`Adding ${src} to image feed`)
|
||||
createImageBtn(src)
|
||||
}
|
||||
}
|
||||
})
|
||||
},
|
||||
//- Hook into the API
|
||||
api.addEventListener('executed', ({ detail }) => {
|
||||
if (detail?.output?.images) {
|
||||
for (const src of detail.output.images) {
|
||||
console.debug(`Adding ${src} to image feed`)
|
||||
createImageBtn(src)
|
||||
}
|
||||
}
|
||||
})
|
||||
},
|
||||
})
|
||||
|
||||
+143
-43
@@ -3,6 +3,7 @@
|
||||
import { app } from '../../scripts/app.js'
|
||||
import { api } from '../../scripts/api.js'
|
||||
|
||||
import * as mtb_ui from './mtb_ui.js'
|
||||
import * as shared from './comfy_shared.js'
|
||||
|
||||
import {
|
||||
@@ -15,7 +16,13 @@ import {
|
||||
} from './mtb_ui.js'
|
||||
|
||||
const offset = 0
|
||||
|
||||
// These are "global" variables mostly meant to sync user settings.
|
||||
let currentWidth = 200
|
||||
let saltUrls =
|
||||
app.extensionManager.setting.get('mtb.io-sidebar.salt_urls') || false
|
||||
let targetWidth =
|
||||
app.extensionManager.setting.get('mtb.io-sidebar.img-size') || 512
|
||||
let currentMode = 'input'
|
||||
let subfolder = ''
|
||||
let currentSort = 'None'
|
||||
@@ -46,15 +53,19 @@ const updateImage = (node, image) => {
|
||||
* @param {ResultItem} resultItem
|
||||
* @returns {string} - The request URL.
|
||||
*/
|
||||
const resultItemToQuery = (resultItem) =>
|
||||
[
|
||||
const resultItemToQuery = (resultItem) => {
|
||||
const res = [
|
||||
`/mtb/view?filename=${resultItem.filename}`,
|
||||
`width=512`,
|
||||
`type=${resultItem.type}`,
|
||||
`subfolder=${resultItem.subfolder}`,
|
||||
`preview=`,
|
||||
].join('&')
|
||||
'preview=',
|
||||
]
|
||||
if (targetWidth > 0) {
|
||||
res.splice(1, 0, `width=${targetWidth}`)
|
||||
}
|
||||
|
||||
return res.join('&')
|
||||
}
|
||||
/**
|
||||
* Retrieves the unique prompt ID from a history task item.
|
||||
* @param {HistoryTaskItem} historyTaskItem
|
||||
@@ -80,7 +91,7 @@ const getNewOutputUrls = (mostRecentTask) => {
|
||||
const imageOutputs = Object.values(nodeOutputs.images)
|
||||
imageOutputs.forEach(
|
||||
(resultItem) =>
|
||||
(urls[resultItem.filename] = resultItemToQuery(resultItem))
|
||||
(urls[resultItem.filename] = resultItemToQuery(resultItem)),
|
||||
)
|
||||
}
|
||||
// Can process `animated` and `audio` outputs here.
|
||||
@@ -209,7 +220,7 @@ const getUrls = async (subfolder) => {
|
||||
if (currentMode === 'video') {
|
||||
const output = await shared.runAction(
|
||||
'getUserVideos',
|
||||
256,
|
||||
targetWidth,
|
||||
count,
|
||||
offset,
|
||||
currentSort,
|
||||
@@ -219,11 +230,13 @@ const getUrls = async (subfolder) => {
|
||||
const output = await shared.runAction(
|
||||
'getUserImages',
|
||||
currentMode,
|
||||
targetWidth,
|
||||
count,
|
||||
offset,
|
||||
currentSort,
|
||||
false,
|
||||
subfolder,
|
||||
saltUrls,
|
||||
)
|
||||
return output || {}
|
||||
}
|
||||
@@ -236,55 +249,110 @@ if (window?.__COMFYUI_FRONTEND_VERSION__) {
|
||||
|
||||
const sidebar_extension = {
|
||||
name: 'mtb.io-sidebar',
|
||||
// init: async () => {
|
||||
// try {
|
||||
// const res = await api.fetchApi('/mtb/server-info')
|
||||
// const msg = await res.json()
|
||||
// exposed = msg.exposed
|
||||
// } catch (e) {
|
||||
// console.error('Error:', e)
|
||||
// }
|
||||
// },
|
||||
init: () => {
|
||||
let handle
|
||||
const version = window?.__COMFYUI_FRONTEND_VERSION__
|
||||
console.log(`%c ${version}`, 'background: orange; color: white;')
|
||||
|
||||
ensureMTBStyles()
|
||||
|
||||
app.ui.settings.addSetting({
|
||||
settings: [
|
||||
{
|
||||
id: 'mtb.io-sidebar.count',
|
||||
category: ['mtb', 'Input & Output Sidebar', 'count'],
|
||||
|
||||
name: 'Number of images to fetch',
|
||||
type: 'number',
|
||||
defaultValue: 1000,
|
||||
|
||||
tooltip:
|
||||
"This setting affects the input/output sidebar to determine how many images to fetch per pagination (pagination is not yet supported so for now it's the static total)",
|
||||
attrs: {
|
||||
style: {
|
||||
// fontFamily: 'monospace',
|
||||
},
|
||||
},
|
||||
{
|
||||
id: 'mtb.io-sidebar.salt_urls',
|
||||
category: ['mtb', 'Input & Output Sidebar', 'salt_urls'],
|
||||
name: 'Salt URLs',
|
||||
type: 'boolean',
|
||||
defaultValue: false,
|
||||
onChange: (n, o) => {
|
||||
saltUrls = n
|
||||
},
|
||||
})
|
||||
|
||||
app.ui.settings.addSetting({
|
||||
tooltip:
|
||||
'Adds a random query parameter to every urls to always invalidate caching.',
|
||||
},
|
||||
{
|
||||
id: 'mtb.io-sidebar.img-size',
|
||||
category: ['mtb', 'Input & Output Sidebar', 'img-size'],
|
||||
|
||||
name: 'Resolution of the images',
|
||||
type: 'number',
|
||||
name: 'Resize width of shown images',
|
||||
defaultValue: 512,
|
||||
type: (name, setter, value, attrs) => {
|
||||
targetWidth = value
|
||||
const container = mtb_ui.makeElement('div', {
|
||||
display: 'flex',
|
||||
alignItems: 'center',
|
||||
gap: '8px',
|
||||
})
|
||||
|
||||
tooltip: "It's recommended to keep it at 512px",
|
||||
attrs: {
|
||||
style: {
|
||||
// fontFamily: 'monospace',
|
||||
},
|
||||
console.log({ name, setter, value, attrs })
|
||||
|
||||
const baseId = name.replace(/[^a-zA-Z0-9]/g, '-').toLowerCase()
|
||||
const checkboxId = `${baseId}-checkbox`
|
||||
const numberInputId = `${baseId}-number`
|
||||
|
||||
const isCheckedInitially = value !== -1
|
||||
|
||||
// TODO: better way to get defaultValue?
|
||||
const defaultValue = 512
|
||||
const initialNumberValue = isCheckedInitially ? value : defaultValue
|
||||
|
||||
console.log('recreate')
|
||||
const checkbox = mtb_ui.makeElement(
|
||||
// harder to match styles (.p-toggleswitch-input)
|
||||
// since it uses a div synced to the input...
|
||||
'input',
|
||||
{},
|
||||
container,
|
||||
)
|
||||
checkbox.type = 'checkbox'
|
||||
checkbox.id = checkboxId
|
||||
checkbox.checked = isCheckedInitially
|
||||
|
||||
const numberInput = mtb_ui.makeElement(
|
||||
'input.p-inputtext',
|
||||
{},
|
||||
container,
|
||||
)
|
||||
numberInput.type = 'number'
|
||||
numberInput.id = numberInputId
|
||||
numberInput.value = initialNumberValue
|
||||
numberInput.disabled = !isCheckedInitially
|
||||
numberInput.min = 128
|
||||
|
||||
checkbox.addEventListener('change', () => {
|
||||
let valToSet = -1
|
||||
if (checkbox.checked) {
|
||||
numberInput.disabled = false
|
||||
|
||||
valToSet = Number.parseInt(numberInput.value, 10)
|
||||
if (Number.isNaN(valToSet) || valToSet < numberInput.min) {
|
||||
valToSet = defaultValue
|
||||
numberInput.value = valToSet
|
||||
}
|
||||
} else {
|
||||
numberInput.disabled = true
|
||||
}
|
||||
setter(valToSet)
|
||||
})
|
||||
|
||||
numberInput.addEventListener('input', () => {
|
||||
if (checkbox.checked) {
|
||||
const numValue = Number.parseInt(numberInput.value, 10)
|
||||
if (!Number.isNaN(numValue) && numberInput.value !== '') {
|
||||
setter(numValue)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
return container
|
||||
},
|
||||
})
|
||||
app.ui.settings.addSetting({
|
||||
|
||||
tooltip:
|
||||
"If browsing large folders it's recommended to use this to avoid overflow/crash of the webpage. Image will get resized to this target width on the server before being sent to the client.",
|
||||
},
|
||||
|
||||
{
|
||||
id: 'mtb.io-sidebar.sort',
|
||||
category: ['mtb', 'Input & Output Sidebar', 'sort'],
|
||||
name: 'Default sort mode',
|
||||
@@ -304,7 +372,39 @@ if (window?.__COMFYUI_FRONTEND_VERSION__) {
|
||||
'Name',
|
||||
'Name-Reverse',
|
||||
],
|
||||
})
|
||||
},
|
||||
{
|
||||
id: 'mtb.io-sidebar.notice',
|
||||
category: ['mtb', 'Input & Output Sidebar', 'sort'],
|
||||
name: ' ',
|
||||
|
||||
type: (name, setter, value, attrs) => {
|
||||
const container = mtb_ui.makeElement('div')
|
||||
const notice =
|
||||
'## Important\nIf you make **any** edits here you need to toggle off and back on the sidebar for it to take effect.'
|
||||
|
||||
if (window.MTB?.mdParser) {
|
||||
MTB.mdParser.parse(notice).then((e) => {
|
||||
container.innerHTML = e
|
||||
})
|
||||
} else {
|
||||
shared.ensureMarkdownParser((p) => {
|
||||
p.parse(notice).then((e) => {
|
||||
container.innerHTML = e
|
||||
})
|
||||
})
|
||||
}
|
||||
return container
|
||||
},
|
||||
},
|
||||
],
|
||||
|
||||
init: () => {
|
||||
let handle
|
||||
const version = window?.__COMFYUI_FRONTEND_VERSION__
|
||||
console.log(`%c ${version}`, 'background: orange; color: white;')
|
||||
|
||||
ensureMTBStyles()
|
||||
|
||||
app.extensionManager.registerSidebarTab({
|
||||
id: 'mtb-inputs-outputs',
|
||||
|
||||
+529
@@ -0,0 +1,529 @@
|
||||
/** 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)
|
||||
@@ -1,28 +0,0 @@
|
||||
// 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 })
|
||||
// },
|
||||
// })
|
||||
// }
|
||||
+4
-1
@@ -203,7 +203,7 @@ export const wrapElement = (element, style = {}) => {
|
||||
* @param {Object} [style] - CSS styles to apply to the element.
|
||||
* @returns {HTMLElement} - The created DOM element.
|
||||
*/
|
||||
export const makeElement = (kind, style) => {
|
||||
export const makeElement = (kind, style, parent) => {
|
||||
let [real_kind, className] = kind.split('.')
|
||||
let id
|
||||
|
||||
@@ -224,6 +224,9 @@ export const makeElement = (kind, style) => {
|
||||
if (id) {
|
||||
el.id = id
|
||||
}
|
||||
if (parent) {
|
||||
parent.appendChild(el)
|
||||
}
|
||||
|
||||
return el
|
||||
}
|
||||
|
||||
+84
-45
@@ -694,7 +694,7 @@ const mtb_widgets = {
|
||||
|
||||
app.ui.settings.addSetting({
|
||||
id: 'mtb.Main.debug-enabled',
|
||||
category: ['mtb', 'Main', 'debug-enabled'],
|
||||
category: ['mtb', ' Main', 'debug-enabled'],
|
||||
name: 'Enable Debug (py and js)',
|
||||
type: 'boolean',
|
||||
defaultValue: false,
|
||||
@@ -1012,12 +1012,15 @@ const mtb_widgets = {
|
||||
)
|
||||
loop_preview.value = 'Iteration: Idle'
|
||||
|
||||
let cancelQueue = false
|
||||
|
||||
const onReset = () => {
|
||||
raw_iteration.value = 0
|
||||
raw_loop.value = 0
|
||||
|
||||
value_preview.value = 'Idle'
|
||||
loop_preview.value = 'Iteration: Idle'
|
||||
cancelQueue = false
|
||||
|
||||
app.canvas.setDirty(true)
|
||||
}
|
||||
@@ -1026,15 +1029,42 @@ const mtb_widgets = {
|
||||
this.addWidget('button', 'Reset', 'reset', onReset)
|
||||
|
||||
// run button
|
||||
this.addWidget('button', 'Queue', 'queue', () => {
|
||||
onReset() // this could maybe be a setting or checkbox
|
||||
app.queuePrompt(0, total_frames.value * loop_count.value)
|
||||
const chunkSize = 10
|
||||
this.addWidget('button', 'Queue', 'queue', async () => {
|
||||
onReset()
|
||||
|
||||
const totalPrompts = total_frames.value * loop_count.value
|
||||
window.MTB?.notify?.(
|
||||
`Started a queue of ${total_frames.value} frames (for ${
|
||||
loop_count.value
|
||||
} loop, so ${total_frames.value * loop_count.value})`,
|
||||
`Starting a queue of ${totalPrompts} frames in chunks of ${chunkSize}...`,
|
||||
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 = () => {
|
||||
@@ -1166,7 +1196,9 @@ const mtb_widgets = {
|
||||
|
||||
//NOTE: dynamic nodes
|
||||
case 'Apply Text Template (mtb)': {
|
||||
shared.setupDynamicConnections(nodeType, 'var', '*')
|
||||
shared.setupDynamicConnections(nodeType, 'var', '*', {
|
||||
rename_menu: 'name',
|
||||
})
|
||||
break
|
||||
}
|
||||
case 'Save Data Bundle (mtb)': {
|
||||
@@ -1298,47 +1330,54 @@ const mtb_widgets = {
|
||||
const related = new Set([this.id])
|
||||
const visited = new Set()
|
||||
if (this.outputs[0].links) {
|
||||
const initLink = this.outputs[0].links[0]
|
||||
const { to: loopEnd } = shared.nodesFromLink(this, initLink)
|
||||
const canReachEnd = (node, visited = new Set()) => {
|
||||
if (node === loopEnd) return true
|
||||
if (visited.has(node.id)) return false
|
||||
visited.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
|
||||
if (canReachEnd(nextNode, visited)) {
|
||||
return true
|
||||
for (const linkId of this.outputs[0].links) {
|
||||
const { to: loopEnd } = shared.nodesFromLink(this, linkId)
|
||||
const canReachEnd = (node, visited = new Set()) => {
|
||||
if (node === loopEnd) return true
|
||||
if (visited.has(node.id)) return false
|
||||
visited.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
|
||||
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.computed_flow = true
|
||||
|
||||
+64
-52
@@ -1,10 +1,13 @@
|
||||
// web/note_plus.constants.js
|
||||
|
||||
export const DEFAULT_CSS = ''
|
||||
export const DEFAULT_CSS = `/** here you can write css**/
|
||||
h1 {
|
||||
color: whitesmoke;
|
||||
}`
|
||||
export const DEFAULT_HTML = `<p style='color:red;font-family:monospace'>
|
||||
Note+
|
||||
</p>`
|
||||
export const DEFAULT_MD = '## Note+'
|
||||
export const DEFAULT_MD = '# 📝 Note+'
|
||||
export const DEFAULT_MODE = 'markdown'
|
||||
export const DEFAULT_THEME = 'one_dark'
|
||||
|
||||
@@ -55,58 +58,57 @@ We also support github callout:
|
||||
`
|
||||
|
||||
export const THEMES = [
|
||||
'ambiance',
|
||||
'chaos',
|
||||
'chrome',
|
||||
'cloud9_day',
|
||||
'cloud9_night',
|
||||
'cloud9_night_low_color',
|
||||
'cloud_editor',
|
||||
'cloud_editor_dark',
|
||||
'clouds',
|
||||
'clouds_midnight',
|
||||
'cobalt',
|
||||
'crimson_editor',
|
||||
'dawn',
|
||||
'dracula',
|
||||
'dreamweaver',
|
||||
'eclipse',
|
||||
'github',
|
||||
'github_dark',
|
||||
'gob',
|
||||
'gruvbox',
|
||||
'gruvbox_dark_hard',
|
||||
'gruvbox_light_hard',
|
||||
'idle_fingers',
|
||||
'iplastic',
|
||||
'katzenmilch',
|
||||
'kr_theme',
|
||||
'kuroir',
|
||||
'merbivore',
|
||||
'merbivore_soft',
|
||||
'mono_industrial',
|
||||
'monokai',
|
||||
'nord_dark',
|
||||
'one_dark',
|
||||
'pastel_on_dark',
|
||||
'solarized_dark',
|
||||
'solarized_light',
|
||||
'sqlserver',
|
||||
'terminal',
|
||||
'textmate',
|
||||
'tomorrow',
|
||||
'tomorrow_night',
|
||||
'tomorrow_night_blue',
|
||||
'tomorrow_night_bright',
|
||||
'tomorrow_night_eighties',
|
||||
'twilight',
|
||||
'vibrant_ink',
|
||||
'vscode',
|
||||
'ambiance',
|
||||
'chaos',
|
||||
'chrome',
|
||||
'cloud9_day',
|
||||
'cloud9_night',
|
||||
'cloud9_night_low_color',
|
||||
'cloud_editor',
|
||||
'cloud_editor_dark',
|
||||
'clouds',
|
||||
'clouds_midnight',
|
||||
'cobalt',
|
||||
'crimson_editor',
|
||||
'dawn',
|
||||
'dracula',
|
||||
'dreamweaver',
|
||||
'eclipse',
|
||||
'github',
|
||||
'github_dark',
|
||||
'gob',
|
||||
'gruvbox',
|
||||
'gruvbox_dark_hard',
|
||||
'gruvbox_light_hard',
|
||||
'idle_fingers',
|
||||
'iplastic',
|
||||
'katzenmilch',
|
||||
'kr_theme',
|
||||
'kuroir',
|
||||
'merbivore',
|
||||
'merbivore_soft',
|
||||
'mono_industrial',
|
||||
'monokai',
|
||||
'nord_dark',
|
||||
'one_dark',
|
||||
'pastel_on_dark',
|
||||
'solarized_dark',
|
||||
'solarized_light',
|
||||
'sqlserver',
|
||||
'terminal',
|
||||
'textmate',
|
||||
'tomorrow',
|
||||
'tomorrow_night',
|
||||
'tomorrow_night_blue',
|
||||
'tomorrow_night_bright',
|
||||
'tomorrow_night_eighties',
|
||||
'twilight',
|
||||
'vibrant_ink',
|
||||
'vscode',
|
||||
]
|
||||
|
||||
export const CSS_RESET = `
|
||||
* {
|
||||
font-family: monospace;
|
||||
line-height: 1.25em;
|
||||
}
|
||||
.shiki{
|
||||
@@ -116,6 +118,8 @@ export const CSS_RESET = `
|
||||
.markdown-callout-title {
|
||||
.octicon{
|
||||
fill:white;
|
||||
width:29px;
|
||||
height:29px;
|
||||
}
|
||||
/* background: var(--current-color); */
|
||||
color: var(--current-color);
|
||||
@@ -124,6 +128,8 @@ export const CSS_RESET = `
|
||||
/* border-start-start-radius: var(--radius); */
|
||||
padding: 0.5em;
|
||||
padding-inline-start: 1em;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
}
|
||||
.markdown-callout-content {
|
||||
padding: 1em;
|
||||
@@ -136,7 +142,12 @@ export const CSS_RESET = `
|
||||
border-left: 3px solid var(--current-color);
|
||||
margin-bottom: 1em;
|
||||
margin-top: 1em;
|
||||
|
||||
}
|
||||
.markdown-callout p:nth-child(2) {
|
||||
padding:1em;
|
||||
}
|
||||
|
||||
|
||||
.markdown-callout-tip {
|
||||
--text-color: whitesmoke;
|
||||
@@ -164,8 +175,9 @@ export const CSS_RESET = `
|
||||
flex-direction:column;
|
||||
align-items: flex-start;
|
||||
width:95%;
|
||||
margin-left: 20px;
|
||||
margin-top:20px;
|
||||
/*margin-left: 20px;*/
|
||||
/*margin-top:20px;*/
|
||||
|
||||
/*background-color: rgba(255,0,0,0.5)!important;*/
|
||||
}
|
||||
|
||||
|
||||
+379
-336
File diff suppressed because it is too large
Load Diff
+12
-3
@@ -41,7 +41,16 @@ const toastStyle = `
|
||||
transition-duration: ${transition_time}ms;
|
||||
`
|
||||
|
||||
function notify(message, timeout = 3000) {
|
||||
function notify(message, timeout = 3000, old_mode = false) {
|
||||
if (!old_mode) {
|
||||
app.extensionManager.toast.add({
|
||||
severity: 'info',
|
||||
summary: 'MTB',
|
||||
detail: message,
|
||||
life: timeout,
|
||||
})
|
||||
return
|
||||
}
|
||||
log('Creating toast')
|
||||
const container = document.getElementById('mtb-notify-container')
|
||||
const toast = document.createElement('div')
|
||||
@@ -59,7 +68,7 @@ function notify(message, timeout = 3000) {
|
||||
log('Transition out')
|
||||
const totalHeight = Array.from(container.children).reduce(
|
||||
(acc, child) => acc + child.offsetHeight + 10, // Add spacing of 10px between toasts
|
||||
0
|
||||
0,
|
||||
)
|
||||
container.style.height = `${totalHeight}px`
|
||||
|
||||
@@ -83,7 +92,7 @@ function notify(message, timeout = 3000) {
|
||||
// Update container's height to fit new toast
|
||||
const totalHeight = Array.from(container.children).reduce(
|
||||
(acc, child) => acc + child.offsetHeight + 10, // Add spacing of 10px between toasts
|
||||
0
|
||||
0,
|
||||
)
|
||||
container.style.height = `${totalHeight}px`
|
||||
|
||||
|
||||
Reference in New Issue
Block a user