Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
821a031bfc | ||
|
|
3d12bd29a8 | ||
|
|
3494a4767e | ||
|
|
41d444ae70 | ||
|
|
41d79d4677 | ||
|
|
54f4963a7b | ||
|
|
d585b16ee7 | ||
|
|
499cd218aa | ||
|
|
fd00e12724 | ||
|
|
e37e4648e5 | ||
|
|
06685a5418 | ||
|
|
40145ddf40 | ||
|
|
c1d74c0d69 | ||
|
|
eca2ea5da9 | ||
|
|
c6d1f73cfc | ||
|
|
4ea7c0b67f | ||
|
|
8316914f02 | ||
|
|
5c0e020c73 | ||
|
|
d00722e9ea | ||
|
|
0106c13250 | ||
|
|
55226058d4 | ||
|
|
50e0f7b357 | ||
|
|
71f601094a | ||
|
|
ea750b5e8b | ||
|
|
ff2e99f73e | ||
|
|
efc6855073 | ||
|
|
0853b7fb6a | ||
|
|
10aa493dd8 | ||
|
|
f038d76748 | ||
|
|
940a781f29 | ||
|
|
c7248344cc | ||
|
|
fab33a40a2 | ||
|
|
8f83e8d4d7 | ||
|
|
6c59d5c32d | ||
|
|
e98f3f626f | ||
|
|
7e89e96e9d | ||
|
|
177b6eeef3 | ||
|
|
502a583409 | ||
|
|
a7966355c1 | ||
|
|
321abea51a | ||
|
|
63be3f26fd | ||
|
|
c4f40e299f | ||
|
|
b541670a5b | ||
|
|
4574c6451c | ||
|
|
7fb27804e1 | ||
|
|
9a7e022df1 | ||
|
|
2c483fd1d2 | ||
|
|
0967d439f5 | ||
|
|
319c02d658 | ||
|
|
265cb953ec | ||
|
|
7e36007933 |
@@ -0,0 +1,7 @@
|
|||||||
|
**/GFPGAN/inputs/**
|
||||||
|
**/GFPGAN/tests/**
|
||||||
|
**/frame_interpolation/photos/*
|
||||||
|
moment.gif
|
||||||
|
node.zip
|
||||||
|
.DS_Store
|
||||||
|
|
||||||
@@ -1,9 +1,6 @@
|
|||||||
name: 📦 Publish to Comfy registry
|
name: 📦 Publish to Comfy registry
|
||||||
on:
|
on:
|
||||||
workflow_dispatch:
|
workflow_dispatch:
|
||||||
push:
|
|
||||||
tags:
|
|
||||||
- '*'
|
|
||||||
|
|
||||||
permissions:
|
permissions:
|
||||||
issues: write
|
issues: write
|
||||||
@@ -21,4 +18,5 @@ jobs:
|
|||||||
- name: 📦 Publish Custom Node
|
- name: 📦 Publish Custom Node
|
||||||
uses: Comfy-Org/publish-node-action@v1
|
uses: Comfy-Org/publish-node-action@v1
|
||||||
with:
|
with:
|
||||||
|
skip_checkout: 'true'
|
||||||
personal_access_token: ${{ secrets.COMFY_REGISTRY_TOKEN }}
|
personal_access_token: ${{ secrets.COMFY_REGISTRY_TOKEN }}
|
||||||
|
|||||||
@@ -1,11 +1,16 @@
|
|||||||
__pycache__
|
__pycache__
|
||||||
*.py[cod]
|
*.py[cod]
|
||||||
*.onnx
|
*.onnx
|
||||||
|
|
||||||
wheels/
|
wheels/
|
||||||
node_modules/
|
node_modules/
|
||||||
compose.yaml
|
compose.yaml
|
||||||
comfy_mtb.wsb
|
comfy_mtb.wsb
|
||||||
Dockerfile
|
Dockerfile
|
||||||
|
|
||||||
|
.DS_Store
|
||||||
|
node.zip
|
||||||
|
|
||||||
# I store the gh-pages worktrees (src & build) there
|
# I store the gh-pages worktrees (src & build) there
|
||||||
.worktrees
|
.worktrees
|
||||||
|
comfy.lock
|
||||||
|
|||||||
@@ -0,0 +1,52 @@
|
|||||||
|
# Code of Conduct
|
||||||
|
|
||||||
|
## Our Commitment
|
||||||
|
|
||||||
|
We are committed to creating a welcoming and inclusive community for everyone. We believe that a diverse and respectful community is essential for fostering creativity and innovation. We expect all members of our community to adhere to this Code of Conduct.
|
||||||
|
|
||||||
|
## Our Expectations
|
||||||
|
|
||||||
|
This Code of Conduct applies to all interactions within the mtb community, including:
|
||||||
|
|
||||||
|
* Public communication channels (e.g., GitHub issues, pull requests, discussions, social media)
|
||||||
|
* Private communication channels (e.g., direct messages, email)
|
||||||
|
* In-person events (if any)
|
||||||
|
|
||||||
|
We expect all members to:
|
||||||
|
|
||||||
|
* **Be respectful and considerate:** Treat others with kindness and empathy.
|
||||||
|
* **Be inclusive:** Welcome and respect people of all backgrounds, identities, and experiences.
|
||||||
|
* **Be constructive:** Focus on providing helpful and positive feedback.
|
||||||
|
* **Be mindful of your language:** Avoid using offensive, discriminatory, or harassing language.
|
||||||
|
* **Respect privacy:** Do not share personal information without consent.
|
||||||
|
|
||||||
|
## Unacceptable Behavior
|
||||||
|
|
||||||
|
The following behaviors are not tolerated:
|
||||||
|
|
||||||
|
* Offensive, discriminatory, or harassing language or conduct
|
||||||
|
* Personal attacks or insults
|
||||||
|
* Spamming or trolling
|
||||||
|
* Sharing of malicious or inappropriate content
|
||||||
|
* Disrupting the community or hindering collaboration
|
||||||
|
* Violating the privacy of others
|
||||||
|
|
||||||
|
## Reporting Violations
|
||||||
|
|
||||||
|
If you experience or witness a violation of this Code of Conduct, please report it to @melmass. All reports will be treated confidentially and investigated promptly.
|
||||||
|
|
||||||
|
## Enforcement
|
||||||
|
|
||||||
|
Violations of this Code of Conduct may result in the following actions:
|
||||||
|
|
||||||
|
* Warning
|
||||||
|
* Removal from the community
|
||||||
|
* Ban from the community
|
||||||
|
|
||||||
|
## License
|
||||||
|
[](code_of_conduct.md)
|
||||||
|
|
||||||
|
## Contact
|
||||||
|
|
||||||
|
If you have any questions or concerns about this Code of Conduct, please contact @melmass.
|
||||||
|
|
||||||
@@ -0,0 +1,62 @@
|
|||||||
|
# Contributing to mtb
|
||||||
|
|
||||||
|
Thank you for your interest in contributing to mtb! We appreciate your help in making this project better. This document outlines how you can contribute to the project.
|
||||||
|
|
||||||
|
## Project Overview
|
||||||
|
|
||||||
|
This project is a collection of custom nodes for ComfyUI, tailored specifically for animation workflows. It aims to provide a streamlined and user-friendly experience for creating animations within the ComfyUI environment.
|
||||||
|
|
||||||
|
## Ways to Contribute
|
||||||
|
|
||||||
|
We welcome all kinds of contributions! Here's how you can get involved:
|
||||||
|
|
||||||
|
* **Bug Reports:** If you encounter any issues, please create a new issue on GitHub. Please include clear steps to reproduce the bug, along with any relevant error messages, workflows or screenshots.
|
||||||
|
* **Feature Requests:** Have an idea for a new node or feature? Create a new issue to discuss it! Please describe the feature in detail, and explain how it would benefit the project.
|
||||||
|
* **Documentation Improvements:** Help us improve the documentation by fixing errors, adding examples, or clarifying explanations.
|
||||||
|
* **Code Contributions:** We welcome contributions to the codebase! Please see the "Development Setup" and "File Structure" sections below for more information.
|
||||||
|
* **Testing:** Help us ensure the stability and reliability of the project by testing new features and bug fixes.
|
||||||
|
* **Refactoring:** Help us improve the codebase by refactoring existing code to improve readability, maintainability, and performance.
|
||||||
|
|
||||||
|
## Development Setup
|
||||||
|
|
||||||
|
```sh
|
||||||
|
git clone --recursive https://github.com/melmass/comfy_mtb
|
||||||
|
```
|
||||||
|
|
||||||
|
## File Structure
|
||||||
|
|
||||||
|
Understanding the project structure is crucial for making effective contributions.
|
||||||
|
|
||||||
|
* **`./nodes/*.py`:** This directory contains the definitions for all custom nodes. Nodes are automatically registered when a file defines an array named `__nodes__` containing the node classes. Make sure your node follows the ComfyUI node definition structure.
|
||||||
|
* **`./web/*.js`:** This directory contains all the frontend JavaScript code for the extension's user interface.
|
||||||
|
* **`./wiki`:** This directory is a Git submodule that contains the project's Wiki documentation, written in Markdown. Node documentation should be created or updated in the corresponding Markdown files within this submodule. This is then referenced by the UI for in-GUI help
|
||||||
|
|
||||||
|
## Coding Style
|
||||||
|
|
||||||
|
We use **Ruff** for code formatting to ensure consistency. Please run Ruff on your code before submitting a pull request. No specific configuration is required, so the default Ruff settings will be used.
|
||||||
|
|
||||||
|
## Contribution Workflow
|
||||||
|
|
||||||
|
1. **Create a Branch:** Create a new branch for your feature or fix. Use a descriptive branch name (e.g., `feature/new-node`, `fix/bug-in-ui`). **Do not fork the main branch directly.**
|
||||||
|
2. **Make Changes:** Implement your changes in your branch.
|
||||||
|
3. **Run Tests:** (Add instructions on how to run tests if available.)
|
||||||
|
4. **Format Code:** Run Ruff on your code to ensure it is properly formatted.
|
||||||
|
5. **Create a Pull Request:** Submit a pull request to the `main` branch. Please provide a clear and concise description of your changes.
|
||||||
|
|
||||||
|
## Code of Conduct
|
||||||
|
|
||||||
|
We are committed to creating a welcoming and inclusive community. We expect all contributors to adhere to a respectful and professional code of conduct. (Consider adding a link to a CODE_OF_CONDUCT.md file or a standard code of conduct.)
|
||||||
|
|
||||||
|
## Tools and Libraries
|
||||||
|
|
||||||
|
* **Python:** The primary programming language for this project.
|
||||||
|
* **ComfyUI:** The underlying framework for the custom nodes.
|
||||||
|
|
||||||
|
## Current Focus
|
||||||
|
|
||||||
|
We are currently focused on a major refactor to clean up the project's codebase. Contributions related to this effort are particularly welcome!
|
||||||
|
|
||||||
|
## Thank You!
|
||||||
|
|
||||||
|
Thank you for considering contributing to mtb! Your contributions are greatly appreciated. We look forward to reviewing your pull requests!
|
||||||
|
|
||||||
+33
-11
@@ -3,11 +3,11 @@
|
|||||||
# File: __init__.py
|
# File: __init__.py
|
||||||
# Project: comfy_mtb
|
# Project: comfy_mtb
|
||||||
# Author: Mel Massadian
|
# Author: Mel Massadian
|
||||||
# Copyright (c) 2023 Mel Massadian
|
# Copyright (c) 2023-2025 Mel Massadian
|
||||||
#
|
#
|
||||||
###
|
###
|
||||||
|
|
||||||
__version__ = "0.3.0"
|
__version__ = "0.6.0"
|
||||||
|
|
||||||
import os
|
import os
|
||||||
|
|
||||||
@@ -34,6 +34,8 @@ from aiohttp import web
|
|||||||
|
|
||||||
IN_COMFY = False
|
IN_COMFY = False
|
||||||
|
|
||||||
|
PromptServer = None
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from server import PromptServer
|
from server import PromptServer
|
||||||
|
|
||||||
@@ -75,7 +77,7 @@ def extract_nodes_from_source(filename: Path):
|
|||||||
)
|
)
|
||||||
break
|
break
|
||||||
except SyntaxError:
|
except SyntaxError:
|
||||||
log.error("Failed to parse")
|
log.error(f"Failed to parse ast from: {filename}")
|
||||||
return nodes
|
return nodes
|
||||||
|
|
||||||
|
|
||||||
@@ -240,10 +242,33 @@ if failed:
|
|||||||
# - ENDPOINT
|
# - ENDPOINT
|
||||||
|
|
||||||
|
|
||||||
if IN_COMFY and hasattr(PromptServer, "instance"):
|
# TODO: move that away and simplify existing endpoints
|
||||||
|
|
||||||
|
|
||||||
|
def register_routes():
|
||||||
|
if not PromptServer:
|
||||||
|
log.error("No prompt server, are you inside comfy?")
|
||||||
|
|
||||||
|
if PromptServer.instance.app.frozen:
|
||||||
|
log.warning(
|
||||||
|
"The router is frozen and cannot be further edited."
|
||||||
|
"If you are hot reloading mtb this is expected."
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
img_cache = None
|
img_cache = None
|
||||||
prompt_cache = None
|
prompt_cache = None
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import os
|
||||||
|
from io import BytesIO
|
||||||
|
|
||||||
|
from PIL import Image
|
||||||
|
|
||||||
|
from .repl import setup_custom_web_routes
|
||||||
|
|
||||||
|
setup_custom_web_routes(PromptServer.instance.app)
|
||||||
|
|
||||||
with contextlib.suppress(ImportError):
|
with contextlib.suppress(ImportError):
|
||||||
from cachetools import TTLCache
|
from cachetools import TTLCache
|
||||||
|
|
||||||
@@ -360,13 +385,6 @@ if IN_COMFY and hasattr(PromptServer, "instance"):
|
|||||||
# Return JSON for other requests
|
# Return JSON for other requests
|
||||||
return web.json_response({"message": "Welcome to MTB!"})
|
return web.json_response({"message": "Welcome to MTB!"})
|
||||||
|
|
||||||
import asyncio
|
|
||||||
import os
|
|
||||||
from io import BytesIO
|
|
||||||
|
|
||||||
from aiohttp import web
|
|
||||||
from PIL import Image
|
|
||||||
|
|
||||||
def get_cached_image(file_path: str, preview_params=None, channel=None):
|
def get_cached_image(file_path: str, preview_params=None, channel=None):
|
||||||
cache_key = (file_path, preview_params, channel)
|
cache_key = (file_path, preview_params, channel)
|
||||||
if img_cache and (cache_key in img_cache):
|
if img_cache and (cache_key in img_cache):
|
||||||
@@ -571,6 +589,10 @@ if IN_COMFY and hasattr(PromptServer, "instance"):
|
|||||||
return await endpoint.do_action(request)
|
return await endpoint.do_action(request)
|
||||||
|
|
||||||
|
|
||||||
|
if IN_COMFY and hasattr(PromptServer, "instance"):
|
||||||
|
register_routes()
|
||||||
|
|
||||||
|
|
||||||
# - WAS Dictionary
|
# - WAS Dictionary
|
||||||
MANIFEST = {
|
MANIFEST = {
|
||||||
"name": "MTB Nodes", # The title that will be displayed on Node Class menu,. and Node Class view
|
"name": "MTB Nodes", # The title that will be displayed on Node Class menu,. and Node Class view
|
||||||
|
|||||||
+13
-6
@@ -1,19 +1,26 @@
|
|||||||
{
|
{
|
||||||
"$schema": "https://biomejs.dev/schemas/1.6.1/schema.json",
|
"$schema": "https://biomejs.dev/schemas/2.0.5/schema.json",
|
||||||
"organizeImports": {
|
"assist": { "actions": { "source": { "organizeImports": "on" } } },
|
||||||
"enabled": true
|
|
||||||
},
|
|
||||||
"linter": {
|
"linter": {
|
||||||
"enabled": true,
|
"enabled": true,
|
||||||
"rules": {
|
"rules": {
|
||||||
"recommended": true,
|
"recommended": true,
|
||||||
"suspicious": {
|
"suspicious": {
|
||||||
"noConsoleLog": "warn"
|
"noConsole": { "level": "warn", "options": { "allow": ["log"] } }
|
||||||
},
|
},
|
||||||
"style": {
|
"style": {
|
||||||
"noParameterAssign": "off",
|
"noParameterAssign": "off",
|
||||||
"noShoutyConstants": "warn",
|
"noShoutyConstants": "warn",
|
||||||
"useNamingConvention": "off"
|
"useNamingConvention": "off",
|
||||||
|
"useAsConstAssertion": "error",
|
||||||
|
"useDefaultParameterLast": "error",
|
||||||
|
"useEnumInitializers": "error",
|
||||||
|
"useSelfClosingElements": "error",
|
||||||
|
"useSingleVarDeclarator": "error",
|
||||||
|
"noUnusedTemplateLiteral": "error",
|
||||||
|
"useNumberNamespace": "error",
|
||||||
|
"noInferrableTypes": "error",
|
||||||
|
"noUselessElse": "error"
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
|
|||||||
+9
-6
@@ -15,7 +15,6 @@ from .utils import (
|
|||||||
backup_file,
|
backup_file,
|
||||||
build_glob_patterns,
|
build_glob_patterns,
|
||||||
glob_multiple,
|
glob_multiple,
|
||||||
import_install,
|
|
||||||
reqs_map,
|
reqs_map,
|
||||||
run_command,
|
run_command,
|
||||||
styles_dir,
|
styles_dir,
|
||||||
@@ -24,7 +23,6 @@ from .utils import (
|
|||||||
endlog = mklog("mtb endpoint")
|
endlog = mklog("mtb endpoint")
|
||||||
|
|
||||||
# - ACTIONS
|
# - ACTIONS
|
||||||
import_install("requirements")
|
|
||||||
|
|
||||||
|
|
||||||
def ACTIONS_installDependency(dependency_names: list[str] | None = None):
|
def ACTIONS_installDependency(dependency_names: list[str] | None = None):
|
||||||
@@ -112,11 +110,15 @@ def ACTIONS_getUserVideos(
|
|||||||
|
|
||||||
def ACTIONS_getUserImages(
|
def ACTIONS_getUserImages(
|
||||||
mode: Literal["input", "output"],
|
mode: Literal["input", "output"],
|
||||||
|
target_width: int | str | None = None,
|
||||||
count=1000,
|
count=1000,
|
||||||
offset=0,
|
offset=0,
|
||||||
sort: str | None = None,
|
sort: str | None = None,
|
||||||
include_subfolders: bool = False,
|
include_subfolders: bool = False,
|
||||||
subfolder=None,
|
subfolder: str | None = None,
|
||||||
|
# IIRC I copied this from Comfy base
|
||||||
|
# just keeping it until I properly checked implications
|
||||||
|
salt_urls=False,
|
||||||
):
|
):
|
||||||
# enabled = "MTB_EXPOSE" in os.environ
|
# enabled = "MTB_EXPOSE" in os.environ
|
||||||
# if not enabled:
|
# if not enabled:
|
||||||
@@ -124,11 +126,12 @@ def ACTIONS_getUserImages(
|
|||||||
|
|
||||||
imgs = {}
|
imgs = {}
|
||||||
count = count or 1000
|
count = count or 1000
|
||||||
|
target_width = int(target_width) if target_width else None
|
||||||
|
|
||||||
input_dir = Path(folder_paths.get_input_directory())
|
input_dir = Path(folder_paths.get_input_directory())
|
||||||
output_dir = Path(folder_paths.get_output_directory())
|
output_dir = Path(folder_paths.get_output_directory())
|
||||||
|
|
||||||
entry_dir = input_dir if mode == "input" else output_dir
|
entry_dir: Path = input_dir if mode == "input" else output_dir
|
||||||
if subfolder:
|
if subfolder:
|
||||||
entry_dir = entry_dir / subfolder
|
entry_dir = entry_dir / subfolder
|
||||||
|
|
||||||
@@ -157,9 +160,9 @@ def ACTIONS_getUserImages(
|
|||||||
|
|
||||||
imgs = {
|
imgs = {
|
||||||
img.name: (
|
img.name: (
|
||||||
f"/mtb/view?filename={img.name}&width=512&type={mode}&subfolder={subfolder or ''}"
|
f"/mtb/view?filename={img.name}{f'&width={target_width}' if target_width and target_width > 0 else ''}&type={mode}&subfolder={subfolder or ''}"
|
||||||
f"{img.parent.relative_to(entry_dir) if include_subfolders else ''}"
|
f"{img.parent.relative_to(entry_dir) if include_subfolders else ''}"
|
||||||
f"&preview=&rand={secrets.randbelow(424242)}"
|
f"&preview={f'&rand={secrets.randbelow(424242)}' if salt_urls else ''}"
|
||||||
)
|
)
|
||||||
for i, img in enumerate(entries)
|
for i, img in enumerate(entries)
|
||||||
if offset <= i < offset + count
|
if offset <= i < offset + count
|
||||||
|
|||||||
@@ -1,85 +1,175 @@
|
|||||||
# NOTE: This file is only use for development you can ignore it
|
# NOTE: This file is only use for development you can ignore it
|
||||||
|
|
||||||
use private/log.nu
|
use log.nu
|
||||||
|
use nssm.nu *
|
||||||
|
use nutils.nu [ make-id upsert-all fwd-slash backup-file ]
|
||||||
|
use os.nu [ link ]
|
||||||
|
|
||||||
def get_root [--clean] {
|
# --- utilities ---
|
||||||
if $clean {
|
def get_root [ --clean] {
|
||||||
$env.COMFY_CLEAN_ROOT
|
if $clean {
|
||||||
} else {
|
$env.COMFY.ROOTS.clean
|
||||||
$env.COMFY_ROOT
|
} else {
|
||||||
}
|
$env.COMFY.ROOTS.main
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
export def "comfy build-web" [] {
|
def --env path-add [pth] {
|
||||||
cd $env.COMFY_MTB
|
$env.PATH = ($env.PATH | append ($pth | path expand))
|
||||||
cd web_source
|
|
||||||
npm run build
|
|
||||||
cp dist/*.js ../web/dist
|
|
||||||
}
|
|
||||||
|
|
||||||
export def "comfy dev-web" [] {
|
|
||||||
cd $env.COMFY_MTB
|
|
||||||
cd web_source
|
|
||||||
npm run dev
|
|
||||||
}
|
|
||||||
|
|
||||||
export def "daily run" [] {
|
|
||||||
let res = (comfy update --rebase)
|
|
||||||
comfy update --clean
|
|
||||||
comfy update_extensions
|
|
||||||
|
|
||||||
daily commit $res.from $res.to
|
|
||||||
}
|
}
|
||||||
|
|
||||||
def short-date [] {
|
def short-date [] {
|
||||||
format date "%Y-%m-%d"
|
format date "%Y-%m-%d"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
export def spawn-for [timeout: duration task: closure] {
|
||||||
|
let input = $in
|
||||||
|
let parent_id = job id
|
||||||
|
let task_id = job spawn {
|
||||||
|
$input | do $task | job send --tag (job id) $parent_id
|
||||||
|
}
|
||||||
|
try {
|
||||||
|
job recv --tag $task_id --timeout $timeout
|
||||||
|
} catch {
|
||||||
|
job kill $task_id
|
||||||
|
error make {
|
||||||
|
msg: "Task timed out."
|
||||||
|
label: {
|
||||||
|
text: "timed out"
|
||||||
|
span: (metadata $task).span
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
# --- exports --
|
||||||
|
export def "comfy profile" [timeout = 60sec] {
|
||||||
|
let to_match = "To see the GUI go to"
|
||||||
|
|
||||||
|
pyinstrument -r html main.py ...($env.COMFY.ARGS)
|
||||||
|
| tee -e {
|
||||||
|
each {
|
||||||
|
let stde = $in
|
||||||
|
print -ne $stde
|
||||||
|
if $to_match in $stde {
|
||||||
|
print $"(ansi gb)Profiling Done!(ansi reset)"
|
||||||
|
let process = (ps -l | where name =~ python | where command =~ pyinstrument | last)
|
||||||
|
kill -f $process.pid
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
| complete
|
||||||
|
| get stdout
|
||||||
|
| save $"profiled_(date now | format date '%s').html"
|
||||||
|
}
|
||||||
|
|
||||||
|
export def "comfy profile-plus" [] {
|
||||||
|
|
||||||
|
let timestamp = (date now | format date "%s")
|
||||||
|
let log_name = $"cprofile_run_($timestamp)"
|
||||||
|
let profiled = (python -m cProfile main.py --port 3000 --preview-method auto | tee -e { print -ne } | complete)
|
||||||
|
|
||||||
|
let out = (
|
||||||
|
$profiled.stdout
|
||||||
|
| lines
|
||||||
|
# skip summary
|
||||||
|
| skip 4
|
||||||
|
| str join "\n"
|
||||||
|
)
|
||||||
|
# save result
|
||||||
|
$out | save $"raw_($log_name).txt"
|
||||||
|
|
||||||
|
# process
|
||||||
|
$out
|
||||||
|
| from ssv
|
||||||
|
| upsert-all { into float } tottime percall cumtime
|
||||||
|
| save $"($log_name).nuon"
|
||||||
|
}
|
||||||
|
|
||||||
|
export def restart-server [] {
|
||||||
|
nssm restart -c comfy
|
||||||
|
}
|
||||||
|
|
||||||
|
# build the web components of mtb
|
||||||
|
export def "comfy build-web" [] {
|
||||||
|
cd $env.COMFY.ROOTS.mtb
|
||||||
|
if ("./web/dist" | path exists) {
|
||||||
|
rm -rt ./web/dist
|
||||||
|
}
|
||||||
|
|
||||||
|
cd web_source
|
||||||
|
^$env.NPM_BINARY run build
|
||||||
|
cp -r dist ../web/dist
|
||||||
|
}
|
||||||
|
|
||||||
|
# start the dev server for web components
|
||||||
|
export def "comfy dev-web" [] {
|
||||||
|
cd $env.COMFY.ROOTS.mtb
|
||||||
|
cd web_source
|
||||||
|
^$env.NPM_BINARY run dev
|
||||||
|
}
|
||||||
|
|
||||||
|
# daily check / update
|
||||||
|
export def "daily run" [] {
|
||||||
|
let res = (comfy update --rebase)
|
||||||
|
comfy update --clean
|
||||||
|
comfy update_extensions
|
||||||
|
|
||||||
|
daily commit $res.from_commit $res.to_commit
|
||||||
|
}
|
||||||
|
|
||||||
# was daily run today?
|
# was daily run today?
|
||||||
export def "daily was-run" [] {
|
export def "daily was-run" [] {
|
||||||
|
|
||||||
let daily = ($env.COMFY_MTB | path join daily.nuon)
|
let daily = ($env.COMFY.ROOTS.mtb | path join daily.nuon)
|
||||||
|
|
||||||
if ($daily | path exists) {
|
if ($daily | path exists) {
|
||||||
let last = (open $daily | sort-by date | get date | last | short-date)
|
let last = (open $daily | sort-by date | get date | last | short-date)
|
||||||
let today = (date now | short-date)
|
let today = (date now | short-date)
|
||||||
return ($last == $today)
|
return ($last == $today)
|
||||||
}
|
}
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
export def "daily commit" [from:string, to:string] {
|
export def "daily commit" [from_commit: string to_commit: string] {
|
||||||
let daily = ($env.COMFY_MTB | path join daily.nuon)
|
let daily = ($env.COMFY.ROOTS.mtb | path join daily.nuon)
|
||||||
let commit = [{date: (date now) from:$from to:$to}]
|
let commit = [{date: (date now) from_commit: $from_commit to_commit: $to_commit}]
|
||||||
|
|
||||||
let dailies = (if ($daily | path exists) {
|
let dailies = (
|
||||||
open $daily | append $commit
|
if ($daily | path exists) {
|
||||||
} else {
|
open $daily | append $commit
|
||||||
|
} else {
|
||||||
$commit
|
$commit
|
||||||
})
|
}
|
||||||
|
)
|
||||||
|
|
||||||
$dailies | save -f $daily
|
$dailies | save -f $daily
|
||||||
log success "Commited daily check"
|
log success "Commited daily check"
|
||||||
}
|
}
|
||||||
|
|
||||||
# start the comfy server
|
# start the comfy server
|
||||||
export def "comfy start" [--clean,--old-ui, --listen, --skip-daily(-s)] {
|
export def "comfy start" [
|
||||||
if (not (daily was-run)) and not $skip_daily {
|
--clean
|
||||||
log info "Running daily checks"
|
--old-ui
|
||||||
daily run
|
--listen
|
||||||
}
|
--skip-daily (-s)
|
||||||
let root = get_root --clean=($clean)
|
] {
|
||||||
cd $root
|
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
|
# update comfy itself and merge master in current branch
|
||||||
export def "comfy update" [
|
export def "comfy update" [
|
||||||
--clean # ??
|
--clean # comfy clean instance
|
||||||
--rebase # Rebase instead of merge
|
--rebase # Rebase instead of merge
|
||||||
] {
|
] {
|
||||||
let root = get_root --clean=$clean
|
let root = get_root --clean=$clean
|
||||||
|
|
||||||
@@ -94,21 +184,28 @@ export def "comfy update" [
|
|||||||
log info "Backing up and removing models symlinks"
|
log info "Backing up and removing models symlinks"
|
||||||
|
|
||||||
# preparing root for pull
|
# preparing root for pull
|
||||||
if not $clean {
|
let pyproject = if not $clean {
|
||||||
|
log info "Backing up the pyproject.toml..."
|
||||||
|
let proj = (backup-file --root pyproject.toml)
|
||||||
|
|
||||||
|
log info "Restoring the original pyproject"
|
||||||
git checkout pyproject.toml
|
git checkout pyproject.toml
|
||||||
cd $models
|
cd $models
|
||||||
# find and store all symlinks
|
# find and store all symlinks
|
||||||
let links = (ls -la |
|
log info "Checking for links in models..."
|
||||||
where not ($it.target | is-empty) |
|
let links = (
|
||||||
select name target |
|
ls -la | where not ($it.target | is-empty) | select name target | sort-by name
|
||||||
sort-by name)
|
)
|
||||||
|
log info $"Found links: ($links)"
|
||||||
|
|
||||||
if not ($links | is-empty) {
|
if not ($links | is-empty) {
|
||||||
|
log info "Backing up the symlinks..."
|
||||||
|
backup-file --root links.nuon
|
||||||
$links | save -f links.nuon
|
$links | save -f links.nuon
|
||||||
# remove them
|
# remove them
|
||||||
open links.nuon | each {|p| rm $p.name }
|
open links.nuon | each {|p| rm $p.name }
|
||||||
}
|
}
|
||||||
|
$proj
|
||||||
} else {
|
} else {
|
||||||
# just remove symlinks
|
# just remove symlinks
|
||||||
rm $models
|
rm $models
|
||||||
@@ -141,7 +238,6 @@ export def "comfy update" [
|
|||||||
if $rebase {
|
if $rebase {
|
||||||
log info "Rebasing changes"
|
log info "Rebasing changes"
|
||||||
git rebase master
|
git rebase master
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
log info "Merging changes"
|
log info "Merging changes"
|
||||||
git merge master
|
git merge master
|
||||||
@@ -152,9 +248,10 @@ export def "comfy update" [
|
|||||||
|
|
||||||
if not $clean {
|
if not $clean {
|
||||||
rm pyproject.toml
|
rm pyproject.toml
|
||||||
cp pyproject-mel.toml pyproject.toml
|
log info "Using our own pyproject..."
|
||||||
|
cp $pyproject pyproject.toml
|
||||||
cd $models
|
cd $models
|
||||||
|
log info "Relinking models..."
|
||||||
# resymlink them
|
# resymlink them
|
||||||
open links.nuon | each {|p| link -a $p.target $p.name }
|
open links.nuon | each {|p| link -a $p.target $p.name }
|
||||||
} else {
|
} else {
|
||||||
@@ -167,72 +264,82 @@ export def "comfy update" [
|
|||||||
|
|
||||||
log success $"Update successful \(($commit_count) new commits\)"
|
log success $"Update successful \(($commit_count) new commits\)"
|
||||||
|
|
||||||
return {from:$current_commit to:$new_commit}
|
return {from_commit: $current_commit to_commit: $new_commit}
|
||||||
|
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
export def "comfy toggle_extensions" [--clean] {
|
export def "comfy toggle_extensions" [
|
||||||
let root = get_root --clean=($clean)
|
--clean
|
||||||
cd $root
|
] {
|
||||||
cd custom_nodes
|
let root = get_root --clean=$clean
|
||||||
let exts = (ls | where type in ["dir","symlink"] | get name)
|
cd $root
|
||||||
let choices = ($exts | input list -m "choose extension to toggle")
|
cd custom_nodes
|
||||||
if ($choices | is-empty) {
|
let exts = (ls | where type in ["dir" "symlink"] | get name)
|
||||||
return
|
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
|
log info "Filtered" $filtered
|
||||||
$filtered | each {|f|
|
$filtered | each {|f|
|
||||||
let new_name = ($f.name | str replace ".disabled" "")
|
let new_name = ($f.name | str replace ".disabled" "")
|
||||||
|
|
||||||
let new_name = if $f.enabled {
|
let new_name = if $f.enabled {
|
||||||
$"($new_name).disabled"
|
$"($new_name).disabled"
|
||||||
} else {
|
} else {
|
||||||
$new_name
|
$new_name
|
||||||
}
|
|
||||||
log info $"Moving ($f.name) to ($new_name)"
|
|
||||||
mv $f.name $new_name
|
|
||||||
}
|
}
|
||||||
|
log info $"Moving ($f.name) to ($new_name)"
|
||||||
|
mv $f.name $new_name
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
# git pull all extensions
|
# git pull all extensions
|
||||||
export def "comfy update_extensions" [--clean] {
|
export def "comfy update_extensions" [ --clean] {
|
||||||
let root = get_root --clean=($clean)
|
let root = get_root --clean=$clean
|
||||||
cd $root
|
cd $root
|
||||||
cd custom_nodes
|
cd custom_nodes
|
||||||
git multipull . -s -q
|
git multipull . -s -q
|
||||||
}
|
}
|
||||||
|
|
||||||
def --env path-add [pth] {
|
# manual set version of mtb
|
||||||
$env.PATH = ($env.PATH | append ($pth | path expand))
|
export def "comfy-mtb set-version" [version: string] {
|
||||||
|
# let pyproject = open pyproject.toml
|
||||||
|
# let current_version = $pyproject.project.version
|
||||||
|
# $pyproject | upsert project.version $version | save -f pyproject.toml
|
||||||
|
# taplo format pyproject.toml
|
||||||
|
sd "(__version__ = )\"(.*)\"" $"${1}\"($version)\"" __init__.py
|
||||||
|
sd "(version = )(.*)" $"${1}\"($version)\"" pyproject.toml
|
||||||
|
# log info $"⬆️ Bump version: ($current_version) → ($version)"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
# -- env
|
||||||
export-env {
|
export-env {
|
||||||
$env.PYTHONUTF8 = 1
|
$env.PYTHONUTF8 = 1
|
||||||
$env.COMFY_MTB = ("." | path expand)
|
$env.COMFY = {
|
||||||
# $env.CUDA_ROOT = 'C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.1\'
|
base_url : "https://mel-pc.tail3c8eb.ts.net"
|
||||||
|
ARGS: [--port 3000 --preview-method auto]
|
||||||
$env.CUDA_HOME = $env.CUDA_ROOT
|
ROOTS: {
|
||||||
|
mtb: ("." | path expand | fwd-slash)
|
||||||
$env.COMFY_ROOT = ("../.." | path expand)
|
main: ("../.." | path expand | fwd-slash)
|
||||||
$env.COMFY_CLEAN_ROOT = ($env.COMFY_ROOT | path dirname | path join ComfyClean)
|
clean: ($env.COMFY_ROOT | path dirname | path join ComfyClean | fwd-slash)
|
||||||
|
}
|
||||||
path-add 'C:/Portable/TensorRT-8.6.0.12/lib'
|
|
||||||
|
|
||||||
if $nu.os-info.family == 'windows' {
|
|
||||||
path-add 'G:\BIN\TensorRT-10.7.0.23\lib'
|
|
||||||
path-add 'G:\BIN\cudnn-windows-x86_64-9.6.0.74_cuda12-archive\bin'
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
$env.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)
|
path-add ($env.CUDA_ROOT | path join bin)
|
||||||
overlay use ../../.venv/Scripts/activate.nu
|
|
||||||
|
overlay use "../../.venv/Scripts/activate.nu"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,3 @@
|
|||||||
|
{
|
||||||
|
"use_repl": false
|
||||||
|
}
|
||||||
@@ -43,7 +43,6 @@ pip_map = {
|
|||||||
"tb-nightly": "tensorboard",
|
"tb-nightly": "tensorboard",
|
||||||
"protobuf": "google.protobuf",
|
"protobuf": "google.protobuf",
|
||||||
"qrcode[pil]": "qrcode",
|
"qrcode[pil]": "qrcode",
|
||||||
"requirements-parser": "requirements",
|
|
||||||
# Add more mappings as needed
|
# Add more mappings as needed
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+14
-13
@@ -1,20 +1,16 @@
|
|||||||
from typing import Any, TypedDict
|
from typing import TYPE_CHECKING, Any, TypedDict
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torchaudio
|
import torchaudio
|
||||||
from comfy.model_management import get_torch_device
|
from comfy.model_management import get_torch_device
|
||||||
from huggingface_hub import snapshot_download
|
from huggingface_hub import snapshot_download
|
||||||
from transformers import (
|
|
||||||
WhisperForConditionalGeneration,
|
|
||||||
WhisperProcessor,
|
|
||||||
)
|
|
||||||
|
|
||||||
# from transformers import (
|
if TYPE_CHECKING:
|
||||||
# AutoFeatureExtractor,
|
from transformers import (
|
||||||
# WhisperForConditionalGeneration,
|
WhisperForConditionalGeneration,
|
||||||
# WhisperModel,
|
WhisperProcessor,
|
||||||
# WhisperProcessor,
|
)
|
||||||
# )
|
|
||||||
from ..log import log
|
from ..log import log
|
||||||
from ..utils import get_model_path
|
from ..utils import get_model_path
|
||||||
|
|
||||||
@@ -101,8 +97,8 @@ class MtbAudio:
|
|||||||
class WhisperPipeline(TypedDict):
|
class WhisperPipeline(TypedDict):
|
||||||
"""Whisper model pipeline."""
|
"""Whisper model pipeline."""
|
||||||
|
|
||||||
processor: WhisperProcessor
|
processor: "WhisperProcessor"
|
||||||
model: WhisperForConditionalGeneration
|
model: "WhisperForConditionalGeneration"
|
||||||
|
|
||||||
|
|
||||||
class MTB_LoadWhisper:
|
class MTB_LoadWhisper:
|
||||||
@@ -148,6 +144,11 @@ class MTB_LoadWhisper:
|
|||||||
|
|
||||||
def load(self, model_size="tiny", download_missing=False):
|
def load(self, model_size="tiny", download_missing=False):
|
||||||
"""Load Whisper model and processor."""
|
"""Load Whisper model and processor."""
|
||||||
|
from transformers import (
|
||||||
|
WhisperForConditionalGeneration,
|
||||||
|
WhisperProcessor,
|
||||||
|
)
|
||||||
|
|
||||||
whisper_dir = get_model_path("whisper")
|
whisper_dir = get_model_path("whisper")
|
||||||
tag = f"whisper-{model_size}"
|
tag = f"whisper-{model_size}"
|
||||||
model_dir = whisper_dir / tag
|
model_dir = whisper_dir / tag
|
||||||
|
|||||||
+190
@@ -0,0 +1,190 @@
|
|||||||
|
import time
|
||||||
|
import uuid
|
||||||
|
from collections import OrderedDict
|
||||||
|
from typing import Any, TypedDict
|
||||||
|
|
||||||
|
from comfy.comfy_types.node_typing import IO as CIO
|
||||||
|
from server import PromptServer
|
||||||
|
|
||||||
|
from ..log import log
|
||||||
|
|
||||||
|
|
||||||
|
class Clock(TypedDict):
|
||||||
|
name: str
|
||||||
|
start: float
|
||||||
|
end: float | None
|
||||||
|
|
||||||
|
|
||||||
|
active_timers: OrderedDict[str, Clock] = OrderedDict()
|
||||||
|
|
||||||
|
# TODO: lower this
|
||||||
|
MAX_CLOCKS = 50
|
||||||
|
|
||||||
|
|
||||||
|
class MTB_StartClock:
|
||||||
|
"""
|
||||||
|
Starts a profiling clock with a given name.
|
||||||
|
|
||||||
|
Outputs a unique ID that must be passed to EndClock.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
pass
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"name": ("STRING", {"default": "Clock A"}),
|
||||||
|
"cache": (
|
||||||
|
"BOOLEAN",
|
||||||
|
{
|
||||||
|
"default": False,
|
||||||
|
"tooltip": "Cache the clock ID, this means the node will follow Comfy's default invalidation system. If False it will always invalidate / mark the node as 'dirty'",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"passthrough": (CIO.ANY,),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = (
|
||||||
|
CIO.ANY,
|
||||||
|
"STRING",
|
||||||
|
)
|
||||||
|
RETURN_NAMES = (
|
||||||
|
"passthrough",
|
||||||
|
"clock_id",
|
||||||
|
)
|
||||||
|
FUNCTION = "start_timer"
|
||||||
|
CATEGORY = "mtb/utils"
|
||||||
|
|
||||||
|
def start_timer(
|
||||||
|
self, *, name: str, passthrough: Any | None = None, **kwargs
|
||||||
|
):
|
||||||
|
global active_timers
|
||||||
|
|
||||||
|
if len(active_timers) >= MAX_CLOCKS:
|
||||||
|
# get oldest clock
|
||||||
|
removed_key = None
|
||||||
|
for key, clock_data in active_timers.items():
|
||||||
|
if clock_data["end"] is not None:
|
||||||
|
removed_key = key
|
||||||
|
break
|
||||||
|
if removed_key:
|
||||||
|
removed_clock = active_timers.pop(removed_key)
|
||||||
|
log.info(
|
||||||
|
f"[Profiling] Evicted finished clock '{removed_clock['name']}' (ID: {removed_key}) due to limit ({MAX_CLOCKS})."
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
removed_key, removed_clock = active_timers.popitem(last=False)
|
||||||
|
log.warning(
|
||||||
|
f"[Profiling] Evicted running clock '{removed_clock['name']}' (ID: {removed_key}) due to limit ({MAX_CLOCKS})."
|
||||||
|
)
|
||||||
|
|
||||||
|
clock_id = str(uuid.uuid4())
|
||||||
|
start_time = time.perf_counter()
|
||||||
|
|
||||||
|
active_timers[clock_id] = {
|
||||||
|
"start": start_time,
|
||||||
|
"name": name,
|
||||||
|
"end": None,
|
||||||
|
}
|
||||||
|
|
||||||
|
active_timers.move_to_end(clock_id)
|
||||||
|
|
||||||
|
log.debug(f"[Profiling] Clock '{name}' (ID: {clock_id}) started.")
|
||||||
|
|
||||||
|
return (
|
||||||
|
passthrough,
|
||||||
|
clock_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def IS_CHANGED(
|
||||||
|
cls, *, name: str, cache: bool = False, passthrough: Any | None = None
|
||||||
|
):
|
||||||
|
if not cache:
|
||||||
|
return float("Nan")
|
||||||
|
|
||||||
|
return {"name": name, "cache": cache, "passthrough": passthrough}
|
||||||
|
|
||||||
|
|
||||||
|
class MTB_EndClock:
|
||||||
|
"""
|
||||||
|
Stops a profiling clock identified by its ID and returns the elapsed time in milliseconds.
|
||||||
|
|
||||||
|
Errors if the clock ID is not found or already stopped.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
pass
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"clock_id": (
|
||||||
|
"STRING",
|
||||||
|
{"forceInput": True},
|
||||||
|
),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"passthrough": (CIO.ANY,),
|
||||||
|
},
|
||||||
|
"hidden": {
|
||||||
|
"unique_id": "UNIQUE_ID",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = (
|
||||||
|
CIO.ANY,
|
||||||
|
"STRING",
|
||||||
|
"FLOAT",
|
||||||
|
"INT",
|
||||||
|
)
|
||||||
|
RETURN_NAMES = (
|
||||||
|
"passthrough",
|
||||||
|
"name",
|
||||||
|
"seconds",
|
||||||
|
"milliseconds",
|
||||||
|
)
|
||||||
|
FUNCTION = "end_timer"
|
||||||
|
CATEGORY = "mtb/utils"
|
||||||
|
|
||||||
|
def end_timer(self, clock_id: str, passthrough, unique_id=None):
|
||||||
|
global active_timers
|
||||||
|
|
||||||
|
if clock_id not in active_timers:
|
||||||
|
raise ValueError(
|
||||||
|
f"Error: Clock with ID '{clock_id}' not found. "
|
||||||
|
"Ensure StartClock was executed for this ID and proper passthrough chaining."
|
||||||
|
)
|
||||||
|
|
||||||
|
clock = active_timers[clock_id]
|
||||||
|
if clock.get("end") is not None:
|
||||||
|
return (passthrough, clock["name"], clock["end"])
|
||||||
|
|
||||||
|
start_time = clock["start"]
|
||||||
|
end_time = time.perf_counter()
|
||||||
|
|
||||||
|
duration_seconds = end_time - start_time
|
||||||
|
duration_ms = int(duration_seconds * 1000)
|
||||||
|
clock["end"] = duration_ms
|
||||||
|
|
||||||
|
active_timers.move_to_end(clock_id)
|
||||||
|
|
||||||
|
log.debug(
|
||||||
|
f"[Profiling] Clock '{clock['name']}' (ID: {clock_id}) stopped. Elapsed: {duration_ms}ms"
|
||||||
|
)
|
||||||
|
if unique_id:
|
||||||
|
PromptServer.instance.send_progress_text(
|
||||||
|
f"Clock '{clock['name']}' took {duration_seconds:.4f} seconds",
|
||||||
|
unique_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
return (passthrough, clock["name"], duration_seconds, duration_ms)
|
||||||
|
|
||||||
|
|
||||||
|
__nodes__ = [MTB_StartClock, MTB_EndClock]
|
||||||
+194
-181
@@ -1,13 +1,22 @@
|
|||||||
import numpy as np
|
from typing import NamedTuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from PIL import Image, ImageDraw, ImageFilter
|
import torchvision.transforms.functional as TF
|
||||||
|
|
||||||
from ..log import log
|
from ..log import log
|
||||||
from ..utils import np2tensor, pil2tensor, tensor2np, tensor2pil
|
|
||||||
|
|
||||||
|
class BoundingBox(NamedTuple):
|
||||||
|
"""The bounding box tuple."""
|
||||||
|
|
||||||
|
x: int
|
||||||
|
y: int
|
||||||
|
width: int
|
||||||
|
height: int
|
||||||
|
|
||||||
|
|
||||||
class MTB_Bbox:
|
class MTB_Bbox:
|
||||||
"""The bounding box (BBOX) custom type used by other nodes"""
|
"""A literal bounding box."""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(cls):
|
def INPUT_TYPES(cls):
|
||||||
@@ -37,12 +46,14 @@ class MTB_Bbox:
|
|||||||
FUNCTION = "do_crop"
|
FUNCTION = "do_crop"
|
||||||
CATEGORY = "mtb/crop"
|
CATEGORY = "mtb/crop"
|
||||||
|
|
||||||
def do_crop(self, x: int, y: int, width: int, height: int): # bbox
|
def do_crop(
|
||||||
return ((x, y, width, height),)
|
self, x: int, y: int, width: int, height: int
|
||||||
|
) -> tuple[BoundingBox]: # bbox
|
||||||
|
return (BoundingBox(x, y, width, height),)
|
||||||
|
|
||||||
|
|
||||||
class MTB_SplitBbox:
|
class MTB_SplitBbox:
|
||||||
"""Split the components of a bbox"""
|
"""Split the components of a bbox."""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(cls):
|
def INPUT_TYPES(cls):
|
||||||
@@ -55,8 +66,8 @@ class MTB_SplitBbox:
|
|||||||
RETURN_TYPES = ("INT", "INT", "INT", "INT")
|
RETURN_TYPES = ("INT", "INT", "INT", "INT")
|
||||||
RETURN_NAMES = ("x", "y", "width", "height")
|
RETURN_NAMES = ("x", "y", "width", "height")
|
||||||
|
|
||||||
def split_bbox(self, bbox):
|
def split_bbox(self, bbox: BoundingBox) -> BoundingBox:
|
||||||
return (bbox[0], bbox[1], bbox[2], bbox[3])
|
return bbox
|
||||||
|
|
||||||
|
|
||||||
class MTB_UpscaleBboxBy:
|
class MTB_UpscaleBboxBy:
|
||||||
@@ -74,26 +85,23 @@ class MTB_UpscaleBboxBy:
|
|||||||
|
|
||||||
FUNCTION = "upscale"
|
FUNCTION = "upscale"
|
||||||
|
|
||||||
def upscale(
|
def upscale(self, bbox: BoundingBox, scale: float) -> tuple[BoundingBox]:
|
||||||
self, bbox: tuple[int, int, int, int], scale: float
|
|
||||||
) -> tuple[tuple[int, int, int, int]]:
|
|
||||||
x, y, width, height = bbox
|
x, y, width, height = bbox
|
||||||
|
|
||||||
center_x = x + width // 2
|
center_x = x + width / 2
|
||||||
center_y = y + height // 2
|
center_y = y + height / 2
|
||||||
|
|
||||||
new_width = int(width * scale)
|
new_width = int(width * scale)
|
||||||
new_height = int(height * scale)
|
new_height = int(height * scale)
|
||||||
|
|
||||||
new_x = center_x - new_width // 2
|
new_x = int(center_x - new_width / 2)
|
||||||
new_y = center_y - new_height // 2
|
new_y = int(center_y - new_height / 2)
|
||||||
|
|
||||||
scaled = (new_x, new_y, new_width, new_height)
|
return (BoundingBox(new_x, new_y, new_width, new_height),)
|
||||||
return (scaled,)
|
|
||||||
|
|
||||||
|
|
||||||
class MTB_BboxFromMask:
|
class MTB_BboxFromMask:
|
||||||
"""From a mask extract the bounding box"""
|
"""From a mask extract the bounding box."""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(cls):
|
def INPUT_TYPES(cls):
|
||||||
@@ -103,7 +111,7 @@ class MTB_BboxFromMask:
|
|||||||
"invert": ("BOOLEAN", {"default": False}),
|
"invert": ("BOOLEAN", {"default": False}),
|
||||||
},
|
},
|
||||||
"optional": {
|
"optional": {
|
||||||
"image": ("IMAGE",),
|
"image": ("IMAGE", {"tooltip": "Optional image"}),
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -119,52 +127,44 @@ class MTB_BboxFromMask:
|
|||||||
CATEGORY = "mtb/crop"
|
CATEGORY = "mtb/crop"
|
||||||
|
|
||||||
def extract_bounding_box(
|
def extract_bounding_box(
|
||||||
self, mask: torch.Tensor, invert: bool, image=None
|
self,
|
||||||
):
|
mask: torch.Tensor,
|
||||||
# if image != None:
|
*,
|
||||||
# if mask.size(0) != image.size(0):
|
invert: bool = False,
|
||||||
# if mask.size(0) != 1:
|
image: torch.Tensor | None = None,
|
||||||
# log.error(
|
) -> tuple[BoundingBox, torch.Tensor | None]:
|
||||||
# f"Batch count mismatch for mask and image, it can either be 1 mask for X images, or X masks for X images (mask: {mask.shape} | image: {image.shape})"
|
mask = 1 - mask if invert else mask
|
||||||
# )
|
non_zero_indices = torch.nonzero(mask)
|
||||||
|
|
||||||
# raise Exception(
|
if non_zero_indices.numel() == 0:
|
||||||
# f"Batch count mismatch for mask and image, it can either be 1 mask for X images, or X masks for X images (mask: {mask.shape} | image: {image.shape})"
|
log.warning(
|
||||||
# )
|
"BboxFromMask: Mask is empty. Returning a (0,0,0,0) bbox."
|
||||||
|
)
|
||||||
|
return (BoundingBox(0, 0, 0, 0), image)
|
||||||
|
|
||||||
# we invert it
|
min_coords = torch.min(non_zero_indices, dim=0).values
|
||||||
_mask = tensor2pil(1.0 - mask)[0] if invert else tensor2pil(mask)[0]
|
max_coords = torch.max(non_zero_indices, dim=0).values
|
||||||
alpha_channel = np.array(_mask)
|
|
||||||
|
|
||||||
non_zero_indices = np.nonzero(alpha_channel)
|
min_y, min_x = min_coords[1].item(), min_coords[2].item()
|
||||||
|
max_y, max_x = max_coords[1].item(), max_coords[2].item()
|
||||||
|
|
||||||
min_x, max_x = np.min(non_zero_indices[1]), np.max(non_zero_indices[1])
|
width = max_x - min_x + 1
|
||||||
min_y, max_y = np.min(non_zero_indices[0]), np.max(non_zero_indices[0])
|
height = max_y - min_y + 1
|
||||||
|
|
||||||
# Create a bounding box tuple
|
bounding_box = BoundingBox(
|
||||||
if image != None:
|
int(min_x), int(min_y), int(width), int(height)
|
||||||
# Convert the image to a NumPy array
|
|
||||||
imgs = tensor2np(image)
|
|
||||||
out = []
|
|
||||||
for img in imgs:
|
|
||||||
# Crop the image from the bounding box
|
|
||||||
img = img[min_y:max_y, min_x:max_x, :]
|
|
||||||
log.debug(f"Cropped image to shape {img.shape}")
|
|
||||||
out.append(img)
|
|
||||||
|
|
||||||
image = np2tensor(out)
|
|
||||||
log.debug(f"Cropped images shape: {image.shape}")
|
|
||||||
bounding_box = (min_x, min_y, max_x - min_x, max_y - min_y)
|
|
||||||
return (
|
|
||||||
bounding_box,
|
|
||||||
image,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
cropped_image = None
|
||||||
|
if image is not None:
|
||||||
|
cropped_image = image[:, min_y : max_y + 1, min_x : max_x + 1, :]
|
||||||
|
|
||||||
|
return (bounding_box, cropped_image)
|
||||||
|
|
||||||
|
|
||||||
class MTB_Crop:
|
class MTB_Crop:
|
||||||
"""Crops an image and an optional mask to a given bounding box
|
"""Crop an image and an optional mask to a given bounding box.
|
||||||
|
|
||||||
The bounding box can be given as a tuple of (x, y, width, height) or as a BBOX type
|
|
||||||
The BBOX input takes precedence over the tuple input
|
The BBOX input takes precedence over the tuple input
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@@ -204,35 +204,38 @@ class MTB_Crop:
|
|||||||
def do_crop(
|
def do_crop(
|
||||||
self,
|
self,
|
||||||
image: torch.Tensor,
|
image: torch.Tensor,
|
||||||
mask=None,
|
*,
|
||||||
x=0,
|
mask: torch.Tensor | None = None,
|
||||||
y=0,
|
x: int = 0,
|
||||||
width=256,
|
y: int = 0,
|
||||||
height=256,
|
width: int = 256,
|
||||||
bbox=None,
|
height: int = 256,
|
||||||
|
bbox: BoundingBox | None = None,
|
||||||
):
|
):
|
||||||
image = image.numpy()
|
|
||||||
if mask is not None:
|
|
||||||
mask = mask.numpy()
|
|
||||||
|
|
||||||
if bbox is not None:
|
if bbox is not None:
|
||||||
x, y, width, height = bbox
|
x, y, width, height = bbox
|
||||||
|
|
||||||
cropped_image = image[:, y : y + height, x : x + width, :]
|
if width <= 0 or height <= 0:
|
||||||
cropped_mask = None
|
log.error(
|
||||||
if mask is not None:
|
"Crop dimensions must be positive. Check the BBOX or widget inputs."
|
||||||
cropped_mask = (
|
|
||||||
mask[:, y : y + height, x : x + width]
|
|
||||||
if mask is not None
|
|
||||||
else None
|
|
||||||
)
|
)
|
||||||
crop_data = (x, y, width, height)
|
return (
|
||||||
|
torch.zeros_like(image),
|
||||||
|
torch.zeros_like(mask) if mask is not None else None,
|
||||||
|
(x, y, width, height),
|
||||||
|
)
|
||||||
|
|
||||||
|
cropped_image = image[:, y : y + height, x : x + width, :]
|
||||||
|
cropped_mask = (
|
||||||
|
mask[:, y : y + height, x : x + width]
|
||||||
|
if mask is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
crop_data = BoundingBox(x, y, width, height)
|
||||||
|
|
||||||
return (
|
return (
|
||||||
torch.from_numpy(cropped_image),
|
cropped_image,
|
||||||
torch.from_numpy(cropped_mask)
|
cropped_mask if cropped_mask is not None else None,
|
||||||
if cropped_mask is not None
|
|
||||||
else None,
|
|
||||||
crop_data,
|
crop_data,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -246,35 +249,33 @@ class MTB_Crop:
|
|||||||
# return (x_left, y_top, x_right, y_bottom)
|
# return (x_left, y_top, x_right, y_bottom)
|
||||||
|
|
||||||
|
|
||||||
def bbox_check(bbox, target_size=None):
|
def bbox_check(bbox: BoundingBox, target_size: tuple[int, int] | None = None):
|
||||||
if not target_size:
|
if not target_size:
|
||||||
return bbox
|
return bbox
|
||||||
|
|
||||||
new_bbox = (
|
new_bbox = BoundingBox(
|
||||||
bbox[0],
|
bbox.x,
|
||||||
bbox[1],
|
bbox.y,
|
||||||
min(target_size[0] - bbox[0], bbox[2]),
|
min(target_size[0] - bbox.x, bbox.width),
|
||||||
min(target_size[1] - bbox[1], bbox[3]),
|
min(target_size[1] - bbox.y, bbox.height),
|
||||||
)
|
)
|
||||||
if new_bbox != bbox:
|
if new_bbox != bbox:
|
||||||
log.warn(f"BBox too big, constrained to {new_bbox}")
|
log.warning(f"BBox too big, constrained to {new_bbox}")
|
||||||
|
|
||||||
return new_bbox
|
return new_bbox
|
||||||
|
|
||||||
|
|
||||||
def bbox_to_region(bbox, target_size=None):
|
def bbox_to_region(
|
||||||
|
bbox: BoundingBox, target_size: tuple[int, int] | None = None
|
||||||
|
):
|
||||||
bbox = bbox_check(bbox, target_size)
|
bbox = bbox_check(bbox, target_size)
|
||||||
|
|
||||||
# to region
|
# to region
|
||||||
return (bbox[0], bbox[1], bbox[0] + bbox[2], bbox[1] + bbox[3])
|
return (bbox.x, bbox.y, bbox.x + bbox.width, bbox.y + bbox.height)
|
||||||
|
|
||||||
|
|
||||||
class MTB_Uncrop:
|
class MTB_Uncrop:
|
||||||
"""Uncrops an image to a given bounding box
|
"""Uncrop an image to a given bounding box."""
|
||||||
|
|
||||||
The bounding box can be given as a tuple of (x, y, width, height) or as a BBOX type
|
|
||||||
The BBOX input takes precedence over the tuple input
|
|
||||||
"""
|
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(cls):
|
def INPUT_TYPES(cls):
|
||||||
@@ -291,91 +292,113 @@ class MTB_Uncrop:
|
|||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ("IMAGE",)
|
RETURN_TYPES = ("IMAGE",)
|
||||||
FUNCTION = "do_crop"
|
FUNCTION = "do_uncrop"
|
||||||
|
|
||||||
CATEGORY = "mtb/crop"
|
CATEGORY = "mtb/crop"
|
||||||
|
|
||||||
def do_crop(self, image, crop_image, bbox, border_blending):
|
def do_uncrop(
|
||||||
def inset_border(image, border_width=20, border_color=(0)):
|
self,
|
||||||
width, height = image.size
|
image: torch.Tensor,
|
||||||
bordered_image = Image.new(
|
crop_image: torch.Tensor,
|
||||||
image.mode, (width, height), border_color
|
bbox: BoundingBox,
|
||||||
|
border_blending: float = 0.25,
|
||||||
|
):
|
||||||
|
if len(image) > 1 and len(image) != len(crop_image):
|
||||||
|
raise ValueError(
|
||||||
|
"Uncrop: Batch size of background 'image' must be 1 or match the 'crop_image' batch size."
|
||||||
)
|
)
|
||||||
bordered_image.paste(image, (0, 0))
|
import comfy.utils
|
||||||
draw = ImageDraw.Draw(bordered_image)
|
|
||||||
draw.rectangle(
|
|
||||||
(0, 0, width - 1, height - 1),
|
|
||||||
outline=border_color,
|
|
||||||
width=border_width,
|
|
||||||
)
|
|
||||||
return bordered_image
|
|
||||||
|
|
||||||
single = image.size(0) == 1
|
pbar = comfy.utils.ProgressBar(4)
|
||||||
if image.size(0) != crop_image.size(0):
|
|
||||||
if not single:
|
|
||||||
raise ValueError(
|
|
||||||
"The Image batch count is greater than 1, but doesn't match the crop_image batch count. If using batches they should either match or only crop_image must be greater than 1"
|
|
||||||
)
|
|
||||||
|
|
||||||
images = tensor2pil(image)
|
device = image.device
|
||||||
crop_imgs = tensor2pil(crop_image)
|
|
||||||
out_images = []
|
|
||||||
for i, crop in enumerate(crop_imgs):
|
|
||||||
if single:
|
|
||||||
img = images[0]
|
|
||||||
else:
|
|
||||||
img = images[i]
|
|
||||||
|
|
||||||
# uncrop the image based on the bounding box
|
log.debug(f"Working on device: {device}")
|
||||||
bb_x, bb_y, bb_width, bb_height = bbox
|
|
||||||
|
|
||||||
paste_region = bbox_to_region(
|
crop_image = crop_image.to(device)
|
||||||
(bb_x, bb_y, bb_width, bb_height), img.size
|
|
||||||
)
|
|
||||||
# log.debug(f"Paste region: {paste_region}")
|
|
||||||
# new_region = adjust_paste_region(img.size, paste_region)
|
|
||||||
# log.debug(f"Adjusted paste region: {new_region}")
|
|
||||||
# # Check if the adjusted paste region is different from the original
|
|
||||||
|
|
||||||
crop_img = crop.convert("RGB")
|
if len(image) == 1 and len(crop_image) > 1:
|
||||||
|
image = image.repeat(len(crop_image), 1, 1, 1)
|
||||||
|
|
||||||
log.debug(f"Crop image size: {crop_img.size}")
|
batch_size, bg_h, bg_w, _ = image.shape
|
||||||
log.debug(f"Image size: {img.size}")
|
_, fg_h, fg_w, _ = crop_image.shape
|
||||||
|
x, y, width, height = bbox
|
||||||
|
|
||||||
if border_blending > 1.0:
|
if (width, height) != (fg_w, fg_h):
|
||||||
border_blending = 1.0
|
log.warning(
|
||||||
elif border_blending < 0.0:
|
f"Uncrop: crop_image size {(fg_w, fg_h)} "
|
||||||
border_blending = 0.0
|
"differs from bbox {(width, height)}. Resizing to fit bbox."
|
||||||
|
|
||||||
blend_ratio = (max(crop_img.size) / 2) * float(border_blending)
|
|
||||||
|
|
||||||
blend = img.convert("RGBA")
|
|
||||||
mask = Image.new("L", img.size, 0)
|
|
||||||
|
|
||||||
mask_block = Image.new("L", (bb_width, bb_height), 255)
|
|
||||||
mask_block = inset_border(mask_block, int(blend_ratio / 2), (0))
|
|
||||||
|
|
||||||
mask.paste(mask_block, paste_region)
|
|
||||||
log.debug(f"Blend size: {blend.size} | kind {blend.mode}")
|
|
||||||
log.debug(
|
|
||||||
f"Crop image size: {crop_img.size} | kind {crop_img.mode}"
|
|
||||||
)
|
|
||||||
log.debug(f"BBox: {paste_region}")
|
|
||||||
blend.paste(crop_img, paste_region)
|
|
||||||
|
|
||||||
mask = mask.filter(ImageFilter.BoxBlur(radius=blend_ratio / 4))
|
|
||||||
mask = mask.filter(
|
|
||||||
ImageFilter.GaussianBlur(radius=blend_ratio / 4)
|
|
||||||
)
|
)
|
||||||
|
|
||||||
blend.putalpha(mask)
|
resized_crop = crop_image.permute(0, 3, 1, 2)
|
||||||
img = Image.alpha_composite(img.convert("RGBA"), blend)
|
resized_crop = torch.nn.functional.interpolate(
|
||||||
out_images.append(img.convert("RGB"))
|
resized_crop,
|
||||||
|
size=(height, width),
|
||||||
|
mode="bicubic",
|
||||||
|
align_corners=False,
|
||||||
|
)
|
||||||
|
resized_crop = resized_crop.permute(0, 2, 3, 1)
|
||||||
|
|
||||||
return (pil2tensor(out_images),)
|
pbar.update(1)
|
||||||
|
# paste coords
|
||||||
|
paste_x1 = max(x, 0)
|
||||||
|
paste_y1 = max(y, 0)
|
||||||
|
paste_x2 = min(x + width, bg_w)
|
||||||
|
paste_y2 = min(y + height, bg_h)
|
||||||
|
|
||||||
|
# region from crop (bound)
|
||||||
|
crop_x1 = max(0, -x)
|
||||||
|
crop_y1 = max(0, -y)
|
||||||
|
crop_x2 = crop_x1 + (paste_x2 - paste_x1)
|
||||||
|
crop_y2 = crop_y1 + (paste_y2 - paste_y1)
|
||||||
|
|
||||||
|
if paste_x1 >= paste_x2 or paste_y1 >= paste_y2:
|
||||||
|
log.warning(
|
||||||
|
"Uncrop: BBOX is entirely outside the image boundaries. Returning original image."
|
||||||
|
)
|
||||||
|
return (image,)
|
||||||
|
|
||||||
|
pbar.update(1)
|
||||||
|
source_slice = resized_crop[:, crop_y1:crop_y2, crop_x1:crop_x2, :]
|
||||||
|
|
||||||
|
final_image = image.clone()
|
||||||
|
final_image[:, paste_y1:paste_y2, paste_x1:paste_x2, :] = source_slice
|
||||||
|
|
||||||
|
pbar.update(1)
|
||||||
|
|
||||||
|
blend_radius = int(max(width, height) * border_blending * 0.5)
|
||||||
|
if blend_radius > 0:
|
||||||
|
_device = device
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
_device = torch.device("cuda")
|
||||||
|
|
||||||
|
log.debug("Processing blending")
|
||||||
|
alpha_mask = torch.zeros((batch_size, bg_h, bg_w), device=_device)
|
||||||
|
alpha_mask[:, paste_y1:paste_y2, paste_x1:paste_x2] = 1.0
|
||||||
|
|
||||||
|
kernel_size = 2 * blend_radius + 1
|
||||||
|
|
||||||
|
log.debug("Gaussian blur...")
|
||||||
|
alpha_mask = TF.gaussian_blur(
|
||||||
|
alpha_mask.unsqueeze(1), kernel_size=[kernel_size, kernel_size]
|
||||||
|
).squeeze(1)
|
||||||
|
alpha_mask = alpha_mask.unsqueeze(-1)
|
||||||
|
|
||||||
|
log.debug("Applying blending")
|
||||||
|
final_image = final_image.to(_device) * alpha_mask + image.to(
|
||||||
|
_device
|
||||||
|
) * (1.0 - alpha_mask)
|
||||||
|
|
||||||
|
pbar.update(1)
|
||||||
|
return (final_image.to(device),)
|
||||||
|
|
||||||
|
|
||||||
class MTB_BBoxForceDimensions:
|
class MTB_BBoxForceDimensions:
|
||||||
|
"""
|
||||||
|
Resize a BBOX to new dimensions while keeping its center.
|
||||||
|
|
||||||
|
Optionally constrains the BBOX to stay within image boundaries.
|
||||||
|
"""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(cls):
|
def INPUT_TYPES(cls):
|
||||||
return {
|
return {
|
||||||
@@ -383,6 +406,7 @@ class MTB_BBoxForceDimensions:
|
|||||||
"bbox": ("BBOX",),
|
"bbox": ("BBOX",),
|
||||||
"width": ("INT", {"default": 512, "min": 1, "max": 8192}),
|
"width": ("INT", {"default": 512, "min": 1, "max": 8192}),
|
||||||
"height": ("INT", {"default": 512, "min": 1, "max": 8192}),
|
"height": ("INT", {"default": 512, "min": 1, "max": 8192}),
|
||||||
|
"constrain_to_image": ("BOOLEAN", {"default": True}),
|
||||||
},
|
},
|
||||||
"optional": {
|
"optional": {
|
||||||
"image": ("IMAGE",),
|
"image": ("IMAGE",),
|
||||||
@@ -395,10 +419,12 @@ class MTB_BBoxForceDimensions:
|
|||||||
|
|
||||||
def force_dimensions(
|
def force_dimensions(
|
||||||
self,
|
self,
|
||||||
|
*,
|
||||||
bbox: tuple[int, int, int, int],
|
bbox: tuple[int, int, int, int],
|
||||||
width: int,
|
width: int,
|
||||||
height: int,
|
height: int,
|
||||||
image: torch.Tensor = None,
|
constrain_to_image: bool = True,
|
||||||
|
image: torch.Tensor | None = None,
|
||||||
) -> tuple[tuple[int, int, int, int]]:
|
) -> tuple[tuple[int, int, int, int]]:
|
||||||
x, y, curr_width, curr_height = bbox
|
x, y, curr_width, curr_height = bbox
|
||||||
|
|
||||||
@@ -408,27 +434,14 @@ class MTB_BBoxForceDimensions:
|
|||||||
new_x = center_x - width // 2
|
new_x = center_x - width // 2
|
||||||
new_y = center_y - height // 2
|
new_y = center_y - height // 2
|
||||||
|
|
||||||
if image is not None:
|
if constrain_to_image and image is not None:
|
||||||
img_height, img_width = image.shape[1:3]
|
img_height, img_width = image.shape[1:3]
|
||||||
x_overflow = max(0, new_x + width - img_width) + min(0, new_x)
|
new_x = max(0, min(new_x, img_width - width))
|
||||||
y_overflow = max(0, new_y + height - img_height) + min(0, new_y)
|
new_y = max(0, min(new_y, img_height - height))
|
||||||
if width > img_width or height > img_height:
|
width = min(width, img_width)
|
||||||
x_exceed = width - img_width if width > img_width else 0
|
height = min(height, img_height)
|
||||||
y_exceed = height - img_height if height > img_height else 0
|
|
||||||
raise ValueError(
|
|
||||||
f"Target bbox dimensions ({width}x{height}) exceed image bounds ({img_width}x{img_height}) "
|
|
||||||
f"by {x_exceed}px horizontally and {y_exceed}px vertically"
|
|
||||||
)
|
|
||||||
|
|
||||||
if x_overflow > 0 or x_overflow < 0:
|
return ((new_x, new_y, width, height),)
|
||||||
new_x -= x_overflow
|
|
||||||
|
|
||||||
if y_overflow > 0:
|
|
||||||
new_y -= y_overflow
|
|
||||||
elif y_overflow < 0:
|
|
||||||
new_y -= y_overflow # Add the negative overflow
|
|
||||||
|
|
||||||
return ((int(new_x), int(new_y), width, height),)
|
|
||||||
|
|
||||||
|
|
||||||
__nodes__ = [
|
__nodes__ = [
|
||||||
|
|||||||
+605
-179
@@ -1,33 +1,70 @@
|
|||||||
import base64
|
import base64
|
||||||
import io
|
import io
|
||||||
import json
|
import textwrap
|
||||||
from pathlib import Path
|
from collections.abc import Callable
|
||||||
|
from functools import wraps
|
||||||
|
from typing import Any, Literal, Protocol, TypedDict, runtime_checkable
|
||||||
|
|
||||||
import folder_paths
|
|
||||||
import torch
|
import torch
|
||||||
|
from rich import inspect
|
||||||
|
from rich.console import Console
|
||||||
|
|
||||||
from ..log import log
|
from ..log import log
|
||||||
from ..utils import tensor2pil
|
from ..utils import LazyProxyTensor, get_torch_tensor_info, tensor2pil
|
||||||
|
|
||||||
|
try:
|
||||||
|
import matplotlib.pyplot as plt
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
plt.style.use("dark_background")
|
||||||
|
MATPLOTLIB_AVAILABLE = True
|
||||||
|
except ImportError:
|
||||||
|
MATPLOTLIB_AVAILABLE = False
|
||||||
|
|
||||||
|
|
||||||
def get_detailed_type_info(obj):
|
# region Decorator
|
||||||
type_info = []
|
def metadata(**meta_kwargs: Any) -> Callable[[Any], Any]:
|
||||||
|
"""Add metadata to method (`__meta__` dict)."""
|
||||||
|
|
||||||
|
def decorator(func: Callable[[Any], Any]) -> Callable[[Any], Any]:
|
||||||
|
@wraps(func)
|
||||||
|
def wrapper(*args, **kwargs):
|
||||||
|
return func(*args, **kwargs)
|
||||||
|
|
||||||
|
wrapper.__meta__ = meta_kwargs
|
||||||
|
return wrapper
|
||||||
|
|
||||||
|
return decorator
|
||||||
|
|
||||||
|
|
||||||
|
# endregion
|
||||||
|
class UIResult(TypedDict):
|
||||||
|
kind: Literal["text", "b64_images"]
|
||||||
|
data: str
|
||||||
|
|
||||||
|
|
||||||
|
def indent_results(results: list[UIResult], by: str = " "):
|
||||||
|
for res in results:
|
||||||
|
if res["kind"] == "text":
|
||||||
|
log.debug(f"Indenting: {res['data']}")
|
||||||
|
res["data"] = textwrap.indent(res["data"], by)
|
||||||
|
|
||||||
|
return results
|
||||||
|
|
||||||
|
|
||||||
|
ProcessorResult = list[UIResult]
|
||||||
|
|
||||||
|
|
||||||
|
def _get_detailed_type_info(obj) -> str:
|
||||||
|
type_info: list[str] = []
|
||||||
|
|
||||||
type_name = type(obj).__name__
|
type_name = type(obj).__name__
|
||||||
type_info.append(f"Type: {type_name}")
|
type_info.append(f"Type: {type_name}")
|
||||||
|
|
||||||
if isinstance(obj, torch.Tensor):
|
if isinstance(obj, torch.Tensor):
|
||||||
type_info.extend(
|
return get_torch_tensor_info(obj)
|
||||||
[
|
|
||||||
f"Shape: {obj.shape}",
|
elif isinstance(obj, list | tuple):
|
||||||
f"Dtype: {obj.dtype}",
|
|
||||||
f"Device: {obj.device}",
|
|
||||||
f"Requires grad: {obj.requires_grad}",
|
|
||||||
f"Stride: {obj.stride()}",
|
|
||||||
f"Contiguous: {obj.is_contiguous()}",
|
|
||||||
]
|
|
||||||
)
|
|
||||||
elif isinstance(obj, (list, tuple)):
|
|
||||||
type_info.extend(
|
type_info.extend(
|
||||||
[
|
[
|
||||||
f"Length: {len(obj)}",
|
f"Length: {len(obj)}",
|
||||||
@@ -47,122 +84,184 @@ def get_detailed_type_info(obj):
|
|||||||
attributes = [attr for attr in dir(obj) if not attr.startswith("_")]
|
attributes = [attr for attr in dir(obj) if not attr.startswith("_")]
|
||||||
type_info.append(f"Attributes: {attributes}")
|
type_info.append(f"Attributes: {attributes}")
|
||||||
|
|
||||||
return type_info
|
return "\n".join(type_info)
|
||||||
|
|
||||||
|
|
||||||
|
def _apply_rich_results(processed, mode="none", title=""):
|
||||||
|
processing_text = False
|
||||||
|
acc = ""
|
||||||
|
reshaped: list[UIResult] = []
|
||||||
|
for i in range(len(processed)):
|
||||||
|
if processed[i]["kind"] == "text":
|
||||||
|
if not processing_text:
|
||||||
|
processing_text = True
|
||||||
|
acc += processed[i]["data"] + "\n"
|
||||||
|
if len(processed) == (i + 1):
|
||||||
|
reshaped.append(
|
||||||
|
UIResult(
|
||||||
|
kind="text", data=_apply_rich(acc, mode, title=title)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
if processing_text:
|
||||||
|
processing_text = False
|
||||||
|
reshaped.append(
|
||||||
|
UIResult(
|
||||||
|
kind="text", data=_apply_rich(acc, mode, title=title)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
acc = ""
|
||||||
|
reshaped.append(processed[i])
|
||||||
|
|
||||||
|
return reshaped
|
||||||
|
# for item in processed:
|
||||||
|
|
||||||
|
|
||||||
# region processors
|
# region processors
|
||||||
def process_tensor(tensor: torch.Tensor, as_type=False):
|
def _apply_rich(
|
||||||
log.debug(f"Tensor: {tensor.shape}")
|
formatted: str | list[str], rich_mode: str | None = None, *, title=""
|
||||||
|
) -> str:
|
||||||
if as_type:
|
if rich_mode is None:
|
||||||
return {
|
return (
|
||||||
"text": [f"Tensor of shape {tensor.shape} of type {tensor.dtype}"]
|
formatted if isinstance(formatted, str) else "\n".join(formatted)
|
||||||
}
|
|
||||||
|
|
||||||
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")
|
|
||||||
)
|
)
|
||||||
|
|
||||||
return {"b64_images": b64_imgs}
|
from rich.console import Console
|
||||||
|
|
||||||
|
console = Console(record=True)
|
||||||
|
|
||||||
def process_list(anything, as_type=False):
|
if isinstance(formatted, list):
|
||||||
text = []
|
for line in formatted:
|
||||||
if not anything:
|
console.print(line)
|
||||||
return {"text": []}
|
|
||||||
|
|
||||||
if as_type:
|
|
||||||
type_info = get_detailed_type_info(anything)
|
|
||||||
type_info.extend(get_detailed_type_info(anything[0]))
|
|
||||||
return {"text": type_info}
|
|
||||||
|
|
||||||
first_element = anything[0]
|
|
||||||
if (
|
|
||||||
isinstance(first_element, list)
|
|
||||||
and first_element
|
|
||||||
and isinstance(first_element[0], torch.Tensor)
|
|
||||||
):
|
|
||||||
text.append(
|
|
||||||
"List of List of Tensors: "
|
|
||||||
f"{first_element[0].shape} (x{len(anything)})"
|
|
||||||
)
|
|
||||||
|
|
||||||
elif isinstance(first_element, torch.Tensor):
|
|
||||||
text.append(
|
|
||||||
f"List of Tensors: {first_element.shape} (x{len(anything)})"
|
|
||||||
)
|
|
||||||
else:
|
else:
|
||||||
text.append(f"Array ({len(anything)}): {anything}")
|
console.print(formatted)
|
||||||
|
|
||||||
return {"text": text}
|
CSV_CODE_FORMAT = """
|
||||||
|
<svg class="rich-terminal" viewBox="0 0 {width} {height}" xmlns="http://www.w3.org/2000/svg">
|
||||||
|
<!-- Generated with Rich https://www.textualize.io -->
|
||||||
|
<style>
|
||||||
|
|
||||||
|
@font-face {{
|
||||||
|
font-family: "Fira Code";
|
||||||
|
src: local("FiraCode-Regular"),
|
||||||
|
url("https://cdnjs.cloudflare.com/ajax/libs/firacode/6.2.0/woff2/FiraCode-Regular.woff2") format("woff2"),
|
||||||
|
url("https://cdnjs.cloudflare.com/ajax/libs/firacode/6.2.0/woff/FiraCode-Regular.woff") format("woff");
|
||||||
|
font-style: normal;
|
||||||
|
font-weight: 400;
|
||||||
|
}}
|
||||||
|
@font-face {{
|
||||||
|
font-family: "Fira Code";
|
||||||
|
src: local("FiraCode-Bold"),
|
||||||
|
url("https://cdnjs.cloudflare.com/ajax/libs/firacode/6.2.0/woff2/FiraCode-Bold.woff2") format("woff2"),
|
||||||
|
url("https://cdnjs.cloudflare.com/ajax/libs/firacode/6.2.0/woff/FiraCode-Bold.woff") format("woff");
|
||||||
|
font-style: bold;
|
||||||
|
font-weight: 700;
|
||||||
|
}}
|
||||||
|
|
||||||
def process_dict(anything, as_type=False):
|
.{unique_id}-matrix {{
|
||||||
text = []
|
font-family: Fira Code, monospace;
|
||||||
if as_type:
|
font-size: {char_height}px;
|
||||||
return {"text": get_detailed_type_info(anything)}
|
line-height: {line_height}px;
|
||||||
|
font-variant-east-asian: full-width;
|
||||||
|
}}
|
||||||
|
|
||||||
if "samples" in anything:
|
.{unique_id}-title {{
|
||||||
is_empty = (
|
font-size: 18px;
|
||||||
"(empty)" if torch.count_nonzero(anything["samples"]) == 0 else ""
|
font-weight: bold;
|
||||||
)
|
font-family: arial;
|
||||||
text.append(f"Latent Samples: {anything['samples'].shape} {is_empty}")
|
}}
|
||||||
|
|
||||||
elif "waveform" in anything:
|
{styles}
|
||||||
is_empty = (
|
</style>
|
||||||
"(empty) " if torch.count_nonzero(anything["samples"]) == 0 else ""
|
|
||||||
|
<defs>
|
||||||
|
<clipPath id="{unique_id}-clip-terminal">
|
||||||
|
<rect x="0" y="0" width="{terminal_width}" height="{terminal_height}" />
|
||||||
|
</clipPath>
|
||||||
|
{lines}
|
||||||
|
</defs>
|
||||||
|
|
||||||
|
{chrome}
|
||||||
|
<g clip-path="url(#{unique_id}-clip-terminal)">
|
||||||
|
{backgrounds}
|
||||||
|
<g class="{unique_id}-matrix">
|
||||||
|
{matrix}
|
||||||
|
</g>
|
||||||
|
</g>
|
||||||
|
</svg>
|
||||||
|
"""
|
||||||
|
|
||||||
|
if rich_mode == "svg-window":
|
||||||
|
return console.export_svg(title=title, code_format=CSV_CODE_FORMAT)
|
||||||
|
elif rich_mode == "svg":
|
||||||
|
return console.export_svg(
|
||||||
|
title=title,
|
||||||
|
code_format=CSV_CODE_FORMAT.replace("{chrome}", ""),
|
||||||
)
|
)
|
||||||
|
|
||||||
text.append(
|
elif rich_mode == "html":
|
||||||
f"Audio Samples: {anything['waveform'].shape}{is_empty} | sample rate {anything['sample_rate']}"
|
CONSOLE_HTML_FORMAT = textwrap.dedent("""
|
||||||
|
<div style="color:{foreground};">
|
||||||
|
<code style="font-family:inherit">{code}</code>
|
||||||
|
</div>
|
||||||
|
""").strip()
|
||||||
|
|
||||||
|
import rich.terminal_theme
|
||||||
|
|
||||||
|
return console.export_html(
|
||||||
|
inline_styles=True,
|
||||||
|
code_format=CONSOLE_HTML_FORMAT,
|
||||||
|
theme=rich.terminal_theme.MONOKAI,
|
||||||
)
|
)
|
||||||
|
|
||||||
else:
|
log.error(f"Unknown rich mode: {rich_mode}")
|
||||||
log.debug(f"Unhandled dict: {anything.keys()}")
|
return formatted if isinstance(formatted, str) else "\n".join(formatted)
|
||||||
text.append(json.dumps(anything, indent=2))
|
|
||||||
|
|
||||||
return {"text": text}
|
|
||||||
|
|
||||||
|
|
||||||
def process_bool(anything, as_type=False):
|
|
||||||
return {"text": ["True" if anything else "False"]}
|
|
||||||
|
|
||||||
|
|
||||||
def process_text(anything, as_type=False):
|
|
||||||
if as_type:
|
|
||||||
return {"text": get_detailed_type_info(anything)}
|
|
||||||
|
|
||||||
return {"text": [str(anything)]}
|
|
||||||
|
|
||||||
|
|
||||||
# endregion
|
# endregion
|
||||||
|
|
||||||
|
|
||||||
class MTB_Debug:
|
# region conditions
|
||||||
"""Experimental node to debug any Comfy values.
|
|
||||||
|
|
||||||
support for more types and widgets is planned.
|
|
||||||
"""
|
# those are pretty dumb there is now probably a better way..
|
||||||
|
def is_condition(item):
|
||||||
|
return (
|
||||||
|
isinstance(item, list)
|
||||||
|
and all(isinstance(i, list) for i in item)
|
||||||
|
and isinstance(item[0][0], torch.Tensor)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# endregion
|
||||||
|
|
||||||
|
RICH_MODE = Literal["none", "html", "svg", "svg-window"]
|
||||||
|
|
||||||
|
|
||||||
|
@runtime_checkable
|
||||||
|
class Processor(Protocol):
|
||||||
|
"""Generic protocol for processor functions."""
|
||||||
|
|
||||||
|
def __call__(
|
||||||
|
self, item: Any, *, as_type: bool = False, deep: bool = False
|
||||||
|
) -> ProcessorResult: ...
|
||||||
|
|
||||||
|
|
||||||
|
class MTB_Debug:
|
||||||
|
"""A debug node."""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(cls):
|
def INPUT_TYPES(cls):
|
||||||
return {
|
return {
|
||||||
"required": {"output_to_console": ("BOOLEAN", {"default": False})},
|
"required": {"output_to_console": ("BOOLEAN", {"default": False})},
|
||||||
"optional": {"as_detailed_types": ("BOOLEAN", {"default": False})},
|
"optional": {
|
||||||
|
"as_detailed_types": ("BOOLEAN", {"default": False}),
|
||||||
|
"deep_inspect": ("BOOLEAN", {"default": False}),
|
||||||
|
"rich_mode": (
|
||||||
|
("none", "html", "svg", "svg-window"),
|
||||||
|
{"default": "none"},
|
||||||
|
),
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ()
|
RETURN_TYPES = ()
|
||||||
@@ -170,99 +269,426 @@ class MTB_Debug:
|
|||||||
CATEGORY = "mtb/debug"
|
CATEGORY = "mtb/debug"
|
||||||
OUTPUT_NODE = True
|
OUTPUT_NODE = True
|
||||||
|
|
||||||
|
_processors: dict[type, Processor]
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self._condition_processors = {is_condition: self._process_condition}
|
||||||
|
self._class_name_processors = {
|
||||||
|
"CLIP": self._process_clip,
|
||||||
|
"VAE": self._process_vae,
|
||||||
|
}
|
||||||
|
self._processors = {
|
||||||
|
torch.nn.Module: self._process_module,
|
||||||
|
torch.Tensor: self._process_tensor,
|
||||||
|
LazyProxyTensor: self._process_repr,
|
||||||
|
list: self._process_container,
|
||||||
|
tuple: self._process_container,
|
||||||
|
dict: self._process_dict,
|
||||||
|
bool: self._process_bool,
|
||||||
|
str: self._process_primitive,
|
||||||
|
int: self._process_primitive,
|
||||||
|
float: self._process_primitive,
|
||||||
|
type(None): self._process_primitive,
|
||||||
|
}
|
||||||
|
|
||||||
|
# - Dispatchers ------------------------------------------------------------
|
||||||
|
def _dispatch_processor(
|
||||||
|
self, item: Any, *, as_type=False, deep=False
|
||||||
|
) -> ProcessorResult:
|
||||||
|
"""Find and calls the appropriate processor for the given item."""
|
||||||
|
# first conditions
|
||||||
|
for c, process in self._condition_processors.items():
|
||||||
|
if c(item):
|
||||||
|
return process(item, as_type=as_type, deep=deep)
|
||||||
|
|
||||||
|
# named class
|
||||||
|
class_name = type(item).__name__
|
||||||
|
if class_name in self._class_name_processors:
|
||||||
|
return self._class_name_processors[class_name](
|
||||||
|
item, as_type=as_type, deep=deep
|
||||||
|
)
|
||||||
|
|
||||||
|
# type based or unknown
|
||||||
|
processor = self._processors.get(type(item), self._process_unknown)
|
||||||
|
res = processor(item, as_type=as_type, deep=deep)
|
||||||
|
|
||||||
|
return res
|
||||||
|
|
||||||
def do_debug(
|
def do_debug(
|
||||||
self, output_to_console: bool, as_detailed_types: bool, **kwargs
|
self,
|
||||||
|
**kwargs,
|
||||||
):
|
):
|
||||||
output = {"ui": {"items": []}}
|
output = {"ui": {"items": []}}
|
||||||
|
|
||||||
if output_to_console:
|
settings = {k: kwargs.pop(k) for k in self.INPUT_TYPES()["optional"]}
|
||||||
for k, v in kwargs.items():
|
output_to_console = kwargs.pop("output_to_console")
|
||||||
log.info(f"{k}: {v}")
|
as_type = settings.get("as_detailed_types", False)
|
||||||
|
deep = settings.get("deep_inspect", False)
|
||||||
|
rich_mode = settings.get("rich_mode", "none")
|
||||||
|
|
||||||
for input_name, anything in kwargs.items():
|
for input_name, item in kwargs.items():
|
||||||
processor = processors.get(type(anything), process_text)
|
processed = self._dispatch_processor(
|
||||||
|
item, as_type=as_type, deep=deep
|
||||||
|
)
|
||||||
|
if processed is None:
|
||||||
|
continue
|
||||||
|
|
||||||
processed = processor(anything, as_detailed_types)
|
if rich_mode != "none":
|
||||||
|
title = f"{input_name} ({type(item).__name__})"
|
||||||
|
processed = _apply_rich_results(processed, rich_mode, title)
|
||||||
|
|
||||||
item = {
|
if output_to_console:
|
||||||
"input": input_name,
|
log.info(f"- Input '{input_name}':")
|
||||||
**processed,
|
for p in processed:
|
||||||
}
|
if p["kind"] == "text":
|
||||||
output["ui"]["items"].append(item)
|
log.info(f" {p['data']}")
|
||||||
|
if p["kind"] == "b64_image":
|
||||||
|
log.info(f" (contains {len(p['data'])} images)")
|
||||||
|
|
||||||
|
output["ui"]["items"].append(
|
||||||
|
{"input": input_name, "items": processed}
|
||||||
|
)
|
||||||
return output
|
return output
|
||||||
|
|
||||||
|
def _process_unknown(
|
||||||
|
self, item: Any, *, as_type=False, deep=False
|
||||||
|
) -> ProcessorResult:
|
||||||
|
console = Console(
|
||||||
|
record=True,
|
||||||
|
width=120,
|
||||||
|
)
|
||||||
|
|
||||||
class MTB_SaveTensors:
|
console.print(f"Generic {type(item).__name__}", emoji=True)
|
||||||
"""Save torch tensors (image, mask or latent) to disk.
|
if as_type:
|
||||||
|
inspect(item, console=console, all=deep, methods=deep, docs=deep)
|
||||||
|
else:
|
||||||
|
console.print(item, emoji=True)
|
||||||
|
|
||||||
useful to debug things outside comfy.
|
text_output = console.export_text(clear=True)
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self):
|
return [UIResult(kind="text", data=text_output.strip())]
|
||||||
self.output_dir = folder_paths.get_output_directory()
|
|
||||||
self.type = "mtb/debug"
|
|
||||||
|
|
||||||
@classmethod
|
def _process_repr(
|
||||||
def INPUT_TYPES(cls):
|
self, item: Any, as_type=False, deep=False
|
||||||
return {
|
) -> ProcessorResult:
|
||||||
"required": {
|
return [{"kind": "text", "data": item.__repr__()}]
|
||||||
"filename_prefix": ("STRING", {"default": "ComfyPickle"}),
|
|
||||||
},
|
|
||||||
"optional": {
|
|
||||||
"image": ("IMAGE",),
|
|
||||||
"mask": ("MASK",),
|
|
||||||
"latent": ("LATENT",),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
FUNCTION = "save"
|
def _process_primitive(
|
||||||
OUTPUT_NODE = True
|
self, item: Any, *, as_type=False, deep=False
|
||||||
RETURN_TYPES = ()
|
) -> ProcessorResult:
|
||||||
CATEGORY = "mtb/debug"
|
if as_type:
|
||||||
|
return self._process_unknown(item, as_type=as_type, deep=deep)
|
||||||
|
|
||||||
def save(
|
return [UIResult(kind="text", data=str(item))]
|
||||||
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:
|
def _process_bool(
|
||||||
mask_file = f"{filename}_mask_{counter:05}.pt"
|
self, item: bool, *, as_type=False, deep=False
|
||||||
torch.save(mask, full_output_folder / mask_file)
|
) -> ProcessorResult: # noqa: FBT001
|
||||||
# np.save(full_output_folder/ mask_file, mask.cpu().numpy())
|
return [{"kind": "text", "data": "True" if item else "False"}]
|
||||||
|
|
||||||
if latent is not None:
|
def _process_clip(
|
||||||
# for latent we must use pickle
|
self, item: Any, *, as_type=False, deep=False
|
||||||
latent_file = f"{filename}_latent_{counter:05}.pt"
|
) -> ProcessorResult:
|
||||||
torch.save(latent, full_output_folder / latent_file)
|
try:
|
||||||
# pickle.dump(latent, open(full_output_folder/ latent_file, "wb"))
|
clip_model = getattr(item, "cond_stage_model", None)
|
||||||
|
tokenizer = getattr(item, "tokenizer", None)
|
||||||
|
|
||||||
# np.save(full_output_folder / latent_file,
|
text = [UIResult(kind="text", data="CLIP")]
|
||||||
# latent[""].cpu().numpy())
|
if clip_model:
|
||||||
|
text.append(UIResult(kind="text", data="CLIP Model:"))
|
||||||
|
model_summary = self._process_module(
|
||||||
|
clip_model, as_type=as_type
|
||||||
|
)
|
||||||
|
if model_summary:
|
||||||
|
text.extend(indent_results(model_summary, " "))
|
||||||
|
else:
|
||||||
|
text.append(
|
||||||
|
UIResult(
|
||||||
|
kind="text",
|
||||||
|
data="[error] failed to get informations about clip model",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
return f"{filename_prefix}_{counter:05}"
|
if tokenizer:
|
||||||
|
text.append(UIResult(kind="text", data="Tokenizer:"))
|
||||||
|
vocab_size = getattr(tokenizer, "vocab_size", "N/A")
|
||||||
|
text.append(
|
||||||
|
UIResult(
|
||||||
|
kind="text",
|
||||||
|
data=f" Class: {type(tokenizer).__name__}\n Vocab Size: {vocab_size}",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
return text
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
log.error(f"Failed to process CLIP object: {e}")
|
||||||
|
return self._process_unknown(item, as_type=as_type, deep=deep)
|
||||||
|
|
||||||
|
def _process_condition(
|
||||||
|
self, item: Any, *, as_type=False, deep=False
|
||||||
|
) -> ProcessorResult:
|
||||||
|
count = len(item)
|
||||||
|
result = [UIResult(kind="text", data=f"Conditions: {count}")]
|
||||||
|
|
||||||
|
for cond in item:
|
||||||
|
result.extend(self._preview_conditioning_tensor(cond[0]))
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
def _process_vae(
|
||||||
|
self, item: Any, *, as_type=False, deep=False
|
||||||
|
) -> ProcessorResult:
|
||||||
|
try:
|
||||||
|
vae_model = getattr(
|
||||||
|
item, "first_stage_model", getattr(item, "vae", item)
|
||||||
|
)
|
||||||
|
text = [
|
||||||
|
UIResult(kind="text", data="VAE"),
|
||||||
|
UIResult(kind="text", data="Internal Model:"),
|
||||||
|
]
|
||||||
|
|
||||||
|
model_summary = self._process_module(
|
||||||
|
vae_model, as_type=as_type, deep=deep
|
||||||
|
)
|
||||||
|
text.extend(indent_results(model_summary, " "))
|
||||||
|
|
||||||
|
return text
|
||||||
|
except Exception as e:
|
||||||
|
log.error(f"Failed to process VAE object: {e}")
|
||||||
|
return self._process_unknown(item, as_type=as_type, deep=deep)
|
||||||
|
|
||||||
|
def _process_module(
|
||||||
|
self, item: torch.nn.Module, *, as_type=False, deep=False
|
||||||
|
) -> ProcessorResult:
|
||||||
|
if as_type and deep:
|
||||||
|
return self._process_unknown(item, as_type=as_type, deep=deep)
|
||||||
|
|
||||||
|
total_params = sum(p.numel() for p in item.parameters())
|
||||||
|
trainable_params = sum(
|
||||||
|
p.numel() for p in item.parameters() if p.requires_grad
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
device = next(item.parameters()).device
|
||||||
|
except StopIteration:
|
||||||
|
device = "cpu (no parameters)"
|
||||||
|
|
||||||
|
train_percent = (
|
||||||
|
f"{trainable_params / total_params:.2%}"
|
||||||
|
if total_params > 0
|
||||||
|
else "0.00%"
|
||||||
|
)
|
||||||
|
|
||||||
|
text = [
|
||||||
|
f"Model: {type(item).__name__} on {device}",
|
||||||
|
textwrap.dedent(f"""
|
||||||
|
- Parameters: {total_params:,}
|
||||||
|
- Trainable: {trainable_params:,} ({train_percent})
|
||||||
|
""").strip(),
|
||||||
|
]
|
||||||
|
return [{"kind": "text", "data": d} for d in text]
|
||||||
|
|
||||||
|
def _process_tensor(
|
||||||
|
self, item: torch.Tensor, *, as_type=False, deep=False
|
||||||
|
) -> ProcessorResult:
|
||||||
|
is_latent = item.ndim == 4 and item.shape[1] == 4
|
||||||
|
is_image = (
|
||||||
|
not is_latent and item.ndim == 4 and item.shape[3] in [1, 3, 4]
|
||||||
|
)
|
||||||
|
is_conditioning = item.ndim == 3 and item.shape[2] in [
|
||||||
|
768,
|
||||||
|
1024,
|
||||||
|
1152,
|
||||||
|
1280,
|
||||||
|
2048,
|
||||||
|
4096,
|
||||||
|
]
|
||||||
|
is_mask = (item.ndim == 2) or (item.ndim == 3 and not is_conditioning)
|
||||||
|
|
||||||
|
if as_type:
|
||||||
|
type_name = "Unknown Tensor"
|
||||||
|
if is_latent:
|
||||||
|
type_name = "Latent Tensor"
|
||||||
|
elif is_image:
|
||||||
|
type_name = "Image Tensor"
|
||||||
|
elif is_conditioning:
|
||||||
|
type_name = "CLIP Conditioning Tensor"
|
||||||
|
elif is_mask:
|
||||||
|
type_name = "Mask Tensor"
|
||||||
|
return [
|
||||||
|
{
|
||||||
|
"kind": "text",
|
||||||
|
"data": get_torch_tensor_info(item, name=type_name),
|
||||||
|
}
|
||||||
|
]
|
||||||
|
|
||||||
|
if is_image or is_mask:
|
||||||
|
return self._render_image_tensor(item)
|
||||||
|
if is_latent:
|
||||||
|
return self._preview_latent_tensor(item)
|
||||||
|
if is_conditioning:
|
||||||
|
return self._preview_conditioning_tensor(item)
|
||||||
|
return self._process_unknown(item, as_type=as_type, deep=deep)
|
||||||
|
|
||||||
|
def _visualize_tensor_heatmap(
|
||||||
|
self, tensor_2d: torch.Tensor, title: str
|
||||||
|
) -> str | None:
|
||||||
|
if not MATPLOTLIB_AVAILABLE:
|
||||||
|
log.warning("Matplotlib not found. Skipping tensor visualization.")
|
||||||
|
return None
|
||||||
|
if tensor_2d.ndim != 2:
|
||||||
|
log.warning(
|
||||||
|
f"Cannot visualize tensor with {tensor_2d.ndim} dimensions. Requires 2."
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
fig, ax = plt.subplots(figsize=(6, 4), dpi=100)
|
||||||
|
im = ax.imshow(tensor_2d.cpu().numpy(), cmap="viridis", aspect="auto")
|
||||||
|
fig.colorbar(im, ax=ax)
|
||||||
|
ax.set_title(title)
|
||||||
|
fig.tight_layout()
|
||||||
|
|
||||||
|
buf = io.BytesIO()
|
||||||
|
fig.savefig(buf, format="png", bbox_inches="tight", pad_inches=0.1)
|
||||||
|
plt.close(fig)
|
||||||
|
buf.seek(0)
|
||||||
|
return "data:image/png;base64," + base64.b64encode(buf.read()).decode(
|
||||||
|
"utf-8"
|
||||||
|
)
|
||||||
|
|
||||||
|
def _render_image_tensor(self, item: torch.Tensor) -> ProcessorResult:
|
||||||
|
is_mask = (item.ndim == 2) or (item.ndim == 3 and item.shape[-1] != 3)
|
||||||
|
img_tensor = (
|
||||||
|
item.unsqueeze(0) if item.ndim == 3 and not is_mask else item
|
||||||
|
)
|
||||||
|
img_tensor = item.unsqueeze(0) if item.ndim == 2 else img_tensor
|
||||||
|
|
||||||
|
images = tensor2pil(img_tensor)
|
||||||
|
b64_imgs = []
|
||||||
|
for im in images:
|
||||||
|
if is_mask:
|
||||||
|
im = im.convert("L")
|
||||||
|
buffered = io.BytesIO()
|
||||||
|
im.save(buffered, format="PNG")
|
||||||
|
b64_imgs.append(
|
||||||
|
"data:image/png;base64,"
|
||||||
|
+ base64.b64encode(buffered.getvalue()).decode("utf-8")
|
||||||
|
)
|
||||||
|
return [UIResult(kind="b64_images", data=b64_imgs)]
|
||||||
|
|
||||||
|
def _preview_latent_tensor(self, item: torch.Tensor) -> ProcessorResult:
|
||||||
|
is_empty = "(empty)" if torch.count_nonzero(item) == 0 else ""
|
||||||
|
stats = [
|
||||||
|
f"Min: {item.min():.4f}",
|
||||||
|
f"Max: {item.max():.4f}",
|
||||||
|
f"Mean: {item.mean():.4f}",
|
||||||
|
]
|
||||||
|
text = [
|
||||||
|
get_torch_tensor_info(item, name="Latent Tensor"),
|
||||||
|
is_empty,
|
||||||
|
] + stats
|
||||||
|
|
||||||
|
result = [UIResult(kind="text", data=t) for t in text]
|
||||||
|
vis_tensor = item[0].mean(dim=0)
|
||||||
|
heatmap_b64 = self._visualize_tensor_heatmap(
|
||||||
|
vis_tensor, "Latent Energy (Channel Mean)"
|
||||||
|
)
|
||||||
|
if heatmap_b64:
|
||||||
|
result.append(UIResult(kind="b64_images", data=[heatmap_b64]))
|
||||||
|
return result
|
||||||
|
|
||||||
|
def _preview_conditioning_tensor(
|
||||||
|
self, item: torch.Tensor
|
||||||
|
) -> ProcessorResult:
|
||||||
|
_batch, tokens, embed_dim = item.shape
|
||||||
|
text = [
|
||||||
|
get_torch_tensor_info(item, name="CLIP Conditioning Tensor"),
|
||||||
|
f"Token Count: {tokens}",
|
||||||
|
f"Embedding Dim: {embed_dim}",
|
||||||
|
]
|
||||||
|
|
||||||
|
result = [UIResult(kind="text", data=d) for d in text]
|
||||||
|
heatmap_b64 = self._visualize_tensor_heatmap(
|
||||||
|
item[0], "Token Embeddings (approx)"
|
||||||
|
)
|
||||||
|
if heatmap_b64:
|
||||||
|
result.append(UIResult(kind="b64_images", data=[heatmap_b64]))
|
||||||
|
return result
|
||||||
|
|
||||||
|
def _process_container(
|
||||||
|
self, item: list | tuple, *, as_type=False, deep=False
|
||||||
|
) -> ProcessorResult:
|
||||||
|
if not item:
|
||||||
|
return [UIResult(kind="text", data=f"Empty {type(item).__name__}")]
|
||||||
|
|
||||||
|
container_type = type(item).__name__
|
||||||
|
element_type = type(item[0]).__name__
|
||||||
|
|
||||||
|
all_match = all(type(i) is type(item[0]) for i in item)
|
||||||
|
|
||||||
|
result = [
|
||||||
|
UIResult(
|
||||||
|
kind="text",
|
||||||
|
data=f"{container_type} of {len(item)} x {element_type}",
|
||||||
|
),
|
||||||
|
UIResult(kind="text", data=f"(mixed types: {not all_match})"),
|
||||||
|
]
|
||||||
|
|
||||||
|
if not as_type or (as_type and deep):
|
||||||
|
for i, sub_item in enumerate(item):
|
||||||
|
res = self._dispatch_processor(
|
||||||
|
sub_item, as_type=as_type, deep=deep
|
||||||
|
)
|
||||||
|
if res:
|
||||||
|
text = res[0].get("data", "Unknown")
|
||||||
|
res[0]["data"] = f"[{i}]: {text}"
|
||||||
|
|
||||||
|
result.extend(res)
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
first_item_result = self._dispatch_processor(
|
||||||
|
item[0], as_type=as_type, deep=deep
|
||||||
|
)
|
||||||
|
if not first_item_result:
|
||||||
|
return result
|
||||||
|
|
||||||
|
return (
|
||||||
|
result
|
||||||
|
+ [UIResult(kind="text", data="Preview of first element:")]
|
||||||
|
+ indent_results(first_item_result, " - ")
|
||||||
|
)
|
||||||
|
|
||||||
|
def _process_dict(
|
||||||
|
self, item: dict, *, as_type=False, deep=False
|
||||||
|
) -> ProcessorResult:
|
||||||
|
if "pooled_output" in item and isinstance(
|
||||||
|
item["pooled_output"], torch.Tensor
|
||||||
|
):
|
||||||
|
return self._dispatch_processor(
|
||||||
|
item["pooled_output"], as_type=as_type, deep=deep
|
||||||
|
)
|
||||||
|
|
||||||
|
if "samples" in item and isinstance(item.get("samples"), torch.Tensor):
|
||||||
|
return self._dispatch_processor(
|
||||||
|
item["samples"], as_type=as_type, deep=deep
|
||||||
|
)
|
||||||
|
|
||||||
|
if "waveform" in item and isinstance(
|
||||||
|
item.get("waveform"), torch.Tensor
|
||||||
|
):
|
||||||
|
waveform = item["waveform"]
|
||||||
|
is_empty = "(empty) " if torch.count_nonzero(waveform) == 0 else ""
|
||||||
|
text = textwrap.dedent(f"""
|
||||||
|
Audio Waveform: {waveform.shape}{is_empty}
|
||||||
|
Sample Rate: {item.get("sample_rate", "N/A")}
|
||||||
|
""").strip()
|
||||||
|
return [{"kind": "text", "data": text}]
|
||||||
|
|
||||||
|
log.debug(
|
||||||
|
f"Processing generic dict with rich inspector: {item.keys()}"
|
||||||
|
)
|
||||||
|
return self._process_unknown(item, as_type=as_type, deep=deep)
|
||||||
|
|
||||||
|
|
||||||
processors = {
|
__nodes__ = [MTB_Debug]
|
||||||
torch.Tensor: process_tensor,
|
|
||||||
list: process_list,
|
|
||||||
dict: process_dict,
|
|
||||||
bool: process_bool,
|
|
||||||
}
|
|
||||||
|
|
||||||
__nodes__ = [MTB_Debug, MTB_SaveTensors]
|
|
||||||
|
|||||||
@@ -0,0 +1,70 @@
|
|||||||
|
import folder_paths
|
||||||
|
import torch
|
||||||
|
|
||||||
|
|
||||||
|
class MTB_SaveTensors:
|
||||||
|
"""Save torch tensors (image, mask or latent) to disk.
|
||||||
|
|
||||||
|
useful to debug things outside comfy.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self.output_dir = folder_paths.get_output_directory()
|
||||||
|
self.type = "mtb/debug"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"filename_prefix": ("STRING", {"default": "ComfyPickle"}),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"image": ("IMAGE",),
|
||||||
|
"mask": ("MASK",),
|
||||||
|
"latent": ("LATENT",),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
FUNCTION = "save"
|
||||||
|
OUTPUT_NODE = True
|
||||||
|
RETURN_TYPES = ()
|
||||||
|
CATEGORY = "mtb/debug"
|
||||||
|
|
||||||
|
def save(
|
||||||
|
self,
|
||||||
|
filename_prefix,
|
||||||
|
image: torch.Tensor | None = None,
|
||||||
|
mask: torch.Tensor | None = None,
|
||||||
|
latent: torch.Tensor | None = None,
|
||||||
|
):
|
||||||
|
(
|
||||||
|
full_output_folder,
|
||||||
|
filename,
|
||||||
|
counter,
|
||||||
|
subfolder,
|
||||||
|
filename_prefix,
|
||||||
|
) = folder_paths.get_save_image_path(filename_prefix, self.output_dir)
|
||||||
|
full_output_folder = Path(full_output_folder)
|
||||||
|
if image is not None:
|
||||||
|
image_file = f"{filename}_image_{counter:05}.pt"
|
||||||
|
torch.save(image, full_output_folder / image_file)
|
||||||
|
# np.save(full_output_folder/ image_file, image.cpu().numpy())
|
||||||
|
|
||||||
|
if mask is not None:
|
||||||
|
mask_file = f"{filename}_mask_{counter:05}.pt"
|
||||||
|
torch.save(mask, full_output_folder / mask_file)
|
||||||
|
# np.save(full_output_folder/ mask_file, mask.cpu().numpy())
|
||||||
|
|
||||||
|
if latent is not None:
|
||||||
|
# for latent we must use pickle
|
||||||
|
latent_file = f"{filename}_latent_{counter:05}.pt"
|
||||||
|
torch.save(latent, full_output_folder / latent_file)
|
||||||
|
# pickle.dump(latent, open(full_output_folder/ latent_file, "wb"))
|
||||||
|
|
||||||
|
# np.save(full_output_folder / latent_file,
|
||||||
|
# latent[""].cpu().numpy())
|
||||||
|
|
||||||
|
return f"{filename_prefix}_{counter:05}"
|
||||||
|
|
||||||
|
|
||||||
|
__nodes__ = [MTB_SaveTensors]
|
||||||
+7
-4
@@ -4,12 +4,8 @@ import sys
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import comfy.model_management as model_management
|
import comfy.model_management as model_management
|
||||||
import cv2
|
|
||||||
import insightface
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import onnxruntime
|
|
||||||
import torch
|
import torch
|
||||||
from insightface.model_zoo.inswapper import INSwapper
|
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
|
|
||||||
from ..errors import ModelNotFound
|
from ..errors import ModelNotFound
|
||||||
@@ -43,6 +39,8 @@ class MTB_LoadFaceAnalysisModel:
|
|||||||
DEPRECATED = True
|
DEPRECATED = True
|
||||||
|
|
||||||
def load_model(self, faceswap_model: str):
|
def load_model(self, faceswap_model: str):
|
||||||
|
import insightface
|
||||||
|
|
||||||
if faceswap_model == "antelopev2":
|
if faceswap_model == "antelopev2":
|
||||||
download_antelopev2()
|
download_antelopev2()
|
||||||
|
|
||||||
@@ -81,6 +79,9 @@ class MTB_LoadFaceSwapModel:
|
|||||||
DEPRECATED = True
|
DEPRECATED = True
|
||||||
|
|
||||||
def load_model(self, faceswap_model: str):
|
def load_model(self, faceswap_model: str):
|
||||||
|
import onnxruntime
|
||||||
|
from insightface.model_zoo.inswapper import INSwapper
|
||||||
|
|
||||||
model_path = get_model_path("insightface", faceswap_model)
|
model_path = get_model_path("insightface", faceswap_model)
|
||||||
if not model_path or not model_path.exists():
|
if not model_path or not model_path.exists():
|
||||||
raise ModelNotFound(f"{faceswap_model} ({model_path})")
|
raise ModelNotFound(f"{faceswap_model} ({model_path})")
|
||||||
@@ -212,6 +213,8 @@ def swap_face(
|
|||||||
face_swapper_model,
|
face_swapper_model,
|
||||||
faces_index: set[int] | None = None,
|
faces_index: set[int] | None = None,
|
||||||
) -> Image.Image:
|
) -> Image.Image:
|
||||||
|
import cv2
|
||||||
|
|
||||||
if faces_index is None:
|
if faces_index is None:
|
||||||
faces_index = {0}
|
faces_index = {0}
|
||||||
log.debug(f"Swapping faces: {faces_index}")
|
log.debug(f"Swapping faces: {faces_index}")
|
||||||
|
|||||||
+23
-9
@@ -7,6 +7,13 @@ from PIL import Image, ImageDraw, ImageFont
|
|||||||
from ..log import log
|
from ..log import log
|
||||||
from ..utils import comfy_dir, font_path, pil2tensor
|
from ..utils import comfy_dir, font_path, pil2tensor
|
||||||
|
|
||||||
|
# try:
|
||||||
|
# from cairosvg import svg2png
|
||||||
|
# HAS_CAIRO = True
|
||||||
|
# except ImportError:
|
||||||
|
# HAS_CAIRO = False
|
||||||
|
|
||||||
|
|
||||||
# class MtbExamples:
|
# class MtbExamples:
|
||||||
# """MTB Example Images"""
|
# """MTB Example Images"""
|
||||||
|
|
||||||
@@ -299,7 +306,7 @@ by default it fallsback to a default font.
|
|||||||
|
|
||||||
def text_to_image(
|
def text_to_image(
|
||||||
self,
|
self,
|
||||||
text: str,
|
text: str | list[str],
|
||||||
font,
|
font,
|
||||||
wrap,
|
wrap,
|
||||||
trim,
|
trim,
|
||||||
@@ -341,11 +348,9 @@ by default it fallsback to a default font.
|
|||||||
color = (255, 255, 255, 255)
|
color = (255, 255, 255, 255)
|
||||||
background = (0, 0, 0, 255)
|
background = (0, 0, 0, 255)
|
||||||
|
|
||||||
def render_text(text_to_render, alpha=None):
|
def render_text(text_to_render: str, alpha=None) -> Image.Image:
|
||||||
if trim:
|
if trim:
|
||||||
text_to_render = (
|
text_to_render = text_to_render.strip()
|
||||||
text_to_render.encode("ascii", "ignore").decode().strip()
|
|
||||||
)
|
|
||||||
if wrap:
|
if wrap:
|
||||||
wrap_width = (((width / 100) * h_coverage) / font_size) * 2
|
wrap_width = (((width / 100) * h_coverage) / font_size) * 2
|
||||||
lines = textwrap.wrap(text_to_render, width=wrap_width)
|
lines = textwrap.wrap(text_to_render, width=wrap_width)
|
||||||
@@ -418,7 +423,9 @@ by default it fallsback to a default font.
|
|||||||
active_chunks.append((chunk["text"], alpha))
|
active_chunks.append((chunk["text"], alpha))
|
||||||
|
|
||||||
for chunk_text, alpha in active_chunks:
|
for chunk_text, alpha in active_chunks:
|
||||||
chunk_img = render_text(chunk_text, alpha)
|
chunk_img = render_text(
|
||||||
|
chunk_text.encode("ascii", "ignore").decode(), alpha
|
||||||
|
)
|
||||||
frame = Image.alpha_composite(frame, chunk_img)
|
frame = Image.alpha_composite(frame, chunk_img)
|
||||||
|
|
||||||
frames.append(frame)
|
frames.append(frame)
|
||||||
@@ -426,9 +433,16 @@ by default it fallsback to a default font.
|
|||||||
frame_tensors = [pil2tensor(frame) for frame in frames]
|
frame_tensors = [pil2tensor(frame) for frame in frames]
|
||||||
return (torch.cat(frame_tensors, dim=0),)
|
return (torch.cat(frame_tensors, dim=0),)
|
||||||
else:
|
else:
|
||||||
text_img = render_text(text)
|
results = []
|
||||||
result = Image.alpha_composite(base_img, text_img)
|
if not isinstance(text, list):
|
||||||
return (pil2tensor(result),)
|
text = [text]
|
||||||
|
|
||||||
|
for t in text:
|
||||||
|
text_img = render_text(t)
|
||||||
|
result = Image.alpha_composite(base_img, text_img)
|
||||||
|
results.append(result)
|
||||||
|
|
||||||
|
return (pil2tensor(results),)
|
||||||
|
|
||||||
|
|
||||||
__nodes__ = [
|
__nodes__ = [
|
||||||
|
|||||||
+168
-7
@@ -4,18 +4,22 @@ import re
|
|||||||
import urllib.parse
|
import urllib.parse
|
||||||
import urllib.request
|
import urllib.request
|
||||||
from math import pi
|
from math import pi
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
import comfy.model_management as model_management
|
import comfy.model_management as mm
|
||||||
import comfy.utils
|
import comfy.utils
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
|
from comfy.comfy_types.node_typing import IO as CIO
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
|
|
||||||
from ..log import log
|
from ..log import log
|
||||||
from ..utils import (
|
from ..utils import (
|
||||||
EASINGS,
|
EASINGS,
|
||||||
|
LazyProxyTensor,
|
||||||
apply_easing,
|
apply_easing,
|
||||||
get_server_info,
|
get_server_info,
|
||||||
|
get_torch_tensor_info,
|
||||||
numpy_NFOV,
|
numpy_NFOV,
|
||||||
pil2tensor,
|
pil2tensor,
|
||||||
tensor2np,
|
tensor2np,
|
||||||
@@ -132,12 +136,64 @@ class MTB_ApplyTextTemplate:
|
|||||||
CATEGORY = "mtb/utils"
|
CATEGORY = "mtb/utils"
|
||||||
FUNCTION = "execute"
|
FUNCTION = "execute"
|
||||||
|
|
||||||
def execute(self, *, template: str, **kwargs):
|
def execute(self, *, template: str, **kwargs) -> tuple[str | list[str]]:
|
||||||
res = f"{template}"
|
keys = list(kwargs.keys())
|
||||||
for k, v in kwargs.items():
|
values = list(kwargs.values())
|
||||||
res = res.replace(f"{{{k}}}", f"{v}")
|
|
||||||
|
|
||||||
return (res,)
|
has_list = any(isinstance(v, list) for v in values)
|
||||||
|
target_length = -1
|
||||||
|
|
||||||
|
if has_list:
|
||||||
|
first_list = next(x for x in values if isinstance(x, list))
|
||||||
|
|
||||||
|
# all_list = all(isinstance(x, list) for x in kwargs.values())
|
||||||
|
# if not all_list:
|
||||||
|
# raise ValueError(
|
||||||
|
# "Text template supports either str or list[str] but not a mix of the two (yet?)"
|
||||||
|
# )
|
||||||
|
target_length = len(first_list)
|
||||||
|
same_length = all(
|
||||||
|
len(v) == target_length for v in values if isinstance(v, list)
|
||||||
|
)
|
||||||
|
if not same_length:
|
||||||
|
raise ValueError(
|
||||||
|
"Text template received multiple list[str] but their size is varying, they should match..."
|
||||||
|
)
|
||||||
|
|
||||||
|
if has_list:
|
||||||
|
results = []
|
||||||
|
|
||||||
|
# do a padded loop, not the most efficient but easy
|
||||||
|
# to handle for now
|
||||||
|
for it in range(target_length):
|
||||||
|
res = f"{template}"
|
||||||
|
for k, v in kwargs.items():
|
||||||
|
if isinstance(v, list):
|
||||||
|
res = self.apply_res(res, k, v[it])
|
||||||
|
else:
|
||||||
|
res = self.apply_res(res, k, v)
|
||||||
|
results.append(res)
|
||||||
|
|
||||||
|
return (results,)
|
||||||
|
|
||||||
|
else:
|
||||||
|
res = f"{template}"
|
||||||
|
for k, v in kwargs.items():
|
||||||
|
res = self.apply_res(res, k, v)
|
||||||
|
|
||||||
|
return (res,)
|
||||||
|
|
||||||
|
def apply_res(self, res, key, value):
|
||||||
|
if isinstance(value, float):
|
||||||
|
value = f"{value:.3f}"
|
||||||
|
elif isinstance(value, torch.Tensor):
|
||||||
|
value = get_torch_tensor_info(value)
|
||||||
|
else:
|
||||||
|
log.debug(
|
||||||
|
f"Falling back to default string conversion for {key} of type {type(value).__name__}"
|
||||||
|
)
|
||||||
|
|
||||||
|
return res.replace(f"{{{key}}}", f"{value}")
|
||||||
|
|
||||||
|
|
||||||
class MTB_MatchDimensions:
|
class MTB_MatchDimensions:
|
||||||
@@ -341,7 +397,7 @@ class MTB_AutoPanEquilateral:
|
|||||||
|
|
||||||
frames.append(frame)
|
frames.append(frame)
|
||||||
|
|
||||||
model_management.throw_exception_if_processing_interrupted()
|
mm.throw_exception_if_processing_interrupted()
|
||||||
pbar.update(1)
|
pbar.update(1)
|
||||||
|
|
||||||
return (pil2tensor(frames),)
|
return (pil2tensor(frames),)
|
||||||
@@ -867,6 +923,108 @@ class MTB_TensorOps:
|
|||||||
return (result,)
|
return (result,)
|
||||||
|
|
||||||
|
|
||||||
|
class MTB_GetItem:
|
||||||
|
"""Generic index based getter for common types"""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"container": (CIO.ANY,),
|
||||||
|
"index": ("INT", {"default": 0}),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = (CIO.ANY,)
|
||||||
|
RETURN_NAMES = ("item",)
|
||||||
|
FUNCTION = "get_item"
|
||||||
|
CATEGORY = "mtb/utils"
|
||||||
|
|
||||||
|
def get_item(self, container: Any, index: int):
|
||||||
|
if "__getitem__" in dir(container):
|
||||||
|
log.debug(f"Container is {type(container)}")
|
||||||
|
res = container[index]
|
||||||
|
if type(res) is torch.Tensor:
|
||||||
|
res = res.unsqueeze(0)
|
||||||
|
|
||||||
|
return (res,)
|
||||||
|
|
||||||
|
|
||||||
|
class MTB_BooleanNot:
|
||||||
|
"""Inverts a boolean."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"bool_in": ("BOOLEAN", {"default": False}),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("BOOLEAN",)
|
||||||
|
RETURN_NAMES = ("inverted_bool",)
|
||||||
|
FUNCTION = "invert"
|
||||||
|
CATEGORY = "mtb/utils"
|
||||||
|
|
||||||
|
def invert(self, bool_in: bool):
|
||||||
|
return (not bool_in,)
|
||||||
|
|
||||||
|
|
||||||
|
class MTB_ProxyTensor:
|
||||||
|
"""Wraps an input tensor into a LazyProxyTensor.
|
||||||
|
|
||||||
|
builds upon an idea by @AustinMroz
|
||||||
|
"""
|
||||||
|
|
||||||
|
NODE_NAME = "ProxyTensor"
|
||||||
|
NODE_DISPLAY_NAME = "Proxy Tensor"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"tensor": ("IMAGE",),
|
||||||
|
"target_dtype": (
|
||||||
|
["float32", "float16", "bfloat16"],
|
||||||
|
{"default": "float32"},
|
||||||
|
),
|
||||||
|
"target_device": (
|
||||||
|
["keep", "cpu", "gpu"],
|
||||||
|
{
|
||||||
|
"default": "keep",
|
||||||
|
"tooltip": "CAUTION: This isn't compatible with most nodes for now",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("IMAGE",)
|
||||||
|
RETURN_NAMES = ("proxy_tensor",)
|
||||||
|
FUNCTION = "execute"
|
||||||
|
CATEGORY = "mtb/utils"
|
||||||
|
|
||||||
|
def execute(
|
||||||
|
self,
|
||||||
|
tensor: torch.Tensor,
|
||||||
|
target_dtype: str = "float32",
|
||||||
|
target_device: str = "keep",
|
||||||
|
):
|
||||||
|
torch_dtype: torch.dtype = getattr(torch, target_dtype)
|
||||||
|
|
||||||
|
if target_device == "gpu":
|
||||||
|
torch_device = mm.get_torch_device()
|
||||||
|
elif target_device == "cpu":
|
||||||
|
torch_device = torch.device("cpu")
|
||||||
|
else:
|
||||||
|
torch_device = tensor.device
|
||||||
|
|
||||||
|
proxy = LazyProxyTensor(tensor, torch_dtype, torch_device)
|
||||||
|
|
||||||
|
log.info(f"Created Proxy Tensor: \n{proxy}")
|
||||||
|
|
||||||
|
return (proxy,)
|
||||||
|
|
||||||
|
|
||||||
__nodes__ = [
|
__nodes__ = [
|
||||||
MTB_StringReplace,
|
MTB_StringReplace,
|
||||||
MTB_FitNumber,
|
MTB_FitNumber,
|
||||||
@@ -882,4 +1040,7 @@ __nodes__ = [
|
|||||||
MTB_FloatToFloats,
|
MTB_FloatToFloats,
|
||||||
MTB_FloatsToInts,
|
MTB_FloatsToInts,
|
||||||
MTB_TensorOps,
|
MTB_TensorOps,
|
||||||
|
MTB_BooleanNot,
|
||||||
|
MTB_GetItem,
|
||||||
|
MTB_ProxyTensor,
|
||||||
]
|
]
|
||||||
|
|||||||
+103
-47
@@ -3,11 +3,12 @@ import json
|
|||||||
import math
|
import math
|
||||||
import os
|
import os
|
||||||
|
|
||||||
import comfy.model_management as model_management
|
import comfy.utils
|
||||||
import folder_paths
|
import folder_paths
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
|
from comfy import model_management
|
||||||
from PIL import Image, ImageOps
|
from PIL import Image, ImageOps
|
||||||
from PIL.PngImagePlugin import PngInfo
|
from PIL.PngImagePlugin import PngInfo
|
||||||
from skimage.filters import gaussian
|
from skimage.filters import gaussian
|
||||||
@@ -74,7 +75,10 @@ class MTB_ExtractCoordinatesFromImage:
|
|||||||
def INPUT_TYPES(cls):
|
def INPUT_TYPES(cls):
|
||||||
return {
|
return {
|
||||||
"required": {
|
"required": {
|
||||||
"threshold": ("FLOAT",),
|
"threshold": (
|
||||||
|
"FLOAT",
|
||||||
|
{"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01},
|
||||||
|
),
|
||||||
"max_points": ("INT", {"default": 50, "min": 0}),
|
"max_points": ("INT", {"default": 50, "min": 0}),
|
||||||
},
|
},
|
||||||
"optional": {"image": ("IMAGE",), "mask": ("MASK",)},
|
"optional": {"image": ("IMAGE",), "mask": ("MASK",)},
|
||||||
@@ -87,72 +91,124 @@ class MTB_ExtractCoordinatesFromImage:
|
|||||||
image: torch.Tensor | None = None,
|
image: torch.Tensor | None = None,
|
||||||
mask: torch.Tensor | None = None,
|
mask: torch.Tensor | None = None,
|
||||||
) -> tuple[list[list[tuple[int, int]]], torch.Tensor]:
|
) -> tuple[list[list[tuple[int, int]]], torch.Tensor]:
|
||||||
if image is not None:
|
if image is None and mask is None:
|
||||||
batch_count, height, width, channel_count = image.shape
|
raise ValueError("Must provide either image or mask")
|
||||||
imgs = image
|
|
||||||
else:
|
|
||||||
if mask is None:
|
|
||||||
raise ValueError("Must provide either image or mask")
|
|
||||||
batch_count, height, width = mask.shape
|
|
||||||
channel_count = 1
|
|
||||||
imgs = mask
|
|
||||||
|
|
||||||
if channel_count not in [1, 2, 3, 4]:
|
if image is not None:
|
||||||
raise ValueError(f"Incorrect channel count: {channel_count}")
|
batch_count, height, width, _channel_count = image.shape
|
||||||
|
input_device = image.device
|
||||||
|
if mask is not None:
|
||||||
|
if mask.ndim == 2:
|
||||||
|
mask = mask.unsqueeze(0)
|
||||||
|
if mask.ndim != 3:
|
||||||
|
raise ValueError(
|
||||||
|
f"Mask has unexpected ndim: {mask.ndim}. Expected 2 or 3."
|
||||||
|
)
|
||||||
|
|
||||||
|
b_mask, h_mask, w_mask = mask.shape
|
||||||
|
if not (h_mask == height and w_mask == width):
|
||||||
|
raise ValueError(
|
||||||
|
f"Image dimensions ({height}x{width}) and mask dimensions ({h_mask}x{w_mask}) are spatially incompatible."
|
||||||
|
)
|
||||||
|
if b_mask == 1 and batch_count > 1:
|
||||||
|
mask = mask.expand(batch_count, height, width)
|
||||||
|
|
||||||
|
elif b_mask != batch_count:
|
||||||
|
raise ValueError(
|
||||||
|
f"Image batch size ({batch_count}) and mask batch size ({b_mask}) are incompatible and mask cannot be broadcast."
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
if mask.ndim == 2:
|
||||||
|
mask = mask.unsqueeze(0)
|
||||||
|
|
||||||
|
if mask.ndim != 3:
|
||||||
|
raise ValueError(
|
||||||
|
f"Mask has unexpected ndim: {mask.ndim} when image is not provided. Expected 2 or 3."
|
||||||
|
)
|
||||||
|
|
||||||
|
batch_count, height, width = mask.shape
|
||||||
|
input_device = mask.device
|
||||||
|
|
||||||
all_points: list[list[tuple[int, int]]] = []
|
all_points: list[list[tuple[int, int]]] = []
|
||||||
debug_images = torch.zeros(
|
debug_images = torch.zeros(
|
||||||
(batch_count, height, width, 3),
|
(batch_count, height, width, 3),
|
||||||
dtype=torch.uint8,
|
dtype=torch.uint8,
|
||||||
device=imgs.device,
|
device=input_device,
|
||||||
)
|
)
|
||||||
|
|
||||||
for i, img in enumerate(imgs):
|
points_tensor = torch.tensor(
|
||||||
if channel_count == 1:
|
[255, 255, 255], dtype=torch.uint8, device=input_device
|
||||||
alpha_channel = img if len(img.shape) == 2 else img[:, :, 0]
|
)
|
||||||
elif channel_count == 2:
|
|
||||||
alpha_channel = img[:, :, 1]
|
for i in range(batch_count):
|
||||||
elif channel_count == 4:
|
value_threshold: torch.Tensor
|
||||||
alpha_channel = img[:, :, 3]
|
if image is not None:
|
||||||
|
img_slice = image[i]
|
||||||
|
img_channels = img_slice.shape[2]
|
||||||
|
if img_channels == 1 or img_channels == 2:
|
||||||
|
value_threshold = img_slice[:, :, 0]
|
||||||
|
elif img_channels == 3 or img_channels == 4:
|
||||||
|
value_threshold = img_slice[:, :, :3].max(dim=2)[0]
|
||||||
|
else:
|
||||||
|
raise ValueError(
|
||||||
|
f"Unsupported image channel count: {img_channels} for image at batch index {i}"
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
# get intensity
|
mask_slice = mask[i]
|
||||||
alpha_channel = img[:, :, :3].max(dim=2)[0]
|
value_threshold = mask_slice
|
||||||
|
|
||||||
points = (alpha_channel > threshold).nonzero(as_tuple=False)
|
condition = value_threshold > threshold
|
||||||
|
if image is not None and mask is not None:
|
||||||
|
mask_slice = mask[i]
|
||||||
|
mask_active_condition = mask_slice > 0.0
|
||||||
|
condition = condition & mask_active_condition
|
||||||
|
|
||||||
if len(points) > max_points:
|
points_yx = condition.nonzero(as_tuple=False)
|
||||||
indices = torch.randperm(points.size(0), device=img.device)[
|
|
||||||
:max_points
|
|
||||||
]
|
|
||||||
points = points[indices]
|
|
||||||
|
|
||||||
points = [(int(y.item()), int(x.item())) for x, y in points]
|
if points_yx.size(0) > max_points:
|
||||||
all_points.append(points)
|
# shuffle and pick max_points randomly
|
||||||
|
indices = torch.randperm(
|
||||||
|
points_yx.size(0), device=input_device
|
||||||
|
)[:max_points]
|
||||||
|
points_yx = points_yx[indices]
|
||||||
|
elif max_points == 0:
|
||||||
|
points_yx = torch.empty(
|
||||||
|
(0, 2), dtype=torch.long, device=input_device
|
||||||
|
)
|
||||||
|
|
||||||
for x, y in points:
|
current_points = [
|
||||||
self._draw_circle(debug_images[i], (x, y), 5)
|
(int(p[1].item()), int(p[0].item())) for p in points_yx
|
||||||
|
]
|
||||||
|
all_points.append(current_points)
|
||||||
|
for x_coord, y_coord in current_points:
|
||||||
|
self._draw_circle(
|
||||||
|
debug_images[i],
|
||||||
|
(x_coord, y_coord),
|
||||||
|
radius=5,
|
||||||
|
color_tensor=points_tensor,
|
||||||
|
)
|
||||||
|
|
||||||
return (all_points, debug_images)
|
return (all_points, debug_images)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _draw_circle(
|
def _draw_circle(
|
||||||
image: torch.Tensor, center: tuple[int, int], radius: int
|
image: torch.Tensor,
|
||||||
|
center: tuple[int, int],
|
||||||
|
radius: int,
|
||||||
|
color_tensor: torch.Tensor,
|
||||||
):
|
):
|
||||||
"""Draw a 5px circle on the image."""
|
"""Draw a 5px circle on the image."""
|
||||||
x0, y0 = center
|
x0, y0 = center
|
||||||
for x in range(-radius, radius + 1):
|
h, w, _ = image.shape
|
||||||
for y in range(-radius, radius + 1):
|
min_x_bbox = max(0, x0 - radius)
|
||||||
in_radius = x**2 + y**2 <= radius**2
|
max_x_bbox = min(w - 1, x0 + radius)
|
||||||
in_bounds = (
|
min_y_bbox = max(0, y0 - radius)
|
||||||
0 <= x0 + x < image.shape[1]
|
max_y_bbox = min(h - 1, y0 + radius)
|
||||||
and 0 <= y0 + y < image.shape[0]
|
|
||||||
)
|
for py in range(min_y_bbox, max_y_bbox + 1):
|
||||||
if in_radius and in_bounds:
|
for px in range(min_x_bbox, max_x_bbox + 1):
|
||||||
image[y0 + y, x0 + x] = torch.tensor(
|
if (px - x0) ** 2 + (py - y0) ** 2 <= radius**2:
|
||||||
[255, 255, 255],
|
image[py, px] = color_tensor
|
||||||
dtype=torch.uint8,
|
|
||||||
device=image.device,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class MTB_ColorCorrectGPU:
|
class MTB_ColorCorrectGPU:
|
||||||
|
|||||||
@@ -21,7 +21,11 @@ class MTB_StackImages:
|
|||||||
"match_method": (
|
"match_method": (
|
||||||
["error", "smallest", "largest"],
|
["error", "smallest", "largest"],
|
||||||
{"default": "error"},
|
{"default": "error"},
|
||||||
)
|
),
|
||||||
|
"output_rgb": (
|
||||||
|
"BOOLEAN",
|
||||||
|
{"default": True, "tooltip": "Output RGB instead of RGBA"},
|
||||||
|
),
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -29,7 +33,7 @@ class MTB_StackImages:
|
|||||||
FUNCTION = "stack"
|
FUNCTION = "stack"
|
||||||
CATEGORY = "mtb/image utils"
|
CATEGORY = "mtb/image utils"
|
||||||
|
|
||||||
def stack(self, vertical, match_method="error", **kwargs):
|
def stack(self, vertical, match_method="error", output_rgb=True, **kwargs):
|
||||||
if not kwargs:
|
if not kwargs:
|
||||||
raise ValueError("At least one tensor must be provided.")
|
raise ValueError("At least one tensor must be provided.")
|
||||||
|
|
||||||
@@ -98,6 +102,9 @@ class MTB_StackImages:
|
|||||||
|
|
||||||
stacked_tensor = torch.cat(normalized_tensors, dim=dim)
|
stacked_tensor = torch.cat(normalized_tensors, dim=dim)
|
||||||
|
|
||||||
|
if output_rgb:
|
||||||
|
stacked_tensor = stacked_tensor[:, :, :, :3]
|
||||||
|
|
||||||
return (stacked_tensor,)
|
return (stacked_tensor,)
|
||||||
|
|
||||||
def normalize_to_rgba(self, tensor):
|
def normalize_to_rgba(self, tensor):
|
||||||
|
|||||||
+5
-38
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
|||||||
|
|
||||||
[project]
|
[project]
|
||||||
name = "comfy-mtb"
|
name = "comfy-mtb"
|
||||||
version = "0.3.0"
|
version = "0.6.0"
|
||||||
description = "Animation oriented nodes pack for ComfyUI."
|
description = "Animation oriented nodes pack for ComfyUI."
|
||||||
license = { text = "MIT" }
|
license = { text = "MIT" }
|
||||||
readme = "README.md"
|
readme = "README.md"
|
||||||
@@ -62,39 +62,6 @@ PublisherId = "mel"
|
|||||||
DisplayName = "comfy-mtb"
|
DisplayName = "comfy-mtb"
|
||||||
Icon = "https://avatars.githubusercontent.com/u/7041726?v=4"
|
Icon = "https://avatars.githubusercontent.com/u/7041726?v=4"
|
||||||
|
|
||||||
[tool.bumpversion]
|
|
||||||
current_version = "0.3.0"
|
|
||||||
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
|
|
||||||
serialize = ["{major}.{minor}.{patch}"]
|
|
||||||
search = "{current_version}"
|
|
||||||
replace = "{new_version}"
|
|
||||||
regex = false
|
|
||||||
ignore_missing_version = false
|
|
||||||
ignore_missing_files = false
|
|
||||||
tag = true
|
|
||||||
sign_tags = true
|
|
||||||
tag_name = "v{new_version}"
|
|
||||||
tag_message = "⬆️ Bump version: {current_version} → {new_version}"
|
|
||||||
allow_dirty = true
|
|
||||||
commit = true
|
|
||||||
message = "⬆️ Bump version: {current_version} → {new_version}"
|
|
||||||
commit_args = ""
|
|
||||||
|
|
||||||
[[tool.bumpversion.files]]
|
|
||||||
filename = "__init__.py"
|
|
||||||
search = "__version__ = \"{current_version}\""
|
|
||||||
replace = "__version__ = \"{new_version}\""
|
|
||||||
|
|
||||||
[[tool.bumpversion.files]]
|
|
||||||
filename = "pyproject.toml"
|
|
||||||
search = "version = \"{current_version}\""
|
|
||||||
replace = "version = \"{new_version}\""
|
|
||||||
|
|
||||||
# [[tool.bumpversion.files]]
|
|
||||||
# filename = "your_package/__init__.py"
|
|
||||||
# search = "__version__ = '{current_version}'"
|
|
||||||
# replace = "__version__ = '{new_version}'"
|
|
||||||
|
|
||||||
# INFO: All those remaining keys are meant for local dev
|
# INFO: All those remaining keys are meant for local dev
|
||||||
[tool.pyright]
|
[tool.pyright]
|
||||||
include = ["."]
|
include = ["."]
|
||||||
@@ -111,18 +78,18 @@ stubPath = "src/stubs"
|
|||||||
|
|
||||||
reportMissingImports = true
|
reportMissingImports = true
|
||||||
reportMissingTypeStubs = false
|
reportMissingTypeStubs = false
|
||||||
|
reportExplicitAny = false
|
||||||
typeCheckingMode = "basic"
|
typeCheckingMode = "basic"
|
||||||
|
pythonVersion = "3.11"
|
||||||
pythonVersion = "3.10"
|
|
||||||
pythonPlatform = "Windows"
|
pythonPlatform = "Windows"
|
||||||
|
|
||||||
|
|
||||||
[tool.pytest.ini_options]
|
[tool.pytest.ini_options]
|
||||||
log_level = "DEBUG"
|
log_level = "DEBUG"
|
||||||
log_cli = true
|
log_cli = true
|
||||||
markers = [
|
markers = [
|
||||||
"wip: tests that aren't fully finished yet",
|
"wip: tests that aren't fully finished yet",
|
||||||
"heavy: marks tests as heavy (deselect with '-m \"not heavy\"')",
|
'''heavy: marks tests as heavy (deselect with '-m "not heavy"')''',
|
||||||
|
|
||||||
]
|
]
|
||||||
filterwarnings = ["ignore::UserWarning", 'ignore::DeprecationWarning']
|
filterwarnings = ["ignore::UserWarning", 'ignore::DeprecationWarning']
|
||||||
|
|
||||||
|
|||||||
@@ -1,18 +0,0 @@
|
|||||||
{
|
|
||||||
"exclude": [
|
|
||||||
"**/node_modules",
|
|
||||||
"**/__pycache__",
|
|
||||||
],
|
|
||||||
"ignore": [
|
|
||||||
"extern"
|
|
||||||
],
|
|
||||||
"defineConstant": {
|
|
||||||
"DEBUG": true
|
|
||||||
},
|
|
||||||
"venvPath": "../../../.venv/",
|
|
||||||
"reportMissingImports": true,
|
|
||||||
"reportMissingTypeStubs": false,
|
|
||||||
"pythonVersion": "3.10",
|
|
||||||
"pythonPlatform": "All",
|
|
||||||
"reportOptionalMemberAccess": "none"
|
|
||||||
}
|
|
||||||
@@ -0,0 +1,637 @@
|
|||||||
|
import base64
|
||||||
|
import code
|
||||||
|
import io
|
||||||
|
import re
|
||||||
|
import sys
|
||||||
|
from contextlib import redirect_stderr, redirect_stdout
|
||||||
|
|
||||||
|
# import matplotlib.pyplot as plt
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
from aiohttp import web
|
||||||
|
from PIL import Image
|
||||||
|
from rich.console import Console
|
||||||
|
from rich.traceback import Traceback
|
||||||
|
|
||||||
|
from .log import log
|
||||||
|
|
||||||
|
try:
|
||||||
|
import pyflakes.api
|
||||||
|
import pyflakes.reporter
|
||||||
|
|
||||||
|
_HAS_LINT = True
|
||||||
|
except ImportError:
|
||||||
|
print(
|
||||||
|
"ComfyREPL: pyflakes not found. Linting will be disabled. Install with 'pip install pyflakes'."
|
||||||
|
)
|
||||||
|
_HAS_LINT = False
|
||||||
|
|
||||||
|
# --- Linting Library ---
|
||||||
|
# try:
|
||||||
|
# import ruff
|
||||||
|
# import ruff.lint
|
||||||
|
# import ruff.lint.linter
|
||||||
|
# import ruff.settings
|
||||||
|
#
|
||||||
|
# _HAS_LINT = True
|
||||||
|
# except ImportError:
|
||||||
|
# print(
|
||||||
|
# "ComfyREPL: ruff not found. Linting will be disabled. Install with 'pip install ruff'."
|
||||||
|
# )
|
||||||
|
# _HAS_LINT = False
|
||||||
|
|
||||||
|
# --- Audio/Video Libraries ---
|
||||||
|
try:
|
||||||
|
import scipy.io.wavfile
|
||||||
|
|
||||||
|
_HAS_SCIPY = True
|
||||||
|
except ImportError:
|
||||||
|
print(
|
||||||
|
"ComfyREPL: SciPy not found. Audio display will be disabled. Install with 'pip install scipy'."
|
||||||
|
)
|
||||||
|
_HAS_SCIPY = False
|
||||||
|
|
||||||
|
try:
|
||||||
|
import imageio
|
||||||
|
import imageio.plugins.ffmpeg # Ensure ffmpeg plugin is available
|
||||||
|
|
||||||
|
_HAS_IMAGEIO = True
|
||||||
|
except ImportError:
|
||||||
|
print(
|
||||||
|
"ComfyREPL: Imageio or imageio-ffmpeg not found. Video display will be disabled. Install with 'pip install imageio imageio-ffmpeg'."
|
||||||
|
)
|
||||||
|
_HAS_IMAGEIO = False
|
||||||
|
|
||||||
|
|
||||||
|
# --- Audio Display ---
|
||||||
|
class AudioDisplay:
|
||||||
|
def __init__(self, samples, sample_rate):
|
||||||
|
if not _HAS_SCIPY:
|
||||||
|
raise ImportError("Audio display requires scipy and numpy.")
|
||||||
|
if not isinstance(samples, (np.ndarray, torch.Tensor)):
|
||||||
|
raise TypeError(
|
||||||
|
"Audio samples must be a numpy array or torch tensor."
|
||||||
|
)
|
||||||
|
if isinstance(samples, torch.Tensor):
|
||||||
|
samples = samples.detach().cpu().numpy()
|
||||||
|
|
||||||
|
# Ensure samples are in a format scipy.io.wavfile can handle (e.g., int16, float32)
|
||||||
|
if samples.dtype == np.float64:
|
||||||
|
samples = samples.astype(np.float32)
|
||||||
|
elif samples.dtype == np.int64:
|
||||||
|
# Or scale to int32 if range requires
|
||||||
|
samples = samples.astype(np.int16)
|
||||||
|
|
||||||
|
self.samples = samples
|
||||||
|
self.sample_rate = sample_rate
|
||||||
|
|
||||||
|
def _to_wav_base64(self):
|
||||||
|
buffer = io.BytesIO()
|
||||||
|
try:
|
||||||
|
scipy.io.wavfile.write(buffer, self.sample_rate, self.samples)
|
||||||
|
audio_base64 = base64.b64encode(buffer.getvalue()).decode("utf-8")
|
||||||
|
return audio_base64
|
||||||
|
except Exception as e:
|
||||||
|
return f"<div style='color: red;'>Error encoding audio: {e}</div>"
|
||||||
|
|
||||||
|
def _repr_html_(self):
|
||||||
|
base64_data = self._to_wav_base64()
|
||||||
|
if base64_data.startswith("<div"):
|
||||||
|
return base64_data
|
||||||
|
return f'<audio controls src="data:audio/wav;base64,{base64_data}" style="margin: 5px 0;"/>'
|
||||||
|
|
||||||
|
|
||||||
|
def render_audio(samples, sample_rate):
|
||||||
|
"""
|
||||||
|
Render audio samples as an HTML audio player.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
samples (np.ndarray or torch.Tensor): Audio samples.
|
||||||
|
sample_rate (int): Sample rate in Hz.
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
AudioDisplay: An object that will render as an HTML audio player.
|
||||||
|
"""
|
||||||
|
return AudioDisplay(samples, sample_rate)
|
||||||
|
|
||||||
|
|
||||||
|
# --- Display Classes ---
|
||||||
|
class VideoDisplay:
|
||||||
|
def __init__(self, frames, fps=24, options=None):
|
||||||
|
if not _HAS_IMAGEIO: # numpy/PIL/torch needed for frames
|
||||||
|
raise ImportError(
|
||||||
|
"Video display requires imageio, imageio-ffmpeg, and image libraries (numpy, Pillow, torch)."
|
||||||
|
)
|
||||||
|
|
||||||
|
self.frames = []
|
||||||
|
for frame in frames:
|
||||||
|
if isinstance(frame, Image.Image):
|
||||||
|
self.frames.append(np.array(frame))
|
||||||
|
elif isinstance(frame, np.ndarray):
|
||||||
|
# Ensure HWC and uint8
|
||||||
|
if frame.ndim == 3 and frame.shape[0] in [1, 3, 4]: # CHW
|
||||||
|
frame = np.transpose(frame, (1, 2, 0))
|
||||||
|
if frame.dtype != np.uint8:
|
||||||
|
frame = (
|
||||||
|
(frame * 255).astype(np.uint8)
|
||||||
|
if frame.max() <= 1.0
|
||||||
|
else frame.astype(np.uint8)
|
||||||
|
)
|
||||||
|
self.frames.append(frame)
|
||||||
|
elif isinstance(frame, torch.Tensor):
|
||||||
|
np_frame = frame.detach().cpu().numpy()
|
||||||
|
if np_frame.ndim == 3 and np_frame.shape[0] in [
|
||||||
|
1,
|
||||||
|
3,
|
||||||
|
4,
|
||||||
|
]: # CHW
|
||||||
|
np_frame = np.transpose(np_frame, (1, 2, 0))
|
||||||
|
if np_frame.dtype != np.uint8:
|
||||||
|
np_frame = (
|
||||||
|
(np_frame * 255).astype(np.uint8)
|
||||||
|
if np_frame.max() <= 1.0
|
||||||
|
else np_frame.astype(np.uint8)
|
||||||
|
)
|
||||||
|
self.frames.append(np_frame)
|
||||||
|
else:
|
||||||
|
raise TypeError(
|
||||||
|
f"Unsupported frame type: {type(frame)}. Must be PIL.Image, numpy.ndarray, or torch.Tensor."
|
||||||
|
)
|
||||||
|
|
||||||
|
self.fps = fps
|
||||||
|
self.options = options if options is not None else {}
|
||||||
|
|
||||||
|
def _to_mp4_base64(self):
|
||||||
|
buffer = io.BytesIO()
|
||||||
|
try:
|
||||||
|
# Use imageio to write frames to an in-memory MP4 file
|
||||||
|
imageio.mimwrite(
|
||||||
|
buffer,
|
||||||
|
self.frames,
|
||||||
|
format="mp4",
|
||||||
|
fps=self.fps,
|
||||||
|
codec="libx264",
|
||||||
|
quality=8,
|
||||||
|
) # quality 1-10
|
||||||
|
video_base64 = base64.b64encode(buffer.getvalue()).decode("utf-8")
|
||||||
|
return video_base64
|
||||||
|
except Exception as e:
|
||||||
|
return f"<div style='color: red;'>Error encoding video: {e}</div>"
|
||||||
|
|
||||||
|
def _repr_html_(self):
|
||||||
|
base64_data = self._to_mp4_base64()
|
||||||
|
if base64_data.startswith("<div"): # Check if it's an error message
|
||||||
|
return base64_data
|
||||||
|
|
||||||
|
# Build HTML options string
|
||||||
|
option_str = ""
|
||||||
|
for key, value in self.options.items():
|
||||||
|
if isinstance(value, bool) and value:
|
||||||
|
option_str += f" {key}"
|
||||||
|
elif isinstance(value, str):
|
||||||
|
option_str += f' {key}="{value}"'
|
||||||
|
else:
|
||||||
|
option_str += f' {key}="{value}"' # Fallback for numbers etc.
|
||||||
|
|
||||||
|
return f'<video controls src="data:video/mp4;base64,{base64_data}" style="max-width: 100%; height: auto; border: 1px solid #555; margin: 5px 0;"{option_str}/>'
|
||||||
|
|
||||||
|
|
||||||
|
def render_video(batch_tensor_or_array_of_pil_images, fps=24, options=None):
|
||||||
|
"""
|
||||||
|
Render video frames as an HTML video player.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
batch_tensor_or_array_of_pil_images (list of PIL.Image, np.ndarray, or torch.Tensor):
|
||||||
|
A list of frames, or a single batch tensor/array (B, H, W, C) or (B, C, H, W).
|
||||||
|
fps (int): Frames per second.
|
||||||
|
options (dict): Dictionary of HTML <video> tag attributes (e.g., {"loop": True, "autoplay": True}).
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
VideoDisplay: An object that will render as an HTML video player.
|
||||||
|
"""
|
||||||
|
frames_list = []
|
||||||
|
if isinstance(
|
||||||
|
batch_tensor_or_array_of_pil_images, (np.ndarray, torch.Tensor)
|
||||||
|
):
|
||||||
|
# Assume it's a batch tensor/array
|
||||||
|
for i in range(batch_tensor_or_array_of_pil_images.shape[0]):
|
||||||
|
frames_list.append(batch_tensor_or_array_of_pil_images[i])
|
||||||
|
elif isinstance(batch_tensor_or_array_of_pil_images, list):
|
||||||
|
frames_list = batch_tensor_or_array_of_pil_images
|
||||||
|
else:
|
||||||
|
raise TypeError(
|
||||||
|
"Input for render_video must be a list of frames or a batch tensor/array."
|
||||||
|
)
|
||||||
|
|
||||||
|
return VideoDisplay(frames_list, fps, options)
|
||||||
|
|
||||||
|
|
||||||
|
class ComfyREPLBackend:
|
||||||
|
def __init__(self):
|
||||||
|
self.repl_consoles: dict[str, code.InteractiveConsole] = {}
|
||||||
|
# self.repl_console = None
|
||||||
|
self.image_outputs = []
|
||||||
|
self.audio_outputs = []
|
||||||
|
self.video_outputs = []
|
||||||
|
self._original_displayhook = sys.displayhook
|
||||||
|
# self._init_repl_console()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _init_repl_console():
|
||||||
|
"""Define the globals that will be available in the REPL session."""
|
||||||
|
repl_globals = {"__builtins__": __builtins__}
|
||||||
|
# repl_globals["plt"] = plt
|
||||||
|
repl_globals["np"] = np
|
||||||
|
repl_globals["Image"] = Image
|
||||||
|
repl_globals["torch"] = torch
|
||||||
|
|
||||||
|
repl_globals["repl_display"] = _repl_display_image
|
||||||
|
if _HAS_SCIPY:
|
||||||
|
repl_globals["render_audio"] = render_audio
|
||||||
|
if _HAS_IMAGEIO:
|
||||||
|
repl_globals["render_video"] = render_video
|
||||||
|
|
||||||
|
return code.InteractiveConsole(locals=repl_globals)
|
||||||
|
|
||||||
|
def _custom_displayhook(self, value):
|
||||||
|
"""
|
||||||
|
Displayhook that capture and process image, audio, video objects.
|
||||||
|
|
||||||
|
For other objects, it fallsback to the original displayhook.
|
||||||
|
"""
|
||||||
|
if value is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
# Attempt to handle as an image
|
||||||
|
if (
|
||||||
|
isinstance(value, (Image.Image, np.ndarray, torch.Tensor))
|
||||||
|
# or (
|
||||||
|
# hasattr(value, "figure")
|
||||||
|
# and isinstance(value.figure, plt.Figure)
|
||||||
|
# )
|
||||||
|
# or isinstance(value, plt.Figure)
|
||||||
|
):
|
||||||
|
img_html = _repl_display_image(value)
|
||||||
|
self.image_outputs.append(img_html)
|
||||||
|
return
|
||||||
|
|
||||||
|
# Attempt to handle as audio
|
||||||
|
elif isinstance(value, AudioDisplay):
|
||||||
|
audio_html = value._repr_html_()
|
||||||
|
self.audio_outputs.append(audio_html)
|
||||||
|
return
|
||||||
|
|
||||||
|
# Attempt to handle as video
|
||||||
|
elif isinstance(value, VideoDisplay):
|
||||||
|
video_html = value._repr_html_()
|
||||||
|
self.video_outputs.append(video_html)
|
||||||
|
return
|
||||||
|
|
||||||
|
else:
|
||||||
|
# If not a special media type, let the original displayhook handle it.
|
||||||
|
self._original_displayhook(value)
|
||||||
|
|
||||||
|
def _console_to_html(
|
||||||
|
self, stream: io.StringIO | Traceback, width: int = 120
|
||||||
|
) -> str:
|
||||||
|
if isinstance(stream, io.StringIO):
|
||||||
|
captured_text_output = stream.getvalue()
|
||||||
|
else:
|
||||||
|
captured_text_output = Traceback
|
||||||
|
|
||||||
|
html_console = Console(
|
||||||
|
file=io.StringIO(), record=True, force_terminal=True, width=width
|
||||||
|
)
|
||||||
|
html_console.print(captured_text_output)
|
||||||
|
|
||||||
|
return html_console.export_html(inline_styles=True)
|
||||||
|
|
||||||
|
def get_console(self, node_name: str):
|
||||||
|
console = self.repl_consoles.get(node_name)
|
||||||
|
if console:
|
||||||
|
return console
|
||||||
|
|
||||||
|
console = self._init_repl_console()
|
||||||
|
self.repl_consoles[node_name] = console
|
||||||
|
return self.repl_consoles[node_name]
|
||||||
|
|
||||||
|
def execute_code(self, node_name: str, code: str):
|
||||||
|
# Clear outputs from previous execution
|
||||||
|
self.image_outputs = []
|
||||||
|
self.audio_outputs = []
|
||||||
|
self.video_outputs = []
|
||||||
|
|
||||||
|
output_html = ""
|
||||||
|
error_message = None
|
||||||
|
|
||||||
|
repl_console = self.get_console(node_name)
|
||||||
|
|
||||||
|
string_io = io.StringIO()
|
||||||
|
|
||||||
|
# Temporarily patch sys.displayhook
|
||||||
|
sys.displayhook = self._custom_displayhook
|
||||||
|
|
||||||
|
try:
|
||||||
|
with redirect_stdout(string_io), redirect_stderr(string_io):
|
||||||
|
for line in code.splitlines():
|
||||||
|
repl_console.push(line)
|
||||||
|
|
||||||
|
full_rich_html = self._console_to_html(string_io)
|
||||||
|
match = re.search(
|
||||||
|
r"<body.*?>(.*?)</body>", full_rich_html, re.DOTALL
|
||||||
|
)
|
||||||
|
if match:
|
||||||
|
output_html = match.group(1)
|
||||||
|
else:
|
||||||
|
output_html = full_rich_html
|
||||||
|
|
||||||
|
# Append any captured media HTML *after* the rich text output
|
||||||
|
for img_html in self.image_outputs:
|
||||||
|
output_html += img_html
|
||||||
|
for audio_html in self.audio_outputs:
|
||||||
|
output_html += audio_html
|
||||||
|
for video_html in self.video_outputs:
|
||||||
|
output_html += video_html
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
exc_type, exc_value, exc_traceback = sys.exc_info()
|
||||||
|
rich_traceback = Traceback.from_exception(
|
||||||
|
exc_type,
|
||||||
|
exc_value,
|
||||||
|
exc_traceback,
|
||||||
|
show_locals=True,
|
||||||
|
suppress=[__file__],
|
||||||
|
)
|
||||||
|
# error_console = Console(
|
||||||
|
# file=io.StringIO(), record=True, force_terminal=True, width=120
|
||||||
|
# )
|
||||||
|
# error_console.print(rich_traceback)
|
||||||
|
# full_error_html = error_console.export_html(inline_styles=True)
|
||||||
|
|
||||||
|
full_error_html = self._console_to_html(rich_traceback)
|
||||||
|
|
||||||
|
match = re.search(
|
||||||
|
r"<body.*?>(.*?)</body>", full_error_html, re.DOTALL
|
||||||
|
)
|
||||||
|
output_html = match.group(1) if match else full_error_html
|
||||||
|
|
||||||
|
error_message = str(e)
|
||||||
|
finally:
|
||||||
|
sys.displayhook = (
|
||||||
|
self._original_displayhook
|
||||||
|
) # Always restore original displayhook
|
||||||
|
|
||||||
|
return {"output_html": output_html, "error": error_message}
|
||||||
|
|
||||||
|
def lint_code(self, node_name: str, code: str):
|
||||||
|
diagnostics = []
|
||||||
|
|
||||||
|
if not _HAS_LINT:
|
||||||
|
diagnostics.append(
|
||||||
|
{
|
||||||
|
"row": 0,
|
||||||
|
"column": 0,
|
||||||
|
"text": "Pyflakes not installed. Linting disabled. Install with 'uv add pyflakes'.",
|
||||||
|
"type": "warning",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
return web.json_response({"diagnostics": diagnostics})
|
||||||
|
|
||||||
|
# Use a custom reporter to capture messages
|
||||||
|
class PyflakesReporter(pyflakes.reporter.Reporter):
|
||||||
|
def __init__(self):
|
||||||
|
self.messages = []
|
||||||
|
# Suppress stdout/stderr from pyflakes itself
|
||||||
|
self._stdout = io.StringIO()
|
||||||
|
self._stderr = io.StringIO()
|
||||||
|
super().__init__(self._stdout, self._stderr)
|
||||||
|
|
||||||
|
def flake(self, message):
|
||||||
|
# Ace editor expects 0-indexed row, pyflakes gives 1-indexed lineno
|
||||||
|
self.messages.append(
|
||||||
|
{
|
||||||
|
"row": message.lineno - 1,
|
||||||
|
"column": message.col,
|
||||||
|
"text": str(message),
|
||||||
|
"type": "warning", # pyflakes usually gives warnings
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
def unexpectedError(self, filename, msg):
|
||||||
|
self.messages.append(
|
||||||
|
{
|
||||||
|
"row": 0,
|
||||||
|
"column": 0,
|
||||||
|
"text": f"Pyflakes internal error: {msg}",
|
||||||
|
"type": "error",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
def syntaxError(self, filename, msg, lineno, offset, text):
|
||||||
|
log.info(f"Received {text} to syntax error")
|
||||||
|
self.messages.append(
|
||||||
|
{
|
||||||
|
"row": lineno - 1, # Ace is 0-indexed
|
||||||
|
"column": offset,
|
||||||
|
"text": f"Syntax Error: {msg}",
|
||||||
|
"type": "error",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
reporter = PyflakesReporter()
|
||||||
|
pyflakes.api.check(code, node_name, reporter)
|
||||||
|
|
||||||
|
return {"diagnostics": reporter.messages}
|
||||||
|
|
||||||
|
def lint_code_ruff(self, code: str):
|
||||||
|
diagnostics = []
|
||||||
|
|
||||||
|
if not _HAS_LINT:
|
||||||
|
diagnostics.append(
|
||||||
|
{
|
||||||
|
"row": 0,
|
||||||
|
"column": 0,
|
||||||
|
"text": "Ruff not installed. Linting disabled. Install with 'pip install ruff'.",
|
||||||
|
"type": "warning",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
return {"diagnostics": diagnostics}
|
||||||
|
|
||||||
|
# Define the builtins/globals that Ruff should recognize
|
||||||
|
# These are the names we inject into the REPL's scope
|
||||||
|
repl_builtins = [
|
||||||
|
"repl_display",
|
||||||
|
"render_audio",
|
||||||
|
"render_video",
|
||||||
|
# "plt",
|
||||||
|
"np",
|
||||||
|
"Image",
|
||||||
|
"torch",
|
||||||
|
]
|
||||||
|
|
||||||
|
try:
|
||||||
|
# Lint the code using Ruff's programmatic API
|
||||||
|
result = ruff.lint.linter.lint_stdin(
|
||||||
|
code.encode("utf-8"),
|
||||||
|
path="<stdin>",
|
||||||
|
builtins=repl_builtins,
|
||||||
|
)
|
||||||
|
|
||||||
|
for diagnostic in result.diagnostics:
|
||||||
|
diag_type = "warning" # Default
|
||||||
|
# Ruff's error codes: F (Pyflakes), E (Pycodestyle), W (Pycodestyle warning), I (isort), N (naming), etc.
|
||||||
|
# F821: Undefined name (often an error)
|
||||||
|
if (
|
||||||
|
diagnostic.kind.code.startswith("E")
|
||||||
|
or diagnostic.kind.code == "F821"
|
||||||
|
):
|
||||||
|
diag_type = "error"
|
||||||
|
elif diagnostic.kind.code.startswith("W"):
|
||||||
|
diag_type = "warning"
|
||||||
|
|
||||||
|
diagnostics.append(
|
||||||
|
{
|
||||||
|
"row": diagnostic.location.row - 1, # Ace is 0-indexed
|
||||||
|
"column": diagnostic.location.column
|
||||||
|
- 1, # Ace is 0-indexed
|
||||||
|
"text": diagnostic.message,
|
||||||
|
"type": diag_type,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
diagnostics.append(
|
||||||
|
{
|
||||||
|
"row": 0,
|
||||||
|
"column": 0,
|
||||||
|
"text": f"Ruff internal error: {e}",
|
||||||
|
"type": "error",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
return {"diagnostics": diagnostics}
|
||||||
|
|
||||||
|
|
||||||
|
def _repl_display_image(img_data):
|
||||||
|
"""
|
||||||
|
Internal function to convert image data (PIL, numpy, torch, matplotlib) to base64 HTML.
|
||||||
|
"""
|
||||||
|
pil_img = None
|
||||||
|
# fig = None
|
||||||
|
|
||||||
|
if isinstance(img_data, Image.Image):
|
||||||
|
pil_img = img_data
|
||||||
|
elif isinstance(img_data, np.ndarray):
|
||||||
|
# Handle different numpy array shapes (HWC, CHW)
|
||||||
|
if img_data.ndim == 3:
|
||||||
|
if img_data.shape[0] in [1, 3, 4]: # Likely CHW
|
||||||
|
if img_data.shape[0] == 1: # Grayscale
|
||||||
|
img_data = img_data.squeeze(0)
|
||||||
|
else: # Color
|
||||||
|
img_data = np.transpose(img_data, (1, 2, 0)) # CHW to HWC
|
||||||
|
# Ensure it's uint8 for PIL, assuming float [0,1] or int [0,255]
|
||||||
|
if img_data.dtype != np.uint8:
|
||||||
|
img_data = (
|
||||||
|
(img_data * 255).astype(np.uint8)
|
||||||
|
if img_data.max() <= 1.0
|
||||||
|
else img_data.astype(np.uint8)
|
||||||
|
)
|
||||||
|
pil_img = Image.fromarray(img_data)
|
||||||
|
elif isinstance(img_data, torch.Tensor):
|
||||||
|
# Move to CPU, convert to numpy
|
||||||
|
np_img = img_data.detach().cpu().numpy()
|
||||||
|
# Handle different tensor shapes (CHW, HWC)
|
||||||
|
if np_img.ndim == 3:
|
||||||
|
if np_img.shape[0] in [1, 3, 4]: # Likely CHW
|
||||||
|
if np_img.shape[0] == 1: # Grayscale
|
||||||
|
np_img = np_img.squeeze(0)
|
||||||
|
else: # Color
|
||||||
|
np_img = np.transpose(np_img, (1, 2, 0)) # CHW to HWC
|
||||||
|
# Ensure it's uint8 for PIL, assuming float [0,1] or int [0,255]
|
||||||
|
if np_img.dtype != np.uint8:
|
||||||
|
np_img = (
|
||||||
|
(np_img * 255).astype(np.uint8)
|
||||||
|
if np_img.max() <= 1.0
|
||||||
|
else np_img.astype(np.uint8)
|
||||||
|
)
|
||||||
|
pil_img = Image.fromarray(np_img)
|
||||||
|
# elif hasattr(img_data, "figure") and isinstance(
|
||||||
|
# img_data.figure, plt.Figure
|
||||||
|
# ):
|
||||||
|
# # If it's a matplotlib Axes object, get its figure
|
||||||
|
# fig = img_data.figure
|
||||||
|
# elif isinstance(img_data, plt.Figure):
|
||||||
|
# fig = img_data
|
||||||
|
else:
|
||||||
|
return f"<div style='color: red;'>Unsupported image type for display: {type(img_data)}</div>"
|
||||||
|
|
||||||
|
buffer = io.BytesIO()
|
||||||
|
try:
|
||||||
|
if pil_img:
|
||||||
|
pil_img.save(buffer, format="PNG")
|
||||||
|
# elif fig:
|
||||||
|
# fig.savefig(
|
||||||
|
# buffer, format="PNG", bbox_inches="tight", pad_inches=0.1
|
||||||
|
# )
|
||||||
|
# plt.close(
|
||||||
|
# fig
|
||||||
|
# ) # Close the figure to prevent it from showing up in other contexts
|
||||||
|
else:
|
||||||
|
return (
|
||||||
|
"<div style='color: red;'>Could not process image data.</div>"
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
return f"<div style='color: red;'>Error saving image: {e}</div>"
|
||||||
|
|
||||||
|
img_base64 = base64.b64encode(buffer.getvalue()).decode("utf-8")
|
||||||
|
return f'<img src="data:image/png;base64,{img_base64}" style="max-width: 100%; height: auto; border: 1px solid #555; margin: 5px 0;"/>'
|
||||||
|
|
||||||
|
|
||||||
|
# Instantiate the backend class globally
|
||||||
|
_comfy_repl_backend = ComfyREPLBackend()
|
||||||
|
|
||||||
|
|
||||||
|
# Update aiohttp handlers to use the backend instance
|
||||||
|
async def repl_execute_code_handler(request):
|
||||||
|
data = await request.json()
|
||||||
|
|
||||||
|
name = data.get("name")
|
||||||
|
|
||||||
|
if name is None: # we send an error
|
||||||
|
return web.Response(
|
||||||
|
status=417, reason="Expectation Failed", text="Missing name key"
|
||||||
|
)
|
||||||
|
code = data.get("code", "")
|
||||||
|
result = _comfy_repl_backend.execute_code(name, code)
|
||||||
|
return web.json_response(result)
|
||||||
|
|
||||||
|
|
||||||
|
async def repl_lint_code_handler(request):
|
||||||
|
data = await request.json()
|
||||||
|
name = data.get("name")
|
||||||
|
|
||||||
|
if name is None: # we send an error
|
||||||
|
return web.Response(
|
||||||
|
status=417, reason="Expectation Failed", text="Missing name key"
|
||||||
|
)
|
||||||
|
# raise web.HTTPExpectationFailed(
|
||||||
|
# reason="Missing name key (reason)", text="Missing name key (text)"
|
||||||
|
# )
|
||||||
|
|
||||||
|
code = data.get("code", "")
|
||||||
|
result = _comfy_repl_backend.lint_code(name, code)
|
||||||
|
return web.json_response(result)
|
||||||
|
|
||||||
|
|
||||||
|
def setup_custom_web_routes(app: web.Application):
|
||||||
|
"""
|
||||||
|
Function to register our custom web routes with the ComfyUI server.
|
||||||
|
"""
|
||||||
|
log.info("ComfyREPL: Registering /mtb/execute route...")
|
||||||
|
app.router.add_post("/mtb/execute", repl_execute_code_handler)
|
||||||
|
app.router.add_post("/mtb/lint", repl_lint_code_handler)
|
||||||
|
|
||||||
|
|
||||||
|
# You can add more routes here if needed, e.g., for clearing state.
|
||||||
@@ -1,6 +1,5 @@
|
|||||||
import contextlib
|
import contextlib
|
||||||
import functools
|
import functools
|
||||||
import importlib
|
|
||||||
import math
|
import math
|
||||||
import operator
|
import operator
|
||||||
import os
|
import os
|
||||||
@@ -9,11 +8,14 @@ import shutil
|
|||||||
import socket
|
import socket
|
||||||
import subprocess
|
import subprocess
|
||||||
import sys
|
import sys
|
||||||
|
import textwrap
|
||||||
import uuid
|
import uuid
|
||||||
|
import warnings
|
||||||
from collections.abc import Callable, Sequence
|
from collections.abc import Callable, Sequence
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from functools import reduce
|
from functools import reduce
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from types import EllipsisType
|
||||||
from typing import TypeVar
|
from typing import TypeVar
|
||||||
from urllib.parse import urlparse
|
from urllib.parse import urlparse
|
||||||
|
|
||||||
@@ -462,25 +464,6 @@ def _run_command(shell_cmd, ignored_lines_start):
|
|||||||
print("Command executed successfully!")
|
print("Command executed successfully!")
|
||||||
|
|
||||||
|
|
||||||
def import_install(package_name):
|
|
||||||
package_spec = reqs_map.get(package_name, package_name)
|
|
||||||
|
|
||||||
try:
|
|
||||||
importlib.import_module(package_name)
|
|
||||||
|
|
||||||
except Exception: # (ImportError, ModuleNotFoundError):
|
|
||||||
run_command(
|
|
||||||
[
|
|
||||||
Path(sys.executable).as_posix(),
|
|
||||||
"-m",
|
|
||||||
"pip",
|
|
||||||
"install",
|
|
||||||
package_spec,
|
|
||||||
]
|
|
||||||
)
|
|
||||||
importlib.import_module(package_name)
|
|
||||||
|
|
||||||
|
|
||||||
# endregion
|
# endregion
|
||||||
|
|
||||||
|
|
||||||
@@ -544,11 +527,199 @@ PIL_FILTER_MAP = {
|
|||||||
|
|
||||||
|
|
||||||
# region TENSOR Utilities
|
# region TENSOR Utilities
|
||||||
|
|
||||||
|
|
||||||
|
class LazyProxyTensor:
|
||||||
|
"""Memory-efficient proxy that wrap a tensor but presents itself as a different dtype (e.g., float32).
|
||||||
|
|
||||||
|
It mimics a torch.Tensor's read-only attributes and methods. Data conversion
|
||||||
|
and normalization happen lazily on access (e.g., via slicing), avoiding
|
||||||
|
the high memory cost of a full conversion.
|
||||||
|
|
||||||
|
Supported source dtypes:
|
||||||
|
- torch.uint8 (normalized from [0, 255])
|
||||||
|
- torch.uint16 (normalized from [0, 65535])
|
||||||
|
- All float types (passed through, assumed to be in [0, 1] range)"
|
||||||
|
"""
|
||||||
|
|
||||||
|
_source_tensor: torch.Tensor
|
||||||
|
_target_dtype: torch.dtype
|
||||||
|
_target_element_size: int
|
||||||
|
_scale_divisor: float
|
||||||
|
_warned_inefficient_access: bool
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self, source_tensor, target_dtype=torch.float32, target_device=None
|
||||||
|
):
|
||||||
|
if not isinstance(source_tensor, torch.Tensor):
|
||||||
|
raise ValueError("Input must be a torch.Tensor.")
|
||||||
|
|
||||||
|
self._source_tensor = source_tensor
|
||||||
|
self._target_dtype = target_dtype
|
||||||
|
self._target_device = (
|
||||||
|
target_device
|
||||||
|
if target_device is not None
|
||||||
|
else source_tensor.device
|
||||||
|
)
|
||||||
|
|
||||||
|
# Determine the normalization divisor based on source dtype
|
||||||
|
# fmt: off
|
||||||
|
if source_tensor.dtype == torch.uint8: self._scale_divisor = 255.0
|
||||||
|
elif source_tensor.dtype == torch.uint16: self._scale_divisor = 65535.0
|
||||||
|
elif torch.is_floating_point(source_tensor): self._scale_divisor = 1.0
|
||||||
|
else: raise ValueError(f"Unsupported source dtype for LazyProxyTensor: {source_tensor.dtype}")
|
||||||
|
# fmt: on
|
||||||
|
|
||||||
|
self._target_element_size = torch.empty(
|
||||||
|
(), dtype=self._target_dtype
|
||||||
|
).element_size()
|
||||||
|
self._warned_inefficient_access = False
|
||||||
|
|
||||||
|
def is_contiguous(self, *args, **kwargs):
|
||||||
|
return self._source_tensor.is_contiguous(*args, **kwargs)
|
||||||
|
|
||||||
|
def stride(self, *args, **kwargs):
|
||||||
|
return self._source_tensor.stride(*args, **kwargs)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def shape(self):
|
||||||
|
return self._source_tensor.shape
|
||||||
|
|
||||||
|
@property
|
||||||
|
def requires_grad(self):
|
||||||
|
return False
|
||||||
|
|
||||||
|
def nelement(self):
|
||||||
|
"""Return the total number of elements in the (pretend) tensor."""
|
||||||
|
return self._source_tensor.nelement()
|
||||||
|
|
||||||
|
def element_size(self):
|
||||||
|
"""Return the size in bytes of an individual (pretend) float element."""
|
||||||
|
return self._target_element_size
|
||||||
|
|
||||||
|
@property
|
||||||
|
def dtype(self):
|
||||||
|
return self._target_dtype
|
||||||
|
|
||||||
|
@property
|
||||||
|
def device(self):
|
||||||
|
return self._target_device
|
||||||
|
|
||||||
|
def __len__(self):
|
||||||
|
return self._source_tensor.shape[0]
|
||||||
|
|
||||||
|
def __getitem__(self, key):
|
||||||
|
if (
|
||||||
|
self._source_tensor.device != self._target_device
|
||||||
|
and not self._warned_inefficient_access
|
||||||
|
):
|
||||||
|
warnings.warn(
|
||||||
|
"Inefficient access pattern detected for LazyProxyTensor. "
|
||||||
|
"You are slicing a device-proxied tensor, which causes slow, "
|
||||||
|
"repeated data transfers. For performance, use the .iter_chunks() method."
|
||||||
|
)
|
||||||
|
self._warned_inefficient_access = True
|
||||||
|
|
||||||
|
subset = self._source_tensor[key]
|
||||||
|
|
||||||
|
return (
|
||||||
|
subset.to(self._target_device).to(self._target_dtype)
|
||||||
|
/ self._scale_divisor
|
||||||
|
)
|
||||||
|
|
||||||
|
# def __iter__(self):
|
||||||
|
# for i in range(len(self)):
|
||||||
|
# yield self[i]
|
||||||
|
|
||||||
|
def iter_chunks(self, chunk_size=16):
|
||||||
|
for i in range(0, len(self), chunk_size):
|
||||||
|
chunk = self._source_tensor[i : i + chunk_size]
|
||||||
|
yield (
|
||||||
|
chunk.to(self._target_device, non_blocking=True).to(
|
||||||
|
self._target_dtype
|
||||||
|
)
|
||||||
|
/ self._scale_divisor
|
||||||
|
)
|
||||||
|
|
||||||
|
def squeeze(self, dim: str | EllipsisType | None = None):
|
||||||
|
squeezed = self._source_tensor.squeeze(dim)
|
||||||
|
return LazyProxyTensor(squeezed, self._target_dtype)
|
||||||
|
|
||||||
|
def unsqueeze(self, dim: int = 0):
|
||||||
|
unsqueezed = self._source_tensor.unsqueeze(dim)
|
||||||
|
return LazyProxyTensor(unsqueezed, self._target_dtype)
|
||||||
|
|
||||||
|
def repeat(self, *sizes):
|
||||||
|
repeated = self._source_tensor.repeat(*sizes)
|
||||||
|
return LazyProxyTensor(repeated, self._target_dtype)
|
||||||
|
|
||||||
|
def _format_mem_size(self, mem_bytes):
|
||||||
|
if mem_bytes > 1e9:
|
||||||
|
return f"{mem_bytes / 1e9:.2f} GB"
|
||||||
|
if mem_bytes > 1e6:
|
||||||
|
return f"{mem_bytes / 1e6:.2f} MB"
|
||||||
|
if mem_bytes > 1e3:
|
||||||
|
return f"{mem_bytes / 1e3:.2f} KB"
|
||||||
|
return f"{mem_bytes} B"
|
||||||
|
|
||||||
|
def __repr__(self):
|
||||||
|
|
||||||
|
actual_info = get_torch_tensor_info(self._source_tensor, name="Source")
|
||||||
|
target_info = get_torch_tensor_info(self, name="Target")
|
||||||
|
|
||||||
|
info = f"""
|
||||||
|
{target_info}
|
||||||
|
|
||||||
|
{actual_info}
|
||||||
|
"""
|
||||||
|
return textwrap.dedent(info).strip()
|
||||||
|
|
||||||
|
|
||||||
|
def get_torch_tensor_info(
|
||||||
|
tensor: torch.Tensor | LazyProxyTensor | np.ndarray,
|
||||||
|
*,
|
||||||
|
name: str | None = None,
|
||||||
|
):
|
||||||
|
mem_str = "N/A"
|
||||||
|
|
||||||
|
is_tensor = isinstance(tensor, torch.Tensor | LazyProxyTensor)
|
||||||
|
if is_tensor:
|
||||||
|
mem_bytes = tensor.element_size() * tensor.nelement()
|
||||||
|
else:
|
||||||
|
mem_bytes = tensor.itemsize * tensor.size
|
||||||
|
|
||||||
|
if mem_bytes > 1e9:
|
||||||
|
mem_str = f"{mem_bytes / 1e9:.2f} GB"
|
||||||
|
elif mem_bytes > 1e6:
|
||||||
|
mem_str = f"{mem_bytes / 1e6:.2f} MB"
|
||||||
|
elif mem_bytes > 1e3:
|
||||||
|
mem_str = f"{mem_bytes / 1e3:.2f} KB"
|
||||||
|
else:
|
||||||
|
mem_str = f"{mem_bytes} B"
|
||||||
|
|
||||||
|
device = "N/A"
|
||||||
|
grad = "False"
|
||||||
|
type_name = name or "Tensor" if is_tensor else "Numpy Array"
|
||||||
|
|
||||||
|
if is_tensor:
|
||||||
|
device = tensor.device
|
||||||
|
grad = str(tensor.requires_grad)
|
||||||
|
|
||||||
|
text = f"""
|
||||||
|
{type_name}
|
||||||
|
shape: {tensor.shape}
|
||||||
|
dtype: {str(tensor.dtype).replace("torch.", "")}
|
||||||
|
device: {device}
|
||||||
|
requires grad: {grad}
|
||||||
|
memory: {mem_str}
|
||||||
|
"""
|
||||||
|
|
||||||
|
return textwrap.dedent(text).strip()
|
||||||
|
|
||||||
|
|
||||||
def to_numpy(image: torch.Tensor) -> npt.NDArray[np.uint8]:
|
def to_numpy(image: torch.Tensor) -> npt.NDArray[np.uint8]:
|
||||||
"""Converts a tensor to a ndarray with proper scaling and type conversion."""
|
"""Converts a tensor to a ndarray with proper scaling and type conversion."""
|
||||||
log.debug(f"Converting tensor to numpy array with shape {image.shape}")
|
|
||||||
np_array = np.clip(255.0 * image.cpu().numpy(), 0, 255).astype(np.uint8)
|
np_array = np.clip(255.0 * image.cpu().numpy(), 0, 255).astype(np.uint8)
|
||||||
log.debug(f"Numpy array shape after conversion: {np_array.shape}")
|
|
||||||
return np_array
|
return np_array
|
||||||
|
|
||||||
|
|
||||||
@@ -556,12 +727,12 @@ def handle_batch(
|
|||||||
tensor: torch.Tensor,
|
tensor: torch.Tensor,
|
||||||
func: Callable[[torch.Tensor], Image.Image | npt.NDArray[np.uint8]],
|
func: Callable[[torch.Tensor], Image.Image | npt.NDArray[np.uint8]],
|
||||||
) -> list[Image.Image] | list[npt.NDArray[np.uint8]]:
|
) -> list[Image.Image] | list[npt.NDArray[np.uint8]]:
|
||||||
"""Handles batch processing for a given tensor and conversion function."""
|
"""Handle batch processing for a given tensor and conversion function."""
|
||||||
return [func(tensor[i]) for i in range(tensor.shape[0])]
|
return [func(tensor[i]) for i in range(tensor.shape[0])]
|
||||||
|
|
||||||
|
|
||||||
def tensor2pil(tensor: torch.Tensor) -> list[Image.Image]:
|
def tensor2pil(tensor: torch.Tensor) -> list[Image.Image]:
|
||||||
"""Converts a batch of tensors to a list of PIL Images."""
|
"""Convert a batch of tensors to a list of PIL Images."""
|
||||||
|
|
||||||
def single_tensor2pil(t: torch.Tensor) -> Image.Image:
|
def single_tensor2pil(t: torch.Tensor) -> Image.Image:
|
||||||
np_array = to_numpy(t)
|
np_array = to_numpy(t)
|
||||||
@@ -578,7 +749,7 @@ def tensor2pil(tensor: torch.Tensor) -> list[Image.Image]:
|
|||||||
|
|
||||||
|
|
||||||
def pil2tensor(images: Image.Image | list[Image.Image]) -> torch.Tensor:
|
def pil2tensor(images: Image.Image | list[Image.Image]) -> torch.Tensor:
|
||||||
"""Converts a PIL Image or a list of PIL Images to a tensor."""
|
"""Convert a PIL Image or a list of PIL Images to a tensor."""
|
||||||
|
|
||||||
def single_pil2tensor(image: Image.Image) -> torch.Tensor:
|
def single_pil2tensor(image: Image.Image) -> torch.Tensor:
|
||||||
np_image = np.array(image).astype(np.float32) / 255.0
|
np_image = np.array(image).astype(np.float32) / 255.0
|
||||||
|
|||||||
+236
-90
@@ -14,6 +14,51 @@ import { api } from '../../scripts/api.js'
|
|||||||
|
|
||||||
// #region base utils
|
// #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
|
// - crude uuid
|
||||||
export function makeUUID() {
|
export function makeUUID() {
|
||||||
let dt = new Date().getTime()
|
let dt = new Date().getTime()
|
||||||
@@ -25,6 +70,19 @@ export function makeUUID() {
|
|||||||
return uuid
|
return uuid
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// - basic debounce decorator
|
||||||
|
export function debounce(func, delay) {
|
||||||
|
let timeout
|
||||||
|
let debounced = function (...args) {
|
||||||
|
clearTimeout(timeout)
|
||||||
|
timeout = setTimeout(() => func.apply(this, args), delay)
|
||||||
|
}
|
||||||
|
debounced.cancel = () => {
|
||||||
|
clearTimeout(timeout)
|
||||||
|
}
|
||||||
|
return debounced
|
||||||
|
}
|
||||||
|
|
||||||
//- local storage manager
|
//- local storage manager
|
||||||
export class LocalStorageManager {
|
export class LocalStorageManager {
|
||||||
constructor(namespace) {
|
constructor(namespace) {
|
||||||
@@ -195,6 +253,7 @@ export function hideWidgetForGood(node, widget, suffix = '') {
|
|||||||
widget.origComputeSize = widget.computeSize
|
widget.origComputeSize = widget.computeSize
|
||||||
widget.origSerializeValue = widget.serializeValue
|
widget.origSerializeValue = widget.serializeValue
|
||||||
widget.computeSize = () => [0, -4] // -4 is due to the gap litegraph adds between widgets automatically
|
widget.computeSize = () => [0, -4] // -4 is due to the gap litegraph adds between widgets automatically
|
||||||
|
widget.hidden = true
|
||||||
widget.type = CONVERTED_TYPE + suffix
|
widget.type = CONVERTED_TYPE + suffix
|
||||||
// widget.serializeValue = () => {
|
// widget.serializeValue = () => {
|
||||||
// // Prevent serializing the widget if we have no input linked
|
// // Prevent serializing the widget if we have no input linked
|
||||||
@@ -276,6 +335,10 @@ export const getNamedWidget = (node, ...names) => {
|
|||||||
* @returns {{to:LGraphNode, from:LGraphNode, type:'error' | 'incoming' | 'outgoing'}}
|
* @returns {{to:LGraphNode, from:LGraphNode, type:'error' | 'incoming' | 'outgoing'}}
|
||||||
*/
|
*/
|
||||||
export const nodesFromLink = (node, link) => {
|
export const nodesFromLink = (node, link) => {
|
||||||
|
if (typeof link === 'number') {
|
||||||
|
link = app.graph.getLink(link)
|
||||||
|
}
|
||||||
|
|
||||||
const fromNode = app.graph.getNodeById(link.origin_id)
|
const fromNode = app.graph.getNodeById(link.origin_id)
|
||||||
const toNode = app.graph.getNodeById(link.target_id)
|
const toNode = app.graph.getNodeById(link.target_id)
|
||||||
|
|
||||||
@@ -366,12 +429,54 @@ export function getWidgetType(config) {
|
|||||||
|
|
||||||
// #endregion
|
// #endregion
|
||||||
|
|
||||||
|
// function to test if input is a dynamic one
|
||||||
|
const isDynamicInput = (input) => {
|
||||||
|
infoLogger('Checking if input dynamic', { input })
|
||||||
|
// return input.name.startsWith(connectionPrefix)
|
||||||
|
return input._isDynamic === true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Add a dynamic input, update node properties and slot colors!
|
||||||
|
const addDynamicInput = (node, name, kind) => {
|
||||||
|
const input = node.addInput(name, kind)
|
||||||
|
input._isDynamic = true
|
||||||
|
|
||||||
|
update_dynamic_properties(node)
|
||||||
|
set_slot_colors(node, ['cyan', undefined], isDynamicInput)
|
||||||
|
|
||||||
|
return input
|
||||||
|
}
|
||||||
|
|
||||||
|
const set_slot_colors = (node, colors, condition) => {
|
||||||
|
if (!condition) {
|
||||||
|
condition = (_s) => true
|
||||||
|
}
|
||||||
|
|
||||||
|
for (const slot of node.slots) {
|
||||||
|
infoLogger('Candidate', { slot, accepted: condition(slot) })
|
||||||
|
if (condition(slot)) {
|
||||||
|
slot.color_off = colors[0]
|
||||||
|
slot.color_on = colors[1]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
const update_dynamic_properties = (node) => {
|
||||||
|
const dyn = []
|
||||||
|
for (const input of node.inputs) {
|
||||||
|
if (isDynamicInput(input)) {
|
||||||
|
dyn.push(input.name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
node.setProperty('dynamic_connections', dyn)
|
||||||
|
}
|
||||||
|
|
||||||
// #region dynamic connections
|
// #region dynamic connections
|
||||||
/**
|
/**
|
||||||
* @param {NodeType} nodeType The nodetype to attach the documentation to
|
* @param {NodeType} nodeType The nodetype to attach the documentation to
|
||||||
* @param {str} prefix A prefix added to each dynamic inputs
|
* @param {str} prefix A prefix added to each dynamic inputs
|
||||||
* @param {str | [str]} inputType The datatype(s) of those dynamic inputs
|
* @param {str | [str]} inputType The datatype(s) of those dynamic inputs
|
||||||
* @param {{separator?:string, start_index?:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} [opts] Extra options
|
* @param {{separator?:string,rename_menu?:'label'|'name', start_index?:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} [opts] Extra options
|
||||||
* @returns
|
* @returns
|
||||||
*/
|
*/
|
||||||
export const setupDynamicConnections = (
|
export const setupDynamicConnections = (
|
||||||
@@ -385,20 +490,115 @@ export const setupDynamicConnections = (
|
|||||||
Object.getOwnPropertyDescriptors(nodeType).title.value,
|
Object.getOwnPropertyDescriptors(nodeType).title.value,
|
||||||
)
|
)
|
||||||
|
|
||||||
/** @type {{separator:string, start_index:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} */
|
/** @type {{separator:string,rename_menu?:"label"|"name" start_index:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} */
|
||||||
const options = Object.assign(
|
const options = Object.assign(
|
||||||
{
|
{
|
||||||
separator: '_',
|
separator: '_',
|
||||||
start_index: 1,
|
start_index: 1,
|
||||||
|
rename_menu: 'label',
|
||||||
},
|
},
|
||||||
opts || {},
|
opts || {},
|
||||||
)
|
)
|
||||||
|
const is_valid_name = (node, val) => {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
nodeType.prototype.getSlotMenuOptions = (slot) => {
|
||||||
|
if (!slot.input) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
infoLogger('Slot Menu', { slot })
|
||||||
|
return [
|
||||||
|
{
|
||||||
|
content: `Rename Input (${options.rename_menu})`,
|
||||||
|
callback: () => {
|
||||||
|
const dialog = app.canvas.createDialog(
|
||||||
|
"<span class='name'>Name</span><input autofocus type='text'/><button>OK</button>",
|
||||||
|
{},
|
||||||
|
)
|
||||||
|
const dialogInput = dialog.querySelector('input')
|
||||||
|
if (dialogInput) {
|
||||||
|
if (options.rename_menu === 'label') {
|
||||||
|
dialogInput.value = slot.input.label || slot.input.name || ''
|
||||||
|
} else if (options.rename_menu === 'name') {
|
||||||
|
dialogInput.value = slot.input.name || ''
|
||||||
|
}
|
||||||
|
}
|
||||||
|
const inner = () => {
|
||||||
|
// TODO: check if name exists or other guards
|
||||||
|
const val = dialogInput.value
|
||||||
|
if (!is_valid_name(slot.node, val)) {
|
||||||
|
dialog.close()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
app.graph.beforeChange()
|
||||||
|
if (options.rename_menu === 'label') {
|
||||||
|
slot.input.label = val
|
||||||
|
} else if (options.rename_menu === 'name') {
|
||||||
|
slot.input.name = val
|
||||||
|
slot.input.label = val
|
||||||
|
}
|
||||||
|
|
||||||
|
app.graph.afterChange()
|
||||||
|
|
||||||
|
dialog.close()
|
||||||
|
}
|
||||||
|
dialog.querySelector('button').addEventListener('click', inner)
|
||||||
|
dialogInput.addEventListener('keydown', (e) => {
|
||||||
|
dialog.is_modified = true
|
||||||
|
if (e.keyCode === 27) {
|
||||||
|
dialog.close()
|
||||||
|
} else if (e.keyCode === 13) {
|
||||||
|
inner()
|
||||||
|
} else if (e.keyCode !== 13 && e.target?.localName !== 'textarea') {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
e.preventDefault()
|
||||||
|
e.stopPropagation()
|
||||||
|
})
|
||||||
|
dialogInput.focus()
|
||||||
|
},
|
||||||
|
},
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
|
const onConfigure = nodeType.prototype.onConfigure
|
||||||
|
|
||||||
|
nodeType.prototype.onConfigure = function (data) {
|
||||||
|
const r = onConfigure ? onConfigure.apply(this, data) : undefined
|
||||||
|
|
||||||
|
// Set or restore serialized properties, lt seems to auto serialize/deserialize to/from string
|
||||||
|
if (!('dynamic_connections' in this.properties)) {
|
||||||
|
// this.addProperty('dynamic_connections', [], 'string')
|
||||||
|
this.setProperty('dynamic_connections', [])
|
||||||
|
} else {
|
||||||
|
const dynamic_connections = this.properties.dynamic_connections
|
||||||
|
if (typeof dynamic_connections !== 'object') {
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
for (const name of dynamic_connections) {
|
||||||
|
infoLogger(`Would dynamize: ${name}`)
|
||||||
|
const input = this.inputs.find((i) => i.name === name)
|
||||||
|
if (input) {
|
||||||
|
infoLogger('Input found', { input })
|
||||||
|
input._isDynamic = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// set color
|
||||||
|
set_slot_colors(this, ['cyan', undefined], isDynamicInput)
|
||||||
|
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
|
||||||
const onNodeCreated = nodeType.prototype.onNodeCreated
|
const onNodeCreated = nodeType.prototype.onNodeCreated
|
||||||
const inputList = typeof inputType === 'object'
|
const inputList = typeof inputType === 'object'
|
||||||
|
|
||||||
nodeType.prototype.onNodeCreated = function () {
|
nodeType.prototype.onNodeCreated = function () {
|
||||||
const r = onNodeCreated ? onNodeCreated.apply(this, []) : undefined
|
const r = onNodeCreated ? onNodeCreated.apply(this, []) : undefined
|
||||||
this.addInput(
|
|
||||||
|
const input = addDynamicInput(
|
||||||
|
this,
|
||||||
`${prefix}${options.separator}${options.start_index}`,
|
`${prefix}${options.separator}${options.start_index}`,
|
||||||
inputList ? '*' : inputType,
|
inputList ? '*' : inputType,
|
||||||
)
|
)
|
||||||
@@ -470,10 +670,7 @@ export const dynamic_connection = (
|
|||||||
opts || {},
|
opts || {},
|
||||||
)
|
)
|
||||||
|
|
||||||
// function to test if input is a dynamic one
|
if (node.inputs.length > 0 && !isDynamicInput(node.inputs[index])) {
|
||||||
const isDynamicInput = (inputName) => inputName.startsWith(connectionPrefix)
|
|
||||||
|
|
||||||
if (node.inputs.length > 0 && !isDynamicInput(node.inputs[index].name)) {
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -483,6 +680,7 @@ export const dynamic_connection = (
|
|||||||
const nameArray = options.nameArray || []
|
const nameArray = options.nameArray || []
|
||||||
|
|
||||||
const clean_inputs = () => {
|
const clean_inputs = () => {
|
||||||
|
if (node.id < 0) return // being duplicated
|
||||||
if (node.inputs.length === 0) return
|
if (node.inputs.length === 0) return
|
||||||
|
|
||||||
let w_count = node.widgets?.length || 0
|
let w_count = node.widgets?.length || 0
|
||||||
@@ -492,7 +690,7 @@ export const dynamic_connection = (
|
|||||||
const to_remove = []
|
const to_remove = []
|
||||||
for (let n = 1; n < node.inputs.length; n++) {
|
for (let n = 1; n < node.inputs.length; n++) {
|
||||||
const element = node.inputs[n]
|
const element = node.inputs[n]
|
||||||
if (!element.link && isDynamicInput(element.name)) {
|
if (!element.link && isDynamicInput(element)) {
|
||||||
if (node.widgets) {
|
if (node.widgets) {
|
||||||
const w = node.widgets.find((w) => w.name === element.name)
|
const w = node.widgets.find((w) => w.name === element.name)
|
||||||
if (w) {
|
if (w) {
|
||||||
@@ -506,9 +704,12 @@ export const dynamic_connection = (
|
|||||||
}
|
}
|
||||||
for (let i = 0; i < to_remove.length; i++) {
|
for (let i = 0; i < to_remove.length; i++) {
|
||||||
const id = to_remove[i]
|
const id = to_remove[i]
|
||||||
|
try {
|
||||||
node.removeInput(id)
|
node.removeInput(id)
|
||||||
i_count -= 1
|
i_count -= 1
|
||||||
|
} catch (err) {
|
||||||
|
errorLogger('Cannot remove input', err)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
node.inputs.length = i_count
|
node.inputs.length = i_count
|
||||||
|
|
||||||
@@ -522,7 +723,7 @@ export const dynamic_connection = (
|
|||||||
for (let i = 0; i < node.inputs.length; i++) {
|
for (let i = 0; i < node.inputs.length; i++) {
|
||||||
let name = ''
|
let name = ''
|
||||||
// rename only prefixed inputs
|
// rename only prefixed inputs
|
||||||
if (isDynamicInput(node.inputs[i].name)) {
|
if (node.inputs[i].name.startsWith(connectionPrefix)) {
|
||||||
// prefixed => rename and increase index
|
// prefixed => rename and increase index
|
||||||
name = `${connectionPrefix}${prefixed_idx}`
|
name = `${connectionPrefix}${prefixed_idx}`
|
||||||
prefixed_idx += 1
|
prefixed_idx += 1
|
||||||
@@ -576,9 +777,8 @@ export const dynamic_connection = (
|
|||||||
if (node.inputs.length === 0) return
|
if (node.inputs.length === 0) return
|
||||||
// add an extra input
|
// add an extra input
|
||||||
if (node.inputs[node.inputs.length - 1].link !== null) {
|
if (node.inputs[node.inputs.length - 1].link !== null) {
|
||||||
// count only the prefixed inputs
|
|
||||||
const nextIndex = node.inputs.reduce(
|
const nextIndex = node.inputs.reduce(
|
||||||
(acc, cur) => (isDynamicInput(cur.name) ? ++acc : acc),
|
(acc, cur) => (isDynamicInput(cur) ? ++acc : acc),
|
||||||
0,
|
0,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -588,7 +788,7 @@ export const dynamic_connection = (
|
|||||||
: `${connectionPrefix}${nextIndex + options.start_index}`
|
: `${connectionPrefix}${nextIndex + options.start_index}`
|
||||||
|
|
||||||
infoLogger(`Adding input ${nextIndex + 1} (${name})`)
|
infoLogger(`Adding input ${nextIndex + 1} (${name})`)
|
||||||
node.addInput(name, conType)
|
addDynamicInput(node, name, conType)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -621,21 +821,21 @@ function getBrightness(rgbObj) {
|
|||||||
export function calculateTotalChildrenHeight(parentElement) {
|
export function calculateTotalChildrenHeight(parentElement) {
|
||||||
let totalHeight = 0
|
let totalHeight = 0
|
||||||
|
|
||||||
|
if (!parentElement || !parentElement.children) {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
for (const child of parentElement.children) {
|
for (const child of parentElement.children) {
|
||||||
const style = window.getComputedStyle(child)
|
const style = window.getComputedStyle(child)
|
||||||
|
|
||||||
// Get height as an integer (without 'px')
|
const height = Number.parseFloat(style.height)
|
||||||
const height = Number.parseInt(style.height, 10)
|
const marginTop = Number.parseFloat(style.marginTop)
|
||||||
|
const marginBottom = Number.parseFloat(style.marginBottom)
|
||||||
|
|
||||||
// Get vertical margin as integers
|
|
||||||
const marginTop = Number.parseInt(style.marginTop, 10)
|
|
||||||
const marginBottom = Number.parseInt(style.marginBottom, 10)
|
|
||||||
|
|
||||||
// Sum up height and vertical margins
|
|
||||||
totalHeight += height + marginTop + marginBottom
|
totalHeight += height + marginTop + marginBottom
|
||||||
}
|
}
|
||||||
|
|
||||||
return totalHeight
|
return Math.ceil(totalHeight)
|
||||||
}
|
}
|
||||||
|
|
||||||
export const loadScript = (
|
export const loadScript = (
|
||||||
@@ -646,13 +846,15 @@ export const loadScript = (
|
|||||||
return new Promise((resolve, reject) => {
|
return new Promise((resolve, reject) => {
|
||||||
try {
|
try {
|
||||||
// Check if the script already exists
|
// Check if the script already exists
|
||||||
const existingScript = document.querySelector(`script[src="${FILE_URL}"]`)
|
let scriptEle = document.querySelector(`script[src="${FILE_URL}"]`)
|
||||||
if (existingScript) {
|
if (scriptEle) {
|
||||||
resolve({ status: true, message: 'Script already loaded' })
|
scriptEle.addEventListener('load', (_ev) => {
|
||||||
|
resolve({ status: true })
|
||||||
|
})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
const scriptEle = document.createElement('script')
|
scriptEle = document.createElement('script')
|
||||||
scriptEle.type = type
|
scriptEle.type = type
|
||||||
scriptEle.async = async
|
scriptEle.async = async
|
||||||
scriptEle.src = FILE_URL
|
scriptEle.src = FILE_URL
|
||||||
@@ -671,6 +873,8 @@ export const loadScript = (
|
|||||||
document.body.appendChild(scriptEle)
|
document.body.appendChild(scriptEle)
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
reject(error)
|
reject(error)
|
||||||
|
} finally {
|
||||||
|
infoLogger(`Finally loaded script: ${FILE_URL}`)
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -784,12 +988,10 @@ function loadParser(shiki) {
|
|||||||
|
|
||||||
export const ensureMarkdownParser = async (callback) => {
|
export const ensureMarkdownParser = async (callback) => {
|
||||||
infoLogger('Ensuring md parser')
|
infoLogger('Ensuring md parser')
|
||||||
let use_shiki = false
|
const use_shiki = app.extensionManager.setting.get(
|
||||||
try {
|
'mtb.noteplus.use-shiki',
|
||||||
use_shiki = await api.getSetting('mtb.Use Shiki')
|
false,
|
||||||
} catch (e) {
|
)
|
||||||
console.warn('Option not available yet', e)
|
|
||||||
}
|
|
||||||
|
|
||||||
if (window.MTB?.mdParser) {
|
if (window.MTB?.mdParser) {
|
||||||
infoLogger('Markdown parser found')
|
infoLogger('Markdown parser found')
|
||||||
@@ -814,8 +1016,7 @@ export const ensureMarkdownParser = async (callback) => {
|
|||||||
callbackQueue.push(callback)
|
callbackQueue.push(callback)
|
||||||
}
|
}
|
||||||
|
|
||||||
await parserPromise
|
await await parserPromise
|
||||||
await parserPromise
|
|
||||||
|
|
||||||
return window.MTB.mdParser
|
return window.MTB.mdParser
|
||||||
}
|
}
|
||||||
@@ -1154,58 +1355,3 @@ export const setServerInfo = async (opts) => {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// #endregion
|
// #endregion
|
||||||
|
|
||||||
// #region Authoring API / graph utilities
|
|
||||||
export const getAPIInputs = () => {
|
|
||||||
const inputs = {}
|
|
||||||
let counter = 1
|
|
||||||
for (const node of getNodes(true)) {
|
|
||||||
const widgets = node.widgets
|
|
||||||
|
|
||||||
if (node.properties.mtb_api && node.properties.useAPI) {
|
|
||||||
if (node.properties.mtb_api.inputs) {
|
|
||||||
for (const currentName in node.properties.mtb_api.inputs) {
|
|
||||||
const current = node.properties.mtb_api.inputs[currentName]
|
|
||||||
if (current.enabled) {
|
|
||||||
const inputName = current.name || currentName
|
|
||||||
const widget = widgets.find((w) => w.name === currentName)
|
|
||||||
if (!widget) continue
|
|
||||||
if (!(inputName in inputs)) {
|
|
||||||
inputs[inputName] = {
|
|
||||||
...current,
|
|
||||||
id: counter,
|
|
||||||
name: inputName,
|
|
||||||
type: current.type,
|
|
||||||
node_id: node.id,
|
|
||||||
widgets: [],
|
|
||||||
}
|
|
||||||
}
|
|
||||||
inputs[inputName].widgets.push(widget)
|
|
||||||
counter = counter + 1
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return inputs
|
|
||||||
}
|
|
||||||
|
|
||||||
export const getNodes = (skip_unused) => {
|
|
||||||
const nodes = []
|
|
||||||
for (const outerNode of app.graph.computeExecutionOrder(false)) {
|
|
||||||
const skipNode =
|
|
||||||
(outerNode.mode === 2 || outerNode.mode === 4) && skip_unused
|
|
||||||
const innerNodes =
|
|
||||||
!skipNode && outerNode.getInnerNodes
|
|
||||||
? outerNode.getInnerNodes()
|
|
||||||
: [outerNode]
|
|
||||||
for (const node of innerNodes) {
|
|
||||||
if ((node.mode === 2 || node.mode === 4) && skip_unused) {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
nodes.push(node)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nodes
|
|
||||||
}
|
|
||||||
// #endregion
|
|
||||||
|
|||||||
+76
-100
@@ -11,7 +11,12 @@
|
|||||||
/// <reference path="../types/typedefs.js" />
|
/// <reference path="../types/typedefs.js" />
|
||||||
|
|
||||||
import { app } from '../../scripts/app.js'
|
import { app } from '../../scripts/app.js'
|
||||||
import * as shared from './comfy_shared.js'
|
|
||||||
|
import {
|
||||||
|
setupDynamicConnections,
|
||||||
|
cleanupNode,
|
||||||
|
infoLogger,
|
||||||
|
} from './comfy_shared.js'
|
||||||
import * as mtb_ui from './mtb_ui.js'
|
import * as mtb_ui from './mtb_ui.js'
|
||||||
|
|
||||||
function escapeHtml(unsafe) {
|
function escapeHtml(unsafe) {
|
||||||
@@ -28,7 +33,7 @@ function createDebugSection(title) {
|
|||||||
margin: '8px 0',
|
margin: '8px 0',
|
||||||
padding: '8px',
|
padding: '8px',
|
||||||
borderRadius: '4px',
|
borderRadius: '4px',
|
||||||
backgroundColor: 'rgba(0,0,0,0.2)'
|
backgroundColor: 'rgba(0,0,0,0.2)',
|
||||||
})
|
})
|
||||||
|
|
||||||
const header = mtb_ui.makeElement('h3', {
|
const header = mtb_ui.makeElement('h3', {
|
||||||
@@ -37,7 +42,7 @@ function createDebugSection(title) {
|
|||||||
borderBottom: '1px solid rgba(255,255,255,0.1)',
|
borderBottom: '1px solid rgba(255,255,255,0.1)',
|
||||||
fontSize: '14px',
|
fontSize: '14px',
|
||||||
fontWeight: 'bold',
|
fontWeight: 'bold',
|
||||||
color: '#9f9'
|
color: '#9f9',
|
||||||
})
|
})
|
||||||
header.textContent = title
|
header.textContent = title
|
||||||
section.appendChild(header)
|
section.appendChild(header)
|
||||||
@@ -45,25 +50,25 @@ function createDebugSection(title) {
|
|||||||
return section
|
return section
|
||||||
}
|
}
|
||||||
|
|
||||||
function createDebugContent(content, type) {
|
function createDebugContent(item) {
|
||||||
const wrapper = mtb_ui.makeElement('div', {
|
const wrapper = mtb_ui.makeElement('div', {
|
||||||
margin: '4px 0'
|
margin: '4px 0',
|
||||||
})
|
})
|
||||||
|
|
||||||
if (type === 'text') {
|
if (item.kind === 'text') {
|
||||||
const text = mtb_ui.makeElement('p', {
|
const text = mtb_ui.makeElement('div', {
|
||||||
margin: '2px 0',
|
margin: '2px 0',
|
||||||
fontFamily: 'monospace',
|
fontFamily: 'monospace',
|
||||||
whiteSpace: 'pre-wrap'
|
whiteSpace: 'pre-wrap',
|
||||||
})
|
})
|
||||||
text.innerHTML = content
|
text.innerHTML = item.data
|
||||||
wrapper.appendChild(text)
|
wrapper.appendChild(text)
|
||||||
} else if (type === 'image') {
|
} else if (item.kind === 'b64_images') {
|
||||||
const img = mtb_ui.makeElement('img', {
|
const img = mtb_ui.makeElement('img', {
|
||||||
width: '100%',
|
width: '100%',
|
||||||
borderRadius: '2px'
|
borderRadius: '2px',
|
||||||
})
|
})
|
||||||
img.src = content
|
img.src = item.data
|
||||||
wrapper.appendChild(img)
|
wrapper.appendChild(img)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -80,115 +85,85 @@ app.registerExtension({
|
|||||||
*/
|
*/
|
||||||
async beforeRegisterNodeDef(nodeType, nodeData, app) {
|
async beforeRegisterNodeDef(nodeType, nodeData, app) {
|
||||||
if (nodeData.name === 'Debug (mtb)') {
|
if (nodeData.name === 'Debug (mtb)') {
|
||||||
const onNodeCreated = nodeType.prototype.onNodeCreated
|
const clear_widgets = (target) => {
|
||||||
nodeType.prototype.onNodeCreated = function (...args) {
|
if (target.widgets) {
|
||||||
this.options = {}
|
let tgt_len = target.widgets.length
|
||||||
const r = onNodeCreated ? onNodeCreated.apply(this, args) : undefined
|
for (let i = 0; i < target.widgets.length; i++) {
|
||||||
this.addInput('anything_1', '*')
|
if (
|
||||||
return r
|
![
|
||||||
|
'output_to_console',
|
||||||
|
'deep_inspect',
|
||||||
|
'as_detailed_types',
|
||||||
|
'rich_mode',
|
||||||
|
].includes(target.widgets[i].name)
|
||||||
|
) {
|
||||||
|
target.widgets[i].onRemove?.()
|
||||||
|
target.widgets[i].onRemoved?.()
|
||||||
|
tgt_len -= 1
|
||||||
|
}
|
||||||
|
}
|
||||||
|
target.widgets.length = tgt_len
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
const onConnectionsChange = nodeType.prototype.onConnectionsChange
|
const original_getExtraMenuOptions =
|
||||||
/**
|
nodeType.prototype.getExtraMenuOptions
|
||||||
* @param {OnConnectionsChangeParams} args
|
nodeType.prototype.getExtraMenuOptions = function (_, options) {
|
||||||
*/
|
original_getExtraMenuOptions?.apply(this, arguments)
|
||||||
nodeType.prototype.onConnectionsChange = function (...args) {
|
options.push({
|
||||||
const [_type, index, connected, link_info, ioSlot] = args
|
content: '🐛 Clear Outputs',
|
||||||
const r = onConnectionsChange
|
callback: async () => {
|
||||||
? onConnectionsChange.apply(this, args)
|
clear_widgets(this)
|
||||||
: undefined
|
},
|
||||||
// TODO: remove all widgets on disconnect once computed
|
|
||||||
shared.dynamic_connection(this, index, connected, 'anything_', '*', {
|
|
||||||
link: link_info,
|
|
||||||
ioSlot: ioSlot,
|
|
||||||
})
|
})
|
||||||
|
|
||||||
//- infer type
|
|
||||||
if (link_info) {
|
|
||||||
// const fromNode = this.graph._nodes.find(
|
|
||||||
// (otherNode) => otherNode.id === link_info.origin_id,
|
|
||||||
// )
|
|
||||||
// const fromNode = app.graph.getNodeById(link_info.origin_id)
|
|
||||||
const { from } = shared.nodesFromLink(this, link_info)
|
|
||||||
if (!from || this.inputs.length === 0) return
|
|
||||||
const type = from.outputs[link_info.origin_slot].type
|
|
||||||
this.inputs[index].type = type
|
|
||||||
// this.inputs[index].label = type.toLowerCase()
|
|
||||||
}
|
|
||||||
//- restore dynamic input
|
|
||||||
if (!connected) {
|
|
||||||
this.inputs[index].type = '*'
|
|
||||||
this.inputs[index].label = `anything_${index + 1}`
|
|
||||||
}
|
|
||||||
return r
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
setupDynamicConnections(nodeType, 'var', '*')
|
||||||
|
|
||||||
const onExecuted = nodeType.prototype.onExecuted
|
const onExecuted = nodeType.prototype.onExecuted
|
||||||
nodeType.prototype.onExecuted = function (...args) {
|
nodeType.prototype.onExecuted = function (...args) {
|
||||||
onExecuted?.apply(this, args)
|
onExecuted?.apply(this, args)
|
||||||
const [data, ..._rest] = args
|
const [data, ..._rest] = args
|
||||||
|
|
||||||
if (this.widgets) {
|
clear_widgets(this)
|
||||||
let tgt_len = this.widgets.length
|
|
||||||
for (let i = 0; i < this.widgets.length; i++) {
|
|
||||||
if (
|
|
||||||
this.widgets[i].name !== 'output_to_console' &&
|
|
||||||
this.widgets[i].name !== 'as_detailed_types'
|
|
||||||
) {
|
|
||||||
this.widgets[i].onRemove?.()
|
|
||||||
this.widgets[i].onRemoved?.()
|
|
||||||
tgt_len -= 1
|
|
||||||
}
|
|
||||||
}
|
|
||||||
this.widgets.length = tgt_len
|
|
||||||
}
|
|
||||||
|
|
||||||
const inputData = {}
|
const inputData = {}
|
||||||
|
|
||||||
const uiData = data.ui || data
|
const uiData = data.ui || data
|
||||||
|
|
||||||
if (uiData.items) {
|
const name_to_label = this.inputs.reduce((acc, input) => {
|
||||||
uiData.items.forEach(item => {
|
acc[input.name] = input.label || input.name
|
||||||
const inputName = item.input
|
return acc
|
||||||
if (!inputData[inputName]) {
|
}, {})
|
||||||
inputData[inputName] = { text: [], b64_images: [] }
|
|
||||||
}
|
|
||||||
if (item.text) {
|
|
||||||
inputData[inputName].text.push(...item.text)
|
|
||||||
}
|
|
||||||
if (item.b64_images) {
|
|
||||||
inputData[inputName].b64_images.push(...item.b64_images)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
let widgetI = 1
|
|
||||||
|
if (uiData.items) {
|
||||||
|
uiData.items.forEach((item) => {
|
||||||
|
const inputName = item.input
|
||||||
|
inputData[inputName] = item.items
|
||||||
|
})
|
||||||
|
}
|
||||||
|
const mainDebugContainer = mtb_ui.makeElement('div', {
|
||||||
|
width: '100%',
|
||||||
|
})
|
||||||
|
let hasContent = false
|
||||||
for (const [inputName, content] of Object.entries(inputData)) {
|
for (const [inputName, content] of Object.entries(inputData)) {
|
||||||
if (content.text.length === 0 && content.b64_images.length === 0) {
|
if (!content || content?.length === 0) {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
hasContent = true
|
||||||
|
|
||||||
const section = createDebugSection(inputName)
|
const section = createDebugSection(name_to_label[inputName])
|
||||||
|
|
||||||
if (content.text.length > 0) {
|
for (const item of content) {
|
||||||
content.text.forEach(text => {
|
section.appendChild(createDebugContent(item))
|
||||||
section.appendChild(createDebugContent(text, 'text'))
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
mainDebugContainer.appendChild(section)
|
||||||
if (content.b64_images.length > 0) {
|
}
|
||||||
content.b64_images.forEach(img => {
|
if (hasContent) {
|
||||||
section.appendChild(createDebugContent(img, 'image'))
|
this.addDOMWidget('debug_output', 'CUSTOM', mainDebugContainer, {
|
||||||
})
|
hideOnZoom: false,
|
||||||
}
|
})
|
||||||
|
|
||||||
this.addDOMWidget(
|
|
||||||
`debug_section_${widgetI}`,
|
|
||||||
'CUSTOM',
|
|
||||||
section,
|
|
||||||
{}
|
|
||||||
)
|
|
||||||
widgetI++
|
|
||||||
}
|
}
|
||||||
|
|
||||||
this.onRemoved = function () {
|
this.onRemoved = function () {
|
||||||
@@ -199,8 +174,9 @@ app.registerExtension({
|
|||||||
widget.onRemoved?.()
|
widget.onRemoved?.()
|
||||||
widget.onRemove?.()
|
widget.onRemove?.()
|
||||||
}
|
}
|
||||||
shared.cleanupNode(this)
|
cleanupNode(this)
|
||||||
}
|
}
|
||||||
|
this.setDirtyCanvas(true, true)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
|
|||||||
+296
-296
@@ -13,40 +13,40 @@ import { api } from '../../scripts/api.js'
|
|||||||
import { app } from '../../scripts/app.js'
|
import { app } from '../../scripts/app.js'
|
||||||
import { LocalStorageManager } from './comfy_shared.js'
|
import { LocalStorageManager } from './comfy_shared.js'
|
||||||
const styles = {
|
const styles = {
|
||||||
lighbox: {
|
lighbox: {
|
||||||
position: 'fixed',
|
position: 'fixed',
|
||||||
top: 0,
|
top: 0,
|
||||||
left: 0,
|
left: 0,
|
||||||
width: '100vw',
|
width: '100vw',
|
||||||
height: '100vh',
|
height: '100vh',
|
||||||
background: 'rgba(0,0,0,0.5)',
|
background: 'rgba(0,0,0,0.5)',
|
||||||
display: 'none',
|
display: 'none',
|
||||||
justifyContent: 'center',
|
justifyContent: 'center',
|
||||||
alignItems: 'center',
|
alignItems: 'center',
|
||||||
zIndex: 999,
|
zIndex: 999,
|
||||||
},
|
},
|
||||||
lightboxBtn: (extra) => ({
|
lightboxBtn: (extra) => ({
|
||||||
position: 'absolute',
|
position: 'absolute',
|
||||||
top: '50%',
|
top: '50%',
|
||||||
background: 'none',
|
background: 'none',
|
||||||
border: 'none',
|
border: 'none',
|
||||||
color: '#fff',
|
color: '#fff',
|
||||||
zIndex: 1000,
|
zIndex: 1000,
|
||||||
fontSize: '30px',
|
fontSize: '30px',
|
||||||
cursor: 'pointer',
|
cursor: 'pointer',
|
||||||
pointerEvents: 'auto',
|
pointerEvents: 'auto',
|
||||||
...extra,
|
...extra,
|
||||||
}),
|
}),
|
||||||
img_list: {
|
img_list: {
|
||||||
minHeight: '30px',
|
minHeight: '30px',
|
||||||
maxHeight: '300px',
|
maxHeight: '300px',
|
||||||
width: '100vw',
|
width: '100vw',
|
||||||
position: 'absolute',
|
position: 'absolute',
|
||||||
bottom: 0,
|
bottom: 0,
|
||||||
zIndex: 10,
|
zIndex: 10,
|
||||||
background: '#333',
|
background: '#333',
|
||||||
overflow: 'auto',
|
overflow: 'auto',
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
let currentImageIndex = 0
|
let currentImageIndex = 0
|
||||||
@@ -58,299 +58,299 @@ const storage = new LocalStorageManager('mtb')
|
|||||||
let activated = storage.get('image_feed', false)
|
let activated = storage.get('image_feed', false)
|
||||||
|
|
||||||
app.registerExtension({
|
app.registerExtension({
|
||||||
name: 'mtb.ImageFeed',
|
name: 'mtb.ImageFeed',
|
||||||
setup: () => {
|
setup: () => {
|
||||||
app.ui.settings.addSetting({
|
app.ui.settings.addSetting({
|
||||||
id: 'mtb.Main.image-feed-enabled',
|
id: 'mtb.Main.image-feed-enabled',
|
||||||
category: ['mtb', 'Main', 'image-feed-enabled'],
|
category: ['mtb', ' Main', 'image-feed-enabled'],
|
||||||
name: 'Enable Image Feed',
|
name: 'Enable Image Feed',
|
||||||
type: 'boolean',
|
type: 'boolean',
|
||||||
defaultValue: false,
|
defaultValue: false,
|
||||||
attrs: {
|
attrs: {
|
||||||
style: {
|
style: {
|
||||||
fontFamily: 'monospace',
|
fontFamily: 'monospace',
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
async onChange(value) {
|
async onChange(value) {
|
||||||
storage.set('image_feed', value)
|
storage.set('image_feed', value)
|
||||||
activated = value
|
activated = value
|
||||||
},
|
},
|
||||||
})
|
})
|
||||||
},
|
},
|
||||||
init: async () => {
|
init: async () => {
|
||||||
if (!activated) {
|
if (!activated) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
const pythongossFeed = app.extensions.find(
|
const pythongossFeed = app.extensions.find(
|
||||||
(e) => e.name === 'pysssss.ImageFeed',
|
(e) => e.name === 'pysssss.ImageFeed',
|
||||||
)
|
)
|
||||||
if (pythongossFeed) {
|
if (pythongossFeed) {
|
||||||
console.warn(
|
console.warn(
|
||||||
"[mtb] - Aborting the loading of mtb's imageFeed in favor of pysssss.ImageFeed",
|
"[mtb] - Aborting the loading of mtb's imageFeed in favor of pysssss.ImageFeed",
|
||||||
)
|
)
|
||||||
activated = false // just in case other methods are added later on
|
activated = false // just in case other methods are added later on
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
// - HTML & CSS
|
// - HTML & CSS
|
||||||
//- lightbox
|
//- lightbox
|
||||||
const lightboxContainer = document.createElement('div')
|
const lightboxContainer = document.createElement('div')
|
||||||
Object.assign(lightboxContainer.style, styles.lighbox)
|
Object.assign(lightboxContainer.style, styles.lighbox)
|
||||||
|
|
||||||
const lightboxImage = document.createElement('img')
|
const lightboxImage = document.createElement('img')
|
||||||
Object.assign(lightboxImage.style, {
|
Object.assign(lightboxImage.style, {
|
||||||
maxHeight: '100%',
|
maxHeight: '100%',
|
||||||
maxWidth: '100%',
|
maxWidth: '100%',
|
||||||
borderRadius: '5px',
|
borderRadius: '5px',
|
||||||
})
|
})
|
||||||
|
|
||||||
// previous and next buttons
|
// previous and next buttons
|
||||||
const lightboxPrevBtn = document.createElement('button')
|
const lightboxPrevBtn = document.createElement('button')
|
||||||
const lightboxNextBtn = document.createElement('button')
|
const lightboxNextBtn = document.createElement('button')
|
||||||
|
|
||||||
lightboxPrevBtn.textContent = '❮'
|
lightboxPrevBtn.textContent = '❮'
|
||||||
lightboxNextBtn.textContent = '❯'
|
lightboxNextBtn.textContent = '❯'
|
||||||
|
|
||||||
Object.assign(lightboxPrevBtn.style, styles.lightboxBtn({ left: '0%' }))
|
Object.assign(lightboxPrevBtn.style, styles.lightboxBtn({ left: '0%' }))
|
||||||
Object.assign(lightboxNextBtn.style, styles.lightboxBtn({ right: '0%' }))
|
Object.assign(lightboxNextBtn.style, styles.lightboxBtn({ right: '0%' }))
|
||||||
|
|
||||||
// close button
|
// close button
|
||||||
const lightboxCloseBtn = document.createElement('button')
|
const lightboxCloseBtn = document.createElement('button')
|
||||||
Object.assign(
|
Object.assign(
|
||||||
lightboxCloseBtn.style,
|
lightboxCloseBtn.style,
|
||||||
styles.lightboxBtn({ right: '0', top: '0' }),
|
styles.lightboxBtn({ right: '0', top: '0' }),
|
||||||
)
|
)
|
||||||
lightboxCloseBtn.textContent = '❌'
|
lightboxCloseBtn.textContent = '❌'
|
||||||
|
|
||||||
const lightboxButtons = document.createElement('div')
|
const lightboxButtons = document.createElement('div')
|
||||||
Object.assign(lightboxButtons.style, {
|
Object.assign(lightboxButtons.style, {
|
||||||
position: 'absolute',
|
position: 'absolute',
|
||||||
top: '0%',
|
top: '0%',
|
||||||
right: '0%',
|
right: '0%',
|
||||||
// transform: "translate(50%, -50%)",
|
// transform: "translate(50%, -50%)",
|
||||||
height: '100%',
|
height: '100%',
|
||||||
width: '100%',
|
width: '100%',
|
||||||
background: 'none',
|
background: 'none',
|
||||||
border: 'none',
|
border: 'none',
|
||||||
color: '#fff',
|
color: '#fff',
|
||||||
fontSize: '30px',
|
fontSize: '30px',
|
||||||
cursor: 'pointer',
|
cursor: 'pointer',
|
||||||
pointerEvents: 'none',
|
pointerEvents: 'none',
|
||||||
})
|
})
|
||||||
|
|
||||||
lightboxButtons.append(lightboxPrevBtn, lightboxNextBtn, lightboxCloseBtn)
|
lightboxButtons.append(lightboxPrevBtn, lightboxNextBtn, lightboxCloseBtn)
|
||||||
lightboxContainer.append(lightboxButtons, lightboxImage)
|
lightboxContainer.append(lightboxButtons, lightboxImage)
|
||||||
|
|
||||||
//- image list
|
//- image list
|
||||||
const imageListContainer = document.createElement('div')
|
const imageListContainer = document.createElement('div')
|
||||||
Object.assign(imageListContainer.style, styles.img_list)
|
Object.assign(imageListContainer.style, styles.img_list)
|
||||||
|
|
||||||
const createImgListBtn = (text, style) => {
|
const createImgListBtn = (text, style) => {
|
||||||
const btn = document.createElement('button')
|
const btn = document.createElement('button')
|
||||||
btn.type = 'button'
|
btn.type = 'button'
|
||||||
btn.textContent = text
|
btn.textContent = text
|
||||||
Object.assign(btn.style, {
|
Object.assign(btn.style, {
|
||||||
...style,
|
...style,
|
||||||
border: 'none',
|
border: 'none',
|
||||||
color: '#fff',
|
color: '#fff',
|
||||||
background: 'none',
|
background: 'none',
|
||||||
height: '20px',
|
height: '20px',
|
||||||
cursor: 'pointer',
|
cursor: 'pointer',
|
||||||
position: 'absolute',
|
position: 'absolute',
|
||||||
top: '5px',
|
top: '5px',
|
||||||
fontSize: '12px',
|
fontSize: '12px',
|
||||||
lineHeight: '12px',
|
lineHeight: '12px',
|
||||||
})
|
})
|
||||||
imageListContainer.append(btn)
|
imageListContainer.append(btn)
|
||||||
return btn
|
return btn
|
||||||
}
|
}
|
||||||
const showBtn = document.createElement('button')
|
const showBtn = document.createElement('button')
|
||||||
const closeBtn = createImgListBtn('❌', {
|
const closeBtn = createImgListBtn('❌', {
|
||||||
width: '20px',
|
width: '20px',
|
||||||
textIndent: '-4px',
|
textIndent: '-4px',
|
||||||
right: '5px',
|
right: '5px',
|
||||||
})
|
})
|
||||||
const loadButton = createImgListBtn('Load Session History', {
|
const loadButton = createImgListBtn('Load Session History', {
|
||||||
right: '90px',
|
right: '90px',
|
||||||
})
|
})
|
||||||
const clearButton = createImgListBtn('Clear', {
|
const clearButton = createImgListBtn('Clear', {
|
||||||
right: '30px',
|
right: '30px',
|
||||||
})
|
})
|
||||||
|
|
||||||
//- tools popup button
|
//- tools popup button
|
||||||
showBtn.classList.add('comfy-settings-btn')
|
showBtn.classList.add('comfy-settings-btn')
|
||||||
Object.assign(showBtn.style, {
|
Object.assign(showBtn.style, {
|
||||||
right: '16px',
|
right: '16px',
|
||||||
cursor: 'pointer',
|
cursor: 'pointer',
|
||||||
display: 'none',
|
display: 'none',
|
||||||
})
|
})
|
||||||
|
|
||||||
//- append to DOM
|
//- append to DOM
|
||||||
document.body.append(imageListContainer)
|
document.body.append(imageListContainer)
|
||||||
|
|
||||||
showBtn.textContent = '🖼'
|
showBtn.textContent = '🖼'
|
||||||
showBtn.onclick = () => {
|
showBtn.onclick = () => {
|
||||||
imageListContainer.style.display = 'block'
|
imageListContainer.style.display = 'block'
|
||||||
showBtn.style.display = 'none'
|
showBtn.style.display = 'none'
|
||||||
}
|
}
|
||||||
document.querySelector('.comfy-settings-btn').after(showBtn)
|
document.querySelector('.comfy-settings-btn').after(showBtn)
|
||||||
document.querySelector('.comfy-settings-btn').after(lightboxContainer)
|
document.querySelector('.comfy-settings-btn').after(lightboxContainer)
|
||||||
|
|
||||||
// for (const { output } of history) {
|
// for (const { output } of history) {
|
||||||
// if (output?.images) {
|
// if (output?.images) {
|
||||||
// for (const src of output.images) {
|
// for (const src of output.images) {
|
||||||
// const img = document.createElement("img");
|
// const img = document.createElement("img");
|
||||||
// const but = document.createElement("button");
|
// const but = document.createElement("button");
|
||||||
|
|
||||||
//- callbacks
|
//- callbacks
|
||||||
closeBtn.onclick = () => {
|
closeBtn.onclick = () => {
|
||||||
imageListContainer.style.display = 'none'
|
imageListContainer.style.display = 'none'
|
||||||
showBtn.style.display = 'unset'
|
showBtn.style.display = 'unset'
|
||||||
}
|
}
|
||||||
|
|
||||||
clearButton.onclick = () => {
|
clearButton.onclick = () => {
|
||||||
imageListContainer.replaceChildren(closeBtn, clearButton, loadButton)
|
imageListContainer.replaceChildren(closeBtn, clearButton, loadButton)
|
||||||
}
|
}
|
||||||
|
|
||||||
lightboxNextBtn.onclick = () => {
|
lightboxNextBtn.onclick = () => {
|
||||||
currentImageIndex = (currentImageIndex + 1) % imageUrls.length
|
currentImageIndex = (currentImageIndex + 1) % imageUrls.length
|
||||||
const imageUrl = imageUrls[currentImageIndex]
|
const imageUrl = imageUrls[currentImageIndex]
|
||||||
lightboxImage.src = imageUrl
|
lightboxImage.src = imageUrl
|
||||||
}
|
}
|
||||||
|
|
||||||
// Modify the lightboxPrevBtn onclick callback
|
// Modify the lightboxPrevBtn onclick callback
|
||||||
lightboxPrevBtn.onclick = () => {
|
lightboxPrevBtn.onclick = () => {
|
||||||
currentImageIndex =
|
currentImageIndex =
|
||||||
(currentImageIndex - 1 + imageUrls.length) % imageUrls.length
|
(currentImageIndex - 1 + imageUrls.length) % imageUrls.length
|
||||||
const imageUrl = imageUrls[currentImageIndex]
|
const imageUrl = imageUrls[currentImageIndex]
|
||||||
lightboxImage.src = imageUrl
|
lightboxImage.src = imageUrl
|
||||||
}
|
}
|
||||||
|
|
||||||
lightboxCloseBtn.onclick = () => {
|
lightboxCloseBtn.onclick = () => {
|
||||||
lightboxContainer.style.display = 'none'
|
lightboxContainer.style.display = 'none'
|
||||||
}
|
}
|
||||||
lightboxImage.onclick = lightboxNextBtn.onclick
|
lightboxImage.onclick = lightboxNextBtn.onclick
|
||||||
/**
|
/**
|
||||||
* This is the function that creates the image buttons for the image list
|
* This is the function that creates the image buttons for the image list
|
||||||
* They are wrapped in a button so that they can be clicked and open
|
* They are wrapped in a button so that they can be clicked and open
|
||||||
* the image in the lightbox.
|
* the image in the lightbox.
|
||||||
* @param {*} src
|
* @param {*} src
|
||||||
*/
|
*/
|
||||||
const createImageBtn = (src) => {
|
const createImageBtn = (src) => {
|
||||||
console.debug(`making image ${src.filename}`)
|
console.debug(`making image ${src.filename}`)
|
||||||
const img = document.createElement('img')
|
const img = document.createElement('img')
|
||||||
const but = document.createElement('button')
|
const but = document.createElement('button')
|
||||||
|
|
||||||
Object.assign(but.style, {
|
Object.assign(but.style, {
|
||||||
height: '120px',
|
height: '120px',
|
||||||
width: '120px',
|
width: '120px',
|
||||||
border: 'none',
|
border: 'none',
|
||||||
padding: 0,
|
padding: 0,
|
||||||
margin: 0,
|
margin: 0,
|
||||||
})
|
})
|
||||||
Object.assign(img.style, {
|
Object.assign(img.style, {
|
||||||
width: '100%',
|
width: '100%',
|
||||||
height: '100%',
|
height: '100%',
|
||||||
objectFit: 'cover',
|
objectFit: 'cover',
|
||||||
})
|
})
|
||||||
|
|
||||||
img.src = `/view?filename=${encodeURIComponent(src.filename)}&type=${
|
img.src = `/view?filename=${encodeURIComponent(src.filename)}&type=${
|
||||||
src.type
|
src.type
|
||||||
}&subfolder=${encodeURIComponent(src.subfolder)}`
|
}&subfolder=${encodeURIComponent(src.subfolder)}`
|
||||||
|
|
||||||
imageUrls.push(img.src)
|
imageUrls.push(img.src)
|
||||||
|
|
||||||
console.debug(img.src)
|
console.debug(img.src)
|
||||||
|
|
||||||
img.onload = () => {
|
img.onload = () => {
|
||||||
but.style.width = `${120 * (img.naturalWidth / img.naturalHeight)}px`
|
but.style.width = `${120 * (img.naturalWidth / img.naturalHeight)}px`
|
||||||
}
|
}
|
||||||
|
|
||||||
but.onclick = () => {
|
but.onclick = () => {
|
||||||
lightboxContainer.style.display = 'flex'
|
lightboxContainer.style.display = 'flex'
|
||||||
// add the same image to the lightbox
|
// add the same image to the lightbox
|
||||||
lightboxImage.src = img.src
|
lightboxImage.src = img.src
|
||||||
// lighboxContainer.replaceChildren(lightboxButtons, img);
|
// lighboxContainer.replaceChildren(lightboxButtons, img);
|
||||||
}
|
}
|
||||||
|
|
||||||
// add right click menu
|
// add right click menu
|
||||||
but.addEventListener('contextmenu', (e) => {
|
but.addEventListener('contextmenu', (e) => {
|
||||||
e.preventDefault()
|
e.preventDefault()
|
||||||
|
|
||||||
if (image_menu) {
|
if (image_menu) {
|
||||||
image_menu.remove()
|
image_menu.remove()
|
||||||
}
|
}
|
||||||
|
|
||||||
image_menu = document.createElement('div')
|
image_menu = document.createElement('div')
|
||||||
Object.assign(image_menu.style, {
|
Object.assign(image_menu.style, {
|
||||||
position: 'absolute',
|
position: 'absolute',
|
||||||
top: `${e.clientY}px`,
|
top: `${e.clientY}px`,
|
||||||
left: `${e.clientX}px`,
|
left: `${e.clientX}px`,
|
||||||
background: '#333',
|
background: '#333',
|
||||||
color: '#fff',
|
color: '#fff',
|
||||||
padding: '5px',
|
padding: '5px',
|
||||||
borderRadius: '5px',
|
borderRadius: '5px',
|
||||||
zIndex: 999,
|
zIndex: 999,
|
||||||
})
|
})
|
||||||
const load_img = document.createElement('button')
|
const load_img = document.createElement('button')
|
||||||
load_img.textContent = 'Load'
|
load_img.textContent = 'Load'
|
||||||
load_img.onclick = () => {
|
load_img.onclick = () => {
|
||||||
app.handleFile(img.src)
|
app.handleFile(img.src)
|
||||||
}
|
}
|
||||||
|
|
||||||
image_menu.appendChild(load_img)
|
image_menu.appendChild(load_img)
|
||||||
document.body.appendChild(image_menu)
|
document.body.appendChild(image_menu)
|
||||||
})
|
})
|
||||||
|
|
||||||
but.append(img)
|
but.append(img)
|
||||||
imageListContainer.prepend(but)
|
imageListContainer.prepend(but)
|
||||||
}
|
}
|
||||||
|
|
||||||
loadButton.onclick = async () => {
|
loadButton.onclick = async () => {
|
||||||
const all_history = await api.getHistory()
|
const all_history = await api.getHistory()
|
||||||
for (const history of all_history.History) {
|
for (const history of all_history.History) {
|
||||||
if (history.outputs) {
|
if (history.outputs) {
|
||||||
for (const key of Object.keys(history.outputs)) {
|
for (const key of Object.keys(history.outputs)) {
|
||||||
console.debug(key)
|
console.debug(key)
|
||||||
if (history.outputs[key].images) {
|
if (history.outputs[key].images) {
|
||||||
for (const im of history.outputs[key].images) {
|
for (const im of history.outputs[key].images) {
|
||||||
console.debug(im)
|
console.debug(im)
|
||||||
createImageBtn(im)
|
createImageBtn(im)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
// for (const src of outputs.outputs.images) {
|
// for (const src of outputs.outputs.images) {
|
||||||
// console.debug(src)
|
// console.debug(src)
|
||||||
// makeImage(`${src.subfolder}/${src.filename}`)
|
// makeImage(`${src.subfolder}/${src.filename}`)
|
||||||
// }
|
// }
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
///////-------
|
///////-------
|
||||||
|
|
||||||
// const all_history = await api.getHistory()
|
// const all_history = await api.getHistory()
|
||||||
// for (const history of all_history.History) {
|
// for (const history of all_history.History) {
|
||||||
// if (history.outputs) {
|
// if (history.outputs) {
|
||||||
// for (const key of Object.keys(history.outputs)) {
|
// for (const key of Object.keys(history.outputs)) {
|
||||||
// for (const im of history.outputs[key].images) {
|
// for (const im of history.outputs[key].images) {
|
||||||
// makeImage(im)
|
// makeImage(im)
|
||||||
// }
|
// }
|
||||||
// }
|
// }
|
||||||
// // for (const src of outputs.outputs.images) {
|
// // for (const src of outputs.outputs.images) {
|
||||||
// // console.debug(src)
|
// // console.debug(src)
|
||||||
// // makeImage(`${src.subfolder}/${src.filename}`)
|
// // makeImage(`${src.subfolder}/${src.filename}`)
|
||||||
// // }
|
// // }
|
||||||
// }
|
// }
|
||||||
// }
|
// }
|
||||||
|
|
||||||
//- Hook into the API
|
//- Hook into the API
|
||||||
api.addEventListener('executed', ({ detail }) => {
|
api.addEventListener('executed', ({ detail }) => {
|
||||||
if (detail?.output?.images) {
|
if (detail?.output?.images) {
|
||||||
for (const src of detail.output.images) {
|
for (const src of detail.output.images) {
|
||||||
console.debug(`Adding ${src} to image feed`)
|
console.debug(`Adding ${src} to image feed`)
|
||||||
createImageBtn(src)
|
createImageBtn(src)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
},
|
},
|
||||||
})
|
})
|
||||||
|
|||||||
+143
-43
@@ -3,6 +3,7 @@
|
|||||||
import { app } from '../../scripts/app.js'
|
import { app } from '../../scripts/app.js'
|
||||||
import { api } from '../../scripts/api.js'
|
import { api } from '../../scripts/api.js'
|
||||||
|
|
||||||
|
import * as mtb_ui from './mtb_ui.js'
|
||||||
import * as shared from './comfy_shared.js'
|
import * as shared from './comfy_shared.js'
|
||||||
|
|
||||||
import {
|
import {
|
||||||
@@ -15,7 +16,13 @@ import {
|
|||||||
} from './mtb_ui.js'
|
} from './mtb_ui.js'
|
||||||
|
|
||||||
const offset = 0
|
const offset = 0
|
||||||
|
|
||||||
|
// These are "global" variables mostly meant to sync user settings.
|
||||||
let currentWidth = 200
|
let currentWidth = 200
|
||||||
|
let saltUrls =
|
||||||
|
app.extensionManager.setting.get('mtb.io-sidebar.salt_urls') || false
|
||||||
|
let targetWidth =
|
||||||
|
app.extensionManager.setting.get('mtb.io-sidebar.img-size') || 512
|
||||||
let currentMode = 'input'
|
let currentMode = 'input'
|
||||||
let subfolder = ''
|
let subfolder = ''
|
||||||
let currentSort = 'None'
|
let currentSort = 'None'
|
||||||
@@ -46,15 +53,19 @@ const updateImage = (node, image) => {
|
|||||||
* @param {ResultItem} resultItem
|
* @param {ResultItem} resultItem
|
||||||
* @returns {string} - The request URL.
|
* @returns {string} - The request URL.
|
||||||
*/
|
*/
|
||||||
const resultItemToQuery = (resultItem) =>
|
const resultItemToQuery = (resultItem) => {
|
||||||
[
|
const res = [
|
||||||
`/mtb/view?filename=${resultItem.filename}`,
|
`/mtb/view?filename=${resultItem.filename}`,
|
||||||
`width=512`,
|
|
||||||
`type=${resultItem.type}`,
|
`type=${resultItem.type}`,
|
||||||
`subfolder=${resultItem.subfolder}`,
|
`subfolder=${resultItem.subfolder}`,
|
||||||
`preview=`,
|
'preview=',
|
||||||
].join('&')
|
]
|
||||||
|
if (targetWidth > 0) {
|
||||||
|
res.splice(1, 0, `width=${targetWidth}`)
|
||||||
|
}
|
||||||
|
|
||||||
|
return res.join('&')
|
||||||
|
}
|
||||||
/**
|
/**
|
||||||
* Retrieves the unique prompt ID from a history task item.
|
* Retrieves the unique prompt ID from a history task item.
|
||||||
* @param {HistoryTaskItem} historyTaskItem
|
* @param {HistoryTaskItem} historyTaskItem
|
||||||
@@ -80,7 +91,7 @@ const getNewOutputUrls = (mostRecentTask) => {
|
|||||||
const imageOutputs = Object.values(nodeOutputs.images)
|
const imageOutputs = Object.values(nodeOutputs.images)
|
||||||
imageOutputs.forEach(
|
imageOutputs.forEach(
|
||||||
(resultItem) =>
|
(resultItem) =>
|
||||||
(urls[resultItem.filename] = resultItemToQuery(resultItem))
|
(urls[resultItem.filename] = resultItemToQuery(resultItem)),
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
// Can process `animated` and `audio` outputs here.
|
// Can process `animated` and `audio` outputs here.
|
||||||
@@ -209,7 +220,7 @@ const getUrls = async (subfolder) => {
|
|||||||
if (currentMode === 'video') {
|
if (currentMode === 'video') {
|
||||||
const output = await shared.runAction(
|
const output = await shared.runAction(
|
||||||
'getUserVideos',
|
'getUserVideos',
|
||||||
256,
|
targetWidth,
|
||||||
count,
|
count,
|
||||||
offset,
|
offset,
|
||||||
currentSort,
|
currentSort,
|
||||||
@@ -219,11 +230,13 @@ const getUrls = async (subfolder) => {
|
|||||||
const output = await shared.runAction(
|
const output = await shared.runAction(
|
||||||
'getUserImages',
|
'getUserImages',
|
||||||
currentMode,
|
currentMode,
|
||||||
|
targetWidth,
|
||||||
count,
|
count,
|
||||||
offset,
|
offset,
|
||||||
currentSort,
|
currentSort,
|
||||||
false,
|
false,
|
||||||
subfolder,
|
subfolder,
|
||||||
|
saltUrls,
|
||||||
)
|
)
|
||||||
return output || {}
|
return output || {}
|
||||||
}
|
}
|
||||||
@@ -236,55 +249,110 @@ if (window?.__COMFYUI_FRONTEND_VERSION__) {
|
|||||||
|
|
||||||
const sidebar_extension = {
|
const sidebar_extension = {
|
||||||
name: 'mtb.io-sidebar',
|
name: 'mtb.io-sidebar',
|
||||||
// init: async () => {
|
settings: [
|
||||||
// try {
|
{
|
||||||
// const res = await api.fetchApi('/mtb/server-info')
|
|
||||||
// const msg = await res.json()
|
|
||||||
// exposed = msg.exposed
|
|
||||||
// } catch (e) {
|
|
||||||
// console.error('Error:', e)
|
|
||||||
// }
|
|
||||||
// },
|
|
||||||
init: () => {
|
|
||||||
let handle
|
|
||||||
const version = window?.__COMFYUI_FRONTEND_VERSION__
|
|
||||||
console.log(`%c ${version}`, 'background: orange; color: white;')
|
|
||||||
|
|
||||||
ensureMTBStyles()
|
|
||||||
|
|
||||||
app.ui.settings.addSetting({
|
|
||||||
id: 'mtb.io-sidebar.count',
|
id: 'mtb.io-sidebar.count',
|
||||||
category: ['mtb', 'Input & Output Sidebar', 'count'],
|
category: ['mtb', 'Input & Output Sidebar', 'count'],
|
||||||
|
|
||||||
name: 'Number of images to fetch',
|
name: 'Number of images to fetch',
|
||||||
type: 'number',
|
type: 'number',
|
||||||
defaultValue: 1000,
|
defaultValue: 1000,
|
||||||
|
|
||||||
tooltip:
|
tooltip:
|
||||||
"This setting affects the input/output sidebar to determine how many images to fetch per pagination (pagination is not yet supported so for now it's the static total)",
|
"This setting affects the input/output sidebar to determine how many images to fetch per pagination (pagination is not yet supported so for now it's the static total)",
|
||||||
attrs: {
|
},
|
||||||
style: {
|
{
|
||||||
// fontFamily: 'monospace',
|
id: 'mtb.io-sidebar.salt_urls',
|
||||||
},
|
category: ['mtb', 'Input & Output Sidebar', 'salt_urls'],
|
||||||
|
name: 'Salt URLs',
|
||||||
|
type: 'boolean',
|
||||||
|
defaultValue: false,
|
||||||
|
onChange: (n, o) => {
|
||||||
|
saltUrls = n
|
||||||
},
|
},
|
||||||
})
|
tooltip:
|
||||||
|
'Adds a random query parameter to every urls to always invalidate caching.',
|
||||||
app.ui.settings.addSetting({
|
},
|
||||||
|
{
|
||||||
id: 'mtb.io-sidebar.img-size',
|
id: 'mtb.io-sidebar.img-size',
|
||||||
category: ['mtb', 'Input & Output Sidebar', 'img-size'],
|
category: ['mtb', 'Input & Output Sidebar', 'img-size'],
|
||||||
|
|
||||||
name: 'Resolution of the images',
|
name: 'Resize width of shown images',
|
||||||
type: 'number',
|
|
||||||
defaultValue: 512,
|
defaultValue: 512,
|
||||||
|
type: (name, setter, value, attrs) => {
|
||||||
|
targetWidth = value
|
||||||
|
const container = mtb_ui.makeElement('div', {
|
||||||
|
display: 'flex',
|
||||||
|
alignItems: 'center',
|
||||||
|
gap: '8px',
|
||||||
|
})
|
||||||
|
|
||||||
tooltip: "It's recommended to keep it at 512px",
|
console.log({ name, setter, value, attrs })
|
||||||
attrs: {
|
|
||||||
style: {
|
const baseId = name.replace(/[^a-zA-Z0-9]/g, '-').toLowerCase()
|
||||||
// fontFamily: 'monospace',
|
const checkboxId = `${baseId}-checkbox`
|
||||||
},
|
const numberInputId = `${baseId}-number`
|
||||||
|
|
||||||
|
const isCheckedInitially = value !== -1
|
||||||
|
|
||||||
|
// TODO: better way to get defaultValue?
|
||||||
|
const defaultValue = 512
|
||||||
|
const initialNumberValue = isCheckedInitially ? value : defaultValue
|
||||||
|
|
||||||
|
console.log('recreate')
|
||||||
|
const checkbox = mtb_ui.makeElement(
|
||||||
|
// harder to match styles (.p-toggleswitch-input)
|
||||||
|
// since it uses a div synced to the input...
|
||||||
|
'input',
|
||||||
|
{},
|
||||||
|
container,
|
||||||
|
)
|
||||||
|
checkbox.type = 'checkbox'
|
||||||
|
checkbox.id = checkboxId
|
||||||
|
checkbox.checked = isCheckedInitially
|
||||||
|
|
||||||
|
const numberInput = mtb_ui.makeElement(
|
||||||
|
'input.p-inputtext',
|
||||||
|
{},
|
||||||
|
container,
|
||||||
|
)
|
||||||
|
numberInput.type = 'number'
|
||||||
|
numberInput.id = numberInputId
|
||||||
|
numberInput.value = initialNumberValue
|
||||||
|
numberInput.disabled = !isCheckedInitially
|
||||||
|
numberInput.min = 128
|
||||||
|
|
||||||
|
checkbox.addEventListener('change', () => {
|
||||||
|
let valToSet = -1
|
||||||
|
if (checkbox.checked) {
|
||||||
|
numberInput.disabled = false
|
||||||
|
|
||||||
|
valToSet = Number.parseInt(numberInput.value, 10)
|
||||||
|
if (Number.isNaN(valToSet) || valToSet < numberInput.min) {
|
||||||
|
valToSet = defaultValue
|
||||||
|
numberInput.value = valToSet
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
numberInput.disabled = true
|
||||||
|
}
|
||||||
|
setter(valToSet)
|
||||||
|
})
|
||||||
|
|
||||||
|
numberInput.addEventListener('input', () => {
|
||||||
|
if (checkbox.checked) {
|
||||||
|
const numValue = Number.parseInt(numberInput.value, 10)
|
||||||
|
if (!Number.isNaN(numValue) && numberInput.value !== '') {
|
||||||
|
setter(numValue)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
return container
|
||||||
},
|
},
|
||||||
})
|
|
||||||
app.ui.settings.addSetting({
|
tooltip:
|
||||||
|
"If browsing large folders it's recommended to use this to avoid overflow/crash of the webpage. Image will get resized to this target width on the server before being sent to the client.",
|
||||||
|
},
|
||||||
|
|
||||||
|
{
|
||||||
id: 'mtb.io-sidebar.sort',
|
id: 'mtb.io-sidebar.sort',
|
||||||
category: ['mtb', 'Input & Output Sidebar', 'sort'],
|
category: ['mtb', 'Input & Output Sidebar', 'sort'],
|
||||||
name: 'Default sort mode',
|
name: 'Default sort mode',
|
||||||
@@ -304,7 +372,39 @@ if (window?.__COMFYUI_FRONTEND_VERSION__) {
|
|||||||
'Name',
|
'Name',
|
||||||
'Name-Reverse',
|
'Name-Reverse',
|
||||||
],
|
],
|
||||||
})
|
},
|
||||||
|
{
|
||||||
|
id: 'mtb.io-sidebar.notice',
|
||||||
|
category: ['mtb', 'Input & Output Sidebar', 'sort'],
|
||||||
|
name: ' ',
|
||||||
|
|
||||||
|
type: (name, setter, value, attrs) => {
|
||||||
|
const container = mtb_ui.makeElement('div')
|
||||||
|
const notice =
|
||||||
|
'## Important\nIf you make **any** edits here you need to toggle off and back on the sidebar for it to take effect.'
|
||||||
|
|
||||||
|
if (window.MTB?.mdParser) {
|
||||||
|
MTB.mdParser.parse(notice).then((e) => {
|
||||||
|
container.innerHTML = e
|
||||||
|
})
|
||||||
|
} else {
|
||||||
|
shared.ensureMarkdownParser((p) => {
|
||||||
|
p.parse(notice).then((e) => {
|
||||||
|
container.innerHTML = e
|
||||||
|
})
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return container
|
||||||
|
},
|
||||||
|
},
|
||||||
|
],
|
||||||
|
|
||||||
|
init: () => {
|
||||||
|
let handle
|
||||||
|
const version = window?.__COMFYUI_FRONTEND_VERSION__
|
||||||
|
console.log(`%c ${version}`, 'background: orange; color: white;')
|
||||||
|
|
||||||
|
ensureMTBStyles()
|
||||||
|
|
||||||
app.extensionManager.registerSidebarTab({
|
app.extensionManager.registerSidebarTab({
|
||||||
id: 'mtb-inputs-outputs',
|
id: 'mtb-inputs-outputs',
|
||||||
|
|||||||
+529
@@ -0,0 +1,529 @@
|
|||||||
|
/** Python REPL for the frontend (uses rich)*/
|
||||||
|
|
||||||
|
import { app } from '../../scripts/app.js'
|
||||||
|
import * as shared from './comfy_shared.js'
|
||||||
|
import * as mtb_ui from './mtb_ui.js'
|
||||||
|
|
||||||
|
class ComfyREPL extends LiteGraph.LGraphNode {
|
||||||
|
constructor() {
|
||||||
|
super()
|
||||||
|
|
||||||
|
this.shape = LiteGraph.BOX_SHAPE
|
||||||
|
this.isVirtualNode = true
|
||||||
|
this.category = 'mtb/repl'
|
||||||
|
this.title = '🐍 REPL (mtb)'
|
||||||
|
|
||||||
|
this.uuid = shared.makeUUID()
|
||||||
|
|
||||||
|
this.size = [600, 400]
|
||||||
|
|
||||||
|
// Create a container for our custom widgets
|
||||||
|
this.widget = this.addDOMWidget('HTML', 'html', this.createREPLWidget())
|
||||||
|
|
||||||
|
this.loadAceEditor()
|
||||||
|
|
||||||
|
// Store input and output for persistence
|
||||||
|
this.properties = {
|
||||||
|
inputCode: '',
|
||||||
|
outputHistory: '',
|
||||||
|
}
|
||||||
|
|
||||||
|
this.outputArea.innerHTML = this.properties.outputHistory
|
||||||
|
this.outputArea.scrollTop = this.outputArea.scrollHeight
|
||||||
|
|
||||||
|
// Debounced linting function
|
||||||
|
this.debouncedLint = shared.debounce(this.lintCode.bind(this), 500)
|
||||||
|
|
||||||
|
// Resizing state variables
|
||||||
|
this.isResizing = false
|
||||||
|
this.initialMouseY = 0
|
||||||
|
this.initialInputHeight = 0
|
||||||
|
this.initialOutputHeight = 0
|
||||||
|
}
|
||||||
|
|
||||||
|
loadAceEditor() {
|
||||||
|
if (window.MTB?.ace_loaded) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
let NEED_PATCH = false
|
||||||
|
if (window.ace) {
|
||||||
|
shared.infoLogger(
|
||||||
|
'A global ace was found in scope, to avoid issues with it we will patch it',
|
||||||
|
)
|
||||||
|
NEED_PATCH = true
|
||||||
|
// window._backupAce = window.ace
|
||||||
|
// window.ace = null
|
||||||
|
}
|
||||||
|
|
||||||
|
shared
|
||||||
|
.loadScript('/mtb_async/ace/ace.js')
|
||||||
|
.then((m) => {
|
||||||
|
shared.infoLogger('ACE was loaded', m)
|
||||||
|
// window.MTB_ACE = window.ace
|
||||||
|
window.MTB.ace_loaded = true
|
||||||
|
this.initAceEditor()
|
||||||
|
|
||||||
|
this.aceEditor.setValue(this.properties.inputCode, -1)
|
||||||
|
})
|
||||||
|
.catch((e) => {
|
||||||
|
shared.errorLogger(e)
|
||||||
|
})
|
||||||
|
.finally(() => {
|
||||||
|
if (NEED_PATCH) {
|
||||||
|
console.log('Patching back window object')
|
||||||
|
window.ace = window._backupAce
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
initAceEditor() {
|
||||||
|
if (!window.MTB.ace_loaded) {
|
||||||
|
console.error('ACE editor not loaded. Cannot set up editors.')
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if (!this.inputDiv) {
|
||||||
|
console.error('Input div not found for Ace editor initialization.')
|
||||||
|
return
|
||||||
|
}
|
||||||
|
this.aceEditor = ace.edit(this.inputDiv)
|
||||||
|
this.aceEditor.setTheme('ace/theme/monokai') //"ace/theme/dracula", "ace/theme/github"
|
||||||
|
this.aceEditor.session.setMode('ace/mode/python')
|
||||||
|
this.aceEditor.setOptions({
|
||||||
|
enableBasicAutocompletion: true,
|
||||||
|
enableLiveAutocompletion: true,
|
||||||
|
enableSnippets: true,
|
||||||
|
fontSize: '14px',
|
||||||
|
fontFamily: 'monospace',
|
||||||
|
showPrintMargin: false,
|
||||||
|
wrap: true,
|
||||||
|
tabSize: 4,
|
||||||
|
useSoftTabs: true,
|
||||||
|
highlightActiveLine: true,
|
||||||
|
highlightSelectedWord: true,
|
||||||
|
cursorStyle: 'ace', // "ace" | "slim" | "smooth" | "wide"
|
||||||
|
behavioursEnabled: true,
|
||||||
|
displayIndentGuides: true,
|
||||||
|
fixedWidthGutter: true,
|
||||||
|
scrollPastEnd: 0.5,
|
||||||
|
})
|
||||||
|
|
||||||
|
// Custom keybinding for Ctrl+Enter
|
||||||
|
this.aceEditor.commands.addCommand({
|
||||||
|
name: 'runCode',
|
||||||
|
bindKey: { win: 'Ctrl-Enter', mac: 'Command-Enter' },
|
||||||
|
exec: () => this.executeCode(),
|
||||||
|
})
|
||||||
|
|
||||||
|
// Listen for changes to trigger linting
|
||||||
|
let lintDisabled = false
|
||||||
|
this.aceEditor.session.on('change', () => {
|
||||||
|
if (!lintDisabled) {
|
||||||
|
this.debouncedLint()
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
this.outputArea.scrollTop = this.outputArea.scrollHeight
|
||||||
|
}
|
||||||
|
|
||||||
|
addOutput(html) {
|
||||||
|
this.outputArea.innerHTML += html
|
||||||
|
this.properties.outputHistory += html
|
||||||
|
this.outputArea.scrollTop = this.outputArea.scrollHeight
|
||||||
|
}
|
||||||
|
|
||||||
|
createREPLWidget() {
|
||||||
|
const container = mtb_ui.makeElement('div', {
|
||||||
|
display: 'flex',
|
||||||
|
flexDirection: 'column',
|
||||||
|
width: '100%',
|
||||||
|
height: '100%',
|
||||||
|
boxSizing: 'border-box',
|
||||||
|
padding: '5px',
|
||||||
|
})
|
||||||
|
|
||||||
|
this.inputDiv = mtb_ui.makeElement(
|
||||||
|
'div',
|
||||||
|
{
|
||||||
|
width: 'calc(100% - 10px)',
|
||||||
|
height: '100px',
|
||||||
|
backgroundColor: '#333',
|
||||||
|
color: '#eee',
|
||||||
|
border: '1px solid #555',
|
||||||
|
borderRadius: '4px',
|
||||||
|
marginBottom: '5px',
|
||||||
|
boxSizing: 'border-box',
|
||||||
|
overflow: 'hidden',
|
||||||
|
},
|
||||||
|
container,
|
||||||
|
)
|
||||||
|
// Resizable Handle
|
||||||
|
this.handleDiv = mtb_ui.makeElement(
|
||||||
|
'div',
|
||||||
|
{
|
||||||
|
width: '100%',
|
||||||
|
height: '5px',
|
||||||
|
backgroundColor: '#666',
|
||||||
|
cursor: 'ns-resize',
|
||||||
|
marginBottom: '5px',
|
||||||
|
borderRadius: '2px',
|
||||||
|
},
|
||||||
|
container,
|
||||||
|
)
|
||||||
|
this.handleDiv.addEventListener('mousedown', this.startResizing.bind(this))
|
||||||
|
|
||||||
|
// Run Button
|
||||||
|
this.runButton = mtb_ui.makeElement(
|
||||||
|
'button',
|
||||||
|
{
|
||||||
|
width: '100%',
|
||||||
|
padding: '8px',
|
||||||
|
backgroundColor: '#555',
|
||||||
|
color: '#fff',
|
||||||
|
border: 'none',
|
||||||
|
borderRadius: '4px',
|
||||||
|
cursor: 'pointer',
|
||||||
|
marginBottom: '5px',
|
||||||
|
fontSize: '14px',
|
||||||
|
},
|
||||||
|
container,
|
||||||
|
)
|
||||||
|
this.runButton.textContent = 'Run Code (Ctrl+Enter)'
|
||||||
|
this.runButton.onclick = () => this.executeCode()
|
||||||
|
|
||||||
|
// Clear Button
|
||||||
|
this.clearButton = mtb_ui.makeElement(
|
||||||
|
'button',
|
||||||
|
{
|
||||||
|
width: '100%',
|
||||||
|
padding: '8px',
|
||||||
|
backgroundColor: '#555',
|
||||||
|
color: '#fff',
|
||||||
|
border: 'none',
|
||||||
|
borderRadius: '4px',
|
||||||
|
cursor: 'pointer',
|
||||||
|
marginBottom: '5px',
|
||||||
|
fontSize: '14px',
|
||||||
|
},
|
||||||
|
container,
|
||||||
|
)
|
||||||
|
this.clearButton.textContent = 'Clear Output'
|
||||||
|
this.clearButton.onclick = () => {
|
||||||
|
this.outputArea.innerHTML = ''
|
||||||
|
this.properties.outputHistory = ''
|
||||||
|
}
|
||||||
|
|
||||||
|
// Output Area
|
||||||
|
this.outputArea = mtb_ui.makeElement(
|
||||||
|
'div',
|
||||||
|
{
|
||||||
|
flexGrow: '1',
|
||||||
|
width: 'calc(100% - 10px)',
|
||||||
|
backgroundColor: '#222',
|
||||||
|
color: '#ddd',
|
||||||
|
border: '1px solid #555',
|
||||||
|
borderRadius: '4px',
|
||||||
|
padding: '5px',
|
||||||
|
fontFamily: 'monospace',
|
||||||
|
fontSize: '14px',
|
||||||
|
overflowY: 'auto',
|
||||||
|
whiteSpace: 'pre-wrap',
|
||||||
|
boxSizing: 'border-box',
|
||||||
|
},
|
||||||
|
container,
|
||||||
|
)
|
||||||
|
|
||||||
|
return container
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- Resizing Logic ---
|
||||||
|
startResizing(e) {
|
||||||
|
if (!this.inputDiv) {
|
||||||
|
shared.infoLogger("The input div isn't ready", this)
|
||||||
|
shared.errorLogger("The input div isn't ready")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
this.isResizing = true
|
||||||
|
this.initialMouseY = e.clientY
|
||||||
|
this.initialInputHeight = this.inputDiv.offsetHeight
|
||||||
|
this.initialOutputHeight = this.outputArea.offsetHeight
|
||||||
|
|
||||||
|
document.addEventListener('mousemove', this.doResize.bind(this))
|
||||||
|
document.addEventListener('mouseup', this.stopResizing.bind(this))
|
||||||
|
document.body.style.cursor = 'ns-resize' // Change cursor globally
|
||||||
|
}
|
||||||
|
doResize(e) {
|
||||||
|
if (!this.isResizing) return
|
||||||
|
|
||||||
|
const deltaY = e.clientY - this.initialMouseY
|
||||||
|
|
||||||
|
let new_input_height = this.initialInputHeight + deltaY
|
||||||
|
let new_output_height = this.initialOutputHeight - deltaY
|
||||||
|
|
||||||
|
const minInputHeight = 50 // Minimum height for Ace editor
|
||||||
|
const minOutputHeight = 50 // Minimum height for output area
|
||||||
|
|
||||||
|
// Clamp heights to minimums
|
||||||
|
if (new_input_height < minInputHeight) {
|
||||||
|
new_input_height = minInputHeight
|
||||||
|
new_output_height =
|
||||||
|
this.initialInputHeight + this.initialOutputHeight - minInputHeight
|
||||||
|
}
|
||||||
|
if (new_output_height < minOutputHeight) {
|
||||||
|
new_output_height = minOutputHeight
|
||||||
|
new_input_height =
|
||||||
|
this.initialInputHeight + this.initialOutputHeight - minOutputHeight
|
||||||
|
}
|
||||||
|
|
||||||
|
this.inputDiv.style.height = `${new_input_height}px`
|
||||||
|
this.outputArea.style.height = `${new_output_height}px`
|
||||||
|
|
||||||
|
// Update the stored ratio for persistence
|
||||||
|
const totalDynamicHeight =
|
||||||
|
this.inputDiv.offsetHeight + this.outputArea.offsetHeight
|
||||||
|
if (totalDynamicHeight > 0) {
|
||||||
|
this.properties.inputHeightRatio = new_input_height / totalDynamicHeight
|
||||||
|
}
|
||||||
|
|
||||||
|
this.aceEditor.resize() // Important for Ace to redraw
|
||||||
|
}
|
||||||
|
|
||||||
|
stopResizing() {
|
||||||
|
this.isResizing = false
|
||||||
|
document.removeEventListener('mousemove', this.doResize)
|
||||||
|
document.removeEventListener('mouseup', this.stopResizing)
|
||||||
|
document.body.style.cursor = '' // Restore default cursor
|
||||||
|
}
|
||||||
|
// --- End Resizing Logic ---
|
||||||
|
|
||||||
|
async executeCode() {
|
||||||
|
const code = this.aceEditor.getValue()
|
||||||
|
|
||||||
|
if (!code.trim()) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
const inputPrompt = `<div style="color:#888; margin-top: 10px;">>>> ${code}</div>`
|
||||||
|
|
||||||
|
this.addOutput(inputPrompt)
|
||||||
|
|
||||||
|
try {
|
||||||
|
const response = await fetch('/mtb/execute', {
|
||||||
|
method: 'POST',
|
||||||
|
headers: {
|
||||||
|
'Content-Type': 'application/json',
|
||||||
|
},
|
||||||
|
body: JSON.stringify({ code: code, name: this.uuid }),
|
||||||
|
})
|
||||||
|
|
||||||
|
if (!response.ok) {
|
||||||
|
throw new Error(`HTTP error! status: ${response.status}`)
|
||||||
|
}
|
||||||
|
|
||||||
|
const result = await response.json()
|
||||||
|
console.debug('Received from backend', result)
|
||||||
|
const outputHtml = result.output_html || ''
|
||||||
|
const error = result.error
|
||||||
|
|
||||||
|
if (error) {
|
||||||
|
this.addOutput(
|
||||||
|
`<div style="color: #f00; font-weight: bold;">Error:</div>${outputHtml}`,
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
this.addOutput(outputHtml)
|
||||||
|
}
|
||||||
|
} catch (e) {
|
||||||
|
const errorMessage = `<div style="color: #f00;">Frontend Error: ${e.message}</div>`
|
||||||
|
|
||||||
|
this.addOutput(errorMessage)
|
||||||
|
console.error('ComfyREPL Frontend Error:', e)
|
||||||
|
} finally {
|
||||||
|
// Not clearing
|
||||||
|
// this.inputArea.value = '' // Clear input after execution
|
||||||
|
// this.properties.inputCode = '' // Clear persisted input
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
async lintCode() {
|
||||||
|
if (!this.aceEditor) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
const code = this.aceEditor.getValue()
|
||||||
|
if (!code.trim()) {
|
||||||
|
this.aceEditor.session.setAnnotations([]) // Clear annotations if empty
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
try {
|
||||||
|
const response = await fetch('/mtb/lint', {
|
||||||
|
// New linting endpoint
|
||||||
|
method: 'POST',
|
||||||
|
headers: {
|
||||||
|
'Content-Type': 'application/json',
|
||||||
|
},
|
||||||
|
body: JSON.stringify({ code: code, name: this.uuid }),
|
||||||
|
})
|
||||||
|
|
||||||
|
if (!response.ok) {
|
||||||
|
console.log(response)
|
||||||
|
throw new Error(
|
||||||
|
`HTTP error! status: ${response.status} ${response.statusText}`,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
const result = await response.json()
|
||||||
|
// result.diagnostics should be an array of {row, column, text, type}
|
||||||
|
this.aceEditor.session.setAnnotations(result.diagnostics)
|
||||||
|
} catch (e) {
|
||||||
|
console.error('ComfyREPL Linting Error:', e)
|
||||||
|
this.aceEditor.session.setAnnotations([
|
||||||
|
{
|
||||||
|
row: 0,
|
||||||
|
column: 0,
|
||||||
|
text: `Linting failed: ${e.message}`,
|
||||||
|
type: 'error',
|
||||||
|
},
|
||||||
|
])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Restore properties when loading a graph
|
||||||
|
onConfigure() {
|
||||||
|
if (this.properties.inputCode && this.aceEditor) {
|
||||||
|
this.aceEditor.setValue(this.properties.inputCode, -1)
|
||||||
|
}
|
||||||
|
// if (this.properties.inputCode) {
|
||||||
|
// this.inputArea.value = this.properties.inputCode
|
||||||
|
// }
|
||||||
|
if (this.properties.outputHistory) {
|
||||||
|
this.outputArea.innerHTML = this.properties.outputHistory
|
||||||
|
this.outputArea.scrollTop = this.outputArea.scrollHeight
|
||||||
|
}
|
||||||
|
if (this.properties.uuid) {
|
||||||
|
this.uuid = this.properties.uuid
|
||||||
|
}
|
||||||
|
this.debouncedLint()
|
||||||
|
this.onResize(this.size)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Save properties when saving a graph
|
||||||
|
onSerialize(o) {
|
||||||
|
if (this.aceEditor) {
|
||||||
|
o.properties.inputCode = this.aceEditor.getValue() //this.inputArea.value
|
||||||
|
}
|
||||||
|
o.properties.outputHistory = this.outputArea.innerHTML
|
||||||
|
o.properties.uuid = this.uuid
|
||||||
|
o.properties.inputHeightRatio = this.properties.inputHeightRatio
|
||||||
|
}
|
||||||
|
|
||||||
|
onRemoved() {
|
||||||
|
// Clean up DOM elements when node is removed
|
||||||
|
if (this.widget?.element?.parentNode) {
|
||||||
|
this.widget.element.parentNode.removeChild(this.widget.element)
|
||||||
|
}
|
||||||
|
// Destroy Ace editor instance to prevent memory leaks
|
||||||
|
if (this.aceEditor) {
|
||||||
|
this.aceEditor.destroy()
|
||||||
|
this.aceEditor.container.remove() // Remove the Ace container div from DOM
|
||||||
|
}
|
||||||
|
// Clean up global event listeners if node is removed while resizing
|
||||||
|
document.removeEventListener('mousemove', this.doResize)
|
||||||
|
document.removeEventListener('mouseup', this.stopResizing)
|
||||||
|
document.body.style.cursor = ''
|
||||||
|
}
|
||||||
|
// LiteGraph method to handle node resizing
|
||||||
|
onResize(size) {
|
||||||
|
// Call parent method if it exists (important for LiteGraph's internal sizing)
|
||||||
|
if (super.onResize) {
|
||||||
|
super.onResize(size)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Adjust container size
|
||||||
|
const container = this.widget.element
|
||||||
|
container.style.width = `${size[0] - 10}px` // Account for padding
|
||||||
|
container.style.height = `${size[1] - 10}px`
|
||||||
|
|
||||||
|
// Adjust input and output area widths
|
||||||
|
this.inputDiv.style.width = 'calc(100% - 10px)'
|
||||||
|
this.outputArea.style.width = 'calc(100% - 10px)'
|
||||||
|
//
|
||||||
|
// const old = () => {
|
||||||
|
// // Calculate remaining height for output area
|
||||||
|
// // Ace editor manages its own height within this.inputDiv, so we use offsetHeight
|
||||||
|
// const inputHeight = this.inputDiv.offsetHeight
|
||||||
|
// const runButtonHeight = this.runButton.offsetHeight
|
||||||
|
// const clearButtonHeight = this.clearButton.offsetHeight
|
||||||
|
// const totalFixedHeight =
|
||||||
|
// inputHeight + runButtonHeight + clearButtonHeight + 15 // 15 for margins/padding
|
||||||
|
//
|
||||||
|
// const remainingHeight = size[1] - 10 - totalFixedHeight
|
||||||
|
// this.outputArea.style.height = `${Math.max(50, remainingHeight)}px` // Min height 50px
|
||||||
|
// }
|
||||||
|
// Calculate dynamic heights
|
||||||
|
const containerHeight = size[1] - 10
|
||||||
|
const handleHeight = this.handleDiv.offsetHeight
|
||||||
|
const buttonHeights =
|
||||||
|
this.runButton.offsetHeight + this.clearButton.offsetHeight + 15 // Sum of button heights + margins
|
||||||
|
|
||||||
|
const dynamicContentHeight = containerHeight - buttonHeights - handleHeight
|
||||||
|
|
||||||
|
const minInputHeight = 50
|
||||||
|
const minOutputHeight = 50
|
||||||
|
|
||||||
|
let inputHeight = Math.max(
|
||||||
|
minInputHeight,
|
||||||
|
dynamicContentHeight * (this.properties.inputHeightRatio || 1.0),
|
||||||
|
)
|
||||||
|
let outputHeight = Math.max(
|
||||||
|
minOutputHeight,
|
||||||
|
dynamicContentHeight - inputHeight,
|
||||||
|
)
|
||||||
|
//
|
||||||
|
// // Re-distribute if one hits its minimum
|
||||||
|
// if (
|
||||||
|
// inputHeight === minInputHeight &&
|
||||||
|
// dynamicContentHeight - minInputHeight > minOutputHeight
|
||||||
|
// ) {
|
||||||
|
// outputHeight = dynamicContentHeight - minInputHeight
|
||||||
|
// } else if (
|
||||||
|
// outputHeight === minOutputHeight &&
|
||||||
|
// dynamicContentHeight - minOutputHeight > minInputHeight
|
||||||
|
// ) {
|
||||||
|
// inputHeight = dynamicContentHeight - minOutputHeight
|
||||||
|
// }
|
||||||
|
//
|
||||||
|
// // Final check to ensure total height matches available dynamic space
|
||||||
|
// const currentTotal = inputHeight + outputHeight
|
||||||
|
// if (currentTotal !== dynamicContentHeight) {
|
||||||
|
// // Adjust one of them if there's a small discrepancy due to rounding
|
||||||
|
// if (inputHeight > minInputHeight) {
|
||||||
|
// inputHeight += dynamicContentHeight - currentTotal
|
||||||
|
// } else if (outputHeight > minOutputHeight) {
|
||||||
|
// outputHeight += dynamicContentHeight - currentTotal
|
||||||
|
// }
|
||||||
|
// }
|
||||||
|
|
||||||
|
this.inputDiv.style.height = `${inputHeight}px`
|
||||||
|
this.outputArea.style.height = `${outputHeight}px`
|
||||||
|
|
||||||
|
// Update the ratio based on the actual heights set
|
||||||
|
if (dynamicContentHeight > 0) {
|
||||||
|
this.properties.inputHeightRatio = inputHeight / dynamicContentHeight
|
||||||
|
}
|
||||||
|
|
||||||
|
// Inform Ace editor about the resize so it can redraw its content
|
||||||
|
if (this.aceEditor) {
|
||||||
|
this.aceEditor.resize()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
const repl = {
|
||||||
|
name: 'mtb.repl',
|
||||||
|
|
||||||
|
registerCustomNodes() {
|
||||||
|
LiteGraph.registerNodeType('Python REPL', ComfyREPL)
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
app.registerExtension(repl)
|
||||||
@@ -1,28 +0,0 @@
|
|||||||
// NOTE: this will be the LT part of mtb API system
|
|
||||||
// I need to properly publish the source and fix a few things before
|
|
||||||
|
|
||||||
// import { app } from '../../scripts/app.js'
|
|
||||||
// // import { api } from '../../scripts/api.js'
|
|
||||||
//
|
|
||||||
// import * as shared from './comfy_shared.js'
|
|
||||||
// import { createOutliner } from './dist/mtb_inspector.js'
|
|
||||||
//
|
|
||||||
// if (window?.__COMFYUI_FRONTEND_VERSION__) {
|
|
||||||
// const version = window?.__COMFYUI_FRONTEND_VERSION__
|
|
||||||
// console.log(`%c ${version}`, 'background: orange; color: white;')
|
|
||||||
//
|
|
||||||
// const panel = app.extensionManager.registerSidebarTab({
|
|
||||||
// id: 'mtb-nodes',
|
|
||||||
// icon: 'pi pi-bolt',
|
|
||||||
// title: 'MTB',
|
|
||||||
// tooltip: 'MTB: API outliner',
|
|
||||||
// type: 'custom',
|
|
||||||
// // this is run everytime the tab's diplay is toggled on.
|
|
||||||
// render: (el) => {
|
|
||||||
// const outliner = createOutliner(el)
|
|
||||||
// const inputs = shared.getAPIInputs()
|
|
||||||
// console.log('INPUTS', inputs)
|
|
||||||
// outliner.$$set({ inputs })
|
|
||||||
// },
|
|
||||||
// })
|
|
||||||
// }
|
|
||||||
+4
-1
@@ -203,7 +203,7 @@ export const wrapElement = (element, style = {}) => {
|
|||||||
* @param {Object} [style] - CSS styles to apply to the element.
|
* @param {Object} [style] - CSS styles to apply to the element.
|
||||||
* @returns {HTMLElement} - The created DOM element.
|
* @returns {HTMLElement} - The created DOM element.
|
||||||
*/
|
*/
|
||||||
export const makeElement = (kind, style) => {
|
export const makeElement = (kind, style, parent) => {
|
||||||
let [real_kind, className] = kind.split('.')
|
let [real_kind, className] = kind.split('.')
|
||||||
let id
|
let id
|
||||||
|
|
||||||
@@ -224,6 +224,9 @@ export const makeElement = (kind, style) => {
|
|||||||
if (id) {
|
if (id) {
|
||||||
el.id = id
|
el.id = id
|
||||||
}
|
}
|
||||||
|
if (parent) {
|
||||||
|
parent.appendChild(el)
|
||||||
|
}
|
||||||
|
|
||||||
return el
|
return el
|
||||||
}
|
}
|
||||||
|
|||||||
+84
-45
@@ -694,7 +694,7 @@ const mtb_widgets = {
|
|||||||
|
|
||||||
app.ui.settings.addSetting({
|
app.ui.settings.addSetting({
|
||||||
id: 'mtb.Main.debug-enabled',
|
id: 'mtb.Main.debug-enabled',
|
||||||
category: ['mtb', 'Main', 'debug-enabled'],
|
category: ['mtb', ' Main', 'debug-enabled'],
|
||||||
name: 'Enable Debug (py and js)',
|
name: 'Enable Debug (py and js)',
|
||||||
type: 'boolean',
|
type: 'boolean',
|
||||||
defaultValue: false,
|
defaultValue: false,
|
||||||
@@ -1012,12 +1012,15 @@ const mtb_widgets = {
|
|||||||
)
|
)
|
||||||
loop_preview.value = 'Iteration: Idle'
|
loop_preview.value = 'Iteration: Idle'
|
||||||
|
|
||||||
|
let cancelQueue = false
|
||||||
|
|
||||||
const onReset = () => {
|
const onReset = () => {
|
||||||
raw_iteration.value = 0
|
raw_iteration.value = 0
|
||||||
raw_loop.value = 0
|
raw_loop.value = 0
|
||||||
|
|
||||||
value_preview.value = 'Idle'
|
value_preview.value = 'Idle'
|
||||||
loop_preview.value = 'Iteration: Idle'
|
loop_preview.value = 'Iteration: Idle'
|
||||||
|
cancelQueue = false
|
||||||
|
|
||||||
app.canvas.setDirty(true)
|
app.canvas.setDirty(true)
|
||||||
}
|
}
|
||||||
@@ -1026,15 +1029,42 @@ const mtb_widgets = {
|
|||||||
this.addWidget('button', 'Reset', 'reset', onReset)
|
this.addWidget('button', 'Reset', 'reset', onReset)
|
||||||
|
|
||||||
// run button
|
// run button
|
||||||
this.addWidget('button', 'Queue', 'queue', () => {
|
const chunkSize = 10
|
||||||
onReset() // this could maybe be a setting or checkbox
|
this.addWidget('button', 'Queue', 'queue', async () => {
|
||||||
app.queuePrompt(0, total_frames.value * loop_count.value)
|
onReset()
|
||||||
|
|
||||||
|
const totalPrompts = total_frames.value * loop_count.value
|
||||||
window.MTB?.notify?.(
|
window.MTB?.notify?.(
|
||||||
`Started a queue of ${total_frames.value} frames (for ${
|
`Starting a queue of ${totalPrompts} frames in chunks of ${chunkSize}...`,
|
||||||
loop_count.value
|
|
||||||
} loop, so ${total_frames.value * loop_count.value})`,
|
|
||||||
5000,
|
5000,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
for (let i = 0; i < totalPrompts; i += chunkSize) {
|
||||||
|
console.log({ cancelQueue })
|
||||||
|
if (cancelQueue) {
|
||||||
|
window.MTB?.notify?.(
|
||||||
|
`Queueing cancelled after ${i} frames.`,
|
||||||
|
3000,
|
||||||
|
)
|
||||||
|
break
|
||||||
|
}
|
||||||
|
const currentChunkSize = Math.min(chunkSize, totalPrompts - i)
|
||||||
|
|
||||||
|
await app.queuePrompt(0, currentChunkSize)
|
||||||
|
}
|
||||||
|
if (!cancelQueue) {
|
||||||
|
window.MTB?.notify?.(
|
||||||
|
`Finished queuing ${totalPrompts} frames.`,
|
||||||
|
5000,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
this.addWidget('button', 'Cancel', 'cancel', () => {
|
||||||
|
cancelQueue = true
|
||||||
|
window.MTB?.notify?.(
|
||||||
|
'Cancellation requested. Waiting for current chunk to finish...',
|
||||||
|
3000,
|
||||||
|
)
|
||||||
})
|
})
|
||||||
|
|
||||||
this.onRemoved = () => {
|
this.onRemoved = () => {
|
||||||
@@ -1166,7 +1196,9 @@ const mtb_widgets = {
|
|||||||
|
|
||||||
//NOTE: dynamic nodes
|
//NOTE: dynamic nodes
|
||||||
case 'Apply Text Template (mtb)': {
|
case 'Apply Text Template (mtb)': {
|
||||||
shared.setupDynamicConnections(nodeType, 'var', '*')
|
shared.setupDynamicConnections(nodeType, 'var', '*', {
|
||||||
|
rename_menu: 'name',
|
||||||
|
})
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
case 'Save Data Bundle (mtb)': {
|
case 'Save Data Bundle (mtb)': {
|
||||||
@@ -1298,47 +1330,54 @@ const mtb_widgets = {
|
|||||||
const related = new Set([this.id])
|
const related = new Set([this.id])
|
||||||
const visited = new Set()
|
const visited = new Set()
|
||||||
if (this.outputs[0].links) {
|
if (this.outputs[0].links) {
|
||||||
const initLink = this.outputs[0].links[0]
|
for (const linkId of this.outputs[0].links) {
|
||||||
const { to: loopEnd } = shared.nodesFromLink(this, initLink)
|
const { to: loopEnd } = shared.nodesFromLink(this, linkId)
|
||||||
const canReachEnd = (node, visited = new Set()) => {
|
const canReachEnd = (node, visited = new Set()) => {
|
||||||
if (node === loopEnd) return true
|
if (node === loopEnd) return true
|
||||||
if (visited.has(node.id)) return false
|
if (visited.has(node.id)) return false
|
||||||
visited.add(node.id)
|
visited.add(node.id)
|
||||||
for (const output of node.outputs || []) {
|
for (const output of node.outputs || []) {
|
||||||
if (!output.links) continue
|
if (!output.links) continue
|
||||||
for (const linkId of output.links) {
|
for (const linkId of output.links) {
|
||||||
const { to: nextNode } = shared.nodesFromLink(node, linkId)
|
const { to: nextNode } = shared.nodesFromLink(
|
||||||
if (!nextNode) continue
|
node,
|
||||||
if (canReachEnd(nextNode, visited)) {
|
linkId,
|
||||||
return true
|
)
|
||||||
|
if (!nextNode) continue
|
||||||
|
if (canReachEnd(nextNode, visited)) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
const traverseNodes = (node) => {
|
||||||
|
if (visited.has(node.id)) return
|
||||||
|
visited.add(node.id)
|
||||||
|
|
||||||
|
// can reach the end
|
||||||
|
if (node !== this && node !== loopEnd && !canReachEnd(node)) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
related.add(node.id)
|
||||||
|
for (const output of node.outputs || []) {
|
||||||
|
if (!output.links) continue
|
||||||
|
|
||||||
|
for (const linkId of output.links) {
|
||||||
|
const { to: nextNode } = shared.nodesFromLink(
|
||||||
|
node,
|
||||||
|
linkId,
|
||||||
|
)
|
||||||
|
if (!nextNode) continue
|
||||||
|
|
||||||
|
traverseNodes(nextNode)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return false
|
|
||||||
|
traverseNodes(this)
|
||||||
}
|
}
|
||||||
const traverseNodes = (node) => {
|
|
||||||
if (visited.has(node.id)) return
|
|
||||||
visited.add(node.id)
|
|
||||||
|
|
||||||
// can reach the end
|
|
||||||
if (node !== this && node !== loopEnd && !canReachEnd(node)) {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
related.add(node.id)
|
|
||||||
for (const output of node.outputs || []) {
|
|
||||||
if (!output.links) continue
|
|
||||||
|
|
||||||
for (const linkId of output.links) {
|
|
||||||
const { to: nextNode } = shared.nodesFromLink(node, linkId)
|
|
||||||
if (!nextNode) continue
|
|
||||||
|
|
||||||
traverseNodes(nextNode)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
traverseNodes(this)
|
|
||||||
}
|
}
|
||||||
this.related_to_flow = Array.from(related)
|
this.related_to_flow = Array.from(related)
|
||||||
this.computed_flow = true
|
this.computed_flow = true
|
||||||
|
|||||||
+64
-52
@@ -1,10 +1,13 @@
|
|||||||
// web/note_plus.constants.js
|
// web/note_plus.constants.js
|
||||||
|
|
||||||
export const DEFAULT_CSS = ''
|
export const DEFAULT_CSS = `/** here you can write css**/
|
||||||
|
h1 {
|
||||||
|
color: whitesmoke;
|
||||||
|
}`
|
||||||
export const DEFAULT_HTML = `<p style='color:red;font-family:monospace'>
|
export const DEFAULT_HTML = `<p style='color:red;font-family:monospace'>
|
||||||
Note+
|
Note+
|
||||||
</p>`
|
</p>`
|
||||||
export const DEFAULT_MD = '## Note+'
|
export const DEFAULT_MD = '# 📝 Note+'
|
||||||
export const DEFAULT_MODE = 'markdown'
|
export const DEFAULT_MODE = 'markdown'
|
||||||
export const DEFAULT_THEME = 'one_dark'
|
export const DEFAULT_THEME = 'one_dark'
|
||||||
|
|
||||||
@@ -55,58 +58,57 @@ We also support github callout:
|
|||||||
`
|
`
|
||||||
|
|
||||||
export const THEMES = [
|
export const THEMES = [
|
||||||
'ambiance',
|
'ambiance',
|
||||||
'chaos',
|
'chaos',
|
||||||
'chrome',
|
'chrome',
|
||||||
'cloud9_day',
|
'cloud9_day',
|
||||||
'cloud9_night',
|
'cloud9_night',
|
||||||
'cloud9_night_low_color',
|
'cloud9_night_low_color',
|
||||||
'cloud_editor',
|
'cloud_editor',
|
||||||
'cloud_editor_dark',
|
'cloud_editor_dark',
|
||||||
'clouds',
|
'clouds',
|
||||||
'clouds_midnight',
|
'clouds_midnight',
|
||||||
'cobalt',
|
'cobalt',
|
||||||
'crimson_editor',
|
'crimson_editor',
|
||||||
'dawn',
|
'dawn',
|
||||||
'dracula',
|
'dracula',
|
||||||
'dreamweaver',
|
'dreamweaver',
|
||||||
'eclipse',
|
'eclipse',
|
||||||
'github',
|
'github',
|
||||||
'github_dark',
|
'github_dark',
|
||||||
'gob',
|
'gob',
|
||||||
'gruvbox',
|
'gruvbox',
|
||||||
'gruvbox_dark_hard',
|
'gruvbox_dark_hard',
|
||||||
'gruvbox_light_hard',
|
'gruvbox_light_hard',
|
||||||
'idle_fingers',
|
'idle_fingers',
|
||||||
'iplastic',
|
'iplastic',
|
||||||
'katzenmilch',
|
'katzenmilch',
|
||||||
'kr_theme',
|
'kr_theme',
|
||||||
'kuroir',
|
'kuroir',
|
||||||
'merbivore',
|
'merbivore',
|
||||||
'merbivore_soft',
|
'merbivore_soft',
|
||||||
'mono_industrial',
|
'mono_industrial',
|
||||||
'monokai',
|
'monokai',
|
||||||
'nord_dark',
|
'nord_dark',
|
||||||
'one_dark',
|
'one_dark',
|
||||||
'pastel_on_dark',
|
'pastel_on_dark',
|
||||||
'solarized_dark',
|
'solarized_dark',
|
||||||
'solarized_light',
|
'solarized_light',
|
||||||
'sqlserver',
|
'sqlserver',
|
||||||
'terminal',
|
'terminal',
|
||||||
'textmate',
|
'textmate',
|
||||||
'tomorrow',
|
'tomorrow',
|
||||||
'tomorrow_night',
|
'tomorrow_night',
|
||||||
'tomorrow_night_blue',
|
'tomorrow_night_blue',
|
||||||
'tomorrow_night_bright',
|
'tomorrow_night_bright',
|
||||||
'tomorrow_night_eighties',
|
'tomorrow_night_eighties',
|
||||||
'twilight',
|
'twilight',
|
||||||
'vibrant_ink',
|
'vibrant_ink',
|
||||||
'vscode',
|
'vscode',
|
||||||
]
|
]
|
||||||
|
|
||||||
export const CSS_RESET = `
|
export const CSS_RESET = `
|
||||||
* {
|
* {
|
||||||
font-family: monospace;
|
|
||||||
line-height: 1.25em;
|
line-height: 1.25em;
|
||||||
}
|
}
|
||||||
.shiki{
|
.shiki{
|
||||||
@@ -116,6 +118,8 @@ export const CSS_RESET = `
|
|||||||
.markdown-callout-title {
|
.markdown-callout-title {
|
||||||
.octicon{
|
.octicon{
|
||||||
fill:white;
|
fill:white;
|
||||||
|
width:29px;
|
||||||
|
height:29px;
|
||||||
}
|
}
|
||||||
/* background: var(--current-color); */
|
/* background: var(--current-color); */
|
||||||
color: var(--current-color);
|
color: var(--current-color);
|
||||||
@@ -124,6 +128,8 @@ export const CSS_RESET = `
|
|||||||
/* border-start-start-radius: var(--radius); */
|
/* border-start-start-radius: var(--radius); */
|
||||||
padding: 0.5em;
|
padding: 0.5em;
|
||||||
padding-inline-start: 1em;
|
padding-inline-start: 1em;
|
||||||
|
display: flex;
|
||||||
|
align-items: center;
|
||||||
}
|
}
|
||||||
.markdown-callout-content {
|
.markdown-callout-content {
|
||||||
padding: 1em;
|
padding: 1em;
|
||||||
@@ -136,7 +142,12 @@ export const CSS_RESET = `
|
|||||||
border-left: 3px solid var(--current-color);
|
border-left: 3px solid var(--current-color);
|
||||||
margin-bottom: 1em;
|
margin-bottom: 1em;
|
||||||
margin-top: 1em;
|
margin-top: 1em;
|
||||||
|
|
||||||
}
|
}
|
||||||
|
.markdown-callout p:nth-child(2) {
|
||||||
|
padding:1em;
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
.markdown-callout-tip {
|
.markdown-callout-tip {
|
||||||
--text-color: whitesmoke;
|
--text-color: whitesmoke;
|
||||||
@@ -164,8 +175,9 @@ export const CSS_RESET = `
|
|||||||
flex-direction:column;
|
flex-direction:column;
|
||||||
align-items: flex-start;
|
align-items: flex-start;
|
||||||
width:95%;
|
width:95%;
|
||||||
margin-left: 20px;
|
/*margin-left: 20px;*/
|
||||||
margin-top:20px;
|
/*margin-top:20px;*/
|
||||||
|
|
||||||
/*background-color: rgba(255,0,0,0.5)!important;*/
|
/*background-color: rgba(255,0,0,0.5)!important;*/
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+379
-336
File diff suppressed because it is too large
Load Diff
+12
-3
@@ -41,7 +41,16 @@ const toastStyle = `
|
|||||||
transition-duration: ${transition_time}ms;
|
transition-duration: ${transition_time}ms;
|
||||||
`
|
`
|
||||||
|
|
||||||
function notify(message, timeout = 3000) {
|
function notify(message, timeout = 3000, old_mode = false) {
|
||||||
|
if (!old_mode) {
|
||||||
|
app.extensionManager.toast.add({
|
||||||
|
severity: 'info',
|
||||||
|
summary: 'MTB',
|
||||||
|
detail: message,
|
||||||
|
life: timeout,
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
log('Creating toast')
|
log('Creating toast')
|
||||||
const container = document.getElementById('mtb-notify-container')
|
const container = document.getElementById('mtb-notify-container')
|
||||||
const toast = document.createElement('div')
|
const toast = document.createElement('div')
|
||||||
@@ -59,7 +68,7 @@ function notify(message, timeout = 3000) {
|
|||||||
log('Transition out')
|
log('Transition out')
|
||||||
const totalHeight = Array.from(container.children).reduce(
|
const totalHeight = Array.from(container.children).reduce(
|
||||||
(acc, child) => acc + child.offsetHeight + 10, // Add spacing of 10px between toasts
|
(acc, child) => acc + child.offsetHeight + 10, // Add spacing of 10px between toasts
|
||||||
0
|
0,
|
||||||
)
|
)
|
||||||
container.style.height = `${totalHeight}px`
|
container.style.height = `${totalHeight}px`
|
||||||
|
|
||||||
@@ -83,7 +92,7 @@ function notify(message, timeout = 3000) {
|
|||||||
// Update container's height to fit new toast
|
// Update container's height to fit new toast
|
||||||
const totalHeight = Array.from(container.children).reduce(
|
const totalHeight = Array.from(container.children).reduce(
|
||||||
(acc, child) => acc + child.offsetHeight + 10, // Add spacing of 10px between toasts
|
(acc, child) => acc + child.offsetHeight + 10, // Add spacing of 10px between toasts
|
||||||
0
|
0,
|
||||||
)
|
)
|
||||||
container.style.height = `${totalHeight}px`
|
container.style.height = `${totalHeight}px`
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user