merge: 🔀 pull request #50 from melMass/dev/august-refactor

This commit is contained in:
Mel Massadian
2023-08-12 23:56:08 +02:00
committed by GitHub
30 changed files with 1266 additions and 3297 deletions
+28 -21
View File
@@ -45,19 +45,15 @@ def extract_nodes_from_source(filename):
if isinstance(target, ast.Name) and target.id == "__nodes__":
value = ast.get_source_segment(source_code, node.value)
node_value = ast.parse(value).body[0].value
if isinstance(node_value, ast.List) or isinstance(
node_value, ast.Tuple
):
for element in node_value.elts:
if isinstance(element, ast.Name):
print(element.id)
nodes.append(element.id)
if isinstance(node_value, (ast.List, ast.Tuple)):
nodes.extend(
element.id
for element in node_value.elts
if isinstance(element, ast.Name)
)
break
except SyntaxError:
log.error("Failed to parse")
pass # File couldn't be parsed
return nodes
@@ -187,6 +183,14 @@ import logging
from .endpoint import endlog
if hasattr(PromptServer, "instance"):
restore_deps = ["basicsr"]
swap_deps = ["insightface", "onnxruntime"]
node_dependency_mapping = {
"FaceSwap": swap_deps,
"LoadFaceSwapModel": swap_deps,
"LoadFaceAnalysisModel": restore_deps,
}
@PromptServer.instance.routes.get("/mtb/status")
async def get_full_library(request):
@@ -202,7 +206,13 @@ if hasattr(PromptServer, "instance"):
NODE_CLASS_MAPPINGS_DEBUG, title="Registered"
)
html_response += endpoint.render_table(
{k: "-" for k in failed}, title="Failed to load"
{
k: {"dependencies": node_dependency_mapping.get(k)}
if node_dependency_mapping.get(k)
else "-"
for k in failed
},
title="Failed to load",
)
return web.Response(
@@ -226,11 +236,10 @@ if hasattr(PromptServer, "instance"):
log.setLevel(logging.DEBUG)
log.debug("Debug mode set from API (/mtb/debug POST route)")
else:
if "MTB_DEBUG" in os.environ:
# del os.environ["MTB_DEBUG"]
os.environ.pop("MTB_DEBUG")
log.setLevel(logging.INFO)
elif "MTB_DEBUG" in os.environ:
# del os.environ["MTB_DEBUG"]
os.environ.pop("MTB_DEBUG")
log.setLevel(logging.INFO)
return web.json_response(
{"message": f"Debug mode {'set' if enabled else 'unset'}"}
@@ -244,7 +253,7 @@ if hasattr(PromptServer, "instance"):
# Check if the request prefers HTML content
if "text/html" in request.headers.get("Accept", ""):
# # Return an HTML page
html_response = f"""
html_response = """
<div class="flex-container menu">
<a href="/mtb/debug">debug</a>
<a href="/mtb/status">status</a>
@@ -263,9 +272,7 @@ if hasattr(PromptServer, "instance"):
from . import endpoint
reload(endpoint)
enabled = False
if "MTB_DEBUG" in os.environ:
enabled = True
enabled = "MTB_DEBUG" in os.environ
# Check if the request prefers HTML content
if "text/html" in request.headers.get("Accept", ""):
# # Return an HTML page
@@ -285,7 +292,7 @@ if hasattr(PromptServer, "instance"):
from . import endpoint
if "text/html" in request.headers.get("Accept", ""):
html_response = f"""
html_response = """
<h1>Actions has no get for now...</h1>
"""
return web.Response(
+77 -13
View File
@@ -1,11 +1,35 @@
from .utils import here
from .utils import here, run_command, comfy_mode
from aiohttp import web
from .log import mklog
import os
import sys
endlog = mklog("mtb endpoint")
#- ACTIONS
# - ACTIONS
import requirements
def ACTIONS_installDependency(dependency_names=None):
if dependency_names is None:
return {"error": "No dependency name provided"}
endlog.debug(f"Received Install Dependency request for {dependency_names}")
reqs = []
if comfy_mode == "embeded":
reqs = list(requirements.parse((here / "reqs_portable.txt").read_text()))
else:
reqs = list(requirements.parse((here / "reqs.txt").read_text()))
print([x.specs for x in reqs])
print(
"\n".join([f"{x.line} {''.join(x.specs[0] if x.specs else '')}" for x in reqs])
)
for dependency_name in dependency_names:
for req in reqs:
if req.name == dependency_name:
endlog.debug(f"Dependency {dependency_name} installed")
break
return {"success": True}
def ACTIONS_getStyles(style_name=None):
from .nodes.conditions import StylesLoader
@@ -19,10 +43,7 @@ def ACTIONS_getStyles(style_name=None):
if not key.startswith("__") and key not in match_list
}
if style_name:
if style_name in filtered_styles:
return filtered_styles[style_name]
else:
return {"error": "Style not found"}
return filtered_styles.get(style_name, {"error": "Style not found"})
return filtered_styles
return {"error": "No styles found"}
@@ -35,7 +56,7 @@ async def do_action(request) -> web.Response:
endlog.debug(f"Received action request: {name} {args}")
method_name = "ACTIONS_" + name
method_name = f"ACTIONS_{name}"
method = globals().get(method_name)
if callable(method):
@@ -53,16 +74,36 @@ async def do_action(request) -> web.Response:
# - HTML UTILS
def dependencies_button(name, dependencies):
deps = ",".join([f"'{x}'" for x in dependencies])
return f"""
<button class="dependency-button" onclick="window.mtb_action('installDependency',[{deps}])">Install {name} deps</button>
"""
def render_table(table_dict, sort=True, title=None):
table_rows = ""
table_dict = sorted(
table_dict.items(), key=lambda item: item[0]
) # Sort the dictionary by keys
for name, description in table_dict:
table_rows += f"<tr><td>{name}</td><td>{description}</td></tr>"
table_rows = ""
for name, item in table_dict:
if isinstance(item, dict):
if "dependencies" in item:
table_rows += f"<tr><td>{name}</td><td>"
table_rows += f"{dependencies_button(name,item['dependencies'])}"
html_response = f"""
table_rows += "</td></tr>"
else:
table_rows += f"<tr><td>{name}</td><td>{render_table(item)}</td></tr>"
# elif isinstance(item, str):
# table_rows += f"<tr><td>{name}</td><td>{item}</td></tr>"
else:
table_rows += f"<tr><td>{name}</td><td>{item}</td></tr>"
return f"""
<div class="table-container">
{"" if title is None else f"<h1>{title}</h1>"}
<table>
@@ -78,7 +119,6 @@ def render_table(table_dict, sort=True, title=None):
</table>
</div>
"""
return html_response
def render_base_template(title, content):
@@ -98,6 +138,29 @@ def render_base_template(title, content):
{css_content}
</style>
</head>
<script type="module">
import {{ api }} from '/scripts/api.js'
const mtb_action = async (action, args) =>{{
console.log(`Sending ${{action}} with args: ${{args}}`)
}}
window.mtb_action = async (action, args) =>{{
console.log(`Sending ${{action}} with args: ${{args}} to the API`)
const res = await api.fetchApi('/actions', {{
method: 'POST',
body: JSON.stringify({{
name: action,
args,
}}),
}})
const output = await res.json()
console.debug(`Received ${{action}} response:`, output)
if (output?.result?.error){{
alert(`An error occured: {{output?.result?.error}}`)
}}
return output?.result
}}
</script>
<body>
<header>
<a href="/">Back to Comfy</a>
@@ -117,5 +180,6 @@ def render_base_template(title, content):
<!-- Shared footer content here -->
</footer>
</body>
</html>
"""
+64 -65
View File
@@ -265,7 +265,7 @@
"Node name for S&R": "CLIPTextEncode"
},
"widgets_values": [
"Closeup portrait of an old bearded Caucasian man smiling, (NYC 1995), trench coat, golden ring, brown eyes"
"Medium cinematic shot of an old Caucasian man smiling, (NYC 1995), trench coat, golden ring, brown eyes, (with a blue light saber)"
]
},
{
@@ -424,65 +424,6 @@
"embedding:EasyNegative, embedding:EasyNegativeV2, watermark, text, deformed, disfigured, blurry"
]
},
{
"id": 3,
"type": "KSampler",
"pos": [
-483,
-21
],
"size": [
315,
474
],
"flags": {},
"order": 14,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "MODEL",
"link": 120
},
{
"name": "positive",
"type": "CONDITIONING",
"link": 4
},
{
"name": "negative",
"type": "CONDITIONING",
"link": 6
},
{
"name": "latent_image",
"type": "LATENT",
"link": 2
}
],
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
138
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "KSampler"
},
"widgets_values": [
542821171533322,
"fixed",
45,
8,
"dpmpp_sde",
"simple",
1
]
},
{
"id": 66,
"type": "VAEDecodeTiled",
@@ -584,7 +525,7 @@
52
],
"size": [
260.3902351585391,
260.3902282714844,
58
],
"flags": {},
@@ -615,8 +556,8 @@
40
],
"size": [
265.97600515853924,
87.31192548828142
265.97601318359375,
87.31192779541016
],
"flags": {},
"order": 11,
@@ -795,8 +736,66 @@
"Node name for S&R": "Face Swap (mtb)"
},
"widgets_values": [
"0",
false
"0"
]
},
{
"id": 3,
"type": "KSampler",
"pos": [
-483,
-21
],
"size": [
315,
474
],
"flags": {},
"order": 14,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "MODEL",
"link": 120
},
{
"name": "positive",
"type": "CONDITIONING",
"link": 4
},
{
"name": "negative",
"type": "CONDITIONING",
"link": 6
},
{
"name": "latent_image",
"type": "LATENT",
"link": 2
}
],
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
138
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "KSampler"
},
"widgets_values": [
542071534529,
"fixed",
32,
9,
"dpmpp_2m",
"normal",
1
]
}
],
+86 -132
View File
@@ -1,6 +1,6 @@
{
"last_node_id": 86,
"last_link_id": 171,
"last_link_id": 172,
"nodes": [
{
"id": 59,
@@ -111,40 +111,6 @@
"horizontal": false
}
},
{
"id": 5,
"type": "EmptyLatentImage",
"pos": [
-1410,
660
],
"size": [
315,
106
],
"flags": {},
"order": 0,
"mode": 0,
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
2,
153
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "EmptyLatentImage"
},
"widgets_values": [
768,
512,
1
]
},
{
"id": 3,
"type": "KSampler",
@@ -211,7 +177,7 @@
"Node name for S&R": "KSampler"
},
"widgets_values": [
45,
1682,
"fixed",
45,
8,
@@ -232,7 +198,7 @@
58
],
"flags": {},
"order": 1,
"order": 0,
"mode": 0,
"outputs": [
{
@@ -263,7 +229,7 @@
246
],
"flags": {},
"order": 23,
"order": 22,
"mode": 0,
"inputs": [
{
@@ -289,7 +255,7 @@
75.28300476074219
],
"flags": {},
"order": 18,
"order": 16,
"mode": 0,
"inputs": [
{
@@ -301,7 +267,7 @@
{
"name": "text",
"type": "STRING",
"link": 171,
"link": 172,
"widget": {
"name": "text",
"config": [
@@ -437,7 +403,7 @@
"flags": {
"collapsed": false
},
"order": 20,
"order": 19,
"mode": 0,
"inputs": [
{
@@ -480,7 +446,7 @@
474
],
"flags": {},
"order": 19,
"order": 18,
"mode": 0,
"inputs": [
{
@@ -536,7 +502,7 @@
"Node name for S&R": "KSampler"
},
"widgets_values": [
45,
1682,
"fixed",
45,
8,
@@ -557,7 +523,7 @@
98
],
"flags": {},
"order": 2,
"order": 1,
"mode": 0,
"outputs": [
{
@@ -648,7 +614,7 @@
46
],
"flags": {},
"order": 21,
"order": 20,
"mode": 0,
"inputs": [
{
@@ -757,10 +723,10 @@
],
"size": [
294,
104.4554216999652
104.4554214477539
],
"flags": {},
"order": 3,
"order": 2,
"mode": 0,
"outputs": [
{
@@ -847,7 +813,7 @@
78
],
"flags": {},
"order": 22,
"order": 21,
"mode": 0,
"inputs": [
{
@@ -882,34 +848,36 @@
]
},
{
"id": 85,
"type": "Save Gif (mtb)",
"id": 83,
"type": "Text box",
"pos": [
1546,
401
-2456,
236
],
"size": [
210,
336
400,
200
],
"flags": {},
"order": 24,
"order": 3,
"mode": 0,
"inputs": [
"outputs": [
{
"name": "image",
"type": "IMAGE",
"link": 169
"name": "STRING",
"type": "STRING",
"links": [
166,
167
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "Save Gif (mtb)"
"Node name for S&R": "Text box"
},
"widgets_values": [
12,
0.7,
true,
"/view?filename=eda71478da.gif&subfolder=&type=output"
"Close up photo of the face of a Caucasian young man (looking down, and frowning), rim lighting, Tokyo 1987, Bernard, over a blue sky, blue eyes shaved close"
]
},
{
@@ -951,49 +919,16 @@
"title": "seed",
"properties": {},
"widgets_values": [
45,
1682,
"fixed"
]
},
{
"id": 83,
"type": "Text box",
"pos": [
-2456,
236
],
"size": [
400,
200
],
"flags": {},
"order": 5,
"mode": 0,
"outputs": [
{
"name": "STRING",
"type": "STRING",
"links": [
166,
167
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "Text box"
},
"widgets_values": [
"Close up photo of the face of a man (looking down, and frowning), rim lighting, Tokyo 1987, Bernard, over a blue sky, blue eyes and a long beard"
]
},
{
"id": 69,
"type": "String Replace (mtb)",
"pos": [
-1442,
-1
-1529,
-3
],
"size": [
315,
@@ -1014,7 +949,7 @@
"name": "STRING",
"type": "STRING",
"links": [
170
172
],
"shape": 3,
"slot_index": 0
@@ -1024,48 +959,75 @@
"Node name for S&R": "String Replace (mtb)"
},
"widgets_values": [
"looking down",
"looking up"
"a Caucasian young",
"an African old"
]
},
{
"id": 86,
"type": "String Replace (mtb)",
"id": 5,
"type": "EmptyLatentImage",
"pos": [
-1082,
1
-1410,
660
],
"size": [
315,
82
106
],
"flags": {},
"order": 16,
"order": 5,
"mode": 0,
"inputs": [
{
"name": "string",
"type": "STRING",
"link": 170
}
],
"outputs": [
{
"name": "STRING",
"type": "STRING",
"name": "LATENT",
"type": "LATENT",
"links": [
171
2,
153
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "String Replace (mtb)"
"Node name for S&R": "EmptyLatentImage"
},
"widgets_values": [
"frowning",
"smiling"
768,
320,
1
]
},
{
"id": 85,
"type": "Save Gif (mtb)",
"pos": [
1519,
364
],
"size": [
862.6054045703117,
496.63413712402314
],
"flags": {},
"order": 23,
"mode": 0,
"inputs": [
{
"name": "image",
"type": "IMAGE",
"link": 169
}
],
"properties": {
"Node name for S&R": "Save Gif (mtb)"
},
"widgets_values": [
12,
0.7,
true,
"/view?filename=eda71478da.gif&subfolder=&type=output",
"nearest",
"/view?filename=421883da47.gif&subfolder=&type=output"
]
}
],
@@ -1327,17 +1289,9 @@
"IMAGE"
],
[
170,
172,
69,
0,
86,
0,
"STRING"
],
[
171,
86,
0,
71,
1,
"STRING"
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
+112 -99
View File
@@ -119,22 +119,13 @@ def print_formatted(text, *formats, color=None, background=None, **kwargs):
encoded_header = header.encode("utf-8", errors="replace").decode("utf-8")
encoded_text = formatted_text.encode("utf-8", errors="replace").decode("utf-8")
if sys.platform == "win32":
output_text = (
" " * len(encoded_header)
if kwargs.get("no_header")
else apply_color(apply_format(encoded_header, "bold"), color="yellow")
)
output_text += encoded_text + "\n"
sys.stdout.buffer.write(output_text.encode("utf-8"))
else:
print(
" " * len(encoded_header)
if kwargs.get("no_header")
else apply_color(apply_format(encoded_header, "bold"), color="yellow"),
encoded_text,
file=file,
)
print(
" " * len(encoded_header)
if kwargs.get("no_header")
else apply_color(apply_format(encoded_header, "bold"), color="yellow"),
encoded_text,
file=file,
)
# endregion
@@ -142,12 +133,15 @@ def print_formatted(text, *formats, color=None, background=None, **kwargs):
# region utils
def enqueue_output(out, queue):
for line in iter(out.readline, b""):
queue.put(line)
for char in iter(lambda: out.read(1), b""):
queue.put(char)
out.close()
def run_command(cmd):
def run_command(cmd, ignored_lines_start=None):
if ignored_lines_start is None:
ignored_lines_start = []
if isinstance(cmd, str):
shell_cmd = cmd
elif isinstance(cmd, list):
@@ -193,18 +187,37 @@ def run_command(cmd):
# Register the signal handler for keyboard interrupts (SIGINT)
signal.signal(signal.SIGINT, signal_handler)
stdout_buffer = ""
stderr_buffer = ""
# Process output from both streams until the process completes or interrupted
while not interrupted and (
process.poll() is None or not stdout_queue.empty() or not stderr_queue.empty()
):
with suppress(Empty):
stdout_line = stdout_queue.get_nowait()
if stdout_line.strip() != "":
print(stdout_line.strip())
stdout_char = stdout_queue.get_nowait()
stdout_buffer += stdout_char
if stdout_char == "\n":
if not any(
stdout_buffer.startswith(ign) for ign in ignored_lines_start
):
print(stdout_buffer.strip())
stdout_buffer = ""
with suppress(Empty):
stderr_line = stderr_queue.get_nowait()
if stderr_line.strip() != "":
print(stderr_line.strip())
stderr_char = stderr_queue.get_nowait()
stderr_buffer += stderr_char
if stderr_char == "\n":
print(stderr_buffer.strip())
stderr_buffer = ""
# Print any remaining content in buffers
if stdout_buffer and not any(
stdout_buffer.startswith(ign) for ign in ignored_lines_start
):
print(stdout_buffer.strip())
if stderr_buffer:
print(stderr_buffer.strip())
return_code = process.returncode
if return_code == 0 and not interrupted:
@@ -231,8 +244,6 @@ except ImportError:
print_formatted("Installing tqdm...", "italic", color="yellow")
run_command([sys.executable, "-m", "pip", "install", "--upgrade", "tqdm"])
from tqdm import tqdm
import importlib
pip_map = {
"onnxruntime-gpu": "onnxruntime",
@@ -462,6 +473,7 @@ if __name__ == "__main__":
# default=get_local_version(),
# help="Version to check against the GitHub API",
# )
print_formatted("mtb install", "bold", color="yellow")
args = parser.parse_args()
@@ -472,7 +484,7 @@ if __name__ == "__main__":
clone_dir = Path(args.path)
if not clone_dir.exists():
print_formatted(
f"The path provided does not exist on disk... It must be pointing to ComfyUI's custom_nodes directory"
"The path provided does not exist on disk... It must be pointing to ComfyUI's custom_nodes directory"
)
sys.exit()
@@ -493,48 +505,49 @@ if __name__ == "__main__":
# Install dependencies from requirements.txt
# if args.requirements or mode == "venv":
if (not args.wheels and mode not in ["colab", "embeded"]) and not full:
print_formatted(
"Skipping wheel installation. Use --wheels to install wheel dependencies. (only needed for Comfy embed)",
"italic",
color="yellow",
)
# if (not args.wheels and mode not in ["colab", "embeded"]) and not full:
# print_formatted(
# "Skipping wheel installation. Use --wheels to install wheel dependencies. (only needed for Comfy embed)",
# "italic",
# color="yellow",
# )
install_dependencies(dry=args.dry)
sys.exit()
# install_dependencies(dry=args.dry)
# sys.exit()
if mode in ["colab", "embeded"]:
print_formatted(
f"Downloading and installing release wheels since we are in a Comfy {apply_color(mode,'cyan')} environment",
"italic",
color="yellow",
)
if full:
print_formatted(
f"Downloading and installing release wheels since no arguments where provided",
"italic",
color="yellow",
)
# if mode in ["colab", "embeded"]:
# print_formatted(
# f"Downloading and installing release wheels since we are in a Comfy {apply_color(mode,'cyan')} environment",
# "italic",
# color="yellow",
# )
# if full:
# print_formatted(
# f"Downloading and installing release wheels since no arguments where provided",
# "italic",
# color="yellow",
# )
# - Check the env before proceeding.
print_formatted("Checking environment...", "italic", color="yellow")
missing_deps = []
parsed_requirements = get_requirements(here / "reqs.txt")
if parsed_requirements:
if parsed_requirements := get_requirements(here / "reqs.txt"):
for requirement in parsed_requirements:
installed, pip_name, pip_spec, import_name = try_import(requirement)
if not installed:
missing_deps.append(pip_name.split("-")[0])
if len(missing_deps) == 0:
if not missing_deps:
print_formatted(
f"All requirements are already installed. Enjoy 🚀", "italic", color="green"
"All requirements are already installed. Enjoy 🚀",
"italic",
color="green",
)
sys.exit()
# - Get the tag version from the GitHub API
tag_data, tag_name = get_github_assets(tag=None)
# # - Get the tag version from the GitHub API
# tag_data, tag_name = get_github_assets(tag=None)
# - keep
# # - keep
# version = args.version
# # Compare the local and tag versions
# if version and tag_name:
@@ -552,58 +565,58 @@ if __name__ == "__main__":
# )
# sys.exit()
matching_assets = [
asset
for asset in tag_data["assets"]
if asset["name"].endswith(".whl")
and (
"any" in asset["name"] or short_platform[current_platform] in asset["name"]
)
]
if not matching_assets:
print_formatted(
f"Unsupported operating system: {current_platform}", color="yellow"
)
wheel_order_asset = next(
(asset for asset in tag_data["assets"] if asset["name"] == "wheel_order.txt"),
None,
)
if wheel_order_asset is not None:
print_formatted(
"⚙️ Sorting the release wheels using wheels order", "italic", color="yellow"
)
response = requests.get(wheel_order_asset["browser_download_url"])
if response.status_code == 200:
wheel_order = [line.strip() for line in response.text.splitlines()]
# matching_assets = [
# asset
# for asset in tag_data["assets"]
# if asset["name"].endswith(".whl")
# and (
# "any" in asset["name"] or short_platform[current_platform] in asset["name"]
# )
# ]
# if not matching_assets:
# print_formatted(
# f"Unsupported operating system: {current_platform}", color="yellow"
# )
# wheel_order_asset = next(
# (asset for asset in tag_data["assets"] if asset["name"] == "wheel_order.txt"),
# None,
# )
# if wheel_order_asset is not None:
# print_formatted(
# "⚙️ Sorting the release wheels using wheels order", "italic", color="yellow"
# )
# response = requests.get(wheel_order_asset["browser_download_url"])
# if response.status_code == 200:
# wheel_order = [line.strip() for line in response.text.splitlines()]
def get_order_index(val):
try:
return wheel_order.index(val)
except ValueError:
return len(wheel_order)
# def get_order_index(val):
# try:
# return wheel_order.index(val)
# except ValueError:
# return len(wheel_order)
matching_assets = sorted(
matching_assets,
key=lambda x: get_order_index(x["name"].split("-")[0]),
)
else:
print("Failed to fetch wheel_order.txt. Status code:", response.status_code)
# matching_assets = sorted(
# matching_assets,
# key=lambda x: get_order_index(x["name"].split("-")[0]),
# )
# else:
# print("Failed to fetch wheel_order.txt. Status code:", response.status_code)
missing_deps_urls = []
for whl_file in matching_assets:
# check if installed
missing_deps_urls.append(whl_file["browser_download_url"])
# missing_deps_urls = []
# for whl_file in matching_assets:
# # check if installed
# missing_deps_urls.append(whl_file["browser_download_url"])
install_cmd = [sys.executable, "-m", "pip", "install"]
wheel_cmd = install_cmd + missing_deps_urls
# - Install all deps
if not args.dry:
# - first install the ordered wheel files
if platform.system() == "Windows":
wheel_cmd = install_cmd + ["-r", (here / "reqs_windows.txt")]
else:
wheel_cmd = install_cmd + ["-r", (here / "reqs.txt")]
run_command(wheel_cmd)
# - then install the rest of the deps
run_command(install_cmd + ["-r", (here / "reqs.txt")])
print_formatted(
"✅ Successfully installed all dependencies.", "italic", color="green"
)
-1
View File
@@ -3,7 +3,6 @@ import re
import os
base_log_level = logging.DEBUG if os.environ.get("MTB_DEBUG") else logging.INFO
print(f"Log level: {base_log_level}")
# Custom object that discards the output
+2 -1
View File
@@ -1,5 +1,6 @@
{
"Animation Builder (mtb)": "Convenient way to manage basic animation maths at the core of many of my workflows",
"Any To String (mtb)": "Tries to take any input and convert it to a string",
"Bbox (mtb)": "The bounding box (BBOX) custom type used by other nodes",
"Bbox From Mask (mtb)": "From a mask extract the bounding box",
"Blur (mtb)": "Blur an image using a Gaussian filter.",
@@ -9,7 +10,7 @@
"Crop (mtb)": "Crops an image and an optional mask to a given bounding box\n\n The bounding box can be given as a tuple of (x, y, width, height) or as a BBOX type\n The BBOX input takes precedence over the tuple input\n ",
"Debug (mtb)": "Experimental node to debug any Comfy values, support for more types and widgets is planned",
"Deep Bump (mtb)": "Normal & height maps generation from single pictures",
"Export To Prores (mtb)": "Export to ProRes 4444 (Experimental)",
"Export With Ffmpeg (mtb)": "Export with FFmpeg (Experimental)",
"Face Swap (mtb)": "Face swap using deepinsight/insightface models",
"Film Interpolation (mtb)": "Google Research FILM frame interpolation for large motion",
"Fit Number (mtb)": "Fit the input float using a source and target range",
+1 -1
View File
@@ -17,7 +17,7 @@ class AnimationBuilder:
},
}
RETURN_TYPES = ("INT", "FLOAT", "INT", "BOOL")
RETURN_TYPES = ("INT", "FLOAT", "INT", "BOOLEAN")
RETURN_NAMES = ("frame", "0-1 (scaled)", "count", "loop_ended")
CATEGORY = "mtb/animation"
FUNCTION = "build_animation"
+2 -109
View File
@@ -1,5 +1,4 @@
from ..utils import pil2tensor
from ..utils import here, comfy_dir
from ..utils import here
from ..log import log
import folder_paths
from pathlib import Path
@@ -97,110 +96,4 @@ class StylesLoader:
return (self.options[style_name][0], self.options[style_name][1])
class TextToImage:
"""Utils to convert text to image using a font
The tool looks for any .ttf file in the Comfy folder hierarchy.
"""
fonts = {}
def __init__(self):
# - This is executed when the graph is executed, we could conditionaly reload fonts there
pass
@classmethod
def CACHE_FONTS(cls):
font_extensions = ["*.ttf", "*.otf", "*.woff", "*.woff2", "*.eot"]
fonts = []
for extension in font_extensions:
fonts.extend(comfy_dir.glob(f"**/{extension}"))
if not fonts:
log.warn(
"> No fonts found in the comfy folder, place at least one font file somewhere in ComfyUI's hierarchy"
)
else:
log.debug(f"> Found {len(fonts)} fonts")
for font in fonts:
log.debug(f"Adding font {font}")
cls.fonts[font.stem] = font.as_posix()
@classmethod
def INPUT_TYPES(cls):
if not cls.fonts:
cls.CACHE_FONTS()
else:
log.debug(f"Using cached fonts (count: {len(cls.fonts)})")
return {
"required": {
"text": (
"STRING",
{"default": "Hello world!"},
),
"font": ((sorted(cls.fonts.keys())),),
"wrap": (
"INT",
{"default": 120, "min": 0, "max": 8096, "step": 1},
),
"font_size": (
"INT",
{"default": 12, "min": 1, "max": 2500, "step": 1},
),
"width": (
"INT",
{"default": 512, "min": 1, "max": 8096, "step": 1},
),
"height": (
"INT",
{"default": 512, "min": 1, "max": 8096, "step": 1},
),
# "position": (["INT"], {"default": 0, "min": 0, "max": 100, "step": 1}),
"color": (
"COLOR",
{"default": "black"},
),
"background": (
"COLOR",
{"default": "white"},
),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
FUNCTION = "text_to_image"
CATEGORY = "mtb/generate"
def text_to_image(
self, text, font, wrap, font_size, width, height, color, background
):
from PIL import Image, ImageDraw, ImageFont
import textwrap
font = self.fonts[font]
font = ImageFont.truetype(font, font_size)
if wrap == 0:
wrap = width / font_size
lines = textwrap.wrap(text, width=wrap)
log.debug(f"Lines: {lines}")
line_height = font.getsize("hg")[1]
img_height = height # line_height * len(lines)
img_width = width # max(font.getsize(line)[0] for line in lines)
img = Image.new("RGBA", (img_width, img_height), background)
draw = ImageDraw.Draw(img)
y_text = 0
for line in lines:
width, height = font.getsize(line)
draw.text((0, y_text), line, color, font=font)
y_text += height
# img.save(os.path.join(folder_paths.base_path, f'{str(uuid.uuid4())}.png'))
return (pil2tensor(img),)
__nodes__ = [SmartStep, TextToImage, StylesLoader]
__nodes__ = [SmartStep, StylesLoader]
+2 -2
View File
@@ -38,9 +38,9 @@ class Debug:
b64_imgs = []
for im in image:
buffered = io.BytesIO()
im.save(buffered, format="JPEG")
im.save(buffered, format="PNG")
b64_imgs.append(
"data:image/jpeg;base64,"
"data:image/png;base64,"
+ base64.b64encode(buffered.getvalue()).decode("utf-8")
)
+1 -1
View File
@@ -261,7 +261,7 @@ class DeepBump:
"LARGEST",
],
),
"normals_to_height_seamless": ("BOOL", {"default": False}),
"normals_to_height_seamless": ("BOOLEAN", {"default": False}),
},
}
+3 -3
View File
@@ -165,12 +165,12 @@ class RestoreFace:
"image": ("IMAGE",),
"model": ("FACEENHANCE_MODEL",),
# Input are aligned faces
"aligned": ("BOOL", {"default": False}),
"aligned": ("BOOLEAN", {"default": False}),
# Only restore the center face
"only_center_face": ("BOOL", {"default": False}),
"only_center_face": ("BOOLEAN", {"default": False}),
# Adjustable weights
"weight": ("FLOAT", {"default": 0.5}),
"save_tmp_steps": ("BOOL", {"default": True}),
"save_tmp_steps": ("BOOLEAN", {"default": True}),
}
}
+1 -2
View File
@@ -2,14 +2,13 @@
import onnxruntime
from pathlib import Path
from PIL import Image
from typing import List, Set, Tuple, Union, Optional
from typing import List, Set, Union, Optional
import cv2
import folder_paths
import glob
import insightface
import numpy as np
import os
import tempfile
import torch
from insightface.model_zoo.inswapper import INSwapper
from ..utils import pil2tensor, tensor2pil, download_antelopev2
+122 -2
View File
@@ -1,5 +1,7 @@
import qrcode
from ..utils import pil2tensor
from ..utils import comfy_dir
from typing import cast
from PIL import Image
from ..log import log
@@ -121,7 +123,7 @@ class QrCode:
"error_correct": (("L", "M", "Q", "H"), {"default": "L"}),
"box_size": ("INT", {"default": 10, "max": 8096, "min": 0, "step": 1}),
"border": ("INT", {"default": 4, "max": 8096, "min": 0, "step": 1}),
"invert": (("BOOL",), {"default": False}),
"invert": (("BOOLEAN",), {"default": False}),
}
}
@@ -130,6 +132,9 @@ class QrCode:
CATEGORY = "mtb/generate"
def do_qr(self, url, width, height, error_correct, box_size, border, invert):
log.warning(
"This node will soon be deprecated, there are much better alternatives like https://github.com/coreyryanhanson/comfy-qr"
)
if error_correct == "L" or error_correct not in ["M", "Q", "H"]:
error_correct = qrcode.constants.ERROR_CORRECT_L
elif error_correct == "M":
@@ -159,8 +164,123 @@ class QrCode:
return (pil2tensor(code),)
def bbox_dim(bbox):
left, upper, right, lower = bbox
width = right - left
height = lower - upper
return width, height
class TextToImage:
"""Utils to convert text to image using a font
The tool looks for any .ttf file in the Comfy folder hierarchy.
"""
fonts = {}
def __init__(self):
# - This is executed when the graph is executed, we could conditionaly reload fonts there
pass
@classmethod
def CACHE_FONTS(cls):
font_extensions = ["*.ttf", "*.otf", "*.woff", "*.woff2", "*.eot"]
fonts = []
for extension in font_extensions:
fonts.extend(comfy_dir.glob(f"**/{extension}"))
if not fonts:
log.warn(
"> No fonts found in the comfy folder, place at least one font file somewhere in ComfyUI's hierarchy"
)
else:
log.debug(f"> Found {len(fonts)} fonts")
for font in fonts:
log.debug(f"Adding font {font}")
cls.fonts[font.stem] = font.as_posix()
@classmethod
def INPUT_TYPES(cls):
if not cls.fonts:
cls.CACHE_FONTS()
else:
log.debug(f"Using cached fonts (count: {len(cls.fonts)})")
return {
"required": {
"text": (
"STRING",
{"default": "Hello world!"},
),
"font": ((sorted(cls.fonts.keys())),),
"wrap": (
"INT",
{"default": 120, "min": 0, "max": 8096, "step": 1},
),
"font_size": (
"INT",
{"default": 12, "min": 1, "max": 2500, "step": 1},
),
"width": (
"INT",
{"default": 512, "min": 1, "max": 8096, "step": 1},
),
"height": (
"INT",
{"default": 512, "min": 1, "max": 8096, "step": 1},
),
# "position": (["INT"], {"default": 0, "min": 0, "max": 100, "step": 1}),
"color": (
"COLOR",
{"default": "black"},
),
"background": (
"COLOR",
{"default": "white"},
),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
FUNCTION = "text_to_image"
CATEGORY = "mtb/generate"
def text_to_image(
self, text, font, wrap, font_size, width, height, color, background
):
from PIL import Image, ImageDraw, ImageFont
import textwrap
font = self.fonts[font]
font = cast(ImageFont.FreeTypeFont, ImageFont.truetype(font, font_size))
if wrap == 0:
wrap = width / font_size
lines = textwrap.wrap(text, width=wrap)
log.debug(f"Lines: {lines}")
line_height = bbox_dim(font.getbbox("hg"))[1]
img_height = height # line_height * len(lines)
img_width = width # max(font.getsize(line)[0] for line in lines)
img = Image.new("RGBA", (img_width, img_height), background)
draw = ImageDraw.Draw(img)
y_text = 0
# - bbox is [left, upper, right, lower]
for line in lines:
width, height = bbox_dim(font.getbbox(line))
draw.text((0, y_text), line, color, font=font)
y_text += height
# img.save(os.path.join(folder_paths.base_path, f'{str(uuid.uuid4())}.png'))
return (pil2tensor(img),)
__nodes__ = [
QrCode,
UnsplashImage
UnsplashImage,
TextToImage
# MtbExamples,
]
+165 -5
View File
@@ -1,4 +1,133 @@
from ..log import log
from PIL import Image
import urllib.request
import urllib.parse
import torch
import json
from comfy.cli_args import args
from ..utils import pil2tensor, apply_easing
import io
import numpy as np
def get_image(filename, subfolder, folder_type):
log.debug(f"Getting image {filename} from {subfolder} of {folder_type}")
data = {"filename": filename, "subfolder": subfolder, "type": folder_type}
url_values = urllib.parse.urlencode(data)
with urllib.request.urlopen(
f"http://{args.listen}:{args.port}/view?{url_values}"
) as response:
return io.BytesIO(response.read())
class GetBatchFromHistory:
"""Very experimental node to load images from the history of the server.
Queue items without output are ignored in the count."""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"enable": ("BOOLEAN", {"default": True}),
"count": ("INT", {"default": 1, "min": 0}),
"offset": ("INT", {"default": 0, "min": -1e9, "max": 1e9}),
"internal_count": ("INT", {"default": 0}),
},
"optional": {
"passthrough_image": ("IMAGE",),
},
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("images",)
CATEGORY = "mtb/animation"
FUNCTION = "load_from_history"
def load_from_history(
self,
enable=True,
count=0,
offset=0,
internal_count=0, # hacky way to invalidate the node
passthrough_image=None,
):
if not enable or count == 0:
if passthrough_image is not None:
log.debug("Using passthrough image")
return (passthrough_image,)
log.debug("Load from history is disabled for this iteration")
return (torch.zeros(0),)
frames = []
with urllib.request.urlopen(
f"http://{args.listen}:{args.port}/history"
) as response:
return self.load_batch_frames(response, offset, count, frames)
def load_batch_frames(self, response, offset, count, frames):
history = json.loads(response.read())
output_images = []
for run in history.values():
for node_output in run["outputs"].values():
if "images" in node_output:
for image in node_output["images"]:
image_data = get_image(
image["filename"], image["subfolder"], image["type"]
)
output_images.append(image_data)
if not output_images:
return (torch.zeros(0),)
# Directly get desired range of images
start_index = max(len(output_images) - offset - count, 0)
end_index = len(output_images) - offset
selected_images = output_images[start_index:end_index]
frames = [Image.open(image) for image in selected_images]
if not frames:
return (torch.zeros(0),)
elif len(frames) != count:
log.warning(f"Expected {count} images, got {len(frames)} instead")
output = pil2tensor(frames)
return (output,)
class AnyToString:
"""Tries to take any input and convert it to a string"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {"input": ("*")},
}
RETURN_TYPES = ("STRING",)
FUNCTION = "do_str"
CATEGORY = "mtb/converters"
def do_str(self, input):
if isinstance(input, str):
return (input,)
elif isinstance(input, torch.Tensor):
return (f"Tensor of shape {input.shape} and dtype {input.dtype}",)
elif isinstance(input, Image.Image):
return (f"PIL Image of size {input.size} and mode {input.mode}",)
elif isinstance(input, np.ndarray):
return (f"Numpy array of shape {input.shape} and dtype {input.dtype}",)
elif isinstance(input, dict):
return (f"Dictionary of {len(input)} items, with keys {input.keys()}",)
else:
log.debug(f"Falling back to string conversion of {input}")
return (str(input),)
class StringReplace:
@@ -38,11 +167,38 @@ class FitNumber:
return {
"required": {
"value": ("FLOAT", {"default": 0, "forceInput": True}),
"clamp": ("BOOL", {"default": False}),
"clamp": ("BOOLEAN", {"default": False}),
"source_min": ("FLOAT", {"default": 0.0}),
"source_max": ("FLOAT", {"default": 1.0}),
"target_min": ("FLOAT", {"default": 0.0}),
"target_max": ("FLOAT", {"default": 1.0}),
"easing": (
[
"Linear",
"Sine In",
"Sine Out",
"Sine In/Out",
"Quart In",
"Quart Out",
"Quart In/Out",
"Cubic In",
"Cubic Out",
"Cubic In/Out",
"Circ In",
"Circ Out",
"Circ In/Out",
"Back In",
"Back Out",
"Back In/Out",
"Elastic In",
"Elastic Out",
"Elastic In/Out",
"Bounce In",
"Bounce Out",
"Bounce In/Out",
],
{"default": "Linear"},
),
}
}
@@ -58,10 +214,14 @@ class FitNumber:
source_max: float,
target_min: float,
target_max: float,
easing: str,
):
res = target_min + (target_max - target_min) * (value - source_min) / (
source_max - source_min
)
normalized_value = (value - source_min) / (source_max - source_min)
eased_value = apply_easing(normalized_value, easing)
# - Convert the eased value to the target range
res = target_min + (target_max - target_min) * eased_value
if clamp:
if target_min > target_max:
@@ -72,4 +232,4 @@ class FitNumber:
return (res,)
__nodes__ = [StringReplace, FitNumber]
__nodes__ = [StringReplace, FitNumber, GetBatchFromHistory, AnyToString]
+2 -103
View File
@@ -6,107 +6,11 @@ import folder_paths
from ..log import log
import torch
from frame_interpolation.eval import util, interpolator
from ..utils import tensor2np
import numpy as np
import comfy
from PIL import Image
import urllib.request
import urllib.parse
import json
import comfy.utils
import tensorflow as tf
import comfy.model_management as model_management
import io
from comfy.cli_args import args
from ..utils import pil2tensor
def get_image(filename, subfolder, folder_type):
data = {"filename": filename, "subfolder": subfolder, "type": folder_type}
url_values = urllib.parse.urlencode(data)
with urllib.request.urlopen(
"http://{}:{}/view?{}".format(args.listen, args.port, url_values)
) as response:
return io.BytesIO(response.read())
class GetBatchFromHistory:
"""Very experimental node to load images from the history of the server.
Queue items without output are ignore in the count."""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"enable": ("BOOL", {"default": True}),
"count": ("INT", {"default": 1, "min": 0}),
"offset": ("INT", {"default": 0, "min": -1e9, "max": 1e9}),
"internal_count": ("INT", {"default": 0}),
},
"optional": {"passthrough_image": ("IMAGE",)},
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = "images"
CATEGORY = "mtb/animation"
FUNCTION = "load_from_history"
def load_from_history(
self,
enable=True,
count=0,
offset=0,
internal_count=0, # hacky way to invalidate the node
passthrough_image=None,
):
if not enable or count == 0:
if passthrough_image is not None:
return (passthrough_image,)
log.debug("Load from history is disabled for this iteration")
return (torch.zeros(0),)
frames = []
with urllib.request.urlopen(
"http://{}:{}/history".format(args.listen, args.port)
) as response:
history = json.loads(response.read())
output_images = []
for k, run in history.items():
for o in run["outputs"]:
for node_id in run["outputs"]:
node_output = run["outputs"][node_id]
if "images" in node_output:
images_output = []
for image in node_output["images"]:
image_data = get_image(
image["filename"], image["subfolder"], image["type"]
)
images_output.append(image_data)
output_images.extend(images_output)
if len(output_images) == 0:
return (torch.zeros(0),)
for i, image in enumerate(list(reversed(output_images))):
if i < offset:
continue
if i >= offset + count:
break
# Decode image as tensor
img = Image.open(image)
log.debug(f"Image from history {i} of shape {img.size}")
frames.append(img)
# Display the shape of the tensor
# print("Tensor shape:", image_tensor.shape)
# return (output_images,)
output = pil2tensor(
list(reversed(frames)),
)
return (output,)
class LoadFilmModel:
@@ -247,9 +151,4 @@ class ConcatImages:
return (self.concatenate_tensors(imageA, imageB),)
__nodes__ = [
LoadFilmModel,
FilmInterpolation,
ConcatImages,
GetBatchFromHistory,
]
__nodes__ = [LoadFilmModel, FilmInterpolation, ConcatImages]
+60 -92
View File
@@ -1,20 +1,16 @@
import torch
from skimage.filters import gaussian
from skimage.restoration import denoise_tv_chambolle
from skimage.util import compare_images
from skimage.color import rgb2hsv, hsv2rgb
import numpy as np
import torchvision.transforms.functional as F
from PIL import Image, ImageChops
from ..utils import tensor2pil, pil2tensor, np2tensor, tensor2np
import cv2
import torch.nn.functional as F
from PIL import Image
from ..utils import tensor2pil, pil2tensor, tensor2np
import torch
from ..log import log
import folder_paths
from PIL.PngImagePlugin import PngInfo
import json
import os
import comfy.model_management as model_management
import math
# try:
@@ -377,7 +373,7 @@ class ImagePremultiply:
"required": {
"image": ("IMAGE",),
"mask": ("MASK",),
"invert": ("BOOL", {"default": False}),
"invert": ("BOOLEAN", {"default": False}),
}
}
@@ -425,10 +421,18 @@ class ImageResizeFactor:
"FLOAT",
{"default": 2, "min": 0.01, "max": 16.0, "step": 0.01},
),
"supersample": ("BOOL", {"default": True}),
"supersample": ("BOOLEAN", {"default": True}),
"resampling": (
["lanczos", "nearest", "bilinear", "bicubic"],
{"default": "lanczos"},
[
"nearest",
"linear",
"bilinear",
"bicubic",
"trilinear",
"area",
"nearest-exact",
],
{"default": "nearest"},
),
},
"optional": {
@@ -440,71 +444,6 @@ class ImageResizeFactor:
RETURN_TYPES = ("IMAGE", "MASK")
FUNCTION = "resize"
def resize_image(
self,
image: torch.Tensor,
factor: float = 0.5,
supersample=False,
resample="lanczos",
mask=None,
) -> torch.Tensor:
model_management.throw_exception_if_processing_interrupted()
batch_count = 1
img = tensor2pil(image)
if isinstance(img, list):
log.debug("Multiple images detected (list)")
out = []
for im in img:
im = self.resize_image(
pil2tensor(im), factor, supersample, resample, mask
)
out.append(im)
return torch.cat(out, dim=0)
elif isinstance(img, torch.Tensor):
if len(image.shape) > 3:
batch_count = image.size(0)
if batch_count > 1:
log.debug("Multiple images detected (batch count)")
out = [
self.resize_image(image[i], factor, supersample, resample, mask)
for i in range(batch_count)
]
return torch.cat(out, dim=0)
log.debug("Resizing image")
# Get the current width and height of the image
current_width, current_height = img.size
log.debug(f"Current width: {current_width}, Current height: {current_height}")
# Calculate the new width and height based on the given mode and parameters
new_width, new_height = int(factor * current_width), int(
factor * current_height
)
log.debug(f"New width: {new_width}, New height: {new_height}")
# Define a dictionary of resampling filters
resample_filters = {"nearest": 0, "bilinear": 2, "bicubic": 3, "lanczos": 1}
# Apply supersample
if supersample:
super_size = (new_width * 8, new_height * 8)
log.debug(f"Applying supersample: {super_size}")
img = img.resize(
super_size, resample=Image.Resampling(resample_filters[resample])
)
# Resize the image using the given resampling filter
resized_image = img.resize(
(new_width, new_height),
resample=Image.Resampling(resample_filters[resample]),
)
return pil2tensor(resized_image)
def resize(
self,
image: torch.Tensor,
@@ -513,24 +452,53 @@ class ImageResizeFactor:
resampling: str,
mask=None,
):
log.debug(f"Resizing image with factor {factor} and resampling {resampling}")
# Check if the tensor has the correct dimension
if len(image.shape) not in [3, 4]: # HxWxC or BxHxWxC
raise ValueError("Expected image tensor of shape (H, W, C) or (B, H, W, C)")
batch_count = image.size(0)
log.debug(f"Batch count: {batch_count}")
if batch_count == 1:
log.debug("Batch count is 1, returning single image")
return (self.resize_image(image, factor, supersample, resampling),)
# Transpose to CxHxW or BxCxHxW for PyTorch
if len(image.shape) == 3:
image = image.permute(2, 0, 1).unsqueeze(0) # CxHxW
else:
log.debug("Batch count is greater than 1, returning multiple images")
images = [
self.resize_image(image[i], factor, supersample, resampling)
for i in range(batch_count)
]
images = torch.cat(images, dim=0)
return (images,)
image = image.permute(0, 3, 1, 2) # BxCxHxW
# Compute new dimensions
B, C, H, W = image.shape
new_H, new_W = int(H * factor), int(W * factor)
import math
align_corner_filters = ("linear", "bilinear", "bicubic", "trilinear")
# Resize the image
resized_image = F.interpolate(
image,
size=(new_H, new_W),
mode=resampling,
align_corners=resampling in align_corner_filters,
)
# Optionally supersample
if supersample:
resized_image = F.interpolate(
resized_image,
scale_factor=2,
mode=resampling,
align_corners=resampling in align_corner_filters,
)
# Transpose back to the original format: BxHxWxC or HxWxC
if len(image.shape) == 4:
resized_image = resized_image.permute(0, 2, 3, 1)
else:
resized_image = resized_image.squeeze(0).permute(1, 2, 0)
# Apply mask if provided
if mask is not None:
if len(mask.shape) != len(resized_image.shape):
raise ValueError(
"Mask tensor should have the same dimensions as the image tensor"
)
resized_image = resized_image * mask
return (resized_image,)
class SaveImageGrid:
@@ -546,7 +514,7 @@ class SaveImageGrid:
"required": {
"images": ("IMAGE",),
"filename_prefix": ("STRING", {"default": "ComfyUI"}),
"save_intermediate": ("BOOL", {"default": False}),
"save_intermediate": ("BOOLEAN", {"default": False}),
},
"hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"},
}
+83 -56
View File
@@ -1,4 +1,4 @@
from ..utils import tensor2np
from ..utils import tensor2np, PIL_FILTER_MAP
import uuid
import folder_paths
from ..log import log
@@ -7,10 +7,12 @@ import subprocess
import torch
from pathlib import Path
import numpy as np
from PIL import Image
from typing import Optional, List
class ExportToProres:
"""Export to ProRes 4444 (Experimental)"""
class ExportWithFfmpeg:
"""Export with FFmpeg (Experimental)"""
@classmethod
def INPUT_TYPES(cls):
@@ -20,6 +22,11 @@ class ExportToProres:
# "frames": ("FRAMES",),
"fps": ("FLOAT", {"default": 24, "min": 1}),
"prefix": ("STRING", {"default": "export"}),
"format": (["mov", "mp4", "mkv", "avi"], {"default": "mov"}),
"codec": (
["prores_ks", "libx264", "libx265"],
{"default": "prores_ks"},
),
}
}
@@ -33,13 +40,17 @@ class ExportToProres:
images: torch.Tensor,
fps: float,
prefix: str,
format: str,
codec: str,
):
if images.size(0) == 0:
return ("",)
output_dir = Path(folder_paths.get_output_directory())
id = f"{prefix}_{uuid.uuid4()}.mov"
pix_fmt = "rgb48le" if codec == "prores_ks" else "yuv420p"
file_ext = format
file_id = f"{prefix}_{uuid.uuid4()}.{file_ext}"
log.debug(f"Exporting to {output_dir / id}")
log.debug(f"Exporting to {output_dir / file_id}")
frames = tensor2np(images)
log.debug(f"Frames type {type(frames[0])}")
@@ -49,7 +60,7 @@ class ExportToProres:
height, width, _ = frames[0].shape
out_path = (output_dir / id).as_posix()
out_path = (output_dir / file_id).as_posix()
# Prepare the FFmpeg command
command = [
@@ -62,17 +73,13 @@ class ExportToProres:
"-s",
f"{width}x{height}",
"-pix_fmt",
"rgb48le",
pix_fmt,
"-r",
str(fps),
"-i",
"-",
"-c:v",
"prores_ks",
"-profile:v",
"4",
"-pix_fmt",
"yuva444p10le",
codec,
"-r",
str(fps),
"-y",
@@ -91,6 +98,37 @@ class ExportToProres:
return (out_path,)
def prepare_animated_batch(
batch: torch.Tensor,
pingpong=False,
resize_by=1.0,
resample_filter: Optional[Image.Resampling] = None,
image_type=np.uint8,
) -> List[Image.Image]:
images = tensor2np(batch)
images = [frame.astype(image_type) for frame in images]
height, width, _ = batch[0].shape
if pingpong:
reversed_frames = images[::-1]
images.extend(reversed_frames)
pil_images = [Image.fromarray(frame) for frame in images]
# Resize frames if necessary
if abs(resize_by - 1.0) > 1e-6:
new_width = int(width * resize_by)
new_height = int(height * resize_by)
pil_images_resized = [
frame.resize((new_width, new_height), resample=resample_filter)
for frame in pil_images
]
pil_images = pil_images_resized
return pil_images
# todo: deprecate for apng
class SaveGif:
"""Save the images from the batch as a GIF"""
@@ -101,8 +139,12 @@ class SaveGif:
"image": ("IMAGE",),
"fps": ("INT", {"default": 12, "min": 1, "max": 120}),
"resize_by": ("FLOAT", {"default": 1.0, "min": 0.1}),
"pingpong": ("BOOL", {"default": False}),
}
"optimize": ("BOOLEAN", {"default": False}),
"pingpong": ("BOOLEAN", {"default": False}),
},
"optional": {
"resample_filter": (list(PIL_FILTER_MAP.keys()),),
},
}
RETURN_TYPES = ()
@@ -110,59 +152,44 @@ class SaveGif:
CATEGORY = "mtb/IO"
FUNCTION = "save_gif"
def save_gif(self, image, fps=12, resize_by=1.0, pingpong=False):
def save_gif(
self,
image,
fps=12,
resize_by=1.0,
optimize=False,
pingpong=False,
resample_filter=None,
):
if image.size(0) == 0:
return ("",)
images = tensor2np(image)
images = [frame.astype(np.uint8) for frame in images]
if pingpong:
reversed_frames = images[::-1]
images.extend(reversed_frames)
if resample_filter is not None:
resample_filter = PIL_FILTER_MAP.get(resample_filter)
height, width, _ = image[0].shape
pil_images = prepare_animated_batch(
image,
pingpong,
resize_by,
resample_filter,
)
ruuid = uuid.uuid4()
ruuid = ruuid.hex[:10]
out_path = f"{folder_paths.output_directory}/{ruuid}.gif"
log.debug(f"Saving a gif file {width}x{height} as {ruuid}.gif")
# Prepare the FFmpeg command
command = [
"ffmpeg",
"-y",
"-f",
"rawvideo",
"-vcodec",
"rawvideo",
"-s",
f"{width}x{height}",
"-pix_fmt",
"rgb24", # GIF only supports rgb24
"-r",
str(fps),
"-i",
"-",
"-vf",
f"fps={fps},scale={width * resize_by}:-1", # Set frame rate and resize if necessary
"-y",
# Create the GIF from PIL images
pil_images[0].save(
out_path,
]
save_all=True,
append_images=pil_images[1:],
optimize=optimize,
duration=int(1000 / fps),
loop=0,
)
process = subprocess.Popen(command, stdin=subprocess.PIPE)
for frame in images:
model_management.throw_exception_if_processing_interrupted()
process.stdin.write(frame.tobytes())
process.stdin.close()
process.wait()
results = []
results.append({"filename": f"{ruuid}.gif", "subfolder": "", "type": "output"})
results = [{"filename": f"{ruuid}.gif", "subfolder": "", "type": "output"}]
return {"ui": {"gif": results}}
__nodes__ = [SaveGif, ExportToProres]
__nodes__ = [SaveGif, ExportWithFfmpeg]
+3 -3
View File
@@ -13,7 +13,7 @@ class ImageRemoveBackgroundRembg:
"required": {
"image": ("IMAGE",),
"alpha_matting": (
"BOOL",
"BOOLEAN",
{"default": False},
),
"alpha_matting_foreground_threshold": (
@@ -29,12 +29,12 @@ class ImageRemoveBackgroundRembg:
{"default": 10, "min": 0, "max": 255},
),
"post_process_mask": (
"BOOL",
"BOOLEAN",
{"default": False},
),
"bgcolor": (
"COLOR",
{"default": "black"},
{"default": "#000000"},
),
},
}
+1 -1
View File
@@ -14,7 +14,7 @@ class IntToBool:
}
}
RETURN_TYPES = ("BOOL",)
RETURN_TYPES = ("BOOLEAN",)
FUNCTION = "int_to_bool"
CATEGORY = "mtb/number"
+71 -15
View File
@@ -1,5 +1,9 @@
import torch
import torchvision.transforms.functional as F
import torchvision.transforms.functional as TF
from ..utils import log, hex_to_rgb, tensor2pil, pil2tensor
from math import sqrt, ceil
from typing import cast
from PIL import Image
class TransformImage:
@@ -14,11 +18,19 @@ class TransformImage:
return {
"required": {
"image": ("IMAGE",),
"x": ("FLOAT", {"default": 0}),
"y": ("FLOAT", {"default": 0}),
"zoom": ("FLOAT", {"default": 1.0, "min": 0.001}),
"angle": ("FLOAT", {"default": 0}),
"shear": ("FLOAT", {"default": 0}),
"x": ("FLOAT", {"default": 0, "step": 1, "min": -4096, "max": 4096}),
"y": ("FLOAT", {"default": 0, "step": 1, "min": -4096, "max": 4096}),
"zoom": ("FLOAT", {"default": 1.0, "min": 0.001, "step": 0.01}),
"angle": ("FLOAT", {"default": 0, "step": 1, "min": -360, "max": 360}),
"shear": (
"FLOAT",
{"default": 0, "step": 1, "min": -4096, "max": 4096},
),
"border_handling": (
["edge", "constant", "reflect", "symmetric"],
{"default": "edge"},
),
"constant_color": ("COLOR", {"default": "#000000"}),
},
}
@@ -32,23 +44,67 @@ class TransformImage:
x: float,
y: float,
zoom: float,
angle: int,
shear,
angle: float,
shear: float,
border_handling="edge",
constant_color=None,
):
x = int(x)
y = int(y)
angle = int(angle)
log.debug(f"Zoom: {zoom} | x: {x}, y: {y}, angle: {angle}, shear: {shear}")
if image.size(0) == 0:
return (torch.zeros(0),)
transformed_images = []
for img in image:
img = img.transpose(0, 2)
frames_count, frame_height, frame_width, frame_channel_count = image.size()
transformed_image = F.affine(
img, angle=angle, scale=zoom, translate=[int(y), int(x)], shear=shear
new_height, new_width = int(frame_height * zoom), int(frame_width * zoom)
log.debug(f"New height: {new_height}, New width: {new_width}")
# - Calculate diagonal of the original image
diagonal = sqrt(frame_width**2 + frame_height**2)
max_padding = ceil(diagonal * zoom - min(frame_width, frame_height))
# Calculate padding for zoom
pw = int(frame_width - new_width)
ph = int(frame_height - new_height)
pw += abs(max_padding)
ph += abs(max_padding)
padding = [max(0, pw + x), max(0, ph + y), max(0, pw - x), max(0, ph - y)]
constant_color = hex_to_rgb(constant_color)
log.debug(f"Fill Tuple: {constant_color}")
for img in tensor2pil(image):
img = TF.pad(
img, # transformed_frame,
padding=padding,
padding_mode=border_handling,
fill=constant_color or 0,
)
transformed_image = transformed_image.transpose(2, 0)
transformed_images.append(transformed_image.unsqueeze(0))
img = cast(
Image.Image,
TF.affine(img, angle=angle, scale=zoom, translate=[x, y], shear=shear),
)
return (torch.cat(transformed_images, dim=0),)
left = abs(padding[0])
upper = abs(padding[1])
right = img.width - abs(padding[2])
bottom = img.height - abs(padding[3])
# log.debug("crop is [:,top:bottom, left:right] for tensors")
log.debug("crop is [left, top, right, bottom] for PIL")
log.debug(f"crop is {left}, {upper}, {right}, {bottom}")
img = img.crop((left, upper, right, bottom))
transformed_images.append(img)
return (pil2tensor(transformed_images),)
__nodes__ = [TransformImage]
+1 -5
View File
@@ -1,11 +1,7 @@
onnxruntime-gpu==1.15.1
qrcode[pil]
rembg==2.0.50
# on windows non WSL 2.10 is the last version with GPU support
tensorflow==2.10.1; platform_system == "Windows"
tb-nightly==2.12.0a20230126; platform_system == "Windows"
tensorflow; platform_system != "Windows"
# the old tf version on windows comes with a breaking protobuf version
tensorflow
facexlib==0.3.0
insightface==0.7.3
basicsr==1.4.2
+18
View File
@@ -0,0 +1,18 @@
https://github.com/melMass/comfy_mtb/releases/download/v0.1.3/pycocotools-2.0.6-cp310-cp310-win_amd64.whl
https://github.com/melMass/comfy_mtb/releases/download/v0.1.3/future-0.18.3-py3-none-any.whl
https://github.com/melMass/comfy_mtb/releases/download/v0.1.3/filterpy-1.4.5-py3-none-any.whl
https://github.com/melMass/comfy_mtb/releases/download/v0.1.3/easydict-1.10-py3-none-any.whl
https://github.com/melMass/comfy_mtb/releases/download/v0.1.3/gdown-4.7.1-py3-none-any.whl
https://github.com/melMass/comfy_mtb/releases/download/v0.1.3/basicsr-1.4.2-py3-none-any.whl
https://github.com/melMass/comfy_mtb/releases/download/v0.1.3/mmcv-2.0.0-py2.py3-none-any.whl
https://github.com/melMass/comfy_mtb/releases/download/v0.1.3/insightface-0.7.3-cp310-cp310-win_amd64.whl
onnxruntime-gpu==1.15.1
qrcode[pil]
rembg==2.0.50
# on windows non WSL 2.10 is the last version with GPU support
tensorflow==2.10.1;
tb-nightly==2.12.0a20230126; platform_system == "Windows"
facexlib==0.3.0
# the old tf version on windows comes with a breaking protobuf version
protobuf==3.19.6
+320 -14
View File
@@ -4,10 +4,38 @@ import torch
from pathlib import Path
import sys
from typing import List
from .log import log
import signal
from contextlib import suppress
from queue import Queue, Empty
import subprocess
import threading
import os
import math
try:
from .log import log
except ImportError:
try:
from log import log
log.warn("Imported log without relative path")
except ImportError:
import logging
log = logging.getLogger("comfy mtb utils")
log.warn("[comfy mtb] You probably called the file outside a module.")
# region MISC Utilities
def hex_to_rgb(hex_color):
try:
hex_color = hex_color.lstrip("#")
return tuple(int(hex_color[i : i + 2], 16) for i in (0, 2, 4))
except ValueError:
log.error(f"Invalid hex color: {hex_color}")
return (0, 0, 0)
def add_path(path, prepend=False):
if isinstance(path, list):
for p in path:
@@ -24,6 +52,79 @@ def add_path(path, prepend=False):
sys.path.append(path)
def enqueue_output(out, queue):
for line in iter(out.readline, b""):
queue.put(line)
out.close()
def run_command(cmd):
if isinstance(cmd, str):
shell_cmd = cmd
elif isinstance(cmd, list):
shell_cmd = ""
for arg in cmd:
if isinstance(arg, Path):
arg = arg.as_posix()
shell_cmd += f"{arg} "
else:
raise ValueError(
"Invalid 'cmd' argument. It must be a string or a list of arguments."
)
process = subprocess.Popen(
shell_cmd,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
universal_newlines=True,
shell=True,
)
# Create separate threads to read standard output and standard error streams
stdout_queue = Queue()
stderr_queue = Queue()
stdout_thread = threading.Thread(
target=enqueue_output, args=(process.stdout, stdout_queue)
)
stderr_thread = threading.Thread(
target=enqueue_output, args=(process.stderr, stderr_queue)
)
stdout_thread.daemon = True
stderr_thread.daemon = True
stdout_thread.start()
stderr_thread.start()
interrupted = False
def signal_handler(signum, frame):
nonlocal interrupted
interrupted = True
print("Command execution interrupted.")
# Register the signal handler for keyboard interrupts (SIGINT)
signal.signal(signal.SIGINT, signal_handler)
# Process output from both streams until the process completes or interrupted
while not interrupted and (
process.poll() is None or not stdout_queue.empty() or not stderr_queue.empty()
):
with suppress(Empty):
stdout_line = stdout_queue.get_nowait()
if stdout_line.strip() != "":
print(stdout_line.strip())
with suppress(Empty):
stderr_line = stderr_queue.get_nowait()
if stderr_line.strip() != "":
print(stderr_line.strip())
return_code = process.returncode
if return_code == 0 and not interrupted:
print("Command executed successfully!")
else:
if not interrupted:
print(f"Command failed with return code: {return_code}")
# todo use the requirements library
reqs_map = {
"onnxruntime": "onnxruntime-gpu==1.15.1",
@@ -50,36 +151,51 @@ def import_install(package_name):
# endregion
# region GLOBAL VARIABLES
# Get the absolute path of the parent directory of the current script
# - detect mode
comfy_mode = None
if os.environ.get("COLAB_GPU"):
comfy_mode = "colab"
elif "python_embeded" in sys.executable:
comfy_mode = "embeded"
elif ".venv" in sys.executable:
comfy_mode = "venv"
# - Get the absolute path of the parent directory of the current script
here = Path(__file__).parent.resolve()
# Construct the absolute path to the ComfyUI directory
# - Construct the absolute path to the ComfyUI directory
comfy_dir = here.parent.parent
# Construct the path to the font file
# - Construct the path to the font file
font_path = here / "font.ttf"
# Add extern folder to path
# - Add extern folder to path
extern_root = here / "extern"
add_path(extern_root)
for pth in extern_root.iterdir():
if pth.is_dir():
add_path(pth)
# Add the ComfyUI directory and custom nodes path to the sys.path list
# - Add the ComfyUI directory and custom nodes path to the sys.path list
add_path(comfy_dir)
add_path((comfy_dir / "custom_nodes"))
PIL_FILTER_MAP = {
"nearest": Image.Resampling.NEAREST,
"box": Image.Resampling.BOX,
"bilinear": Image.Resampling.BILINEAR,
"hamming": Image.Resampling.HAMMING,
"bicubic": Image.Resampling.BICUBIC,
"lanczos": Image.Resampling.LANCZOS,
}
# endregion
# region TENSOR UTILITIES
def tensor2pil(image: torch.Tensor) -> List[Image.Image]:
batch_count = 1
if len(image.shape) > 3:
batch_count = image.size(0)
batch_count = image.size(0) if len(image.shape) > 3 else 1
if batch_count > 1:
out = []
for i in range(batch_count):
@@ -108,9 +224,7 @@ def np2tensor(img_np: np.ndarray | List[np.ndarray]) -> torch.Tensor:
def tensor2np(tensor: torch.Tensor) -> List[np.ndarray]:
batch_count = 1
if len(tensor.shape) > 3:
batch_count = tensor.size(0)
batch_count = tensor.size(0) if len(tensor.shape) > 3 else 1
if batch_count > 1:
out = []
for i in range(batch_count):
@@ -122,6 +236,7 @@ def tensor2np(tensor: torch.Tensor) -> List[np.ndarray]:
# endregion
# region MODEL Utilities
def download_antelopev2():
antelopev2_url = "https://drive.google.com/uc?id=18wEUfMNohBJ4K3Ly5wpTejPfDzp-8fI8"
@@ -161,3 +276,194 @@ def download_antelopev2():
# endregion
# region UV Utilities
def create_uv_map_tensor(width=512, height=512):
u = torch.linspace(0.0, 1.0, steps=width)
v = torch.linspace(0.0, 1.0, steps=height)
U, V = torch.meshgrid(u, v)
uv_map = torch.zeros(height, width, 3, dtype=torch.float32)
uv_map[:, :, 0] = U.t()
uv_map[:, :, 1] = V.t()
return uv_map.unsqueeze(0)
# endregion
# region ANIMATION Utilities
def apply_easing(value, easing_type):
if value < 0 or value > 1:
raise ValueError("The value should be between 0 and 1.")
if easing_type == "Linear":
return value
# Back easing functions
def easeInBack(t):
s = 1.70158
return t * t * ((s + 1) * t - s)
def easeOutBack(t):
s = 1.70158
return ((t - 1) * t * ((s + 1) * t + s)) + 1
def easeInOutBack(t):
s = 1.70158 * 1.525
if t < 0.5:
return (t * t * (t * (s + 1) - s)) * 2
return ((t - 2) * t * ((s + 1) * t + s) + 2) * 2
# Elastic easing functions
def easeInElastic(t):
if t == 0:
return 0
if t == 1:
return 1
p = 0.3
s = p / 4
return -(math.pow(2, 10 * (t - 1)) * math.sin((t - 1 - s) * (2 * math.pi) / p))
def easeOutElastic(t):
if t == 0:
return 0
if t == 1:
return 1
p = 0.3
s = p / 4
return math.pow(2, -10 * t) * math.sin((t - s) * (2 * math.pi) / p) + 1
def easeInOutElastic(t):
if t == 0:
return 0
if t == 1:
return 1
p = 0.3 * 1.5
s = p / 4
t = t * 2
if t < 1:
return -0.5 * (
math.pow(2, 10 * (t - 1)) * math.sin((t - 1 - s) * (2 * math.pi) / p)
)
return (
0.5 * math.pow(2, -10 * (t - 1)) * math.sin((t - 1 - s) * (2 * math.pi) / p)
+ 1
)
# Bounce easing functions
def easeInBounce(t):
return 1 - easeOutBounce(1 - t)
def easeOutBounce(t):
if t < (1 / 2.75):
return 7.5625 * t * t
elif t < (2 / 2.75):
t -= 1.5 / 2.75
return 7.5625 * t * t + 0.75
elif t < (2.5 / 2.75):
t -= 2.25 / 2.75
return 7.5625 * t * t + 0.9375
else:
t -= 2.625 / 2.75
return 7.5625 * t * t + 0.984375
def easeInOutBounce(t):
if t < 0.5:
return easeInBounce(t * 2) * 0.5
return easeOutBounce(t * 2 - 1) * 0.5 + 0.5
# Quart easing functions
def easeInQuart(t):
return t * t * t * t
def easeOutQuart(t):
t -= 1
return -(t**2 * t * t - 1)
def easeInOutQuart(t):
t *= 2
if t < 1:
return 0.5 * t * t * t * t
t -= 2
return -0.5 * (t**2 * t * t - 2)
# Cubic easing functions
def easeInCubic(t):
return t * t * t
def easeOutCubic(t):
t -= 1
return t**2 * t + 1
def easeInOutCubic(t):
t *= 2
if t < 1:
return 0.5 * t * t * t
t -= 2
return 0.5 * (t**2 * t + 2)
# Circ easing functions
def easeInCirc(t):
return -(math.sqrt(1 - t * t) - 1)
def easeOutCirc(t):
t -= 1
return math.sqrt(1 - t**2)
def easeInOutCirc(t):
t *= 2
if t < 1:
return -0.5 * (math.sqrt(1 - t**2) - 1)
t -= 2
return 0.5 * (math.sqrt(1 - t**2) + 1)
# Sine easing functions
def easeInSine(t):
return -math.cos(t * (math.pi / 2)) + 1
def easeOutSine(t):
return math.sin(t * (math.pi / 2))
def easeInOutSine(t):
return -0.5 * (math.cos(math.pi * t) - 1)
easing_functions = {
"Sine In": easeInSine,
"Sine Out": easeOutSine,
"Sine In/Out": easeInOutSine,
"Quart In": easeInQuart,
"Quart Out": easeOutQuart,
"Quart In/Out": easeInOutQuart,
"Cubic In": easeInCubic,
"Cubic Out": easeOutCubic,
"Cubic In/Out": easeInOutCubic,
"Circ In": easeInCirc,
"Circ Out": easeOutCirc,
"Circ In/Out": easeInOutCirc,
"Back In": easeInBack,
"Back Out": easeOutBack,
"Back In/Out": easeInOutBack,
"Elastic In": easeInElastic,
"Elastic Out": easeOutElastic,
"Elastic In/Out": easeInOutElastic,
"Bounce In": easeInBounce,
"Bounce Out": easeOutBounce,
"Bounce In/Out": easeInOutBounce,
}
function_ease = easing_functions.get(easing_type)
if function_ease:
return function_ease(value)
log.error(f"Unknown easing type: {easing_type}")
log.error(f"Available easing types: {list(easing_functions.keys())}")
raise ValueError(f"Unknown easing type: {easing_type}")
# endregion
+2 -2
View File
@@ -49,7 +49,7 @@ export function offsetDOMWidget(
position: 'absolute',
background: !node.color ? '' : node.color,
color: !node.color ? '' : 'white',
zIndex: app.graph._nodes.indexOf(node),
zIndex: 5, //app.graph._nodes.indexOf(node),
})
}
@@ -60,7 +60,7 @@ export function offsetDOMWidget(
*/
export function getWidgetType(config) {
// Special handling for COMBO so we restrict links based on the entries
let type = config[0]
let type = config?.[0]
let linkType = type
if (type instanceof Array) {
type = 'COMBO'
+2 -2
View File
@@ -31,7 +31,7 @@ const styles = {
background: 'none',
border: 'none',
color: '#fff',
zIndex: 9999999,
zIndex: 1000,
fontSize: '30px',
cursor: 'pointer',
pointerEvents: 'auto',
@@ -43,7 +43,7 @@ const styles = {
width: '100vw',
position: 'absolute',
bottom: 0,
zIndex: 9999999,
zIndex: 10,
background: '#333',
overflow: 'auto',
},
+35 -110
View File
@@ -13,7 +13,7 @@ import * as shared from '/extensions/mtb/comfy_shared.js'
import { log } from '/extensions/mtb/comfy_shared.js'
import { api } from '/scripts/api.js'
const newTypes = ['BOOL', 'COLOR', 'BBOX']
const newTypes = [, /*'BOOL'*/ 'COLOR', 'BBOX']
export const MtbWidgets = {
BBOX: (key, val) => {
@@ -204,116 +204,7 @@ export const MtbWidgets = {
widget.desc = 'Represents a Bounding Box with x, y, width, and height.'
return widget
},
BOOL: (key, val, compute = false) => {
/** @type {import("/types/litegraph").IWidget} */
const widget = {
name: key,
type: 'BOOL',
options: { default: false },
y: 0,
draw: function (ctx, node, widget_width, widgetY, height) {
const hide = this.type !== 'BOOL' && app.canvas.ds.scale > 0.5
if (hide) {
return
}
const outline_color = LiteGraph.WIDGET_OUTLINE_COLOR
const background_color = LiteGraph.WIDGET_BGCOLOR
const text_color = LiteGraph.WIDGET_TEXT_COLOR
const H = LiteGraph.NODE_WIDGET_HEIGHT
// const arrowSize = 8
let margin = 15
if (hide) return
let currentY = widgetY
ctx.textAlign = 'left'
ctx.strokeStyle = outline_color
ctx.fillStyle = background_color
ctx.beginPath()
// ctx.roundRect(margin, currentY, widget_width - margin * 2, H, [H * 0.5]);
ctx.rect(margin, currentY, H, H) // Draw checkbox square
ctx.fill()
ctx.stroke()
ctx.fillStyle = text_color
// ctx.fillText(this.label || this.name, margin * 2 + 5, currentY + H * 0.7);
ctx.fillText(
this.label || this.name,
H + margin * 2,
currentY + H * 0.7
)
// Draw arrow if the value is true
// Draw checkmark if the value is true
if (this.value) {
ctx.fillStyle = text_color
ctx.beginPath()
ctx.moveTo(margin + H * 0.15, currentY + H * 0.5)
ctx.lineTo(margin + H * 0.4, currentY + H * 0.8)
ctx.lineTo(margin + H * 0.85, currentY + H * 0.2)
ctx.stroke()
}
},
get value() {
return this.inputEl.value === 'true'
},
set value(x) {
this.inputEl.value = x
},
computeSize: function (width) {
return [width, 32]
},
mouse: function (event, pos, node) {
// let x = pos[0] - node.pos[0];
// let y = pos[1] - node.pos[1];
// let width = node.size[0];
// let H = LiteGraph.NODE_WIDGET_HEIGHT;
// let margin = 15;
// if (event.type == LiteGraph.pointerevents_method + "down") {
// if (x > margin && x < widget_width - margin && y > widgetY && y < widgetY + H) {
// this.value = !this.value; // Toggle checkbox value
// shared.inner_value_change(this, this.value, event);
// app.canvas.setDirty(true);
// }
// }
if (event.type === 'pointerdown') {
// get widgets of type type : "COLOR"
const widgets = node.widgets.filter((w) => w.type === 'BOOL')
for (const w of widgets) {
// color picker
const rect = [w.last_y, w.last_y + 32]
if (pos[1] > rect[0] && pos[1] < rect[1]) {
// picker.style.position = "absolute";
// picker.style.left = ( pos[0]) + "px";
// picker.style.top = ( pos[1]) + "px";
// place at screen center
// picker.style.position = "absolute";
// picker.style.left = (window.innerWidth / 2) + "px";
// picker.style.top = (window.innerHeight / 2) + "px";
// picker.style.transform = "translate(-50%, -50%)";
// picker.style.zIndex = 1000;
this.value = this.value ? false : true
}
}
}
},
}
// create a checkbox
widget.inputEl = document.createElement('input')
widget.inputEl.type = 'checkbox'
widget.value = val || false
document.body.appendChild(widget.inputEl)
return widget
},
COLOR: (key, val, compute = false) => {
/** @type {import("/types/litegraph").IWidget} */
const widget = {}
@@ -890,6 +781,40 @@ const mtb_widgets = {
}
break
}
case 'Interpolate Clip Sequential (mtb)': {
const onNodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
const r = onNodeCreated
? onNodeCreated.apply(this, arguments)
: undefined
const addReplacement = () => {
const input = this.addInput(
`replacement_${this.widgets.length}`,
'STRING',
''
)
console.log(input)
this.addWidget('STRING', `replacement_${this.widgets.length}`, '')
}
//- add
this.addWidget('button', '+', 'add', function (value, widget, node) {
console.log('Button clicked', value, widget, node)
addReplacement()
})
//- remove
this.addWidget(
'button',
'-',
'remove',
function (value, widget, node) {
console.log(`Button clicked: ${value}`, widget, node)
}
)
return r
}
break
}
case 'Styles Loader (mtb)': {
const origGetExtraMenuOptions = nodeType.prototype.getExtraMenuOptions
nodeType.prototype.getExtraMenuOptions = function (_, options) {