Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
821a031bfc | ||
|
|
3d12bd29a8 | ||
|
|
3494a4767e | ||
|
|
41d444ae70 | ||
|
|
41d79d4677 | ||
|
|
54f4963a7b | ||
|
|
d585b16ee7 | ||
|
|
499cd218aa | ||
|
|
fd00e12724 | ||
|
|
e37e4648e5 | ||
|
|
06685a5418 | ||
|
|
40145ddf40 | ||
|
|
c1d74c0d69 | ||
|
|
eca2ea5da9 | ||
|
|
c6d1f73cfc | ||
|
|
4ea7c0b67f | ||
|
|
8316914f02 |
+5
-1
@@ -7,7 +7,7 @@
|
||||
#
|
||||
###
|
||||
|
||||
__version__ = "0.5.4"
|
||||
__version__ = "0.6.0"
|
||||
|
||||
import os
|
||||
|
||||
@@ -265,6 +265,10 @@ def register_routes():
|
||||
|
||||
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
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
+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]
|
||||
+19
-5
@@ -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,7 +348,7 @@ 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.strip()
|
||||
if wrap:
|
||||
@@ -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__ = [
|
||||
|
||||
+117
-7
@@ -6,7 +6,7 @@ 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
|
||||
@@ -16,8 +16,10 @@ 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,
|
||||
@@ -134,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:
|
||||
@@ -343,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),)
|
||||
@@ -916,6 +970,61 @@ class MTB_BooleanNot:
|
||||
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,
|
||||
@@ -933,4 +1042,5 @@ __nodes__ = [
|
||||
MTB_TensorOps,
|
||||
MTB_BooleanNot,
|
||||
MTB_GetItem,
|
||||
MTB_ProxyTensor,
|
||||
]
|
||||
|
||||
+5
-38
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "comfy-mtb"
|
||||
version = "0.5.4"
|
||||
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.5.1"
|
||||
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.
|
||||
@@ -8,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
|
||||
|
||||
@@ -461,8 +464,6 @@ def _run_command(shell_cmd, ignored_lines_start):
|
||||
print("Command executed successfully!")
|
||||
|
||||
|
||||
|
||||
|
||||
# endregion
|
||||
|
||||
|
||||
@@ -526,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)
|
||||
@@ -536,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)
|
||||
@@ -558,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
|
||||
|
||||
+201
-15
@@ -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()
|
||||
@@ -290,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)
|
||||
|
||||
@@ -380,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 = (
|
||||
@@ -399,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,
|
||||
)
|
||||
@@ -484,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
|
||||
}
|
||||
|
||||
@@ -497,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
|
||||
@@ -506,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) {
|
||||
@@ -520,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
|
||||
|
||||
@@ -536,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
|
||||
@@ -590,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,
|
||||
)
|
||||
|
||||
@@ -602,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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+67
-86
@@ -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) {
|
||||
@@ -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',
|
||||
})
|
||||
|
||||
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',
|
||||
})
|
||||
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',
|
||||
})
|
||||
img.src = content
|
||||
img.src = item.data
|
||||
wrapper.appendChild(img)
|
||||
}
|
||||
|
||||
@@ -80,110 +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
|
||||
|
||||
const name_to_label = this.inputs.reduce((acc, input) => {
|
||||
acc[input.name] = input.label || input.name
|
||||
return acc
|
||||
}, {})
|
||||
|
||||
|
||||
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)
|
||||
}
|
||||
inputData[inputName] = item.items
|
||||
})
|
||||
}
|
||||
|
||||
let widgetI = 1
|
||||
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 () {
|
||||
@@ -194,8 +174,9 @@ app.registerExtension({
|
||||
widget.onRemoved?.()
|
||||
widget.onRemove?.()
|
||||
}
|
||||
shared.cleanupNode(this)
|
||||
cleanupNode(this)
|
||||
}
|
||||
this.setDirtyCanvas(true, true)
|
||||
}
|
||||
}
|
||||
},
|
||||
|
||||
+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)
|
||||
+47
-39
@@ -1051,7 +1051,6 @@ const mtb_widgets = {
|
||||
const currentChunkSize = Math.min(chunkSize, totalPrompts - i)
|
||||
|
||||
await app.queuePrompt(0, currentChunkSize)
|
||||
|
||||
}
|
||||
if (!cancelQueue) {
|
||||
window.MTB?.notify?.(
|
||||
@@ -1197,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)': {
|
||||
@@ -1329,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
|
||||
|
||||
Reference in New Issue
Block a user