Compare commits

..
9 Commits
Author SHA1 Message Date
drbaph b35b5d8a17 fix: whisper transcription compatibility with newer transformers (#274)
- Use getattr for max_length to handle removed WhisperConfig attribute
- Cast input_features to model dtype to fix float16 mismatch
2026-07-04 21:40:35 +02:00
carlostsai 6d5fd74333 fix: typo on vitmatte torch script name (#276) 2026-06-27 21:15:11 +02:00
Anderson Yan b705a177d3 fix: add retry to LoadImageFromURL 2026-03-19 08:39:38 +01:00
Mel Massadian 00fbad37c5 docs: remove deprecation
Updated caution and note sections regarding recent changes and versioning.
2026-01-10 10:32:45 +01:00
Benjamin Gregg 6cbe294c1b Fix Deepcopy Error
Fix Deepcopy Error in new comfy versions
2026-01-10 10:30:32 +01:00
Mel Massadian eabe43db79 fix: 🐛 add missing widgetTypes for COLOR 2025-09-07 11:54:27 +00:00
Austin Mroz 1c99a1c63c Set widgetType for COLOR widgets 2025-09-06 16:56:35 +02:00
Mel Massadian 426cdf5f9f fix: 🐛 temporary fix for COLOR 2025-09-06 12:38:41 +00:00
Mel Massadian 5fa3791559 📚 docs: add caution about project status 2025-09-06 11:03:42 +02:00
24 changed files with 581 additions and 2733 deletions
+6
View File
@@ -1,4 +1,10 @@
# MTB Nodes
> [!NOTE]
> master/main is outdated for now to keep backward compatibility, the next version is being worked on in
> [`dev/0.6.0`](https://github.com/melMass/comfy_mtb/tree/dev/0.6.0)
[![embedded test](https://github.com/melMass/comfy_mtb/actions/workflows/test_embedded.yml/badge.svg)](https://github.com/melMass/comfy_mtb/actions/workflows/test_embedded.yml)
![home](https://repository-images.githubusercontent.com/649047066/a3eef9a7-20dd-4ef9-b839-884502d4e873)
+1 -5
View File
@@ -7,7 +7,7 @@
#
###
__version__ = "0.6.0"
__version__ = "0.5.4"
import os
@@ -265,10 +265,6 @@ 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
+104 -211
View File
@@ -1,175 +1,85 @@
# NOTE: This file is only use for development you can ignore it
use log.nu
use nssm.nu *
use nutils.nu [ make-id upsert-all fwd-slash backup-file ]
use os.nu [ link ]
use private/log.nu
# --- utilities ---
def get_root [ --clean] {
if $clean {
$env.COMFY.ROOTS.clean
} else {
$env.COMFY.ROOTS.main
}
def get_root [--clean] {
if $clean {
$env.COMFY_CLEAN_ROOT
} else {
$env.COMFY_ROOT
}
}
def --env path-add [pth] {
$env.PATH = ($env.PATH | append ($pth | path expand))
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 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.ROOTS.mtb | path join daily.nuon)
let daily = ($env.COMFY_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_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}]
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}]
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 # comfy clean instance
--rebase # Rebase instead of merge
--clean # ??
--rebase # Rebase instead of merge
] {
let root = get_root --clean=$clean
@@ -184,28 +94,21 @@ export def "comfy update" [
log info "Backing up and removing models symlinks"
# preparing root for pull
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"
if not $clean {
git checkout pyproject.toml
cd $models
# find and store all symlinks
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)"
let links = (ls -la |
where not ($it.target | is-empty) |
select name target |
sort-by name)
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
@@ -238,6 +141,7 @@ export def "comfy update" [
if $rebase {
log info "Rebasing changes"
git rebase master
} else {
log info "Merging changes"
git merge master
@@ -248,10 +152,9 @@ export def "comfy update" [
if not $clean {
rm pyproject.toml
log info "Using our own pyproject..."
cp $pyproject pyproject.toml
cp pyproject-mel.toml pyproject.toml
cd $models
log info "Relinking models..."
# resymlink them
open links.nuon | each {|p| link -a $p.target $p.name }
} else {
@@ -264,82 +167,72 @@ export def "comfy update" [
log success $"Update successful \(($commit_count) new commits\)"
return {from_commit: $current_commit to_commit: $new_commit}
return {from:$current_commit to:$new_commit}
}
export def "comfy toggle_extensions" [
--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
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
}
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
}
# 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)"
def --env path-add [pth] {
$env.PATH = ($env.PATH | append ($pth | path expand))
}
# -- env
export-env {
$env.PYTHONUTF8 = 1
$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.COMFY_MTB = ("." | path expand)
# $env.CUDA_ROOT = 'C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.1\'
$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"
$env.COMFY_ROOT = ("../.." | path expand)
$env.COMFY_CLEAN_ROOT = ($env.COMFY_ROOT | path dirname | path join ComfyClean)
path-add 'C:/Portable/TensorRT-8.6.0.12/lib'
if $nu.os-info.family == 'windows' {
path-add 'G:\BIN\TensorRT-10.7.0.23\lib'
path-add 'G:\BIN\cudnn-windows-x86_64-9.6.0.74_cuda12-archive\bin'
}
path-add ($env.CUDA_ROOT | path join bin)
overlay use ../../.venv/Scripts/activate.nu
}
-3
View File
@@ -1,3 +0,0 @@
{
"use_repl": false
}
+2 -2
View File
@@ -277,14 +277,14 @@ class MTB_AudioToText(MtbAudio):
f"Processing chunk {chunk_offset:.1f}s - {chunk_end / sample_rate:.1f}s"
)
max_length = model.config.max_length or 448
max_length = getattr(model.config, "max_length", None) or 448
attention_mask = torch.ones((1, max_length))
input_features = processor(
chunk_waveform,
sampling_rate=sample_rate,
return_tensors="pt",
).input_features.to(device)
).input_features.to(device=device, dtype=model.dtype)
with torch.no_grad():
predicted_ids = model.generate(
+4 -4
View File
@@ -335,9 +335,9 @@ class MTB_BatchShape:
"image_width": ("INT", {"default": 512}),
"image_height": ("INT", {"default": 512}),
"shape_size": ("INT", {"default": 100}),
"color": ("COLOR", {"default": "#ffffff"}),
"bg_color": ("COLOR", {"default": "#000000"}),
"shade_color": ("COLOR", {"default": "#000000"}),
"color": ("COLOR", {"default": "#ffffff","widgetType": "MTB_COLOR"}),
"bg_color": ("COLOR", {"default": "#000000","widgetType": "MTB_COLOR"}),
"shade_color": ("COLOR", {"default": "#000000","widgetType": "MTB_COLOR"}),
"thickness": ("INT", {"default": 5}),
"shadex": ("FLOAT", {"default": 0.0}),
"shadey": ("FLOAT", {"default": 0.0}),
@@ -842,7 +842,7 @@ class MTB_Batch2dTransform:
["edge", "constant", "reflect", "symmetric"],
{"default": "edge"},
),
"constant_color": ("COLOR", {"default": "#000000"}),
"constant_color": ("COLOR", {"default": "#000000","widgetType": "MTB_COLOR"}),
},
"optional": {
"x": ("FLOATS",),
+179 -605
View File
@@ -1,70 +1,33 @@
import base64
import io
import textwrap
from collections.abc import Callable
from functools import wraps
from typing import Any, Literal, Protocol, TypedDict, runtime_checkable
import json
from pathlib import Path
import folder_paths
import torch
from rich import inspect
from rich.console import Console
from ..log import log
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
from ..utils import tensor2pil
# 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] = []
def get_detailed_type_info(obj):
type_info = []
type_name = type(obj).__name__
type_info.append(f"Type: {type_name}")
if isinstance(obj, torch.Tensor):
return get_torch_tensor_info(obj)
elif isinstance(obj, list | tuple):
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)):
type_info.extend(
[
f"Length: {len(obj)}",
@@ -84,184 +47,122 @@ def _get_detailed_type_info(obj) -> str:
attributes = [attr for attr in dir(obj) if not attr.startswith("_")]
type_info.append(f"Attributes: {attributes}")
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:
return type_info
# region processors
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)
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")
)
from rich.console import Console
return {"b64_images": b64_imgs}
console = Console(record=True)
if isinstance(formatted, list):
for line in formatted:
console.print(line)
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)})"
)
else:
console.print(formatted)
text.append(f"Array ({len(anything)}): {anything}")
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>
return {"text": text}
@font-face {{
font-family: "Fira Code";
src: local("FiraCode-Regular"),
url("https://cdnjs.cloudflare.com/ajax/libs/firacode/6.2.0/woff2/FiraCode-Regular.woff2") format("woff2"),
url("https://cdnjs.cloudflare.com/ajax/libs/firacode/6.2.0/woff/FiraCode-Regular.woff") format("woff");
font-style: normal;
font-weight: 400;
}}
@font-face {{
font-family: "Fira Code";
src: local("FiraCode-Bold"),
url("https://cdnjs.cloudflare.com/ajax/libs/firacode/6.2.0/woff2/FiraCode-Bold.woff2") format("woff2"),
url("https://cdnjs.cloudflare.com/ajax/libs/firacode/6.2.0/woff/FiraCode-Bold.woff") format("woff");
font-style: bold;
font-weight: 700;
}}
.{unique_id}-matrix {{
font-family: Fira Code, monospace;
font-size: {char_height}px;
line-height: {line_height}px;
font-variant-east-asian: full-width;
}}
def process_dict(anything, as_type=False):
text = []
if as_type:
return {"text": get_detailed_type_info(anything)}
.{unique_id}-title {{
font-size: 18px;
font-weight: bold;
font-family: arial;
}}
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}")
{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}", ""),
elif "waveform" in anything:
is_empty = (
"(empty) " if torch.count_nonzero(anything["samples"]) == 0 else ""
)
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,
text.append(
f"Audio Samples: {anything['waveform'].shape}{is_empty} | sample rate {anything['sample_rate']}"
)
log.error(f"Unknown rich mode: {rich_mode}")
return formatted if isinstance(formatted, str) else "\n".join(formatted)
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)]}
# endregion
# region conditions
# those are pretty dumb there is now probably a better way..
def is_condition(item):
return (
isinstance(item, list)
and all(isinstance(i, list) for i in item)
and isinstance(item[0][0], torch.Tensor)
)
# endregion
RICH_MODE = Literal["none", "html", "svg", "svg-window"]
@runtime_checkable
class Processor(Protocol):
"""Generic protocol for processor functions."""
def __call__(
self, item: Any, *, as_type: bool = False, deep: bool = False
) -> ProcessorResult: ...
class MTB_Debug:
"""A debug node."""
"""Experimental node to debug any Comfy values.
support for more types and widgets is planned.
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {"output_to_console": ("BOOLEAN", {"default": False})},
"optional": {
"as_detailed_types": ("BOOLEAN", {"default": False}),
"deep_inspect": ("BOOLEAN", {"default": False}),
"rich_mode": (
("none", "html", "svg", "svg-window"),
{"default": "none"},
),
},
"optional": {"as_detailed_types": ("BOOLEAN", {"default": False})},
}
RETURN_TYPES = ()
@@ -269,426 +170,99 @@ 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,
**kwargs,
self, output_to_console: bool, as_detailed_types: bool, **kwargs
):
output = {"ui": {"items": []}}
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")
if output_to_console:
for k, v in kwargs.items():
log.info(f"{k}: {v}")
for input_name, item in kwargs.items():
processed = self._dispatch_processor(
item, as_type=as_type, deep=deep
)
if processed is None:
continue
for input_name, anything in kwargs.items():
processor = processors.get(type(anything), process_text)
if rich_mode != "none":
title = f"{input_name} ({type(item).__name__})"
processed = _apply_rich_results(processed, rich_mode, title)
processed = processor(anything, as_detailed_types)
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)")
item = {
"input": input_name,
**processed,
}
output["ui"]["items"].append(item)
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,
)
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)
class MTB_SaveTensors:
"""Save torch tensors (image, mask or latent) to disk.
text_output = console.export_text(clear=True)
useful to debug things outside comfy.
"""
return [UIResult(kind="text", data=text_output.strip())]
def __init__(self):
self.output_dir = folder_paths.get_output_directory()
self.type = "mtb/debug"
def _process_repr(
self, item: Any, as_type=False, deep=False
) -> ProcessorResult:
return [{"kind": "text", "data": item.__repr__()}]
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"filename_prefix": ("STRING", {"default": "ComfyPickle"}),
},
"optional": {
"image": ("IMAGE",),
"mask": ("MASK",),
"latent": ("LATENT",),
},
}
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)
FUNCTION = "save"
OUTPUT_NODE = True
RETURN_TYPES = ()
CATEGORY = "mtb/debug"
return [UIResult(kind="text", data=str(item))]
def save(
self,
filename_prefix,
image: torch.Tensor | None = None,
mask: torch.Tensor | None = None,
latent: torch.Tensor | None = None,
):
(
full_output_folder,
filename,
counter,
subfolder,
filename_prefix,
) = folder_paths.get_save_image_path(filename_prefix, self.output_dir)
full_output_folder = Path(full_output_folder)
if image is not None:
image_file = f"{filename}_image_{counter:05}.pt"
torch.save(image, full_output_folder / image_file)
# np.save(full_output_folder/ image_file, image.cpu().numpy())
def _process_bool(
self, item: bool, *, as_type=False, deep=False
) -> ProcessorResult: # noqa: FBT001
return [{"kind": "text", "data": "True" if item else "False"}]
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_clip(
self, item: Any, *, as_type=False, deep=False
) -> ProcessorResult:
try:
clip_model = getattr(item, "cond_stage_model", None)
tokenizer = getattr(item, "tokenizer", None)
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"))
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",
)
)
# np.save(full_output_folder / latent_file,
# latent[""].cpu().numpy())
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)
return f"{filename_prefix}_{counter:05}"
__nodes__ = [MTB_Debug]
processors = {
torch.Tensor: process_tensor,
list: process_list,
dict: process_dict,
bool: process_bool,
}
__nodes__ = [MTB_Debug, MTB_SaveTensors]
-70
View File
@@ -1,70 +0,0 @@
import folder_paths
import torch
class MTB_SaveTensors:
"""Save torch tensors (image, mask or latent) to disk.
useful to debug things outside comfy.
"""
def __init__(self):
self.output_dir = folder_paths.get_output_directory()
self.type = "mtb/debug"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"filename_prefix": ("STRING", {"default": "ComfyPickle"}),
},
"optional": {
"image": ("IMAGE",),
"mask": ("MASK",),
"latent": ("LATENT",),
},
}
FUNCTION = "save"
OUTPUT_NODE = True
RETURN_TYPES = ()
CATEGORY = "mtb/debug"
def save(
self,
filename_prefix,
image: torch.Tensor | None = None,
mask: torch.Tensor | None = None,
latent: torch.Tensor | None = None,
):
(
full_output_folder,
filename,
counter,
subfolder,
filename_prefix,
) = folder_paths.get_save_image_path(filename_prefix, self.output_dir)
full_output_folder = Path(full_output_folder)
if image is not None:
image_file = f"{filename}_image_{counter:05}.pt"
torch.save(image, full_output_folder / image_file)
# np.save(full_output_folder/ image_file, image.cpu().numpy())
if mask is not None:
mask_file = f"{filename}_mask_{counter:05}.pt"
torch.save(mask, full_output_folder / mask_file)
# np.save(full_output_folder/ mask_file, mask.cpu().numpy())
if latent is not None:
# for latent we must use pickle
latent_file = f"{filename}_latent_{counter:05}.pt"
torch.save(latent, full_output_folder / latent_file)
# pickle.dump(latent, open(full_output_folder/ latent_file, "wb"))
# np.save(full_output_folder / latent_file,
# latent[""].cpu().numpy())
return f"{filename_prefix}_{counter:05}"
__nodes__ = [MTB_SaveTensors]
+7 -21
View File
@@ -7,13 +7,6 @@ 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"""
@@ -200,11 +193,11 @@ by default it fallsback to a default font.
),
"color": (
"COLOR",
{"default": "black"},
{"default": "black", "widgetType": "MTB_COLOR"},
),
"background": (
"COLOR",
{"default": "white"},
{"default": "white", "widgetType": "MTB_COLOR"},
),
"h_align": (("left", "center", "right"), {"default": "left"}),
"v_align": (("top", "center", "bottom"), {"default": "top"}),
@@ -306,7 +299,7 @@ by default it fallsback to a default font.
def text_to_image(
self,
text: str | list[str],
text: str,
font,
wrap,
trim,
@@ -348,7 +341,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: str, alpha=None) -> Image.Image:
def render_text(text_to_render, alpha=None):
if trim:
text_to_render = text_to_render.strip()
if wrap:
@@ -433,16 +426,9 @@ by default it fallsback to a default font.
frame_tensors = [pil2tensor(frame) for frame in frames]
return (torch.cat(frame_tensors, dim=0),)
else:
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),)
text_img = render_text(text)
result = Image.alpha_composite(base_img, text_img)
return (pil2tensor(result),)
__nodes__ = [
+7 -117
View File
@@ -6,7 +6,7 @@ import urllib.request
from math import pi
from typing import Any
import comfy.model_management as mm
import comfy.model_management as model_management
import comfy.utils
import numpy as np
import torch
@@ -16,10 +16,8 @@ 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,
@@ -136,64 +134,12 @@ class MTB_ApplyTextTemplate:
CATEGORY = "mtb/utils"
FUNCTION = "execute"
def execute(self, *, template: str, **kwargs) -> tuple[str | list[str]]:
keys = list(kwargs.keys())
values = list(kwargs.values())
def execute(self, *, template: str, **kwargs):
res = f"{template}"
for k, v in kwargs.items():
res = res.replace(f"{{{k}}}", f"{v}")
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}")
return (res,)
class MTB_MatchDimensions:
@@ -397,7 +343,7 @@ class MTB_AutoPanEquilateral:
frames.append(frame)
mm.throw_exception_if_processing_interrupted()
model_management.throw_exception_if_processing_interrupted()
pbar.update(1)
return (pil2tensor(frames),)
@@ -970,61 +916,6 @@ 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,
@@ -1042,5 +933,4 @@ __nodes__ = [
MTB_TensorOps,
MTB_BooleanNot,
MTB_GetItem,
MTB_ProxyTensor,
]
+35 -7
View File
@@ -683,6 +683,7 @@ class MTB_ImageCompare:
import requests
import time
class MTB_LoadImageFromUrl:
@@ -698,6 +699,14 @@ class MTB_LoadImageFromUrl:
"default": "https://upload.wikimedia.org/wikipedia/commons/thumb/a/a7/Example.jpg/800px-Example.jpg"
},
),
"retry_count": (
"INT",
{"default": 3, "min": 1, "max": 20, "step": 1},
),
"retry_interval": (
"FLOAT",
{"default": 1.0, "min": 0.0, "max": 60.0, "step": 0.1},
),
}
}
@@ -705,11 +714,27 @@ class MTB_LoadImageFromUrl:
FUNCTION = "load"
CATEGORY = "mtb/IO"
def load(self, url):
# get the image from the url
image = Image.open(requests.get(url, stream=True).raw)
image = ImageOps.exif_transpose(image)
return (pil2tensor(image),)
def load(self, url, retry_count, retry_interval):
# get the image from the url with retry + exponential backoff
last_error = None
for attempt in range(retry_count):
try:
response = requests.get(url, stream=True)
response.raise_for_status()
image = Image.open(response.raw)
image = ImageOps.exif_transpose(image)
return (pil2tensor(image),)
except Exception as e:
last_error = e
if attempt == retry_count - 1:
raise
wait_seconds = retry_interval * (2**attempt)
if wait_seconds > 0:
time.sleep(wait_seconds)
if last_error is not None:
raise last_error
raise RuntimeError("Failed to load image from URL without captured exception")
class MTB_Blur:
@@ -879,8 +904,11 @@ class MTB_MaskToImage:
return {
"required": {
"mask": ("MASK",),
"color": ("COLOR",),
"background": ("COLOR", {"default": "#000000"}),
"color": ("COLOR", {"widgetType": "MTB_COLOR"}),
"background": (
"COLOR",
{"default": "#000000", "widgetType": "MTB_COLOR"},
),
},
"optional": {
"invert": ("BOOLEAN", {"default": False}),
+17
View File
@@ -0,0 +1,17 @@
# from ..utils import hex_to_rgb
class MTB_ColorInput:
RETURN_TYPES = ("COLOR",)
FUNCTION = "color"
CATEGORY = "mtb/color"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {"color": ("MTB_COLOR", {"default": "#ffffff"})},
}
def color(self, color):
return (color,)
__nodes__ = [MTB_ColorInput]
+1 -1
View File
@@ -34,7 +34,7 @@ class MTB_ImageRemoveBackgroundRembg:
),
"bgcolor": (
"COLOR",
{"default": "#000000"},
{"default": "#000000","widgetType": "MTB_COLOR"},
),
},
}
+1 -1
View File
@@ -145,7 +145,7 @@ class MTB_ModelPatchSeamless:
tilingX,
tilingY,
):
hacked_model = copy.deepcopy(model)
hacked_model = model.clone()
self.apply_circular(
hacked_model.model, startStep, stopStep, tilingX, tilingY
)
+4 -1
View File
@@ -43,7 +43,10 @@ class MTB_TransformImage:
["edge", "constant", "reflect", "symmetric"],
{"default": "edge"},
),
"constant_color": ("COLOR", {"default": "#000000"}),
"constant_color": (
"COLOR",
{"default": "#000000", "widgetType": "MTB_COLOR"},
),
},
"optional": {
"filter_type": (
+1 -1
View File
@@ -27,7 +27,7 @@ class MTB_LoadVitMatteModel:
def execute(self, *, kind: str, autodownload: bool):
dest = models_dir / "vitmatte"
dest.mkdir(exist_ok=True)
name = "dist" if kind == "Distinctions-646" else "com"
name = "dis" if kind == "Distinctions-646" else "com"
file = hf_hub_download(
repo_id="melmass/pytorch-scripts",
+38 -5
View File
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
[project]
name = "comfy-mtb"
version = "0.6.0"
version = "0.5.4"
description = "Animation oriented nodes pack for ComfyUI."
license = { text = "MIT" }
readme = "README.md"
@@ -62,6 +62,39 @@ 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 = ["."]
@@ -78,18 +111,18 @@ stubPath = "src/stubs"
reportMissingImports = true
reportMissingTypeStubs = false
reportExplicitAny = false
typeCheckingMode = "basic"
pythonVersion = "3.11"
pythonPlatform = "Windows"
pythonVersion = "3.10"
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']
+18
View File
@@ -0,0 +1,18 @@
{
"exclude": [
"**/node_modules",
"**/__pycache__",
],
"ignore": [
"extern"
],
"defineConstant": {
"DEBUG": true
},
"venvPath": "../../../.venv/",
"reportMissingImports": true,
"reportMissingTypeStubs": false,
"pythonVersion": "3.10",
"pythonPlatform": "All",
"reportOptionalMemberAccess": "none"
}
-637
View File
@@ -1,637 +0,0 @@
import base64
import code
import io
import re
import sys
from contextlib import redirect_stderr, redirect_stdout
# import matplotlib.pyplot as plt
import numpy as np
import torch
from aiohttp import web
from PIL import Image
from rich.console import Console
from rich.traceback import Traceback
from .log import log
try:
import pyflakes.api
import pyflakes.reporter
_HAS_LINT = True
except ImportError:
print(
"ComfyREPL: pyflakes not found. Linting will be disabled. Install with 'pip install pyflakes'."
)
_HAS_LINT = False
# --- Linting Library ---
# try:
# import ruff
# import ruff.lint
# import ruff.lint.linter
# import ruff.settings
#
# _HAS_LINT = True
# except ImportError:
# print(
# "ComfyREPL: ruff not found. Linting will be disabled. Install with 'pip install ruff'."
# )
# _HAS_LINT = False
# --- Audio/Video Libraries ---
try:
import scipy.io.wavfile
_HAS_SCIPY = True
except ImportError:
print(
"ComfyREPL: SciPy not found. Audio display will be disabled. Install with 'pip install scipy'."
)
_HAS_SCIPY = False
try:
import imageio
import imageio.plugins.ffmpeg # Ensure ffmpeg plugin is available
_HAS_IMAGEIO = True
except ImportError:
print(
"ComfyREPL: Imageio or imageio-ffmpeg not found. Video display will be disabled. Install with 'pip install imageio imageio-ffmpeg'."
)
_HAS_IMAGEIO = False
# --- Audio Display ---
class AudioDisplay:
def __init__(self, samples, sample_rate):
if not _HAS_SCIPY:
raise ImportError("Audio display requires scipy and numpy.")
if not isinstance(samples, (np.ndarray, torch.Tensor)):
raise TypeError(
"Audio samples must be a numpy array or torch tensor."
)
if isinstance(samples, torch.Tensor):
samples = samples.detach().cpu().numpy()
# Ensure samples are in a format scipy.io.wavfile can handle (e.g., int16, float32)
if samples.dtype == np.float64:
samples = samples.astype(np.float32)
elif samples.dtype == np.int64:
# Or scale to int32 if range requires
samples = samples.astype(np.int16)
self.samples = samples
self.sample_rate = sample_rate
def _to_wav_base64(self):
buffer = io.BytesIO()
try:
scipy.io.wavfile.write(buffer, self.sample_rate, self.samples)
audio_base64 = base64.b64encode(buffer.getvalue()).decode("utf-8")
return audio_base64
except Exception as e:
return f"<div style='color: red;'>Error encoding audio: {e}</div>"
def _repr_html_(self):
base64_data = self._to_wav_base64()
if base64_data.startswith("<div"):
return base64_data
return f'<audio controls src="data:audio/wav;base64,{base64_data}" style="margin: 5px 0;"/>'
def render_audio(samples, sample_rate):
"""
Render audio samples as an HTML audio player.
Args:
samples (np.ndarray or torch.Tensor): Audio samples.
sample_rate (int): Sample rate in Hz.
Returns
-------
AudioDisplay: An object that will render as an HTML audio player.
"""
return AudioDisplay(samples, sample_rate)
# --- Display Classes ---
class VideoDisplay:
def __init__(self, frames, fps=24, options=None):
if not _HAS_IMAGEIO: # numpy/PIL/torch needed for frames
raise ImportError(
"Video display requires imageio, imageio-ffmpeg, and image libraries (numpy, Pillow, torch)."
)
self.frames = []
for frame in frames:
if isinstance(frame, Image.Image):
self.frames.append(np.array(frame))
elif isinstance(frame, np.ndarray):
# Ensure HWC and uint8
if frame.ndim == 3 and frame.shape[0] in [1, 3, 4]: # CHW
frame = np.transpose(frame, (1, 2, 0))
if frame.dtype != np.uint8:
frame = (
(frame * 255).astype(np.uint8)
if frame.max() <= 1.0
else frame.astype(np.uint8)
)
self.frames.append(frame)
elif isinstance(frame, torch.Tensor):
np_frame = frame.detach().cpu().numpy()
if np_frame.ndim == 3 and np_frame.shape[0] in [
1,
3,
4,
]: # CHW
np_frame = np.transpose(np_frame, (1, 2, 0))
if np_frame.dtype != np.uint8:
np_frame = (
(np_frame * 255).astype(np.uint8)
if np_frame.max() <= 1.0
else np_frame.astype(np.uint8)
)
self.frames.append(np_frame)
else:
raise TypeError(
f"Unsupported frame type: {type(frame)}. Must be PIL.Image, numpy.ndarray, or torch.Tensor."
)
self.fps = fps
self.options = options if options is not None else {}
def _to_mp4_base64(self):
buffer = io.BytesIO()
try:
# Use imageio to write frames to an in-memory MP4 file
imageio.mimwrite(
buffer,
self.frames,
format="mp4",
fps=self.fps,
codec="libx264",
quality=8,
) # quality 1-10
video_base64 = base64.b64encode(buffer.getvalue()).decode("utf-8")
return video_base64
except Exception as e:
return f"<div style='color: red;'>Error encoding video: {e}</div>"
def _repr_html_(self):
base64_data = self._to_mp4_base64()
if base64_data.startswith("<div"): # Check if it's an error message
return base64_data
# Build HTML options string
option_str = ""
for key, value in self.options.items():
if isinstance(value, bool) and value:
option_str += f" {key}"
elif isinstance(value, str):
option_str += f' {key}="{value}"'
else:
option_str += f' {key}="{value}"' # Fallback for numbers etc.
return f'<video controls src="data:video/mp4;base64,{base64_data}" style="max-width: 100%; height: auto; border: 1px solid #555; margin: 5px 0;"{option_str}/>'
def render_video(batch_tensor_or_array_of_pil_images, fps=24, options=None):
"""
Render video frames as an HTML video player.
Args:
batch_tensor_or_array_of_pil_images (list of PIL.Image, np.ndarray, or torch.Tensor):
A list of frames, or a single batch tensor/array (B, H, W, C) or (B, C, H, W).
fps (int): Frames per second.
options (dict): Dictionary of HTML <video> tag attributes (e.g., {"loop": True, "autoplay": True}).
Returns
-------
VideoDisplay: An object that will render as an HTML video player.
"""
frames_list = []
if isinstance(
batch_tensor_or_array_of_pil_images, (np.ndarray, torch.Tensor)
):
# Assume it's a batch tensor/array
for i in range(batch_tensor_or_array_of_pil_images.shape[0]):
frames_list.append(batch_tensor_or_array_of_pil_images[i])
elif isinstance(batch_tensor_or_array_of_pil_images, list):
frames_list = batch_tensor_or_array_of_pil_images
else:
raise TypeError(
"Input for render_video must be a list of frames or a batch tensor/array."
)
return VideoDisplay(frames_list, fps, options)
class ComfyREPLBackend:
def __init__(self):
self.repl_consoles: dict[str, code.InteractiveConsole] = {}
# self.repl_console = None
self.image_outputs = []
self.audio_outputs = []
self.video_outputs = []
self._original_displayhook = sys.displayhook
# self._init_repl_console()
@staticmethod
def _init_repl_console():
"""Define the globals that will be available in the REPL session."""
repl_globals = {"__builtins__": __builtins__}
# repl_globals["plt"] = plt
repl_globals["np"] = np
repl_globals["Image"] = Image
repl_globals["torch"] = torch
repl_globals["repl_display"] = _repl_display_image
if _HAS_SCIPY:
repl_globals["render_audio"] = render_audio
if _HAS_IMAGEIO:
repl_globals["render_video"] = render_video
return code.InteractiveConsole(locals=repl_globals)
def _custom_displayhook(self, value):
"""
Displayhook that capture and process image, audio, video objects.
For other objects, it fallsback to the original displayhook.
"""
if value is None:
return
# Attempt to handle as an image
if (
isinstance(value, (Image.Image, np.ndarray, torch.Tensor))
# or (
# hasattr(value, "figure")
# and isinstance(value.figure, plt.Figure)
# )
# or isinstance(value, plt.Figure)
):
img_html = _repl_display_image(value)
self.image_outputs.append(img_html)
return
# Attempt to handle as audio
elif isinstance(value, AudioDisplay):
audio_html = value._repr_html_()
self.audio_outputs.append(audio_html)
return
# Attempt to handle as video
elif isinstance(value, VideoDisplay):
video_html = value._repr_html_()
self.video_outputs.append(video_html)
return
else:
# If not a special media type, let the original displayhook handle it.
self._original_displayhook(value)
def _console_to_html(
self, stream: io.StringIO | Traceback, width: int = 120
) -> str:
if isinstance(stream, io.StringIO):
captured_text_output = stream.getvalue()
else:
captured_text_output = Traceback
html_console = Console(
file=io.StringIO(), record=True, force_terminal=True, width=width
)
html_console.print(captured_text_output)
return html_console.export_html(inline_styles=True)
def get_console(self, node_name: str):
console = self.repl_consoles.get(node_name)
if console:
return console
console = self._init_repl_console()
self.repl_consoles[node_name] = console
return self.repl_consoles[node_name]
def execute_code(self, node_name: str, code: str):
# Clear outputs from previous execution
self.image_outputs = []
self.audio_outputs = []
self.video_outputs = []
output_html = ""
error_message = None
repl_console = self.get_console(node_name)
string_io = io.StringIO()
# Temporarily patch sys.displayhook
sys.displayhook = self._custom_displayhook
try:
with redirect_stdout(string_io), redirect_stderr(string_io):
for line in code.splitlines():
repl_console.push(line)
full_rich_html = self._console_to_html(string_io)
match = re.search(
r"<body.*?>(.*?)</body>", full_rich_html, re.DOTALL
)
if match:
output_html = match.group(1)
else:
output_html = full_rich_html
# Append any captured media HTML *after* the rich text output
for img_html in self.image_outputs:
output_html += img_html
for audio_html in self.audio_outputs:
output_html += audio_html
for video_html in self.video_outputs:
output_html += video_html
except Exception as e:
exc_type, exc_value, exc_traceback = sys.exc_info()
rich_traceback = Traceback.from_exception(
exc_type,
exc_value,
exc_traceback,
show_locals=True,
suppress=[__file__],
)
# error_console = Console(
# file=io.StringIO(), record=True, force_terminal=True, width=120
# )
# error_console.print(rich_traceback)
# full_error_html = error_console.export_html(inline_styles=True)
full_error_html = self._console_to_html(rich_traceback)
match = re.search(
r"<body.*?>(.*?)</body>", full_error_html, re.DOTALL
)
output_html = match.group(1) if match else full_error_html
error_message = str(e)
finally:
sys.displayhook = (
self._original_displayhook
) # Always restore original displayhook
return {"output_html": output_html, "error": error_message}
def lint_code(self, node_name: str, code: str):
diagnostics = []
if not _HAS_LINT:
diagnostics.append(
{
"row": 0,
"column": 0,
"text": "Pyflakes not installed. Linting disabled. Install with 'uv add pyflakes'.",
"type": "warning",
}
)
return web.json_response({"diagnostics": diagnostics})
# Use a custom reporter to capture messages
class PyflakesReporter(pyflakes.reporter.Reporter):
def __init__(self):
self.messages = []
# Suppress stdout/stderr from pyflakes itself
self._stdout = io.StringIO()
self._stderr = io.StringIO()
super().__init__(self._stdout, self._stderr)
def flake(self, message):
# Ace editor expects 0-indexed row, pyflakes gives 1-indexed lineno
self.messages.append(
{
"row": message.lineno - 1,
"column": message.col,
"text": str(message),
"type": "warning", # pyflakes usually gives warnings
}
)
def unexpectedError(self, filename, msg):
self.messages.append(
{
"row": 0,
"column": 0,
"text": f"Pyflakes internal error: {msg}",
"type": "error",
}
)
def syntaxError(self, filename, msg, lineno, offset, text):
log.info(f"Received {text} to syntax error")
self.messages.append(
{
"row": lineno - 1, # Ace is 0-indexed
"column": offset,
"text": f"Syntax Error: {msg}",
"type": "error",
}
)
reporter = PyflakesReporter()
pyflakes.api.check(code, node_name, reporter)
return {"diagnostics": reporter.messages}
def lint_code_ruff(self, code: str):
diagnostics = []
if not _HAS_LINT:
diagnostics.append(
{
"row": 0,
"column": 0,
"text": "Ruff not installed. Linting disabled. Install with 'pip install ruff'.",
"type": "warning",
}
)
return {"diagnostics": diagnostics}
# Define the builtins/globals that Ruff should recognize
# These are the names we inject into the REPL's scope
repl_builtins = [
"repl_display",
"render_audio",
"render_video",
# "plt",
"np",
"Image",
"torch",
]
try:
# Lint the code using Ruff's programmatic API
result = ruff.lint.linter.lint_stdin(
code.encode("utf-8"),
path="<stdin>",
builtins=repl_builtins,
)
for diagnostic in result.diagnostics:
diag_type = "warning" # Default
# Ruff's error codes: F (Pyflakes), E (Pycodestyle), W (Pycodestyle warning), I (isort), N (naming), etc.
# F821: Undefined name (often an error)
if (
diagnostic.kind.code.startswith("E")
or diagnostic.kind.code == "F821"
):
diag_type = "error"
elif diagnostic.kind.code.startswith("W"):
diag_type = "warning"
diagnostics.append(
{
"row": diagnostic.location.row - 1, # Ace is 0-indexed
"column": diagnostic.location.column
- 1, # Ace is 0-indexed
"text": diagnostic.message,
"type": diag_type,
}
)
except Exception as e:
diagnostics.append(
{
"row": 0,
"column": 0,
"text": f"Ruff internal error: {e}",
"type": "error",
}
)
return {"diagnostics": diagnostics}
def _repl_display_image(img_data):
"""
Internal function to convert image data (PIL, numpy, torch, matplotlib) to base64 HTML.
"""
pil_img = None
# fig = None
if isinstance(img_data, Image.Image):
pil_img = img_data
elif isinstance(img_data, np.ndarray):
# Handle different numpy array shapes (HWC, CHW)
if img_data.ndim == 3:
if img_data.shape[0] in [1, 3, 4]: # Likely CHW
if img_data.shape[0] == 1: # Grayscale
img_data = img_data.squeeze(0)
else: # Color
img_data = np.transpose(img_data, (1, 2, 0)) # CHW to HWC
# Ensure it's uint8 for PIL, assuming float [0,1] or int [0,255]
if img_data.dtype != np.uint8:
img_data = (
(img_data * 255).astype(np.uint8)
if img_data.max() <= 1.0
else img_data.astype(np.uint8)
)
pil_img = Image.fromarray(img_data)
elif isinstance(img_data, torch.Tensor):
# Move to CPU, convert to numpy
np_img = img_data.detach().cpu().numpy()
# Handle different tensor shapes (CHW, HWC)
if np_img.ndim == 3:
if np_img.shape[0] in [1, 3, 4]: # Likely CHW
if np_img.shape[0] == 1: # Grayscale
np_img = np_img.squeeze(0)
else: # Color
np_img = np.transpose(np_img, (1, 2, 0)) # CHW to HWC
# Ensure it's uint8 for PIL, assuming float [0,1] or int [0,255]
if np_img.dtype != np.uint8:
np_img = (
(np_img * 255).astype(np.uint8)
if np_img.max() <= 1.0
else np_img.astype(np.uint8)
)
pil_img = Image.fromarray(np_img)
# elif hasattr(img_data, "figure") and isinstance(
# img_data.figure, plt.Figure
# ):
# # If it's a matplotlib Axes object, get its figure
# fig = img_data.figure
# elif isinstance(img_data, plt.Figure):
# fig = img_data
else:
return f"<div style='color: red;'>Unsupported image type for display: {type(img_data)}</div>"
buffer = io.BytesIO()
try:
if pil_img:
pil_img.save(buffer, format="PNG")
# elif fig:
# fig.savefig(
# buffer, format="PNG", bbox_inches="tight", pad_inches=0.1
# )
# plt.close(
# fig
# ) # Close the figure to prevent it from showing up in other contexts
else:
return (
"<div style='color: red;'>Could not process image data.</div>"
)
except Exception as e:
return f"<div style='color: red;'>Error saving image: {e}</div>"
img_base64 = base64.b64encode(buffer.getvalue()).decode("utf-8")
return f'<img src="data:image/png;base64,{img_base64}" style="max-width: 100%; height: auto; border: 1px solid #555; margin: 5px 0;"/>'
# Instantiate the backend class globally
_comfy_repl_backend = ComfyREPLBackend()
# Update aiohttp handlers to use the backend instance
async def repl_execute_code_handler(request):
data = await request.json()
name = data.get("name")
if name is None: # we send an error
return web.Response(
status=417, reason="Expectation Failed", text="Missing name key"
)
code = data.get("code", "")
result = _comfy_repl_backend.execute_code(name, code)
return web.json_response(result)
async def repl_lint_code_handler(request):
data = await request.json()
name = data.get("name")
if name is None: # we send an error
return web.Response(
status=417, reason="Expectation Failed", text="Missing name key"
)
# raise web.HTTPExpectationFailed(
# reason="Missing name key (reason)", text="Missing name key (text)"
# )
code = data.get("code", "")
result = _comfy_repl_backend.lint_code(name, code)
return web.json_response(result)
def setup_custom_web_routes(app: web.Application):
"""
Function to register our custom web routes with the ComfyUI server.
"""
log.info("ComfyREPL: Registering /mtb/execute route...")
app.router.add_post("/mtb/execute", repl_execute_code_handler)
app.router.add_post("/mtb/lint", repl_lint_code_handler)
# You can add more routes here if needed, e.g., for clearing state.
+5 -196
View File
@@ -8,14 +8,11 @@ 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
@@ -464,6 +461,8 @@ def _run_command(shell_cmd, ignored_lines_start):
print("Command executed successfully!")
# endregion
@@ -527,196 +526,6 @@ 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)
@@ -727,12 +536,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]]:
"""Handle batch processing for a given tensor and conversion function."""
"""Handles batch processing for a given tensor and conversion function."""
return [func(tensor[i]) for i in range(tensor.shape[0])]
def tensor2pil(tensor: torch.Tensor) -> list[Image.Image]:
"""Convert a batch of tensors to a list of PIL Images."""
"""Converts a batch of tensors to a list of PIL Images."""
def single_tensor2pil(t: torch.Tensor) -> Image.Image:
np_array = to_numpy(t)
@@ -749,7 +558,7 @@ def tensor2pil(tensor: torch.Tensor) -> list[Image.Image]:
def pil2tensor(images: Image.Image | list[Image.Image]) -> torch.Tensor:
"""Convert a PIL Image or a list of PIL Images to a tensor."""
"""Converts a PIL Image or a list of PIL Images to a tensor."""
def single_pil2tensor(image: Image.Image) -> torch.Tensor:
np_image = np.array(image).astype(np.float32) / 255.0
+15 -201
View File
@@ -14,51 +14,6 @@ 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()
@@ -335,10 +290,6 @@ 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)
@@ -429,54 +380,12 @@ 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,rename_menu?:'label'|'name', start_index?:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} [opts] Extra options
* @param {{separator?:string, start_index?:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} [opts] Extra options
* @returns
*/
export const setupDynamicConnections = (
@@ -490,115 +399,20 @@ export const setupDynamicConnections = (
Object.getOwnPropertyDescriptors(nodeType).title.value,
)
/** @type {{separator:string,rename_menu?:"label"|"name" start_index:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} */
/** @type {{separator:string, start_index:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} */
const options = Object.assign(
{
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
const input = addDynamicInput(
this,
this.addInput(
`${prefix}${options.separator}${options.start_index}`,
inputList ? '*' : inputType,
)
@@ -670,7 +484,10 @@ export const dynamic_connection = (
opts || {},
)
if (node.inputs.length > 0 && !isDynamicInput(node.inputs[index])) {
// function to test if input is a dynamic one
const isDynamicInput = (inputName) => inputName.startsWith(connectionPrefix)
if (node.inputs.length > 0 && !isDynamicInput(node.inputs[index].name)) {
return
}
@@ -680,7 +497,6 @@ 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
@@ -690,7 +506,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)) {
if (!element.link && isDynamicInput(element.name)) {
if (node.widgets) {
const w = node.widgets.find((w) => w.name === element.name)
if (w) {
@@ -704,12 +520,9 @@ export const dynamic_connection = (
}
for (let i = 0; i < to_remove.length; i++) {
const id = to_remove[i]
try {
node.removeInput(id)
i_count -= 1
} catch (err) {
errorLogger('Cannot remove input', err)
}
node.removeInput(id)
i_count -= 1
}
node.inputs.length = i_count
@@ -723,7 +536,7 @@ export const dynamic_connection = (
for (let i = 0; i < node.inputs.length; i++) {
let name = ''
// rename only prefixed inputs
if (node.inputs[i].name.startsWith(connectionPrefix)) {
if (isDynamicInput(node.inputs[i].name)) {
// prefixed => rename and increase index
name = `${connectionPrefix}${prefixed_idx}`
prefixed_idx += 1
@@ -777,8 +590,9 @@ 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) ? ++acc : acc),
(acc, cur) => (isDynamicInput(cur.name) ? ++acc : acc),
0,
)
@@ -788,7 +602,7 @@ export const dynamic_connection = (
: `${connectionPrefix}${nextIndex + options.start_index}`
infoLogger(`Adding input ${nextIndex + 1} (${name})`)
addDynamicInput(node, name, conType)
node.addInput(name, conType)
}
}
}
+86 -67
View File
@@ -11,12 +11,7 @@
/// <reference path="../types/typedefs.js" />
import { app } from '../../scripts/app.js'
import {
setupDynamicConnections,
cleanupNode,
infoLogger,
} from './comfy_shared.js'
import * as shared from './comfy_shared.js'
import * as mtb_ui from './mtb_ui.js'
function escapeHtml(unsafe) {
@@ -50,25 +45,25 @@ function createDebugSection(title) {
return section
}
function createDebugContent(item) {
function createDebugContent(content, type) {
const wrapper = mtb_ui.makeElement('div', {
margin: '4px 0',
})
if (item.kind === 'text') {
const text = mtb_ui.makeElement('div', {
if (type === 'text') {
const text = mtb_ui.makeElement('p', {
margin: '2px 0',
fontFamily: 'monospace',
whiteSpace: 'pre-wrap',
})
text.innerHTML = item.data
text.innerHTML = content
wrapper.appendChild(text)
} else if (item.kind === 'b64_images') {
} else if (type === 'image') {
const img = mtb_ui.makeElement('img', {
width: '100%',
borderRadius: '2px',
})
img.src = item.data
img.src = content
wrapper.appendChild(img)
}
@@ -85,85 +80,110 @@ app.registerExtension({
*/
async beforeRegisterNodeDef(nodeType, nodeData, app) {
if (nodeData.name === 'Debug (mtb)') {
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 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 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)
},
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,
})
}
setupDynamicConnections(nodeType, 'var', '*')
//- infer type
if (link_info) {
// const fromNode = this.graph._nodes.find(
// (otherNode) => otherNode.id === link_info.origin_id,
// )
// const fromNode = app.graph.getNodeById(link_info.origin_id)
const { from } = shared.nodesFromLink(this, link_info)
if (!from || this.inputs.length === 0) return
const type = from.outputs[link_info.origin_slot].type
this.inputs[index].type = type
// this.inputs[index].label = type.toLowerCase()
}
//- restore dynamic input
if (!connected) {
this.inputs[index].type = '*'
this.inputs[index].label = `anything_${index + 1}`
}
return r
}
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (...args) {
onExecuted?.apply(this, args)
const [data, ..._rest] = args
clear_widgets(this)
if (this.widgets) {
let tgt_len = this.widgets.length
for (let i = 0; i < this.widgets.length; i++) {
if (
this.widgets[i].name !== 'output_to_console' &&
this.widgets[i].name !== 'as_detailed_types'
) {
this.widgets[i].onRemove?.()
this.widgets[i].onRemoved?.()
tgt_len -= 1
}
}
this.widgets.length = tgt_len
}
const inputData = {}
const 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
inputData[inputName] = item.items
if (!inputData[inputName]) {
inputData[inputName] = { text: [], b64_images: [] }
}
if (item.text) {
inputData[inputName].text.push(...item.text)
}
if (item.b64_images) {
inputData[inputName].b64_images.push(...item.b64_images)
}
})
}
const mainDebugContainer = mtb_ui.makeElement('div', {
width: '100%',
})
let hasContent = false
let widgetI = 1
for (const [inputName, content] of Object.entries(inputData)) {
if (!content || content?.length === 0) {
if (content.text.length === 0 && content.b64_images.length === 0) {
continue
}
hasContent = true
const section = createDebugSection(name_to_label[inputName])
const section = createDebugSection(inputName)
for (const item of content) {
section.appendChild(createDebugContent(item))
if (content.text.length > 0) {
content.text.forEach((text) => {
section.appendChild(createDebugContent(text, 'text'))
})
}
mainDebugContainer.appendChild(section)
}
if (hasContent) {
this.addDOMWidget('debug_output', 'CUSTOM', mainDebugContainer, {
hideOnZoom: false,
})
if (content.b64_images.length > 0) {
content.b64_images.forEach((img) => {
section.appendChild(createDebugContent(img, 'image'))
})
}
this.addDOMWidget(`debug_section_${widgetI}`, 'CUSTOM', section, {})
widgetI++
}
this.onRemoved = function () {
@@ -174,9 +194,8 @@ app.registerExtension({
widget.onRemoved?.()
widget.onRemove?.()
}
cleanupNode(this)
shared.cleanupNode(this)
}
this.setDirtyCanvas(true, true)
}
}
},
-529
View File
@@ -1,529 +0,0 @@
/** Python REPL for the frontend (uses rich)*/
import { app } from '../../scripts/app.js'
import * as shared from './comfy_shared.js'
import * as mtb_ui from './mtb_ui.js'
class ComfyREPL extends LiteGraph.LGraphNode {
constructor() {
super()
this.shape = LiteGraph.BOX_SHAPE
this.isVirtualNode = true
this.category = 'mtb/repl'
this.title = '🐍 REPL (mtb)'
this.uuid = shared.makeUUID()
this.size = [600, 400]
// Create a container for our custom widgets
this.widget = this.addDOMWidget('HTML', 'html', this.createREPLWidget())
this.loadAceEditor()
// Store input and output for persistence
this.properties = {
inputCode: '',
outputHistory: '',
}
this.outputArea.innerHTML = this.properties.outputHistory
this.outputArea.scrollTop = this.outputArea.scrollHeight
// Debounced linting function
this.debouncedLint = shared.debounce(this.lintCode.bind(this), 500)
// Resizing state variables
this.isResizing = false
this.initialMouseY = 0
this.initialInputHeight = 0
this.initialOutputHeight = 0
}
loadAceEditor() {
if (window.MTB?.ace_loaded) {
return
}
let NEED_PATCH = false
if (window.ace) {
shared.infoLogger(
'A global ace was found in scope, to avoid issues with it we will patch it',
)
NEED_PATCH = true
// window._backupAce = window.ace
// window.ace = null
}
shared
.loadScript('/mtb_async/ace/ace.js')
.then((m) => {
shared.infoLogger('ACE was loaded', m)
// window.MTB_ACE = window.ace
window.MTB.ace_loaded = true
this.initAceEditor()
this.aceEditor.setValue(this.properties.inputCode, -1)
})
.catch((e) => {
shared.errorLogger(e)
})
.finally(() => {
if (NEED_PATCH) {
console.log('Patching back window object')
window.ace = window._backupAce
}
})
}
initAceEditor() {
if (!window.MTB.ace_loaded) {
console.error('ACE editor not loaded. Cannot set up editors.')
return
}
if (!this.inputDiv) {
console.error('Input div not found for Ace editor initialization.')
return
}
this.aceEditor = ace.edit(this.inputDiv)
this.aceEditor.setTheme('ace/theme/monokai') //"ace/theme/dracula", "ace/theme/github"
this.aceEditor.session.setMode('ace/mode/python')
this.aceEditor.setOptions({
enableBasicAutocompletion: true,
enableLiveAutocompletion: true,
enableSnippets: true,
fontSize: '14px',
fontFamily: 'monospace',
showPrintMargin: false,
wrap: true,
tabSize: 4,
useSoftTabs: true,
highlightActiveLine: true,
highlightSelectedWord: true,
cursorStyle: 'ace', // "ace" | "slim" | "smooth" | "wide"
behavioursEnabled: true,
displayIndentGuides: true,
fixedWidthGutter: true,
scrollPastEnd: 0.5,
})
// Custom keybinding for Ctrl+Enter
this.aceEditor.commands.addCommand({
name: 'runCode',
bindKey: { win: 'Ctrl-Enter', mac: 'Command-Enter' },
exec: () => this.executeCode(),
})
// Listen for changes to trigger linting
let lintDisabled = false
this.aceEditor.session.on('change', () => {
if (!lintDisabled) {
this.debouncedLint()
}
})
this.outputArea.scrollTop = this.outputArea.scrollHeight
}
addOutput(html) {
this.outputArea.innerHTML += html
this.properties.outputHistory += html
this.outputArea.scrollTop = this.outputArea.scrollHeight
}
createREPLWidget() {
const container = mtb_ui.makeElement('div', {
display: 'flex',
flexDirection: 'column',
width: '100%',
height: '100%',
boxSizing: 'border-box',
padding: '5px',
})
this.inputDiv = mtb_ui.makeElement(
'div',
{
width: 'calc(100% - 10px)',
height: '100px',
backgroundColor: '#333',
color: '#eee',
border: '1px solid #555',
borderRadius: '4px',
marginBottom: '5px',
boxSizing: 'border-box',
overflow: 'hidden',
},
container,
)
// Resizable Handle
this.handleDiv = mtb_ui.makeElement(
'div',
{
width: '100%',
height: '5px',
backgroundColor: '#666',
cursor: 'ns-resize',
marginBottom: '5px',
borderRadius: '2px',
},
container,
)
this.handleDiv.addEventListener('mousedown', this.startResizing.bind(this))
// Run Button
this.runButton = mtb_ui.makeElement(
'button',
{
width: '100%',
padding: '8px',
backgroundColor: '#555',
color: '#fff',
border: 'none',
borderRadius: '4px',
cursor: 'pointer',
marginBottom: '5px',
fontSize: '14px',
},
container,
)
this.runButton.textContent = 'Run Code (Ctrl+Enter)'
this.runButton.onclick = () => this.executeCode()
// Clear Button
this.clearButton = mtb_ui.makeElement(
'button',
{
width: '100%',
padding: '8px',
backgroundColor: '#555',
color: '#fff',
border: 'none',
borderRadius: '4px',
cursor: 'pointer',
marginBottom: '5px',
fontSize: '14px',
},
container,
)
this.clearButton.textContent = 'Clear Output'
this.clearButton.onclick = () => {
this.outputArea.innerHTML = ''
this.properties.outputHistory = ''
}
// Output Area
this.outputArea = mtb_ui.makeElement(
'div',
{
flexGrow: '1',
width: 'calc(100% - 10px)',
backgroundColor: '#222',
color: '#ddd',
border: '1px solid #555',
borderRadius: '4px',
padding: '5px',
fontFamily: 'monospace',
fontSize: '14px',
overflowY: 'auto',
whiteSpace: 'pre-wrap',
boxSizing: 'border-box',
},
container,
)
return container
}
// --- Resizing Logic ---
startResizing(e) {
if (!this.inputDiv) {
shared.infoLogger("The input div isn't ready", this)
shared.errorLogger("The input div isn't ready")
return
}
this.isResizing = true
this.initialMouseY = e.clientY
this.initialInputHeight = this.inputDiv.offsetHeight
this.initialOutputHeight = this.outputArea.offsetHeight
document.addEventListener('mousemove', this.doResize.bind(this))
document.addEventListener('mouseup', this.stopResizing.bind(this))
document.body.style.cursor = 'ns-resize' // Change cursor globally
}
doResize(e) {
if (!this.isResizing) return
const deltaY = e.clientY - this.initialMouseY
let new_input_height = this.initialInputHeight + deltaY
let new_output_height = this.initialOutputHeight - deltaY
const minInputHeight = 50 // Minimum height for Ace editor
const minOutputHeight = 50 // Minimum height for output area
// Clamp heights to minimums
if (new_input_height < minInputHeight) {
new_input_height = minInputHeight
new_output_height =
this.initialInputHeight + this.initialOutputHeight - minInputHeight
}
if (new_output_height < minOutputHeight) {
new_output_height = minOutputHeight
new_input_height =
this.initialInputHeight + this.initialOutputHeight - minOutputHeight
}
this.inputDiv.style.height = `${new_input_height}px`
this.outputArea.style.height = `${new_output_height}px`
// Update the stored ratio for persistence
const totalDynamicHeight =
this.inputDiv.offsetHeight + this.outputArea.offsetHeight
if (totalDynamicHeight > 0) {
this.properties.inputHeightRatio = new_input_height / totalDynamicHeight
}
this.aceEditor.resize() // Important for Ace to redraw
}
stopResizing() {
this.isResizing = false
document.removeEventListener('mousemove', this.doResize)
document.removeEventListener('mouseup', this.stopResizing)
document.body.style.cursor = '' // Restore default cursor
}
// --- End Resizing Logic ---
async executeCode() {
const code = this.aceEditor.getValue()
if (!code.trim()) {
return
}
const inputPrompt = `<div style="color:#888; margin-top: 10px;">>>> ${code}</div>`
this.addOutput(inputPrompt)
try {
const response = await fetch('/mtb/execute', {
method: 'POST',
headers: {
'Content-Type': 'application/json',
},
body: JSON.stringify({ code: code, name: this.uuid }),
})
if (!response.ok) {
throw new Error(`HTTP error! status: ${response.status}`)
}
const result = await response.json()
console.debug('Received from backend', result)
const outputHtml = result.output_html || ''
const error = result.error
if (error) {
this.addOutput(
`<div style="color: #f00; font-weight: bold;">Error:</div>${outputHtml}`,
)
} else {
this.addOutput(outputHtml)
}
} catch (e) {
const errorMessage = `<div style="color: #f00;">Frontend Error: ${e.message}</div>`
this.addOutput(errorMessage)
console.error('ComfyREPL Frontend Error:', e)
} finally {
// Not clearing
// this.inputArea.value = '' // Clear input after execution
// this.properties.inputCode = '' // Clear persisted input
}
}
async lintCode() {
if (!this.aceEditor) {
return
}
const code = this.aceEditor.getValue()
if (!code.trim()) {
this.aceEditor.session.setAnnotations([]) // Clear annotations if empty
return
}
try {
const response = await fetch('/mtb/lint', {
// New linting endpoint
method: 'POST',
headers: {
'Content-Type': 'application/json',
},
body: JSON.stringify({ code: code, name: this.uuid }),
})
if (!response.ok) {
console.log(response)
throw new Error(
`HTTP error! status: ${response.status} ${response.statusText}`,
)
}
const result = await response.json()
// result.diagnostics should be an array of {row, column, text, type}
this.aceEditor.session.setAnnotations(result.diagnostics)
} catch (e) {
console.error('ComfyREPL Linting Error:', e)
this.aceEditor.session.setAnnotations([
{
row: 0,
column: 0,
text: `Linting failed: ${e.message}`,
type: 'error',
},
])
}
}
// Restore properties when loading a graph
onConfigure() {
if (this.properties.inputCode && this.aceEditor) {
this.aceEditor.setValue(this.properties.inputCode, -1)
}
// if (this.properties.inputCode) {
// this.inputArea.value = this.properties.inputCode
// }
if (this.properties.outputHistory) {
this.outputArea.innerHTML = this.properties.outputHistory
this.outputArea.scrollTop = this.outputArea.scrollHeight
}
if (this.properties.uuid) {
this.uuid = this.properties.uuid
}
this.debouncedLint()
this.onResize(this.size)
}
// Save properties when saving a graph
onSerialize(o) {
if (this.aceEditor) {
o.properties.inputCode = this.aceEditor.getValue() //this.inputArea.value
}
o.properties.outputHistory = this.outputArea.innerHTML
o.properties.uuid = this.uuid
o.properties.inputHeightRatio = this.properties.inputHeightRatio
}
onRemoved() {
// Clean up DOM elements when node is removed
if (this.widget?.element?.parentNode) {
this.widget.element.parentNode.removeChild(this.widget.element)
}
// Destroy Ace editor instance to prevent memory leaks
if (this.aceEditor) {
this.aceEditor.destroy()
this.aceEditor.container.remove() // Remove the Ace container div from DOM
}
// Clean up global event listeners if node is removed while resizing
document.removeEventListener('mousemove', this.doResize)
document.removeEventListener('mouseup', this.stopResizing)
document.body.style.cursor = ''
}
// LiteGraph method to handle node resizing
onResize(size) {
// Call parent method if it exists (important for LiteGraph's internal sizing)
if (super.onResize) {
super.onResize(size)
}
// Adjust container size
const container = this.widget.element
container.style.width = `${size[0] - 10}px` // Account for padding
container.style.height = `${size[1] - 10}px`
// Adjust input and output area widths
this.inputDiv.style.width = 'calc(100% - 10px)'
this.outputArea.style.width = 'calc(100% - 10px)'
//
// const old = () => {
// // Calculate remaining height for output area
// // Ace editor manages its own height within this.inputDiv, so we use offsetHeight
// const inputHeight = this.inputDiv.offsetHeight
// const runButtonHeight = this.runButton.offsetHeight
// const clearButtonHeight = this.clearButton.offsetHeight
// const totalFixedHeight =
// inputHeight + runButtonHeight + clearButtonHeight + 15 // 15 for margins/padding
//
// const remainingHeight = size[1] - 10 - totalFixedHeight
// this.outputArea.style.height = `${Math.max(50, remainingHeight)}px` // Min height 50px
// }
// Calculate dynamic heights
const containerHeight = size[1] - 10
const handleHeight = this.handleDiv.offsetHeight
const buttonHeights =
this.runButton.offsetHeight + this.clearButton.offsetHeight + 15 // Sum of button heights + margins
const dynamicContentHeight = containerHeight - buttonHeights - handleHeight
const minInputHeight = 50
const minOutputHeight = 50
let inputHeight = Math.max(
minInputHeight,
dynamicContentHeight * (this.properties.inputHeightRatio || 1.0),
)
let outputHeight = Math.max(
minOutputHeight,
dynamicContentHeight - inputHeight,
)
//
// // Re-distribute if one hits its minimum
// if (
// inputHeight === minInputHeight &&
// dynamicContentHeight - minInputHeight > minOutputHeight
// ) {
// outputHeight = dynamicContentHeight - minInputHeight
// } else if (
// outputHeight === minOutputHeight &&
// dynamicContentHeight - minOutputHeight > minInputHeight
// ) {
// inputHeight = dynamicContentHeight - minOutputHeight
// }
//
// // Final check to ensure total height matches available dynamic space
// const currentTotal = inputHeight + outputHeight
// if (currentTotal !== dynamicContentHeight) {
// // Adjust one of them if there's a small discrepancy due to rounding
// if (inputHeight > minInputHeight) {
// inputHeight += dynamicContentHeight - currentTotal
// } else if (outputHeight > minOutputHeight) {
// outputHeight += dynamicContentHeight - currentTotal
// }
// }
this.inputDiv.style.height = `${inputHeight}px`
this.outputArea.style.height = `${outputHeight}px`
// Update the ratio based on the actual heights set
if (dynamicContentHeight > 0) {
this.properties.inputHeightRatio = inputHeight / dynamicContentHeight
}
// Inform Ace editor about the resize so it can redraw its content
if (this.aceEditor) {
this.aceEditor.resize()
}
}
}
const repl = {
name: 'mtb.repl',
registerCustomNodes() {
LiteGraph.registerNodeType('Python REPL', ComfyREPL)
},
}
app.registerExtension(repl)
+50 -49
View File
@@ -21,7 +21,7 @@ import { infoLogger } from './comfy_shared.js'
import { NumberInputWidget } from './numberInput.js'
// NOTE: new widget types registered by MTB Widgets
const newTypes = [/*'BOOL'*/ 'COLOR', 'BBOX']
const newTypes = [/*'BOOL'*/ 'COLOR','MTB_COLOR', 'BBOX']
const deprecated_nodes = {
// 'Animation Builder':
@@ -739,7 +739,6 @@ const mtb_widgets = {
// },
COLOR: (node, inputName, inputData, _app) => {
console.debug('Registering color')
return {
widget: node.addCustomWidget(
MtbWidgets.COLOR(inputName, inputData[1]?.default || '#ff0000'),
@@ -748,6 +747,16 @@ const mtb_widgets = {
minHeight: 30,
}
},
MTB_COLOR: (node, inputName, inputData, _app) => {
return {
widget: node.addCustomWidget(
MtbWidgets.COLOR(inputName, inputData[1]?.default || '#ff0000'),
),
minWidth: 150,
minHeight: 30,
}
},
// BBOX: (node, inputName, inputData, app) => {
// console.debug("Registering bbox")
// return {
@@ -1051,6 +1060,7 @@ const mtb_widgets = {
const currentChunkSize = Math.min(chunkSize, totalPrompts - i)
await app.queuePrompt(0, currentChunkSize)
}
if (!cancelQueue) {
window.MTB?.notify?.(
@@ -1196,9 +1206,7 @@ const mtb_widgets = {
//NOTE: dynamic nodes
case 'Apply Text Template (mtb)': {
shared.setupDynamicConnections(nodeType, 'var', '*', {
rename_menu: 'name',
})
shared.setupDynamicConnections(nodeType, 'var', '*')
break
}
case 'Save Data Bundle (mtb)': {
@@ -1330,54 +1338,47 @@ const mtb_widgets = {
const related = new Set([this.id])
const visited = new Set()
if (this.outputs[0].links) {
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)
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
}
}
}
traverseNodes(this)
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)
}
}
}
traverseNodes(this)
}
this.related_to_flow = Array.from(related)
this.computed_flow = true