commit abb98953a39592cf4b38f4d245b84e9696c904b7 Author: numz Date: Sun Oct 13 17:17:01 2024 +0200 Initial commit for Comfyui-FlowChain diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..14e19a0 --- /dev/null +++ b/.gitignore @@ -0,0 +1,3 @@ +docs/assets/demo.gif +.git/* +**/__pycache__/ \ No newline at end of file diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md new file mode 100644 index 0000000..88132cb --- /dev/null +++ b/CONTRIBUTING.md @@ -0,0 +1,29 @@ +# Contributing to Comfyui-FlowChain + +Thank you for your interest in contributing to sd-wav2lip-uhq! We appreciate your effort and to help us incorporate your contribution in the best way possible, please follow the following contribution guidelines. + +## Reporting Bugs + +If you find a bug in the project, we encourage you to report it. Here's how: + +1. First, check the [existing Issues](url_of_issues) to see if the issue has already been reported. If it has, please add a comment to the existing issue rather than creating a new one. +2. If you can't find an existing issue that matches your bug, create a new issue. Make sure to include as many details as possible so we can understand and reproduce the problem. + +## Proposing Changes + +We welcome code contributions from the community. Here's how to propose changes: + +1. Fork this repository to your own GitHub account. +2. Create a new branch on your fork for your changes. +3. Make your changes in this branch. +4. When you are ready, submit a pull request to the `main` branch of this repository. + +Please note that we use the GitHub Flow workflow, so all pull requests should be made to the `main` branch. + +Before submitting a pull request, please make sure your code adheres to the project's coding conventions and it has passed all tests. If you are adding features, please also add appropriate tests. + +## Contact + +If you have any questions or need help, please ping the developer via discord NumZ#7184 to make sure your addition will fit well into such a large project and to get help if needed. + +Thank you again for your contribution! diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000..1ae0738 --- /dev/null +++ b/LICENSE @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2024 the comfyui-FlowChain + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/README.md b/README.md new file mode 100644 index 0000000..0c399ad --- /dev/null +++ b/README.md @@ -0,0 +1,193 @@ +# ⛓️ Comfyui-FlowChain + +## πŸ’‘ Description +This repository includes a set of custom nodes for ComfyUI that allow you to: + + - Convert your workflows into nodes + - Chain your workflows together + - Bonus: a node to integrate [LipSync Studio v0.6](https://www.patreon.com/Wav2LipStudio) via API (third-party application) + + + +## πŸš€ All Nodes + + + +## πŸ“– Quick Index +* [πŸš€ Updates](#-updates) +* [πŸ’» Installation](#-installation) +* [πŸ•ΈοΈ Nodes](#-nodes) +* [πŸ“Ί Tutorial](#-tutorial) +* [🐍 Usage](#-usage) +* [πŸ’ͺ Special things to know](#-special-things-to-know) +* [πŸ“Ί Examples](#-examples) +* [😎 Contributing](#-contributing) +* [πŸ™ Appreciation](#-appreciation) +* [πŸ“œ License](#-license) +* [β˜• Support](#-support) + +## πŸš€ Updates +**2024.11.01 Initial version features :** +- πŸ’ͺ Convert your workflows into nodes +- ⛓️ Chain your workflow +- πŸ‘„ Extra Node that use [LipSync Studio v0.6](https://www.patreon.com/Wav2LipStudio) + +## πŸ’» Installation + +1. Install [Git](https://git-scm.com/) +2. Go to folder ..\ComfyUI\custom_nodes +3. Run cmd.exe + > **Windows**: + > + > > **Variant 1:** In folder click panel current path and input **cmd** and press **Enter** on keyboard + > > + > > **Variant 2:** Press on keyboard Windows+R, and enter cmd.exe open window cmd, enter **cd /d your_path_to_custom_nodes**, **Enter** on keyboard +4. Then do : + +```git clone https://github.com/numz/Comfyui-FlowChain.git``` + +After this command be created folder Comfyui-FlowChain + +8. Go to the folder: + +```cd Comfyui-FlowChain``` + +8. Then do: + +```pip install -r requirements.txt``` + +7. Run Comfyui... + +## πŸ•ΈοΈ Nodes: + +| | Name | Description | ComfyUI category | +|:-------------------------------------------------:|:--------------------|:------------------------------------------------------------------------------------------------------------:|:----------------:| +| | _Workflow_ | Node that allows loading workflows in API format. It will show Inputs and Outputs into the loaded Workflows | FlowChain ⛓️ | +| | _Workflow Input_ | Node used to declare the inputs of your workflows. | FlowChain ⛓️ | +| | _Workflow Output_ | Node used to declare the outputs of your workflows. | FlowChain ⛓️ | +| | _Workflow Continue_ | Node to stop/Continue the workflow process. | FlowChain ⛓️ | +| | _Workflow Lipsync_ | Extra Node to use LipSync Studio via API | FlowChain ⛓️ | + + +## πŸ“Ί Tutorial +- [Here](https://youtu.be/B84A5alpPDc) + +# 🐍 Usage + +## ⛓️ Workflow Node +![Illustration](docs/assets/workflow2.png) + + - Load a workflow in **workflows** list. This field will show all workflows saved in the comfyui user folder: **ComfyUI\user\default\workflows\api**, if you add a new workflow in this folder you have to refresh UI (F5 to refresh web page) to see it in the **workflows** list. + - Workflows have to be saved as **API format** of comfyui, but save it also in normal format because "api formal file" can't be loaded in comfyui as usually. + + + +If you don't see **"Export (API format)"** options in Comfyui do this : + - go to Settings + - Activate the **Dev Mode** options + +![Illustration](docs/assets/devmode.png) + +-You can also Import the file by "copy/paste" your workflow path in "workflow_api_path" and click import, that will add your workflow in the comfyui api path. + +## ⛓️ Input Node +![Illustration](docs/assets/input2.png) + +- Allow to declare inputs in your workflow. +- Types available : **"IMAGE", "MASK", "STRING", "INT", "FLOAT", "LATENT", "BOOLEAN", "CLIP", "CONDITIONING", "MODEL", "VAE"** +- Give a Name and select the type. +- **Default** value is used when debugging your workflow or if you don't plug an input into the **Workflow** node. + +- ![Illustration](docs/assets/workflow5.png) + +## ⛓️ Output Node +![Illustration](docs/assets/output1.png) + +- Allow to declare outputs in your workflow. +- Types available : **"IMAGE", "MASK", "STRING", "INT", "FLOAT", "LATENT", "BOOLEAN", "CLIP", "CONDITIONING", "MODEL", "VAE"** +- Give a Name and select the type. +- **Default** value is used to connect the output. + +![Illustration](docs/assets/output2.png) + +## ⛓️ Continue Node +![Illustration](docs/assets/continue1.png) + +- Usually associated with a **boolean** input plugged on **"continue_workflow"**, allow to "Stop" a workflow if **"continue_workflow"** is False. +- Types available : **"IMAGE", "LATENT"** +- Give a Name and select the type. +- During development of your workflow, If **continue_worflow" is False it will let pass only 1 image/latent, and if True it will let pass all images/latents. + +![Illustration](docs/assets/continue3.png) + +- But When a workflow is loaded into the **"workflow"** Node, which contain a **"Workflow Continue"** node, it will be delete if **continue_workflow** is False. That allow to create conditional situation where you want to prevent computation of some parts. + +![Illustration](docs/assets/continue5.png) + +## πŸ”‰πŸ‘„ Workflow LipSync Node +![Illustration](docs/assets/lipsync2.png) + +- Extra Node that allow to use third-party app **[Lipsync Studio v0.6](https://www.patreon.com/Wav2LipStudio)** Via it's API +- Inputs: + - **frames**: Images to compute. + - **audio**: Audio to add. + - **faceswap_image**: An image with a face to swap. + - **lipsync_studio_url**: usually http://127.0.0.1:7860 + - **project_name**: name of your project. + - **face_id**: id of the face you want to lipsync and faceswap. + - **fps**: frame per second. + - **avatar**: Will be used create a driving video, 10 avatars are available, each give different output result. + - **close mouth before Lip sync**: Allow to close the mouth before create the lip sync. + - **quality**: Can be **Low, Medium, High**, in High gfpgan will be used to enhance quality output. + - **skip_first_frame**: number of frames to remove at the beginning of the video. + - **load_cap**: number of frames to load. + - **low vram**: allow to decrease VRAM consumption for low pc configuration. + +Project will be automatically created into your Lipsync Studio **projects** folder. You can then load it into studio and work directly from studio if the output not good enough for you. + +![Illustration](docs/assets/lipsync3.png) + +## πŸ’ͺ Special things to know + + the **"πŸͺ› Switch"** nodes from [Crystools](https://github.com/crystian/ComfyUI-Crystools) have a particular place in **workflow Node** + +![Illustration](docs/assets/crystools.png) + +Let's illustrate this with an example: + +![Illustration](docs/assets/switch.png) + +Here we want to choose between video1 or video2. It depends on the **boolean** value in **Switch Image Node**. The issue here is that both videos will be loaded before Switch. To prevent both videos from being loaded, the **"workflow node"** will check the boolean value, remove the unused node, and directly connect the correct value to the preview image. + +![Illustration](docs/assets/switch2.png) + +This gives you the ability to create truly conditional cases in your workflows, without computing irrelevant nodes. + +# πŸ“Ί Examples + +https://user-images.githubusercontent.com/800903/262439441-bb9d888a-d33e-4246-9f0a-1ddeac062d35.mp4 + +https://user-images.githubusercontent.com/800903/262442794-61b1e32f-3f87-4b36-98d6-f711822bdb1e.mp4 + +https://user-images.githubusercontent.com/800903/262449305-901086a3-22cb-42d2-b5be-a5f38db4549a.mp4 + +https://user-images.githubusercontent.com/800903/267808494-300f8cc3-9136-4810-86e2-92f2114a5f9a.mp4 + +# 😎 Contributing + +We welcome contributions to this project. When submitting pull requests, please provide a detailed description of the changes. see [CONTRIBUTING](CONTRIBUTING.md) for more information. + +# πŸ™ Appreciation +- [Jedrzej Kosinski](https://github.com/Kosinkadink/ComfyUI-VideoHelperSuite) : For the code quality that really inspired me during development. + + +# β˜• Support + +this project is open-source effort that is free to use and modify. I rely on the support of users to keep this project going and help improve it. If you'd like to support me, you can make a donation on my [Patreon page](https://www.patreon.com/Wav2LipStudio). Any contribution, large or small, is greatly appreciated! + +Your support helps me cover the costs of development and maintenance, and allows me to allocate more time and resources to enhancing this project. Thank you for your support! + +[patreon page](https://www.patreon.com/Wav2LipStudio) + +# πŸ“œ License +* The code in this repository is released under the MIT license as found in the [LICENSE file](LICENSE). diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..3b982b7 --- /dev/null +++ b/__init__.py @@ -0,0 +1,46 @@ +import os +import importlib.util +import sys +import traceback +from .lipsync_studio import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS +from .workflow_nodes import NODE_CLASS_MAPPINGS_NODES, NODE_DISPLAY_NAME_MAPPINGS_NODES +from .workflow import NODE_CLASS_MAPPINGS_WORKFLOW, NODE_DISPLAY_NAME_MAPPINGS_WORKFLOW +from pathlib import Path + +NODE_CLASS_MAPPINGS.update(NODE_CLASS_MAPPINGS_NODES) +NODE_CLASS_MAPPINGS.update(NODE_CLASS_MAPPINGS_WORKFLOW) +NODE_DISPLAY_NAME_MAPPINGS.update(NODE_DISPLAY_NAME_MAPPINGS_NODES) +NODE_DISPLAY_NAME_MAPPINGS.update(NODE_DISPLAY_NAME_MAPPINGS_WORKFLOW) + +def get_ext_dir(subpath=None, mkdir=False): + dir = os.path.dirname(__file__) + if subpath is not None: + dir = os.path.join(dir, subpath) + + dir = os.path.abspath(dir) + + if mkdir and not os.path.exists(dir): + os.makedirs(dir) + return dir + + +py = Path(get_ext_dir("py")) +files = list(py.glob("*.py")) +for file in files: + try: + name = os.path.splitext(file)[0] + spec = importlib.util.spec_from_file_location(name, os.path.join(py, file)) + module = importlib.util.module_from_spec(spec) + sys.modules[name] = module + spec.loader.exec_module(module) + if hasattr(module, "NODE_CLASS_MAPPINGS") and getattr(module, "NODE_CLASS_MAPPINGS") is not None: + NODE_CLASS_MAPPINGS.update(module.NODE_CLASS_MAPPINGS) + if hasattr(module, "NODE_DISPLAY_NAME_MAPPINGS") and getattr(module, + "NODE_DISPLAY_NAME_MAPPINGS") is not None: + NODE_DISPLAY_NAME_MAPPINGS.update(module.NODE_DISPLAY_NAME_MAPPINGS) + except Exception as e: + traceback.print_exc() +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS", "WEB_DIRECTORY"] + +WEB_DIRECTORY = "./web" + diff --git a/core_old/__init__.py b/core_old/__init__.py new file mode 100644 index 0000000..c097037 --- /dev/null +++ b/core_old/__init__.py @@ -0,0 +1,6 @@ +from .logger import * +from .keys import * +from .types import * +from .config import * +from .common import * +from .version import * diff --git a/core_old/common.py b/core_old/common.py new file mode 100644 index 0000000..a3d0252 --- /dev/null +++ b/core_old/common.py @@ -0,0 +1,107 @@ +import os +import json +import torch +from deepdiff import DeepDiff +from ..core_old import CONFIG, logger + + +# just a helper function to set the widget values (or clear them) +def setWidgetValues(value=None, unique_id=None, extra_pnginfo=None) -> None: + if unique_id and extra_pnginfo: + workflow = extra_pnginfo["workflow"] + node = next((x for x in workflow["nodes"] if str(x["id"]) == unique_id), None) + + if node: + node["widgets_values"] = value + + return None + + +# find difference between two jsons +def findJsonStrDiff(json1, json2): + msgError = "Could not compare jsons" + returnJson = {"error": msgError} + try: + # TODO review this + # dict1 = json.loads(json1) + # dict2 = json.loads(json2) + + returnJson = findJsonsDiff(json1, json2) + + returnJson = json.dumps(returnJson, indent=CONFIG["indent"]) + except Exception as e: + logger.warn(f"{msgError}: {e}") + + return returnJson + + +def findJsonsDiff(json1, json2): + msgError = "Could not compare jsons" + returnJson = {"error": msgError} + + try: + diff = DeepDiff(json1, json2, ignore_order=True, verbose_level=2) + + returnJson = {k: v for k, v in diff.items() if + k in ('dictionary_item_added', 'dictionary_item_removed', 'values_changed')} + + # just for print "values_changed" at first + returnJson = dict(reversed(returnJson.items())) + + except Exception as e: + logger.warn(f"{msgError}: {e}") + + return returnJson + + +# powered by: +# https://github.com/WASasquatch/was-node-suite-comfyui/blob/main/WAS_Node_Suite.py +# class: WAS_Samples_Passthrough_Stat_System +def get_system_stats(): + import psutil + + # RAM + ram = psutil.virtual_memory() + ram_used = ram.used / (1024 ** 3) + ram_total = ram.total / (1024 ** 3) + ram_stats = f"Used RAM: {ram_used:.2f} GB / Total RAM: {ram_total:.2f} GB" + + # VRAM (with PyTorch) + device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + vram_used = torch.cuda.memory_allocated(device) / (1024 ** 3) + vram_total = torch.cuda.get_device_properties(device).total_memory / (1024 ** 3) + vram_stats = f"Used VRAM: {vram_used:.2f} GB / Total VRAM: {vram_total:.2f} GB" + + # Hard Drive Space + hard_drive = psutil.disk_usage("/") + used_space = hard_drive.used / (1024 ** 3) + total_space = hard_drive.total / (1024 ** 3) + hard_drive_stats = f"Used Space: {used_space:.2f} GB / Total Space: {total_space:.2f} GB" + + return [ram_stats, vram_stats, hard_drive_stats] + + +# return x and y resolution of an image (torch tensor) +def getResolutionByTensor(image=None) -> dict: + res = {"x": 0, "y": 0} + + if image is not None: + img = image.movedim(-1, 1) + + res["x"] = img.shape[3] + res["y"] = img.shape[2] + + return res + + +# by https://stackoverflow.com/questions/6080477/how-to-get-the-size-of-tar-gz-in-mb-file-in-python +def get_size(path): + size = os.path.getsize(path) + if size < 1024: + return f"{size} bytes" + elif size < pow(1024, 2): + return f"{round(size / 1024, 2)} KB" + elif size < pow(1024, 3): + return f"{round(size / (pow(1024, 2)), 2)} MB" + elif size < pow(1024, 4): + return f"{round(size / (pow(1024, 3)), 2)} GB" diff --git a/core_old/config.py b/core_old/config.py new file mode 100644 index 0000000..13b2e71 --- /dev/null +++ b/core_old/config.py @@ -0,0 +1,7 @@ +import os +import logging + +CONFIG = { + "loglevel": int(os.environ.get("CRYSTOOLS_LOGLEVEL", logging.INFO)), + "indent": int(os.environ.get("CRYSTOOLS_INDENT", 2)) +} diff --git a/core_old/keys.py b/core_old/keys.py new file mode 100644 index 0000000..68ac0bc --- /dev/null +++ b/core_old/keys.py @@ -0,0 +1,29 @@ +from enum import Enum + + +class TEXTS(Enum): + CUSTOM_NODE_NAME = "Crystools" + LOGGER_PREFIX = "Crystools" + CONCAT = "concatenated" + INACTIVE_MSG = "inactive" + INVALID_METADATA_MSG = "Invalid metadata raw" + FILE_NOT_FOUND = "File not found!" + + +class CATEGORY(Enum): + TESTING = "_for_testing" + MAIN = "crystools πŸͺ›" + PRIMITIVE = "/Primitive" + DEBUGGER = "/Debugger" + LIST = "/List" + SWITCH = "/Switch" + PIPE = "/Pipe" + IMAGE = "/Image" + UTILS = "/Utils" + METADATA = "/Metadata" + + +# remember, all keys should be in lowercase! +class KEYS(Enum): + LIST = "list_string" + PREFIX = "prefix" diff --git a/core_old/logger.py b/core_old/logger.py new file mode 100644 index 0000000..4c21af5 --- /dev/null +++ b/core_old/logger.py @@ -0,0 +1,39 @@ +# by https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/blob/main/control/logger.py +import sys +import copy +import logging +from .keys import TEXTS +from .config import CONFIG + + +class ColoredFormatter(logging.Formatter): + COLORS = { + "DEBUG": "\033[0;36m", # CYAN + "INFO": "\033[0;32m", # GREEN + "WARNING": "\033[0;33m", # YELLOW + "ERROR": "\033[0;31m", # RED + "CRITICAL": "\033[0;37;41m", # WHITE ON RED + "RESET": "\033[0m", # RESET COLOR + } + + def format(self, record): + colored_record = copy.copy(record) + levelname = colored_record.levelname + seq = self.COLORS.get(levelname, self.COLORS["RESET"]) + colored_record.levelname = f"{seq}{levelname}{self.COLORS['RESET']}" + return super().format(colored_record) + + +# Create a new logger +logger = logging.getLogger(TEXTS.LOGGER_PREFIX.value) +logger.propagate = False + +# Add handler if we don't have one. +if not logger.handlers: + handler = logging.StreamHandler(sys.stdout) + handler.setFormatter(ColoredFormatter("[%(name)s %(levelname)s] %(message)s")) + logger.addHandler(handler) + +# Configure logger +loglevel = CONFIG["loglevel"] +logger.setLevel(loglevel) diff --git a/core_old/types.py b/core_old/types.py new file mode 100644 index 0000000..499ae93 --- /dev/null +++ b/core_old/types.py @@ -0,0 +1,36 @@ +import sys + +FLOAT = ("FLOAT", {"default": 1, + "min": -sys.float_info.max, + "max": sys.float_info.max, + "step": 0.01}) + +BOOLEAN = ("BOOLEAN", {"default": True}) +BOOLEAN_FALSE = ("BOOLEAN", {"default": False}) + +INT = ("INT", {"default": 1, + "min": -sys.maxsize, + "max": sys.maxsize, + "step": 1}) + +STRING = ("STRING", {"default": ""}) + +STRING_ML = ("STRING", {"multiline": True, "default": ""}) + +STRING_WIDGET = ("STRING", {"forceInput": True}) + +JSON_WIDGET = ("JSON", {"forceInput": True}) + +METADATA_RAW = ("METADATA_RAW", {"forceInput": True}) + +class AnyType(str): + """A special class that is always equal in not equal comparisons. Credit to pythongosssss""" + + def __eq__(self, _) -> bool: + return True + + def __ne__(self, __value: object) -> bool: + return False + + +any = AnyType("*") diff --git a/core_old/version.py b/core_old/version.py new file mode 100644 index 0000000..af16f8d --- /dev/null +++ b/core_old/version.py @@ -0,0 +1 @@ +version = "1.15.0" diff --git a/docs/assets/Continue.png b/docs/assets/Continue.png new file mode 100644 index 0000000..b49b3c8 Binary files /dev/null and b/docs/assets/Continue.png differ diff --git a/docs/assets/Input.png b/docs/assets/Input.png new file mode 100644 index 0000000..0bbdf32 Binary files /dev/null and b/docs/assets/Input.png differ diff --git a/docs/assets/allnodes.png b/docs/assets/allnodes.png new file mode 100644 index 0000000..1f43ea9 Binary files /dev/null and b/docs/assets/allnodes.png differ diff --git a/docs/assets/continue1.png b/docs/assets/continue1.png new file mode 100644 index 0000000..3af27ae Binary files /dev/null and b/docs/assets/continue1.png differ diff --git a/docs/assets/continue2.png b/docs/assets/continue2.png new file mode 100644 index 0000000..186175a Binary files /dev/null and b/docs/assets/continue2.png differ diff --git a/docs/assets/continue3.png b/docs/assets/continue3.png new file mode 100644 index 0000000..9b66571 Binary files /dev/null and b/docs/assets/continue3.png differ diff --git a/docs/assets/continue4.png b/docs/assets/continue4.png new file mode 100644 index 0000000..eeb7d2c Binary files /dev/null and b/docs/assets/continue4.png differ diff --git a/docs/assets/continue5.png b/docs/assets/continue5.png new file mode 100644 index 0000000..414e508 Binary files /dev/null and b/docs/assets/continue5.png differ diff --git a/docs/assets/crystools.png b/docs/assets/crystools.png new file mode 100644 index 0000000..d168e63 Binary files /dev/null and b/docs/assets/crystools.png differ diff --git a/docs/assets/devmode.png b/docs/assets/devmode.png new file mode 100644 index 0000000..83e0a88 Binary files /dev/null and b/docs/assets/devmode.png differ diff --git a/docs/assets/input2.png b/docs/assets/input2.png new file mode 100644 index 0000000..6e7ada5 Binary files /dev/null and b/docs/assets/input2.png differ diff --git a/docs/assets/lipsync.png b/docs/assets/lipsync.png new file mode 100644 index 0000000..6354f97 Binary files /dev/null and b/docs/assets/lipsync.png differ diff --git a/docs/assets/lipsync1.png b/docs/assets/lipsync1.png new file mode 100644 index 0000000..1471553 Binary files /dev/null and b/docs/assets/lipsync1.png differ diff --git a/docs/assets/lipsync2.png b/docs/assets/lipsync2.png new file mode 100644 index 0000000..aff5b3f Binary files /dev/null and b/docs/assets/lipsync2.png differ diff --git a/docs/assets/lipsync3.png b/docs/assets/lipsync3.png new file mode 100644 index 0000000..5c21966 Binary files /dev/null and b/docs/assets/lipsync3.png differ diff --git a/docs/assets/output.png b/docs/assets/output.png new file mode 100644 index 0000000..f45ee26 Binary files /dev/null and b/docs/assets/output.png differ diff --git a/docs/assets/output1.png b/docs/assets/output1.png new file mode 100644 index 0000000..8cedc09 Binary files /dev/null and b/docs/assets/output1.png differ diff --git a/docs/assets/output2.png b/docs/assets/output2.png new file mode 100644 index 0000000..c5f3346 Binary files /dev/null and b/docs/assets/output2.png differ diff --git a/docs/assets/save_as_api.png b/docs/assets/save_as_api.png new file mode 100644 index 0000000..369058d Binary files /dev/null and b/docs/assets/save_as_api.png differ diff --git a/docs/assets/switch.png b/docs/assets/switch.png new file mode 100644 index 0000000..98b621a Binary files /dev/null and b/docs/assets/switch.png differ diff --git a/docs/assets/switch2.png b/docs/assets/switch2.png new file mode 100644 index 0000000..6589e5c Binary files /dev/null and b/docs/assets/switch2.png differ diff --git a/docs/assets/workflow.png b/docs/assets/workflow.png new file mode 100644 index 0000000..0a426d4 Binary files /dev/null and b/docs/assets/workflow.png differ diff --git a/docs/assets/workflow2.png b/docs/assets/workflow2.png new file mode 100644 index 0000000..79f95f4 Binary files /dev/null and b/docs/assets/workflow2.png differ diff --git a/docs/assets/workflow3.png b/docs/assets/workflow3.png new file mode 100644 index 0000000..cd6bc95 Binary files /dev/null and b/docs/assets/workflow3.png differ diff --git a/docs/assets/workflow4.png b/docs/assets/workflow4.png new file mode 100644 index 0000000..e8a9522 Binary files /dev/null and b/docs/assets/workflow4.png differ diff --git a/docs/assets/workflow5.png b/docs/assets/workflow5.png new file mode 100644 index 0000000..1751161 Binary files /dev/null and b/docs/assets/workflow5.png differ diff --git a/lipsync_studio.py b/lipsync_studio.py new file mode 100644 index 0000000..1e0f9e4 --- /dev/null +++ b/lipsync_studio.py @@ -0,0 +1,221 @@ +import shutil +from gradio_client import Client +import os +import subprocess +import folder_paths +import numpy as np +import hashlib +from .utils.utils import ffmpeg_path +from .utils.logger import Logger +import sys +from PIL import Image + + +class WorkflowLipSync: + def __init__(self): + self.logger = Logger() + self.ws = None + + @classmethod + def INPUT_TYPES(cls): + return {"required": { + "lipsync_studio_url": ("STRING", {"default": "http://127.0.0.1:7860/"}), + "project_name": ("STRING", {"default": "project1"}), + "frames": ("IMAGE",), + "face_id": ("INT", {"default": 0, "min": 0, "max": 10, "step": 1}), + "fps": ("FLOAT", {"default": 25., "min": 0., "max": 60., "step": 1}), + "audio": ("AUDIO",), + "avatar": (["Avatar 1", "Avatar 2", "Avatar 3", "Avatar 4", "Avatar 5", "Avatar 6", "Avatar 7", "Avatar 8", "Avatar 9", "Avatar 10"],), + "close_mouth_before_lipsync": ("BOOLEAN", {"default": True}), + "quality": (["Low", "Medium", "High"],), + "skip_first_frames": ("INT", {"default": 0, "min": 0, "max": 10000, "step": 1}), + "load_cap": ("INT", {"default": 0, "min": 0, "max": 10000, "step": 1}), + "low_vram": ("BOOLEAN", {"default": False}), + + }, + "optional": { + "faceswap_image": ("IMAGE",), + }} + + # RETURN_TYPES = ("STRING", "STRING") + RETURN_TYPES = () + # RETURN_NAMES = ("faceswap_video_path", "lipsync_video_path") + RETURN_NAMES = () + FUNCTION = "generate" + CATEGORY = "FlowChain ⛓️" + + OUTPUT_NODE = True + + @classmethod + def IS_CHANGED(s, project_name, **kworgs): + m = hashlib.sha256() + m.update(project_name.encode()) + return m.digest().hex() + + def generate(self, lipsync_studio_url, project_name, frames, fps, face_id, audio, avatar, close_mouth_before_lipsync, quality, skip_first_frames, + load_cap, low_vram, faceswap_image=None, **kwargs): + client = Client(lipsync_studio_url, verbose=False) + full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path( + project_name, folder_paths.get_output_directory(), frames[0].shape[1], frames[0].shape[0]) + # Set project name + client.predict(project_name, api_name="/set_project_name") + frame_list = [] + counter = 0 + if not os.path.exists(os.path.join(full_output_folder, project_name)): + os.makedirs(os.path.join(full_output_folder, project_name)) + for (batch_number, image) in enumerate(frames): + i = 255. * image.cpu().numpy() + filename_with_batch_num = filename.replace("%batch_num%", str(batch_number)) + file = f"{filename_with_batch_num}_{counter:05}_.png" + img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8)) + img.save(os.path.join(full_output_folder, project_name, file), compress_level=4) + img_info = { + 'path': os.path.join(full_output_folder, project_name, file) + } + frame_list.append(img_info) + counter += 1 + + client.predict( + frame_list, + fps, + api_name="/new_frames" + ) + if load_cap == 0: + load_cap = len(frames) + + client.predict( + skip_first_frames + 1, # float (numeric value between 1 and 1) in 'Trim Video Start' Slider component + api_name="/video_start_frame" + ) + client.predict( + load_cap + 1, # float (numeric value between 1 and 1) in 'Trim Video Start' Slider component + api_name="/video_stop_frame" + ) + + if faceswap_image is not None: + i = 255. * faceswap_image[0].cpu().numpy() + file = f"faceswap_{counter:05}_.png" + img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8)) + img.save(os.path.join(full_output_folder, project_name, file), compress_level=4) + client.predict( + os.path.join(full_output_folder, project_name, file), + # filepath in 'Face Swap' Image component + api_name="/new_face_swap_img" + ) + else: + client.predict( + None, + # filepath in 'Face Swap' Image component + api_name="/new_face_swap_img" + ) + + client.predict( + 1, + # float (numeric value between 1 and 4) in 'Resolution Divide Factor' Slider component + 30, # float (numeric value between 0 and 100) in 'Min Face Width Detection' Slider component + True, # bool in 'Keyframes On Speaker Change' Checkbox component + True, # bool in 'Keyframes On Scene Change' Checkbox component + skip_first_frames + 1, # int 'Trim Video Start' Slider component + load_cap, # float (numeric value between 1 and 1) in 'Trim Video Stop' Slider component + 4, # float (numeric value between 1 and 64) in 'Number of CPU' Slider component + 1000, + api_name="/analyse_video" + ) + # Set Audio Type + client.predict( + # config["audio_path"] if config["audio_path"] else "Input Video",# Literal['File', 'Generate', 'Input Video'] in 'Audio Input' Radio component + "File", # Literal['File', 'Generate', 'Input Video'] in 'Audio Input' Radio component + api_name="/set_audio_type" + ) + output_file_audio = f"{filename}_{counter:05}.wav" + output_file_audio_path = os.path.join(full_output_folder, project_name, output_file_audio) + + # FFmpeg command to save audio in WAV format + channels = audio['waveform'].size(1) + + wav_args = [ffmpeg_path, "-v", "error", "-n", + "-ar", str(audio['sample_rate']), # Sample rate + "-ac", str(channels), # Number of channels + "-f", "f32le", "-i", "-", # Audio format and input from stdin + "-c:a", "pcm_s16le", # Encode as 16-bit PCM WAV + output_file_audio_path] + env = os.environ.copy() + audio_data = audio['waveform'].squeeze(0).transpose(0, 1) \ + .numpy().tobytes() + + try: + res = subprocess.run(wav_args, input=audio_data, + env=env, capture_output=True, check=True) + except subprocess.CalledProcessError as e: + raise Exception("An error occurred in the ffmpeg subprocess:\n" \ + + e.stderr.decode("utf-8")) + + if res.stderr: + print(res.stderr.decode("utf-8"), end="", file=sys.stderr) + + client.predict( + output_file_audio_path, + # filepath in 'Speech' Audio component + api_name="/set_audio_file" + ) + client.predict( + avatar, + # Literal['None', 'Avatar 1', 'Avatar 2', 'Avatar 3', 'Avatar 4', 'Avatar 5', 'Avatar 6', 'Avatar 7', 'Avatar 8', 'Avatar 9', 'Avatar 10'] in 'Avatar' Dropdown component + api_name="/change_avatar" + ) + client.predict( + low_vram, # bool in 'Low VRAM' Checkbox component + api_name="/set_low_vram" + ) + client.predict( + avatar, + api_name="/generate_driving_video" + ) + client.predict( + quality, # Literal['Low', 'Medium', 'High', 'Best'] in 'Video Quality' Radio component + api_name="/set_video_quality" + ) + faceswap_video = None + if faceswap_image is not None: + result = client.predict( + api_name="/generate_faceswap" + ) + faceswap_video = result["value"]["video"] + client.predict( + face_id, # Literal[] in 'Face Id' Dropdown component + False, # bool in 'Show wav2lip Output' Checkbox component + api_name="/set_face_id" + ) + client.predict( + True, # bool in 'Stop video' Checkbox component + api_name="/set_stop_video" + ) + client.predict( + close_mouth_before_lipsync, # bool in 'Stop video' Checkbox component + api_name="/set_face_zero" + ) + + # Generate Wav2lip + result = client.predict( + 1, # float (numeric value between 1 and 100) in 'Volume Amplifier' Slider component + api_name="/generate_w2l" + ) + output_dir = folder_paths.get_output_directory() + video_path = result["value"]["video"] + new_path = os.path.join(output_dir, project_name, os.path.split(video_path)[-1]) + if not os.path.exists(new_path): + shutil.copy(video_path, new_path) + return {"ui": {"video_path": [new_path, project_name]}} + # return (video_path, faceswap_video) + + +# A dictionary that contains all nodes you want to export with their names +# NOTE: names should be globally unique +NODE_CLASS_MAPPINGS = { + "WorkflowLipSync": WorkflowLipSync, +} + +# A dictionary that contains the friendly/humanly readable titles for the nodes +NODE_DISPLAY_NAME_MAPPINGS = { + "WorkflowLipSync": "Workflow LipSync (FlowChain ⛓️)", +} diff --git a/py/endpoints.py b/py/endpoints.py new file mode 100644 index 0000000..face6eb --- /dev/null +++ b/py/endpoints.py @@ -0,0 +1,178 @@ +import sys +import uuid + +import server +from aiohttp import web +import shutil +import os +import subprocess +import json +import urllib.request +import copy +import folder_paths +from app.user_manager import UserManager +import multiprocessing as mp +import time +import queue +from multiprocessing import Process, Queue +import websocket +from nacl import hashlib + +client_id = '5b49a023-b05a-4c53-8dc9-addc3a749911' +server_address = "127.0.0.1:8188" + + +@server.PromptServer.instance.routes.get("/flowchain/workflows") +async def workflows(request): + user = UserManager().get_request_user_id(request) + json_path = folder_paths.user_directory + "/" + user + "/workflows/api/" + result = {} + if os.path.exists(json_path): + files = os.listdir(json_path) + for idx, file in enumerate(files): + with open(json_path + file, "r", encoding="utf-8") as f: + json_content = json.load(f) + nodes_input = {k: v for k, v in json_content.items() if v["class_type"] == "WorkflowInput"} + nodes_output = {k: v for k, v in json_content.items() if v["class_type"] == "WorkflowOutput"} + result[file] = {"inputs": nodes_input, "outputs": nodes_output} + else: + os.makedirs(json_path) + result["No file in worflows/api folder"] = {"inputs": {}, "outputs": {}} + + return web.json_response(result, content_type='application/json') + + +@server.PromptServer.instance.routes.get("/flowchain/workflow") +async def workflow(request): + user = UserManager().get_request_user_id(request) + + original_path = request.query.get("workflow_path") + json_path = original_path.replace("\\", "/").split("/") + if ".json" in json_path[0]: + file_name = json_path[0] + json_path = folder_paths.user_directory + "/" + user + "/workflows/api/" + file_name + else: + file_name = json_path[-1] + json_path = folder_paths.user_directory + "/" + user + "/workflows/api/" + file_name + shutil.copy(original_path, json_path) + if os.path.exists(json_path): + with open(json_path, "r", encoding="utf-8") as f: + json_content = json.load(f) + err = "none" + if "nodes" in json_content: + err = "Not a Json API format workflow" + result = {"error": err, "workflow": json_content, "file_name": file_name} + else: + result = {"error": "File not found"} + + return web.json_response(result, content_type='application/json') + + +""" +def generate(workflow_path, kwargs): + workflow = json.load(open(workflow_path, "r", encoding="utf-8")) + outputs = get_outputs(workflow) + workflow_optimized = copy.deepcopy(workflow) + for idx, field in enumerate(kwargs): + for node_id, node in workflow.items(): + if "input_" + field["name"] in node["_meta"]["title"]: + # get first key of workflow[node_id]["inputs"] + key = list(workflow[node_id]["inputs"].keys())[0] + workflow[node_id]["inputs"][key] = field["value"] + + boolean_values = [] + for node_id, node in workflow.items(): + if "boolean" in node["inputs"] and "input_" in node["_meta"]["title"]: + boolean_values.append((node_id, node["inputs"]["boolean"])) + + for node_id, active in boolean_values: + for node_id2, value2 in workflow_optimized.items(): + if "boolean" in value2["inputs"] and ( + "on_true" in value2["inputs"] or "on_false" in value2["inputs"]): + if node_id2 in workflow: + + if workflow[node_id2]["inputs"]["boolean"] == [node_id, 0]: + input_to_replace = None + if active: + if "on_true" in value2["inputs"]: + input_to_replace = workflow[node_id2]["inputs"]["on_true"] + else: + if "on_false" in value2["inputs"]: + input_to_replace = workflow[node_id2]["inputs"]["on_false"] + worflow_value_to_change = [] + for key3, value3 in workflow.items(): + for k, v in value3["inputs"].items(): + if v == [node_id2, 0]: + worflow_value_to_change.append((key3, k)) + # workflow[key3]["inputs"][k] = input_to_replace + for key3, k in worflow_value_to_change: + if input_to_replace: + workflow[key3]["inputs"][k] = input_to_replace + else: + del workflow[key3]["inputs"][k] + del workflow[node_id2] + + boolean_values = [] + for node_id, value in workflow.items(): + if value["class_type"] == "Continue Workflow": + boolean_values.append((node_id, value["inputs"]["boolean"], value["inputs"]["line"])) + + for node_id, active, line in boolean_values: + for node_id2, value2 in workflow_optimized.items(): + worflow_value_to_change = [] + for inp, val in value2["inputs"].items(): + if val == [node_id, 0]: + if type(active) == list: + continue_workflow = workflow_optimized[active[0]]['inputs']['boolean'] + else: + continue_workflow = active + if continue_workflow: + worflow_value_to_change.append((node_id2, inp, line)) + else: + worflow_value_to_change.append((node_id2, inp, None)) + + for key3, k, line2 in worflow_value_to_change: + if line2: + workflow[key3]["inputs"][k] = line2 + else: + del workflow[key3]["inputs"][k] + queue_prompt(workflow, outputs) + return True + +def get_history(prompt_id): + with urllib.request.urlopen("http://{}/history/{}".format(server_address, prompt_id)) as response: + return json.loads(response.read()) + + +def get_outputs(workflow): + output_images_path = [] + for node_id, node in workflow.items(): + if "output_" in node["_meta"]["title"]: + output_images_path.append(node["_meta"]["title"]) + return output_images_path + + +def queue_prompt(prompt, outputs): + root_folder = os.path.dirname(__file__) + if not os.path.exists(root_folder + "/../queue"): + os.makedirs(root_folder + "/../queue") + + queues = {} + if os.path.exists(root_folder + "/../queue/queue.json"): + queues = json.loads(open(root_folder + "/../queue/queue.json", "r", encoding="utf-8").read()) + + uid = str(uuid.uuid4()) + + queues[uid] = {"prompt": prompt, "client_id": client_id, "output_fields": outputs, "status": {"completed": "false"}} + print(uid) + with open(root_folder + "/../queue/queue.json", "w", encoding="utf-8") as f: + json.dump(queues, f) + time.sleep(0.5) + commands = [sys.executable, root_folder + "/../queue/queue.py", uid] + try: + subprocess.Popen(commands, stderr=subprocess.PIPE) + return True + except subprocess.CalledProcessError as exception: + print(exception.stderr.decode().strip(), __name__.upper()) + return False +""" diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..abd7706 --- /dev/null +++ b/requirements.txt @@ -0,0 +1 @@ +gradio_client==0.8.0 \ No newline at end of file diff --git a/utils/caching.py b/utils/caching.py new file mode 100644 index 0000000..2cee037 --- /dev/null +++ b/utils/caching.py @@ -0,0 +1,315 @@ +import itertools +from typing import Sequence, Mapping +from comfy_execution.graph import DynamicPrompt + +import nodes + +from comfy_execution.graph_utils import is_link + + +class CacheKeySet: + def __init__(self, dynprompt, node_ids, is_changed_cache): + self.keys = {} + self.subcache_keys = {} + + def add_keys(self, node_ids): + raise NotImplementedError() + + def all_node_ids(self): + return set(self.keys.keys()) + + def get_used_keys(self): + return self.keys.values() + + def get_used_subcache_keys(self): + return self.subcache_keys.values() + + def get_data_key(self, node_id): + return self.keys.get(node_id, None) + + def get_subcache_key(self, node_id): + return self.subcache_keys.get(node_id, None) + + +class Unhashable: + def __init__(self): + self.value = float("NaN") + + +def to_hashable(obj): + # So that we don't infinitely recurse since frozenset and tuples + # are Sequences. + if isinstance(obj, (int, float, str, bool, type(None))): + return obj + elif isinstance(obj, Mapping): + return frozenset([(to_hashable(k), to_hashable(v)) for k, v in sorted(obj.items())]) + elif isinstance(obj, Sequence): + return frozenset(zip(itertools.count(), [to_hashable(i) for i in obj])) + else: + # TODO - Support other objects like tensors? + return Unhashable() + + +class CacheKeySetID(CacheKeySet): + def __init__(self, dynprompt, node_ids, is_changed_cache): + super().__init__(dynprompt, node_ids, is_changed_cache) + self.dynprompt = dynprompt + self.add_keys(node_ids) + + def add_keys(self, node_ids): + for node_id in node_ids: + if node_id in self.keys: + continue + if not self.dynprompt.has_node(node_id): + continue + node = self.dynprompt.get_node(node_id) + self.keys[node_id] = (node_id, node["class_type"]) + self.subcache_keys[node_id] = (node_id, node["class_type"]) + + +class CacheKeySetInputSignature(CacheKeySet): + def __init__(self, dynprompt, node_ids, is_changed_cache): + super().__init__(dynprompt, node_ids, is_changed_cache) + self.dynprompt = dynprompt + self.is_changed_cache = is_changed_cache + self.add_keys(node_ids) + + def include_node_id_in_input(self) -> bool: + return False + + def add_keys(self, node_ids): + for node_id in node_ids: + if node_id in self.keys: + continue + if not self.dynprompt.has_node(node_id): + continue + node = self.dynprompt.get_node(node_id) + self.keys[node_id] = self.get_node_signature(self.dynprompt, node_id) + self.subcache_keys[node_id] = (node_id, node["class_type"]) + + def get_node_signature(self, dynprompt, node_id): + signature = [] + ancestors, order_mapping = self.get_ordered_ancestry(dynprompt, node_id) + signature.append(self.get_immediate_node_signature(dynprompt, node_id, order_mapping)) + for ancestor_id in ancestors: + signature.append(self.get_immediate_node_signature(dynprompt, ancestor_id, order_mapping)) + return to_hashable(signature) + + def get_immediate_node_signature(self, dynprompt, node_id, ancestor_order_mapping): + if not dynprompt.has_node(node_id): + # This node doesn't exist -- we can't cache it. + return [float("NaN")] + node = dynprompt.get_node(node_id) + class_type = node["class_type"] + class_def = nodes.NODE_CLASS_MAPPINGS[class_type] + signature = [class_type, self.is_changed_cache.get(node_id)] + if self.include_node_id_in_input() or (hasattr(class_def, "NOT_IDEMPOTENT") and class_def.NOT_IDEMPOTENT): + signature.append(node_id) + inputs = node["inputs"] + for key in sorted(inputs.keys()): + if is_link(inputs[key]): + (ancestor_id, ancestor_socket) = inputs[key] + ancestor_index = ancestor_order_mapping[ancestor_id] + signature.append((key, ("ANCESTOR", ancestor_index, ancestor_socket))) + else: + signature.append((key, inputs[key])) + return signature + + # This function returns a list of all ancestors of the given node. The order of the list is + # deterministic based on which specific inputs the ancestor is connected by. + def get_ordered_ancestry(self, dynprompt, node_id): + ancestors = [] + order_mapping = {} + self.get_ordered_ancestry_internal(dynprompt, node_id, ancestors, order_mapping) + return ancestors, order_mapping + + def get_ordered_ancestry_internal(self, dynprompt, node_id, ancestors, order_mapping): + if not dynprompt.has_node(node_id): + return + inputs = dynprompt.get_node(node_id)["inputs"] + input_keys = sorted(inputs.keys()) + for key in input_keys: + if is_link(inputs[key]): + ancestor_id = inputs[key][0] + if ancestor_id not in order_mapping: + ancestors.append(ancestor_id) + order_mapping[ancestor_id] = len(ancestors) - 1 + self.get_ordered_ancestry_internal(dynprompt, ancestor_id, ancestors, order_mapping) + + +class BasicCache: + def __init__(self, key_class): + self.key_class = key_class + self.initialized = False + self.dynprompt: DynamicPrompt + self.cache_key_set: CacheKeySet + self.cache = {} + self.subcaches = {} + + def set_prompt(self, dynprompt, node_ids, is_changed_cache): + self.dynprompt = dynprompt + self.cache_key_set = self.key_class(dynprompt, node_ids, is_changed_cache) + self.is_changed_cache = is_changed_cache + self.initialized = True + + def all_node_ids(self): + assert self.initialized + node_ids = self.cache_key_set.all_node_ids() + for subcache in self.subcaches.values(): + node_ids = node_ids.union(subcache.all_node_ids()) + return node_ids + + def _clean_cache(self): + preserve_keys = set(self.cache_key_set.get_used_keys()) + to_remove = [] + for key in self.cache: + if key not in preserve_keys: + to_remove.append(key) + for key in to_remove: + del self.cache[key] + + def _clean_subcaches(self): + preserve_subcaches = set(self.cache_key_set.get_used_subcache_keys()) + + to_remove = [] + for key in self.subcaches: + if key not in preserve_subcaches: + to_remove.append(key) + for key in to_remove: + del self.subcaches[key] + + def clean_unused(self): + assert self.initialized + self._clean_cache() + self._clean_subcaches() + + def _set_immediate(self, node_id, value): + assert self.initialized + cache_key = self.cache_key_set.get_data_key(node_id) + self.cache[cache_key] = value + + def _get_immediate(self, node_id): + if not self.initialized: + return None + cache_key = self.cache_key_set.get_data_key(node_id) + if cache_key in self.cache: + return self.cache[cache_key] + else: + return None + + def _ensure_subcache(self, node_id, children_ids): + subcache_key = self.cache_key_set.get_subcache_key(node_id) + subcache = self.subcaches.get(subcache_key, None) + if subcache is None: + subcache = BasicCache(self.key_class) + self.subcaches[subcache_key] = subcache + subcache.set_prompt(self.dynprompt, children_ids, self.is_changed_cache) + return subcache + + def _get_subcache(self, node_id): + assert self.initialized + subcache_key = self.cache_key_set.get_subcache_key(node_id) + if subcache_key in self.subcaches: + return self.subcaches[subcache_key] + else: + return None + + def recursive_debug_dump(self): + result = [] + for key in self.cache: + result.append({"key": key, "value": self.cache[key]}) + for key in self.subcaches: + result.append({"subcache_key": key, "subcache": self.subcaches[key].recursive_debug_dump()}) + return result + + +class HierarchicalCache(BasicCache): + def __init__(self, key_class): + super().__init__(key_class) + + def _get_cache_for(self, node_id): + assert self.dynprompt is not None + parent_id = self.dynprompt.get_parent_node_id(node_id) + if parent_id is None: + return self + + hierarchy = [] + while parent_id is not None: + hierarchy.append(parent_id) + parent_id = self.dynprompt.get_parent_node_id(parent_id) + + cache = self + for parent_id in reversed(hierarchy): + cache = cache._get_subcache(parent_id) + if cache is None: + return None + return cache + + def get(self, node_id): + cache = self._get_cache_for(node_id) + if cache is None: + return None + return cache._get_immediate(node_id) + + def set(self, node_id, value): + cache = self._get_cache_for(node_id) + assert cache is not None + cache._set_immediate(node_id, value) + + def ensure_subcache_for(self, node_id, children_ids): + cache = self._get_cache_for(node_id) + assert cache is not None + return cache._ensure_subcache(node_id, children_ids) + + +class LRUCache(BasicCache): + def __init__(self, key_class, max_size=100): + super().__init__(key_class) + self.max_size = max_size + self.min_generation = 0 + self.generation = 0 + self.used_generation = {} + self.children = {} + + def set_prompt(self, dynprompt, node_ids, is_changed_cache): + super().set_prompt(dynprompt, node_ids, is_changed_cache) + self.generation += 1 + for node_id in node_ids: + self._mark_used(node_id) + + def clean_unused(self): + while len(self.cache) > self.max_size and self.min_generation < self.generation: + self.min_generation += 1 + to_remove = [key for key in self.cache if self.used_generation[key] < self.min_generation] + for key in to_remove: + del self.cache[key] + del self.used_generation[key] + if key in self.children: + del self.children[key] + self._clean_subcaches() + + def get(self, node_id): + self._mark_used(node_id) + return self._get_immediate(node_id) + + def _mark_used(self, node_id): + cache_key = self.cache_key_set.get_data_key(node_id) + if cache_key is not None: + self.used_generation[cache_key] = self.generation + + def set(self, node_id, value): + self._mark_used(node_id) + return self._set_immediate(node_id, value) + + def ensure_subcache_for(self, node_id, children_ids): + # Just uses subcaches for tracking 'live' nodes + super()._ensure_subcache(node_id, children_ids) + + self.cache_key_set.add_keys(children_ids) + self._mark_used(node_id) + cache_key = self.cache_key_set.get_data_key(node_id) + self.children[cache_key] = [] + for child_id in children_ids: + self._mark_used(child_id) + self.children[cache_key].append(self.cache_key_set.get_data_key(child_id)) + return self diff --git a/utils/logger.py b/utils/logger.py new file mode 100644 index 0000000..c60215f --- /dev/null +++ b/utils/logger.py @@ -0,0 +1,194 @@ +import logging +import re +import sys +from dataclasses import dataclass + + + +def Logger(): + + try: + # Create logger + logger = logging.getLogger(__name__) + # set log level to no print + #logger.setLevel(logging.CRITICAL) + logger.setLevel(logging.DEBUG) + if len(logger.handlers) > 0: + logger.handlers.clear() + console_level = "DEBUG" + console_handler = logging.StreamHandler(stream=sys.stdout) + console_handler.setLevel(console_level) + console_format = "%(asctime)s %(levelname)-8s - %(message)s" + colored_formatter = ColorizedArgsFormatter(console_format) + console_handler.setFormatter(colored_formatter) + logger.addHandler(console_handler) + + """ + file_handler = logging.FileHandler(log_filename) + file_level = "DEBUG" + file_handler.setLevel(file_level) + file_format = "%(asctime)s %(levelname)-8s - %(lineno)-5s - %(filename)-20s - %(message)s" + file_handler.setFormatter(BraceFormatStyleFormatter(file_format)) + logger.addHandler(file_handler) + """ + return logger + + except Exception: + err = sys.exc_info() + # print("Error : %s" % (err)) + + +@dataclass +class Log: + _logger: logging.Logger = None + _log_level: int = logging.DEBUG + + @property + def log_level(self): + return self._log_level + + @log_level.setter + def log_level(self, log_level: int): + self._log_level = log_level + + @property + def logger(self): + return self._logger + + @logger.setter + def logger(self, logger: Logger): + self._logger = logger + self.set_level() + + def set_level(self): + self.logger.setLevel(self.log_level) + for handler in self.logger.handlers: + handler.setLevel(self.log_level) + + +class ColorCodes: + grey = "\x1b[38;21m" + green = "\x1b[1;32m" + yellow = "\x1b[33;21m" + red = "\x1b[31;21m" + bold_red = "\x1b[31;1m" + blue = "\x1b[1;34m" + light_blue = "\x1b[1;36m" + purple = "\x1b[1;35m" + reset = "\x1b[0m" + + +class ColorizedArgsFormatter(logging.Formatter): + arg_colors = [ColorCodes.purple, ColorCodes.light_blue, ColorCodes.green, ColorCodes.yellow, ColorCodes.red] + level_fields = ["levelname", "levelno"] + level_to_color = { + logging.DEBUG: ColorCodes.red, + logging.INFO: ColorCodes.green, + logging.WARNING: ColorCodes.yellow, + logging.ERROR: ColorCodes.red, + logging.CRITICAL: ColorCodes.bold_red, + } + + def __init__(self, fmt: str): + super().__init__() + self.level_to_formatter = {} + + def add_color_format(level: int): + color = ColorizedArgsFormatter.level_to_color[level] + _format = fmt + for fld in ColorizedArgsFormatter.level_fields: + search = "(%\(" + fld + "\).*?s)" + _format = re.sub(search, f"{color}\\1{ColorCodes.reset}", _format) + + formatter = logging.Formatter(_format) + self.level_to_formatter[level] = formatter + + add_color_format(logging.DEBUG) + add_color_format(logging.INFO) + add_color_format(logging.WARNING) + add_color_format(logging.ERROR) + add_color_format(logging.CRITICAL) + + @staticmethod + def rewrite_record(record: logging.LogRecord): + if not BraceFormatStyleFormatter.is_brace_format_style(record): + return + + msg = record.msg + msg = msg.replace("{", "_{{") + msg = msg.replace("}", "_}}") + placeholder_count = 0 + # add ANSI escape code for next alternating color before each formatting parameter + # and reset color after it. + while True: + if "_{{" not in msg: + break + color_index = placeholder_count % len(ColorizedArgsFormatter.arg_colors) + color = ColorizedArgsFormatter.arg_colors[color_index] + msg = msg.replace("_{{", color + "{", 1) + msg = msg.replace("_}}", "}" + ColorCodes.reset, 1) + placeholder_count += 1 + + record.msg = msg.format(*record.args) + record.args = [] + + def format(self, record): + + orig_msg = record.msg + orig_args = record.args + formatter = self.level_to_formatter.get(record.levelno) + + self.rewrite_record(record) + formatted = formatter.format(record) + record.msg = orig_msg + record.args = orig_args + return formatted + + +class BraceFormatStyleFormatter(logging.Formatter): + def __init__(self, fmt: str): + super().__init__() + self.formatter = logging.Formatter(fmt) + + @staticmethod + def is_brace_format_style(record: logging.LogRecord): + if len(record.args) == 0: + return False + + msg = record.msg + if '%' in msg: + return False + count_of_start_param = msg.count("{") + count_of_end_param = msg.count("}") + + if count_of_start_param != count_of_end_param: + return False + + if count_of_start_param != len(record.args): + return False + + return True + + @staticmethod + def rewrite_record(record: logging.LogRecord): + if not BraceFormatStyleFormatter.is_brace_format_style(record): + return + record.msg = record.msg.format(*record.args) + record.args = [] + + def format(self, record): + + orig_msg = record.msg + orig_args = record.args + self.rewrite_record(record) + formatted = self.formatter.format(record) + + # formatted = re.sub(r"\'(.*?)\': \'(.*?)\'", f"{ColorCodes.light_blue}\\1{ColorCodes.reset}: {ColorCodes.bold_red}\\2{ColorCodes.reset}", formatted) + + # restore log record to original state for other handlers + record.msg = orig_msg + record.args = orig_args + return formatted + +# logger = Logger('sdf') +# logger.info("{0} {1} {2}", "sdf", "sdf", "sdf") diff --git a/utils/utils.py b/utils/utils.py new file mode 100644 index 0000000..cdbc979 --- /dev/null +++ b/utils/utils.py @@ -0,0 +1,92 @@ +import os +import shutil +import subprocess +from .caching import HierarchicalCache, LRUCache, CacheKeySetInputSignature, CacheKeySetID + + +class CacheSet: + def __init__(self, lru_size=None): + if lru_size is None or lru_size == 0: + self.init_classic_cache() + else: + self.init_lru_cache(lru_size) + self.all = [self.outputs, self.ui, self.objects] + + # Useful for those with ample RAM/VRAM -- allows experimenting without + # blowing away the cache every time + def init_lru_cache(self, cache_size): + self.outputs = LRUCache(CacheKeySetInputSignature, max_size=cache_size) + self.ui = LRUCache(CacheKeySetInputSignature, max_size=cache_size) + self.objects = HierarchicalCache(CacheKeySetID) + + # Performs like the old cache -- dump data ASAP + def init_classic_cache(self): + self.outputs = HierarchicalCache(CacheKeySetInputSignature) + self.ui = HierarchicalCache(CacheKeySetInputSignature) + self.objects = HierarchicalCache(CacheKeySetID) + + def recursive_debug_dump(self): + result = { + "outputs": self.outputs.recursive_debug_dump(), + "ui": self.ui.recursive_debug_dump(), + } + return result + + +caches = CacheSet(None) + + +def ffmpeg_suitability(path): + try: + version = subprocess.run([path, "-version"], check=True, + capture_output=True).stdout.decode("utf-8") + except: + return 0 + score = 0 + # rough layout of the importance of various features + simple_criterion = [("libvpx", 20), ("264", 10), ("265", 3), + ("svtav1", 5), ("libopus", 1)] + for criterion in simple_criterion: + if version.find(criterion[0]) >= 0: + score += criterion[1] + # obtain rough compile year from copyright information + copyright_index = version.find('2000-2') + if copyright_index >= 0: + copyright_year = version[copyright_index + 6:copyright_index + 9] + if copyright_year.isnumeric(): + score += int(copyright_year) + return score + + +if "VHS_FORCE_FFMPEG_PATH" in os.environ: + ffmpeg_path = os.environ.get("VHS_FORCE_FFMPEG_PATH") +else: + ffmpeg_paths = [] + try: + from imageio_ffmpeg import get_ffmpeg_exe + + imageio_ffmpeg_path = get_ffmpeg_exe() + ffmpeg_paths.append(imageio_ffmpeg_path) + except: + if "VHS_USE_IMAGEIO_FFMPEG" in os.environ: + raise + + if "VHS_USE_IMAGEIO_FFMPEG" in os.environ: + ffmpeg_path = imageio_ffmpeg_path + else: + system_ffmpeg = shutil.which("ffmpeg") + if system_ffmpeg is not None: + ffmpeg_paths.append(system_ffmpeg) + if os.path.isfile("ffmpeg"): + ffmpeg_paths.append(os.path.abspath("ffmpeg")) + if os.path.isfile("ffmpeg.exe"): + ffmpeg_paths.append(os.path.abspath("ffmpeg.exe")) + if len(ffmpeg_paths) == 0: + + ffmpeg_path = None + elif len(ffmpeg_paths) == 1: + # Evaluation of suitability isn't required, can take sole option + # to reduce startup time + ffmpeg_path = ffmpeg_paths[0] + else: + ffmpeg_path = max(ffmpeg_paths, key=ffmpeg_suitability) diff --git a/web/js/jsnodes.js b/web/js/jsnodes.js new file mode 100644 index 0000000..d863268 --- /dev/null +++ b/web/js/jsnodes.js @@ -0,0 +1,1015 @@ +import { app } from "../../../scripts/app.js"; +import { api } from '../../../scripts/api.js' +import { ComfyWidgets } from '../../../scripts/widgets.js' +const client_id = '5b49a023-b05a-4c53-8dc9-addc3a749911' + +const colors = ["#222222", "#5940bb", "#FFFFFF", "#7cbb1a", "#29699c", "#777788", "#268bd2", "#2ab7ca", "#d33682", "#dc322f", "#facfad","#77ff77", "#5940bb"] +const bg_colors = ["#000000", "#392978", "#89888d", "#496c12", "#19466a", "#4b4b56", "#165481", "#176974", "#851f50", "#911e1c", "#9f826b", "#499f49", "#392978"] +const node_type_list = ["none", "IMAGE", "MASK", "STRING", "INT", "FLOAT", "LATENT", "CLIP", "CONDITIONING", "MODEL", "VAE", "BOOLEAN", "SWITCH"] + +function chainCallback(object, property, callback) { + if (object == undefined) { + console.error("Tried to add callback to non-existant object") + return; + } + if (property in object) { + const callback_orig = object[property] + object[property] = function () { + const r = callback_orig.apply(this, arguments); + callback.apply(this, arguments); + return r + }; + } else { + object[property] = callback; + } +} + +function useKVState(nodeType) { + chainCallback(nodeType.prototype, "onNodeCreated", function () { + chainCallback(this, "onConfigure", function(info) { + if (!this.widgets) { + //Node has no widgets, there is nothing to restore + return + } + if (typeof(info.widgets_values) != "object") { + //widgets_values is in some unknown inactionable format + return + } + let widgetDict = info.widgets_values + + if (widgetDict.length == undefined) { + for (let w of this.widgets) { + if (w.name in widgetDict) { + w.value = widgetDict[w.name]; + if (w.name == "videopreview") { + w.updateSource(); + } + } + + } + } + }); + chainCallback(this, "onSerialize", function(info) { + info.widgets_values = {}; + if (!this.widgets) { + //object has no widgets, there is nothing to store + return; + } + for (let w of this.widgets) { + info.widgets_values[w.name] = w.value; + } + }); + }) +} + + +function addVideoPreview(nodeType) { + chainCallback(nodeType.prototype, "onNodeCreated", function() { + var element = document.createElement("div"); + const previewNode = this; + var previewWidget = this.addDOMWidget("videopreview", "preview", element, { + serialize: false, + hideOnZoom: false, + getValue() { + return element.value; + }, + setValue(v) { + element.value = v; + }, + }); + previewWidget.computeSize = function(width) { + if (this.aspectRatio && !this.parentEl.hidden) { + let height = (previewNode.size[0]-20)/ this.aspectRatio + 10; + if (!(height > 0)) { + height = 0; + } + this.computedHeight = height + 10; + return [width, height]; + } + return [width, -4];//no loaded src, widget should not display + } + element.addEventListener('contextmenu', (e) => { + e.preventDefault() + return app.canvas._mousedown_callback(e) + }, true); + element.addEventListener('pointerdown', (e) => { + e.preventDefault() + return app.canvas._mousedown_callback(e) + }, true); + element.addEventListener('mousewheel', (e) => { + e.preventDefault() + return app.canvas._mousewheel_callback(e) + }, true); + previewWidget.value = {hidden: false, paused: false, params: {}} + previewWidget.parentEl = document.createElement("div"); + previewWidget.parentEl.className = "vhs_preview"; + previewWidget.parentEl.style['width'] = "100%" + element.appendChild(previewWidget.parentEl); + previewWidget.videoEl = document.createElement("video"); + previewWidget.videoEl.controls = false; + previewWidget.videoEl.loop = true; + previewWidget.videoEl.muted = true; + previewWidget.videoEl.style['width'] = "100%" + previewWidget.videoEl.addEventListener("loadedmetadata", () => { + + previewWidget.aspectRatio = previewWidget.videoEl.videoWidth / previewWidget.videoEl.videoHeight; + fitHeight(this); + }); + previewWidget.videoEl.addEventListener("error", () => { + //TODO: consider a way to properly notify the user why a preview isn't shown. + previewWidget.parentEl.hidden = true; + fitHeight(this); + }); + previewWidget.videoEl.onmouseenter = () => { + previewWidget.videoEl.muted = false; + }; + previewWidget.videoEl.onmouseleave = () => { + previewWidget.videoEl.muted = true; + }; + + previewWidget.imgEl = document.createElement("img"); + previewWidget.imgEl.style['width'] = "100%" + previewWidget.imgEl.hidden = true; + previewWidget.imgEl.onload = () => { + previewWidget.aspectRatio = previewWidget.imgEl.naturalWidth / previewWidget.imgEl.naturalHeight; + fitHeight(this); + }; + + var timeout = null; + this.updateParameters = (params, force_update) => { + if (!previewWidget.value.params) { + if(typeof(previewWidget.value != 'object')) { + previewWidget.value = {hidden: false, paused: false} + } + previewWidget.value.params = {} + } + Object.assign(previewWidget.value.params, params) + timeout = setTimeout(() => previewWidget.updateSource(),100); + }; + previewWidget.updateSource = function () { + if (this.value.params == undefined) { + return; + } + let params = {} + Object.assign(params, this.value.params);//shallow copy + this.parentEl.hidden = this.value.hidden; + if (params.format?.split('/')[0] == 'video' || + app.ui.settings.getSettingValue("VHS.AdvancedPreviews", false) && + (params.format?.split('/')[1] == 'gif') || params.format == 'folder') { + this.videoEl.autoplay = !this.value.paused && !this.value.hidden; + let target_width = 256 + if (element.style?.width) { + //overscale to allow scrolling. Endpoint won't return higher than native + target_width = element.style.width.slice(0,-2)*2; + } + if (!params.force_size || params.force_size.includes("?") || params.force_size == "Disabled") { + params.force_size = target_width+"x?" + } else { + let size = params.force_size.split("x") + let ar = parseInt(size[0])/parseInt(size[1]) + params.force_size = target_width+"x"+(target_width/ar) + } + if (app.ui.settings.getSettingValue("VHS.AdvancedPreviews", false)) { + this.videoEl.src = api.apiURL('/viewvideo?' + new URLSearchParams(params)); + } else { + previewWidget.videoEl.src = api.apiURL('/view?' + new URLSearchParams(params)); + } + this.videoEl.hidden = false; + this.imgEl.hidden = true; + } else if (params.format?.split('/')[0] == 'image'){ + //Is animated image + this.imgEl.src = api.apiURL('/view?' + new URLSearchParams(params)); + this.videoEl.hidden = true; + this.imgEl.hidden = false; + } + } + previewWidget.parentEl.appendChild(previewWidget.videoEl) + previewWidget.parentEl.appendChild(previewWidget.imgEl) + }); +} +function fitHeight(node) { + node.setSize([node.size[0], node.computeSize([node.size[0], node.size[1]])[1]]) + node?.graph?.setDirtyCanvas(true); +} +function addPreviewOptions(nodeType) { + chainCallback(nodeType.prototype, "getExtraMenuOptions", function(_, options) { + let optNew = [] + const previewWidget = this.widgets.find((w) => w.name === "videopreview"); + + let url = null + if (previewWidget.videoEl?.hidden == false && previewWidget.videoEl.src) { + url = api.apiURL('/view?' + new URLSearchParams(previewWidget.value.params)); + url = url.replace('%2503d', '001') + } else if (previewWidget.imgEl?.hidden == false && previewWidget.imgEl.src) { + url = previewWidget.imgEl.src; + url = new URL(url); + } + if (url) { + optNew.push( + { + content: "Open preview", + callback: () => { + window.open(url, "_blank") + }, + }, + { + content: "Save preview", + callback: () => { + const a = document.createElement("a"); + a.href = url; + a.setAttribute("download", new URLSearchParams(previewWidget.value.params).get("filename")); + document.body.append(a); + a.click(); + requestAnimationFrame(() => a.remove()); + }, + } + ); + } + const PauseDesc = (previewWidget.value.paused ? "Resume" : "Pause") + " preview"; + if(previewWidget.videoEl.hidden == false) { + optNew.push({content: PauseDesc, callback: () => { + if(previewWidget.value.paused) { + previewWidget.videoEl?.play(); + } else { + previewWidget.videoEl?.pause(); + } + previewWidget.value.paused = !previewWidget.value.paused; + }}); + } + //TODO: Consider hiding elements if no video preview is available yet. + //It would reduce confusion at the cost of functionality + //(if a video preview lags the computer, the user should be able to hide in advance) + const visDesc = (previewWidget.value.hidden ? "Show" : "Hide") + " preview"; + optNew.push({content: visDesc, callback: () => { + if (!previewWidget.videoEl.hidden && !previewWidget.value.hidden) { + previewWidget.videoEl.pause(); + } else if (previewWidget.value.hidden && !previewWidget.videoEl.hidden && !previewWidget.value.paused) { + previewWidget.videoEl.play(); + } + previewWidget.value.hidden = !previewWidget.value.hidden; + previewWidget.parentEl.hidden = previewWidget.value.hidden; + fitHeight(this); + + }}); + optNew.push({content: "Sync preview", callback: () => { + //TODO: address case where videos have varying length + //Consider a system of sync groups which are opt-in? + for (let p of document.getElementsByClassName("vhs_preview")) { + for (let child of p.children) { + if (child.tagName == "VIDEO") { + child.currentTime=0; + } else if (child.tagName == "IMG") { + child.src = child.src; + } + } + } + }}); + if(options.length > 0 && options[0] != null && optNew.length > 0) { + optNew.push(null); + } + options.unshift(...optNew); + }); +} +function addLoadVideoCommon(nodeType, nodeData) { + addVideoPreview(nodeType); + addPreviewOptions(nodeType); +} +function cleanInputs(root_obj, reset_value=true) { + if (!root_obj.inputs) { + root_obj.inputs = []; + } + if (!root_obj.outputs) { + root_obj.outputs = []; + } + if (!root_obj.widgets) { + root_obj.widgets = []; + //root_obj.widgets_values = []; + } + if (!root_obj.widgets_values) { + root_obj.widgets_values = []; + } + + root_obj.widgets = root_obj.widgets.splice(0,3) + if(reset_value){ + for (let key in root_obj.widgets_values) { + if (key != "workflows" && key != "workflow_api_path" && key != "Import Workflow"){ + delete root_obj.widgets_values[key]; + } + } + const max_node_output = root_obj.outputs.length; + for(let i = 0; i i.name === field_name).length == 0){ + root_obj.addInput(field_name, "IMAGE"); + } + + const input_value = root_obj.widgets.length + 1 < root_obj.widgets_values.length?root_obj.widgets_values[root_obj.widgets.length]:value.default; + if (type == "STRING" || type =="text"){ + ComfyWidgets.STRING(root_obj, field_name, ['STRING',{default: value.default,callback: () => {},},],app,) + } + + if (type == "INT"){ + ComfyWidgets.INT( + root_obj, + field_name, + ['',{default: input_value,callback: () => {},},], app,) + } + if (type == "FLOAT"){ + ComfyWidgets.FLOAT( + root_obj, + field_name, + ['',{default: input_value,callback: (val) => console.log('VALUE', val), "min": 0.00, "max": 1.00, "step": 0.01},], + app, + ) + const widget = root_obj.inputs.filter(i => i.name === field_name); + if(widget.length > 0) + app.convertToWidget(root_obj, widget[0]); + } + + if (type == "BOOLEAN"){ + + root_obj.addWidget("toggle",field_name, input_value, ()=>{}); + const widget = root_obj.inputs.filter(i => i.name === field_name); + if(widget.length > 0) + app.convertToWidget(root_obj, widget[0]); + } + + if (type == "LATENT"){ + root_obj.addInput(field_name, "LATENT"); + } + + if (type == "MODEL"){ + root_obj.addInput(field_name, "MODEL"); + } + + + if (type == "CLIP"){ + root_obj.addInput(field_name, "CLIP"); + } + + if (type == "MASK"){ + root_obj.addInput(field_name, "MASK"); + } + + if (type == "CONDITIONING"){ + root_obj.addInput(field_name, "CONDITIONING"); + } + if (type == "VAE"){ + root_obj.addInput(field_name, "VAE"); + } + +} + +function importWorkflow(root_obj, workflow_path, app, reset_values=true){ + const filename = workflow_path.replace(/\\/g, '/').split("/"); + root_obj.title = "Workflow: "+filename[filename.length-1].replace(".json", "").replace(/_/g, " "); + api.fetchApi("/flowchain/workflow?workflow_path="+workflow_path) + .then(response => response.json()) + .then(data => { + cleanInputs(root_obj, reset_values); + if (data.error != "none"){ + return false + }else{ + const workflow = data.workflow; + console.log('Workflow:', workflow); + const nodes_input = Object.fromEntries( + Object.entries(workflow).filter(([k, v]) => v.class_type == "WorkflowInput") + ); + + const nodes_output = Object.fromEntries( + Object.entries(workflow).filter(([k, v]) => v.class_type == "WorkflowOutput") + ); + + Object.fromEntries( + Object.entries(nodes_input).filter((node, idx) => addWidgetType(root_obj, node[1].inputs)) + ); + + Object.fromEntries( + Object.entries(nodes_output).filter((node, idx) =>root_obj.addOutput(`${node[1].inputs.Name}`, node[1].inputs.type)) + ); + + console.log('Nodes:', nodes_input); + + root_obj.size[0] = 400; + return true + } + }) + .catch(error => { + console.error('Error:', error); + throw error; // Rilancia l'errore per consentire al chiamante di gestirlo + }); +} + +function addWidgetInfo(root_obj, field_name, value, app){ + let type = value.type; + if (type == "converted-widget"){ + type = value.origType; + } + if (type == "STRING" || type =="text"){ + ComfyWidgets.STRING(root_obj, field_name, ['STRING',{default: value.value,callback: () => {},},],app,) + } + if (type == "INT" || type == "number"){ + ComfyWidgets.INT( + root_obj, + field_name, + ['',{default: value.value, callback: () => {},},], + app, + ) + } + if (type == "FLOAT"){ + ComfyWidgets.FLOAT( + root_obj, + field_name, + ['',{default: value.value, callback: (val) => console.log('VALUE', val),},], + app, + ) + } + if (type == "BOOLEAN" || type == "toggle"){ + root_obj.addWidget("toggle",field_name, value.value, ()=>{}); + } +} + + +app.registerExtension({ + name: "FlowChain.jsnodes", + async beforeRegisterNodeDef(nodeType, nodeData, app) { + if(!nodeData?.category?.startsWith("FlowChain")) { + return; + } + + switch (nodeData.name) { + case "Workflow": + nodeType.prototype.onNodeCreated = function() { + + chainCallback(this, "onConfigure", function(info) { + let widgetDict = info.widgets_values + if (info.widgets_values.length == undefined) { + if(info.widgets_values.workflows.value != "None"){ + const workflow_name = info.widgets_values.workflows.value; + console.log("workflow_name", workflow_name) + console.log(app.lipsync_studio[workflow_name]) + const inputs = app.lipsync_studio[workflow_name].inputs; + + for (let w of this.widgets) { + if (w.name in widgetDict) { + w.value = widgetDict[w.name].value; + } + } + for (let [key, value] of Object.entries(widgetDict)) { + let widget = this.widgets.find(w => w.name === key); + if(!widget){ + addWidgetInfo(this, key, value, app); + widget = this.widgets.find(w => w.name === key); + } + widget.options = info.widgets_values[key].options; + widget.value = info.widgets_values[key].value; + for (let input of this.inputs) + if (input.name == key){ + for (let [key2, value2] of Object.entries(inputs)){ + if (value2.inputs.Name == key){ + input.type = value2.inputs.type; + widget.type = "converted-widget" + widget.origType = info.widgets_values[key].origType; + widget.origComputeSize = undefined; + widget.last_y = info.widgets_values[key].last_y; + widget.origSerializeValue = nodeType.prototype.serializeValue; + widget.value = info.widgets_values[key].value; + break; + } + } + break; + } + } + } + } + for(let i = this.outputs.length - 1; i>0; i--){ + if (this.outputs[i].name == "*"){ + this.removeOutput(i); + } + } + }); + chainCallback(this, "onSerialize", function(info) { + let inps = {}; + if (info.widgets_values[2] != "None"){ + const workflow_name = info.widgets_values[2]; + inps = app.lipsync_studio[workflow_name].inputs + } + info.widgets_values = {}; + if (!this.widgets) { + return; + } + + for (let w of this.widgets) { + info.widgets_values[w.name] = {name: w.name, options : w.options, value: w.value, type: w.type, origType: w.origType, last_y: w.last_y}; + } + + for (let w of this.inputs){ + for (let [key, value] of Object.entries(inps)){ + if (value.inputs.Name == w.name){ + w.type = value.inputs.type; + break; + } + } + } + + + }); + const workflow_reload = this.title.startsWith("Workflow: ")?true:false; + + const filename = this.title.replace("Workflow: ", ""); + this.addWidget("STRING", "workflow_api_path", "", ()=>{}); + this.addWidget("button", "Import Workflow", null, () => { + const workflow_path = this.widgets.find(w => w.name === "workflow_api_path")["value"]; + const filename = workflow_path.replace(/\\/g, '/').split("/"); + this.title = "Workflow: "+filename[filename.length-1].replace(".json", "").replace(/_/g, " "); + cleanInputs(this); + + if (workflow_path != "" && workflow_path != "None") + api.fetchApi("/flowchain/workflow?workflow_path="+workflow_path) + .then(response => response.json()) + .then(data => { + // Eseguire l'elaborazione dei dati + const workflow = data.workflow; + //console.log('Workflow:', workflow); + if (data.error == "none"){ + const combo = this.widgets.find(w => w.name === "workflows"); + combo["values"] += data.file_name; + combo["value"] = data.file_name; + importWorkflow(this, data.file_name, app) + }else{ + alert(data.error) + } + }) + .catch(error => { + console.error('Error:', error); + throw error; // Rilancia l'errore per consentire al chiamante di gestirlo + }); + }); + this.addWidget("combo", "workflows", "None", (value) => { + if (value == "None" || value == ""){ + this.title = "Workflow (FlowChain ⛓️)"; + cleanInputs(this); + }else{ + importWorkflow(this, value, app) + } + },{ + values: ["None", ...Object.keys(app.lipsync_studio)] + }); + if(!workflow_reload || !filename in app.lipsync_studio){ + cleanInputs(this); + } + this.color = "#004670"; + this.bgcolor = "#002942"; + } + break; + case "WorkflowInput": + nodeType.prototype.onNodeCreated = function() { + chainCallback(this, "onConfigure", function(info) { + let widgetDict = info.widgets_values + if (info.widgets_values.length == undefined) { + + for (let w of this.widgets) { + if (w.name in widgetDict) { + w.value = widgetDict[w.name].value; + } + } + // check if widgetDict in this.widgets + for (let [key, value] of Object.entries(widgetDict)) { + let widget = this.widgets.find(w => w.name === key); + let type = this.widgets.find(w => w.name === "type"); + if(!widget) + addWidgetInfo(this, key, value, app); + widget = this.widgets.find(w => w.name === key); + //this.widgets.push(value); + widget.options = info.widgets_values[key].options; + widget.value = info.widgets_values[key].value; + //if value exists in inputs + for (let input of this.inputs) + if (input.name == key){ + //find if key exists in inputs array in inputs.Name + if (info.widgets_values[key].type == "converted-widget"){ + input.type = info.widgets_values.type.value; + widget.type = "converted-widget" + widget.origType = info.widgets_values[key].origType; + widget.origComputeSize = undefined; + widget.last_y = info.widgets_values[key].last_y; + widget.origSerializeValue = nodeType.prototype.serializeValue; + }else{ + this.removeInput(this.inputs.indexOf(input)); + } + break; + } + } + + } + for(let i = this.outputs.length - 1; i>0; i--){ + if (this.outputs[i].name == "*"){ + this.removeOutput(i); + } + } + }); + chainCallback(this, "onSerialize", function(info) { + info.widgets_values = {}; + if (!this.widgets) { + return; + } + + for (let w of this.widgets) { + info.widgets_values[w.name] = {name: w.name, options : w.options, value: w.value, type: w.type, origType: w.origType, last_y: w.last_y}; + } + for (let w of this.inputs){ + // if w.name exists in info.widgets_values + if (info.widgets_values[w.name]){ + if(info.widgets_values[w.name].type == "converted-widget"){ + if(info.widgets_values[w.name].origType == "toggle"){ + w.type = "BOOLEAN"; + }else if(info.widgets_values[w.name].origType == "text"){ + w.type = "STRING"; + } + } + } + } + + }); + + this.widgets[1].callback = ( value ) => { + clearInputs(this); + switch(value){ + case "IMAGE": + this.addOutput("output", "IMAGE"); + this.addInput("default", "IMAGE"); + + break; + case "MASK": + this.addOutput("output", "MASK"); + this.addInput("default", "MASK"); + break; + case "STRING": + this.addOutput("output", "STRING"); + ComfyWidgets.STRING( + this, + "default", + ["STRING",{default: "",callback: (val) => console.log('VALUE', val),},], + app, + ) + break; + case "INT": + this.addOutput("output", "INT"); + ComfyWidgets.INT( + this, + "default", + ['',{default: 0,callback: (val) => console.log('VALUE', val),},], + app, + ) + break; + case "FLOAT": + this.addOutput("output", "FLOAT"); + ComfyWidgets.FLOAT( + this, + "default", + ['',{default: 0,callback: (val) => console.log('VALUE', val), "min": 0.00, "max": 1.00, "step": 0.01},], + app, + ) + break; + + case "BOOLEAN": + this.addOutput("output", "BOOLEAN"); + this.addWidget("toggle", "default", false, ()=>{}); + break; + case "LATENT": + this.addOutput("output", "LATENT"); + this.addInput("default", "LATENT"); + break; + case "MODEL": + this.addOutput("output", "MODEL"); + this.addInput("default", "MODEL"); + break; + case "CLIP": + this.addOutput("output", "CLIP"); + this.addInput("default", "CLIP"); + break; + case "CONDITIONING": + this.addOutput("output", "CONDITIONING"); + this.addInput("default", "CONDITIONING"); + break; + case "VAE": + this.addOutput("output", "VAE"); + this.addInput("default", "VAE"); + break; + } + this.color = colors[node_type_list.indexOf(value)]; + this.bgcolor = bg_colors[node_type_list.indexOf(value)]; + }; + this.color = colors[node_type_list.indexOf("none")]; + this.bgcolor = bg_colors[node_type_list.indexOf("none")]; + } + break; + case "WorkflowContinue": + nodeType.prototype.onNodeCreated = function() { + chainCallback(this, "onConfigure", function(info) { + let widgetDict = info.widgets_values + if (info.widgets_values.length == undefined) { + + for (let w of this.widgets) { + if (w.name in widgetDict) { + w.value = widgetDict[w.name].value; + } + } + // check if widgetDict in this.widgets + for (let [key, value] of Object.entries(widgetDict)) { + let widget = this.widgets.find(w => w.name === key); + let type = this.widgets.find(w => w.name === "type"); + if(!widget) + addWidgetInfo(this, key, value, app); + widget = this.widgets.find(w => w.name === key); + //this.widgets.push(value); + widget.options = info.widgets_values[key].options; + widget.value = info.widgets_values[key].value; + //if value exists in inputs + for (let input of this.inputs) + if (input.name == key){ + //find if key exists in inputs array in inputs.Name + if (info.widgets_values[key].type == "converted-widget"){ + input.type = info.widgets_values.type.value; + widget.type = "converted-widget" + widget.origType = info.widgets_values[key].origType; + widget.origComputeSize = undefined; + widget.last_y = info.widgets_values[key].last_y; + widget.origSerializeValue = nodeType.prototype.serializeValue; + }else{ + this.removeInput(this.inputs.indexOf(input)); + } + break; + } + } + + } + for(let i = this.outputs.length - 1; i>0; i--){ + if (this.outputs[i].name == "*"){ + this.removeOutput(i); + } + } + }); + chainCallback(this, "onSerialize", function(info) { + info.widgets_values = {}; + if (!this.widgets) { + return; + } + + for (let w of this.widgets) { + info.widgets_values[w.name] = {name: w.name, options : w.options, value: w.value, type: w.type, origType: w.origType, last_y: w.last_y}; + } + for (let w of this.inputs){ + // if w.name exists in info.widgets_values + if (info.widgets_values[w.name]){ + if(info.widgets_values[w.name].type == "converted-widget"){ + if(info.widgets_values[w.name].origType == "toggle"){ + w.type = "BOOLEAN"; + }else if(info.widgets_values[w.name].origType == "combo"){ + w.type = "COMBO"; + } + } + } + } + + }); + + this.widgets[0].callback = ( value ) => { + clearInputs(this); + switch(value){ + case "IMAGE": + this.addOutput("output", "IMAGE"); + this.addInput("input", "IMAGE"); + break; + case "LATENT": + this.addOutput("output", "LATENT"); + this.addInput("input", "LATENT"); + break; + } + this.color = colors[node_type_list.indexOf(value)]; + this.bgcolor = bg_colors[node_type_list.indexOf(value)]; + }; + this.color = colors[node_type_list.indexOf("none")]; + this.bgcolor = bg_colors[node_type_list.indexOf("none")]; + } + break; + case "WorkflowOutput": + nodeType.prototype.onNodeCreated = function() { + chainCallback(this, "onConfigure", function(info) { + let widgetDict = info.widgets_values + if (info.widgets_values.length == undefined) { + + for (let w of this.widgets) { + if (w.name in widgetDict) { + w.value = widgetDict[w.name].value; + } + } + // check if widgetDict in this.widgets + for (let [key, value] of Object.entries(widgetDict)) { + let widget = this.widgets.find(w => w.name === key); + let type = this.widgets.find(w => w.name === "type"); + if(!widget) + addWidgetInfo(this, key, value, app); + widget = this.widgets.find(w => w.name === key); + //this.widgets.push(value); + widget.options = info.widgets_values[key].options; + widget.value = info.widgets_values[key].value; + //if value exists in inputs + for (let input of this.inputs) + if (input.name == key){ + //find if key exists in inputs array in inputs.Name + if (info.widgets_values[key].type == "converted-widget"){ + input.type = info.widgets_values.type.value; + widget.type = "converted-widget" + widget.origType = info.widgets_values[key].origType; + widget.origComputeSize = undefined; + widget.last_y = info.widgets_values[key].last_y; + widget.origSerializeValue = nodeType.prototype.serializeValue; + }else{ + this.removeInput(this.inputs.indexOf(input)); + } + break; + } + } + + } + for(let i = this.outputs.length - 1; i>0; i--){ + if (this.outputs[i].name == "*"){ + this.removeOutput(i); + } + } + }); + chainCallback(this, "onSerialize", function(info) { + info.widgets_values = {}; + if (!this.widgets) { + return; + } + + for (let w of this.widgets) { + info.widgets_values[w.name] = {name: w.name, options : w.options, value: w.value, type: w.type, origType: w.origType, last_y: w.last_y}; + } + for (let w of this.inputs){ + // if w.name exists in info.widgets_values + if (info.widgets_values[w.name]){ + if(info.widgets_values[w.name].type == "converted-widget"){ + if(info.widgets_values[w.name].origType == "toggle"){ + w.type = "BOOLEAN"; + }else if(info.widgets_values[w.name].origType == "text"){ + w.type = "STRING"; + } + } + } + } + + }); + this.widgets[1].callback = ( value ) => { + clearInputs(this); + switch(value){ + case "IMAGE": + this.addOutput("output", "IMAGE"); + this.addInput("default", "IMAGE"); + break; + case "MASK": + this.addOutput("output", "MASK"); + this.addInput("default", "MASK"); + break; + case "STRING": + this.addOutput("output", "STRING"); + this.addInput("default","STRING"); + break; + case "INT": + this.addOutput("output", "INT"); + this.addInput("default","INT"); + break; + case "FLOAT": + this.addOutput("output", "FLOAT"); + this.addInput("default","FLOAT"); + break; + case "BOOLEAN": + this.addOutput("output", "BOOLEAN"); + this.addInput("default","BOOLEAN"); + break; + case "LATENT": + this.addOutput("output", "LATENT"); + this.addInput("default", "LATENT"); + break; + case "MODEL": + this.addOutput("output", "MODEL"); + this.addInput("default", "MODEL"); + break; + case "CLIP": + this.addOutput("output", "CLIP"); + this.addInput("default", "CLIP"); + break; + case "CONDITIONING": + this.addOutput("output", "CONDITIONING"); + this.addInput("default", "CONDITIONING"); + break; + case "VAE": + this.addOutput("output", "VAE"); + this.addInput("default", "VAE"); + break; + } + this.color = colors[node_type_list.indexOf(value)]; + this.bgcolor = bg_colors[node_type_list.indexOf(value)]; + }; + clearInputs(this); + this.color = colors[node_type_list.indexOf("none")]; + this.bgcolor = bg_colors[node_type_list.indexOf("none")]; + } + break; + case "WorkflowLipSync": + useKVState(nodeType); + chainCallback(nodeType.prototype, "onNodeCreated", function () { + let new_widgets = [] + if (this.widgets) { + for (let w of this.widgets) { + let input = this.constructor.nodeData.input + let config = input?.required[w.name] ?? input.optional[w.name] + if (!config) { + continue + } + if (w?.type == "text" && config[1].vhs_path_extensions) { + new_widgets.push(app.widgets.VHSPATH({}, w.name, ["VHSPATH", config[1]])); + } else { + new_widgets.push(w) + } + } + this.widgets = new_widgets; + } + }); + addLoadVideoCommon(nodeType, nodeData); + const onGetImageSizeExecuted = nodeType.prototype.onExecuted; + nodeType.prototype.onExecuted = function(message) { + const r = onGetImageSizeExecuted? onGetImageSizeExecuted.apply(this,arguments): undefined + let video = message["video_path"][0]; + if(video){ + this.updateParameters({format: "video/mp4", filename: message["video_path"][0], subfolder: message["video_path"][1], "type": "output"}); + } + return r + } + break; + } + }, + async init(app) { + api.fetchApi("/flowchain/workflows") + .then(response => response.json()) + .then(data => { + app.lipsync_studio = data + }) + .catch(error => { + console.error('Error:', error); + throw error; + }); + } +}); \ No newline at end of file diff --git a/workflow.py b/workflow.py new file mode 100644 index 0000000..cbe54a6 --- /dev/null +++ b/workflow.py @@ -0,0 +1,907 @@ +import json +import urllib.request +import urllib.parse +import torch +import logging +import time +import uuid +import traceback +import nodes +import copy +import asyncio +from enum import Enum +import numpy as np +import server +import hashlib +from torchvision import transforms +from .utils.logger import Logger +from .utils.utils import caches +from comfy_execution.graph import get_input_info, ExecutionList, DynamicPrompt, ExecutionBlocker +import comfy.model_management +import sys +from PIL import Image +from comfy_execution.graph_utils import is_link, GraphBuilder +from nodes import SaveImage +import gc + +class ExecutionResult(Enum): + SUCCESS = 0 + FAILURE = 1 + PENDING = 2 + + +class AnyType(str): + """A special class that is always equal in not equal comparisons. Credit to pythongosssss""" + + def __eq__(self, _) -> bool: + return True + + def __ne__(self, __value: object) -> bool: + return False + + +client_id = '5b49a023-b05a-4c53-8dc9-addc3a749911' +server_address = "127.0.0.1:8188" + + +def _map_node_over_list(obj, input_data_all, func, allow_interrupt=False, execution_block_cb=None, pre_execute_cb=None): + # check if node wants the lists + input_is_list = getattr(obj, "INPUT_IS_LIST", False) + + if len(input_data_all) == 0: + max_len_input = 0 + else: + max_len_input = max(len(x) for x in input_data_all.values()) + + # get a slice of inputs, repeat last input when list isn't long enough + def slice_dict(d, i): + return {k: v[i if len(v) > i else -1] for k, v in d.items()} + + results = [] + + def process_inputs(inputs, index=None): + if allow_interrupt: + nodes.before_node_execution() + execution_block = None + for k, v in inputs.items(): + if isinstance(v, ExecutionBlocker): + execution_block = execution_block_cb(v) if execution_block_cb else v + break + if execution_block is None: + if pre_execute_cb is not None and index is not None: + pre_execute_cb(index) + results.append(getattr(obj, func)(**inputs)) + else: + results.append(execution_block) + + if input_is_list: + process_inputs(input_data_all, 0) + elif max_len_input == 0: + process_inputs({}) + else: + for i in range(max_len_input): + input_dict = slice_dict(input_data_all, i) + process_inputs(input_dict, i) + return results + + +def merge_result_data(results, obj): + # check which outputs need concatenating + output = [] + output_is_list = [False] * len(results[0]) + if hasattr(obj, "OUTPUT_IS_LIST"): + output_is_list = obj.OUTPUT_IS_LIST + + # merge node execution results + for i, is_list in zip(range(len(results[0])), output_is_list): + if is_list: + output.append([x for o in results for x in o[i]]) + else: + output.append([o[i] for o in results]) + return output + + +def get_output_data(obj, input_data_all, execution_block_cb=None, pre_execute_cb=None): + results = [] + uis = [] + subgraph_results = [] + return_values = _map_node_over_list(obj, input_data_all, obj.FUNCTION, allow_interrupt=True, + execution_block_cb=execution_block_cb, pre_execute_cb=pre_execute_cb) + has_subgraph = False + for i in range(len(return_values)): + r = return_values[i] + if isinstance(r, dict): + if 'ui' in r: + uis.append(r['ui']) + if 'expand' in r: + # Perform an expansion, but do not append results + has_subgraph = True + new_graph = r['expand'] + result = r.get("result", None) + if isinstance(result, ExecutionBlocker): + result = tuple([result] * len(obj.RETURN_TYPES)) + subgraph_results.append((new_graph, result)) + elif 'result' in r: + result = r.get("result", None) + if isinstance(result, ExecutionBlocker): + result = tuple([result] * len(obj.RETURN_TYPES)) + results.append(result) + subgraph_results.append((None, result)) + else: + if isinstance(r, ExecutionBlocker): + r = tuple([r] * len(obj.RETURN_TYPES)) + results.append(r) + subgraph_results.append((None, r)) + + if has_subgraph: + output = subgraph_results + elif len(results) > 0: + output = merge_result_data(results, obj) + else: + output = [] + ui = dict() + if len(uis) > 0: + # ui = {k: [y for x in uis for y in x[k]] for k in uis[0].keys()} + for k in uis[0].keys(): + for x in uis: + ui[k] = x[k] + # ui = {k: uis[0]["images"] for k in uis[0].keys()} + return output, ui, has_subgraph + + +def get_input_data(inputs, class_def, unique_id, outputs=None, dynprompt=None, extra_data=None): + if extra_data is None: + extra_data = {} + valid_inputs = class_def.INPUT_TYPES() + input_data_all = {} + missing_keys = {} + for x in inputs: + input_data = inputs[x] + input_type, input_category, input_info = get_input_info(class_def, x) + + def mark_missing(): + missing_keys[x] = True + input_data_all[x] = (None,) + + if is_link(input_data) and (not input_info or not input_info.get("rawLink", False)): + input_unique_id = input_data[0] + output_index = input_data[1] + if outputs is None: + mark_missing() + continue # This might be a lazily-evaluated input + cached_output = outputs.get(input_unique_id) + if cached_output is None: + mark_missing() + continue + if output_index >= len(cached_output): + mark_missing() + continue + obj = cached_output[output_index] + input_data_all[x] = obj + elif input_category is not None: + input_data_all[x] = [input_data] + + if "hidden" in valid_inputs: + h = valid_inputs["hidden"] + for x in h: + if h[x] == "PROMPT": + input_data_all[x] = [dynprompt.get_original_prompt() if dynprompt is not None else {}] + if h[x] == "DYNPROMPT": + input_data_all[x] = [dynprompt] + if h[x] == "EXTRA_PNGINFO": + input_data_all[x] = [extra_data.get('extra_pnginfo', None)] + if h[x] == "UNIQUE_ID": + input_data_all[x] = [unique_id] + return input_data_all, missing_keys + + +def full_type_name(klass): + module = klass.__module__ + if module == 'builtins': + return klass.__qualname__ + return module + '.' + klass.__qualname__ + + +def format_value(x): + if x is None: + return None + elif isinstance(x, (int, float, bool, str)): + return x + else: + return str(x) + + +def executes(server, dynprompt, caches, current_item, extra_data, executed, prompt_id, execution_list, + pending_subgraph_results): + unique_id = current_item + real_node_id = dynprompt.get_real_node_id(unique_id) + display_node_id = dynprompt.get_display_node_id(unique_id) + parent_node_id = dynprompt.get_parent_node_id(unique_id) + inputs = dynprompt.get_node(unique_id)['inputs'] + class_type = dynprompt.get_node(unique_id)['class_type'] + class_def = nodes.NODE_CLASS_MAPPINGS[class_type] + if caches.outputs.get(unique_id) is not None: + if server.client_id is not None: + cached_output = caches.ui.get(unique_id) or {} + server.send_sync("executed", {"node": unique_id, "display_node": display_node_id, + "output": cached_output.get("output", None), "prompt_id": prompt_id}, + server.client_id) + return (ExecutionResult.SUCCESS, None, None) + + input_data_all = None + try: + if unique_id in pending_subgraph_results: + cached_results = pending_subgraph_results[unique_id] + resolved_outputs = [] + for is_subgraph, result in cached_results: + if not is_subgraph: + resolved_outputs.append(result) + else: + resolved_output = [] + for r in result: + if is_link(r): + source_node, source_output = r[0], r[1] + node_output = caches.outputs.get(source_node)[source_output] + for o in node_output: + resolved_output.append(o) + + else: + resolved_output.append(r) + resolved_outputs.append(tuple(resolved_output)) + output_data = merge_result_data(resolved_outputs, class_def) + output_ui = [] + has_subgraph = False + else: + input_data_all, missing_keys = get_input_data(inputs, class_def, unique_id, caches.outputs, dynprompt, + extra_data) + if server.client_id is not None: + server.last_node_id = display_node_id + server.send_sync("executing", + {"node": unique_id, "display_node": display_node_id, "prompt_id": prompt_id}, + server.client_id) + + obj = caches.objects.get(unique_id) + if obj is None: + obj = class_def() + caches.objects.set(unique_id, obj) + + if hasattr(obj, "check_lazy_status"): + required_inputs = _map_node_over_list(obj, input_data_all, "check_lazy_status", allow_interrupt=True) + required_inputs = set(sum([r for r in required_inputs if isinstance(r, list)], [])) + required_inputs = [x for x in required_inputs if isinstance(x, str) and ( + x not in input_data_all or x in missing_keys + )] + if len(required_inputs) > 0: + for i in required_inputs: + execution_list.make_input_strong_link(unique_id, i) + return (ExecutionResult.PENDING, None, None) + + def execution_block_cb(block): + if block.message is not None: + """mes = { + "prompt_id": prompt_id, + "node_id": unique_id, + "node_type": class_type, + "executed": list(executed), + + "exception_message": f"Execution Blocked: {block.message}", + "exception_type": "ExecutionBlocked", + "traceback": [], + "current_inputs": [], + "current_outputs": [], + }""" + """server.send_sync("execution_error", mes, server.client_id)""" + return ExecutionBlocker(None) + else: + return block + + def pre_execute_cb(call_index): + GraphBuilder.set_default_prefix(unique_id, call_index, 0) + + output_data, output_ui, has_subgraph = get_output_data(obj, input_data_all, + execution_block_cb=execution_block_cb, + pre_execute_cb=pre_execute_cb) + if len(output_ui) > 0: + caches.ui.set(unique_id, { + "meta": { + "node_id": unique_id, + "display_node": display_node_id, + "parent_node": parent_node_id, + "real_node_id": real_node_id, + }, + "output": output_ui + }) + if server.client_id is not None: + server.send_sync("executed", {"node": unique_id, "display_node": display_node_id, "output": output_ui, + "prompt_id": prompt_id}, server.client_id) + if has_subgraph: + cached_outputs = [] + new_node_ids = [] + new_output_ids = [] + new_output_links = [] + for i in range(len(output_data)): + new_graph, node_outputs = output_data[i] + if new_graph is None: + cached_outputs.append((False, node_outputs)) + else: + # Check for conflicts + + for node_id, node_info in new_graph.items(): + new_node_ids.append(node_id) + display_id = node_info.get("override_display_id", unique_id) + dynprompt.add_ephemeral_node(node_id, node_info, unique_id, display_id) + # Figure out if the newly created node is an output node + class_type = node_info["class_type"] + class_def = nodes.NODE_CLASS_MAPPINGS[class_type] + if hasattr(class_def, 'OUTPUT_NODE') and class_def.OUTPUT_NODE == True: + new_output_ids.append(node_id) + for i in range(len(node_outputs)): + if is_link(node_outputs[i]): + from_node_id, from_socket = node_outputs[i][0], node_outputs[i][1] + new_output_links.append((from_node_id, from_socket)) + cached_outputs.append((True, node_outputs)) + new_node_ids = set(new_node_ids) + for cache in caches.all: + cache.ensure_subcache_for(unique_id, new_node_ids).clean_unused() + for node_id in new_output_ids: + execution_list.add_node(node_id) + for link in new_output_links: + execution_list.add_strong_link(link[0], link[1], unique_id) + pending_subgraph_results[unique_id] = cached_outputs + return (ExecutionResult.PENDING, None, None) + caches.outputs.set(unique_id, output_data) + except comfy.model_management.InterruptProcessingException as iex: + logging.info("Processing interrupted") + + # skip formatting inputs/outputs + error_details = { + "node_id": real_node_id, + } + + return (ExecutionResult.FAILURE, error_details, iex) + except Exception as ex: + typ, _, tb = sys.exc_info() + exception_type = full_type_name(typ) + input_data_formatted = {} + if input_data_all is not None: + input_data_formatted = {} + for name, inputs in input_data_all.items(): + input_data_formatted[name] = [format_value(x) for x in inputs] + + logging.error(f"!!! Exception during processing !!! {ex}") + logging.error(traceback.format_exc()) + + error_details = { + "node_id": real_node_id, + "exception_message": str(ex), + "exception_type": exception_type, + "traceback": traceback.format_tb(tb), + "current_inputs": input_data_formatted + } + if isinstance(ex, comfy.model_management.OOM_EXCEPTION): + logging.error("Got an OOM, unloading all loaded models.") + comfy.model_management.unload_all_models() + + return (ExecutionResult.FAILURE, error_details, ex) + + executed.add(unique_id) + + return (ExecutionResult.SUCCESS, None, None) + + +class IsChangedCache: + def __init__(self, dynprompt, outputs_cache): + self.dynprompt = dynprompt + self.outputs_cache = outputs_cache + self.is_changed = {} + + def get(self, node_id): + if node_id in self.is_changed: + return self.is_changed[node_id] + + node = self.dynprompt.get_node(node_id) + class_type = node["class_type"] + class_def = nodes.NODE_CLASS_MAPPINGS[class_type] + if not hasattr(class_def, "IS_CHANGED"): + self.is_changed[node_id] = False + return self.is_changed[node_id] + + if "is_changed" in node: + self.is_changed[node_id] = node["is_changed"] + return self.is_changed[node_id] + + # Intentionally do not use cached outputs here. We only want constants in IS_CHANGED + input_data_all, _ = get_input_data(node["inputs"], class_def, node_id, None) + try: + is_changed = _map_node_over_list(class_def, input_data_all, "IS_CHANGED") + node["is_changed"] = [None if isinstance(x, ExecutionBlocker) else x for x in is_changed] + except Exception as e: + logging.warning("WARNING: {}".format(e)) + node["is_changed"] = float("NaN") + finally: + self.is_changed[node_id] = node["is_changed"] + return self.is_changed[node_id] + + +status_messages = [] + + +def add_message(servers, event, data: dict, broadcast: bool): + data = { + **data, + "timestamp": int(time.time() * 1000), + } + status_messages.append((event, data)) + """if servers.client_id is not None or broadcast: + servers.send_sync(event, data, servers.client_id)""" + + +def handle_execution_error(servers, prompt_id, prompt, current_outputs, executed, error, ex): + node_id = error["node_id"] + class_type = prompt[node_id]["class_type"] + + # First, send back the status to the frontend depending + # on the exception type + if isinstance(ex, comfy.model_management.InterruptProcessingException): + mes = { + "prompt_id": prompt_id, + "node_id": node_id, + "node_type": class_type, + "executed": list(executed), + } + add_message(servers, "execution_interrupted", mes, broadcast=True) + else: + mes = { + "prompt_id": prompt_id, + "node_id": node_id, + "node_type": class_type, + "executed": list(executed), + "exception_message": error["exception_message"], + "exception_type": error["exception_type"], + "traceback": error["traceback"], + "current_inputs": error["current_inputs"], + "current_outputs": list(current_outputs), + } + add_message(servers, "execution_error", mes, broadcast=False) + + +def execute(server, prompt, prompt_id, extra_data={}, execute_outputs=[]): + nodes.interrupt_processing(False) + + if "client_id" in extra_data: + server.client_id = extra_data["client_id"] + + status_messages = [] + add_message(server,"execution_start", {"prompt_id": prompt_id}, broadcast=False) + + with torch.inference_mode(): + dynamic_prompt = DynamicPrompt(prompt) + is_changed_cache = IsChangedCache(dynamic_prompt, caches.outputs) + for cache in caches.all: + cache.set_prompt(dynamic_prompt, prompt.keys(), is_changed_cache) + cache.clean_unused() + + cached_nodes = [] + for node_id in prompt: + if caches.outputs.get(node_id) is not None: + cached_nodes.append(node_id) + + comfy.model_management.cleanup_models(keep_clone_weights_loaded=True) + add_message(server, "execution_cached",{"nodes": cached_nodes, "prompt_id": prompt_id}, broadcast=False) + pending_subgraph_results = {} + executed = set() + execution_list = ExecutionList(dynamic_prompt, caches.outputs) + current_outputs = caches.outputs.all_node_ids() + for node_id in list(execute_outputs): + execution_list.add_node(node_id) + + while not execution_list.is_empty(): + node_id, error, ex = execution_list.stage_node_execution() + if error is not None: + handle_execution_error(server, prompt_id, dynamic_prompt.original_prompt, current_outputs, executed, + error, ex) + break + if "type" in prompt[node_id]["inputs"] and prompt[node_id]["inputs"]["type"] in ["IMAGE", "LATENT"]: + logging.info("node : {} {} image_count => {}".format(node_id, prompt[node_id]["class_type"], + len(prompt[node_id]["inputs"]["default"]))) + else: + logging.info( + "node : {} {} {}".format(node_id, prompt[node_id]["class_type"], prompt[node_id]["inputs"])) + + result, error, ex = executes(server, dynamic_prompt, caches, node_id, extra_data, executed, + prompt_id, execution_list, pending_subgraph_results) + success = result != ExecutionResult.FAILURE + if result == ExecutionResult.FAILURE: + handle_execution_error(server, prompt_id, dynamic_prompt.original_prompt, current_outputs, executed, + error, ex) + break + elif result == ExecutionResult.PENDING: + execution_list.unstage_node_execution() + else: # result == ExecutionResult.SUCCESS: + execution_list.complete_node_execution() + else: + # Only execute when the while-loop ends without break + #print("execution_success", prompt_id) + add_message(server, "execution_success", {"prompt_id": prompt_id}, broadcast=False) + + ui_outputs = {} + meta_outputs = {} + all_node_ids = caches.ui.all_node_ids() + for node_id in all_node_ids: + ui_info = caches.ui.get(node_id) + if ui_info is not None: + ui_outputs[node_id] = ui_info["output"] + meta_outputs[node_id] = ui_info["meta"] + history_result = {"outputs": ui_outputs, "meta": meta_outputs,} + for node_id in history_result["outputs"]: + for output in history_result["outputs"][node_id]: + if type(history_result["outputs"][node_id][output]) == torch.Tensor: + logging.info("output : {} {} image_count => {}".format(node_id, prompt[node_id]["class_type"], + len(history_result["outputs"][node_id][output]))) + elif len(str(history_result["outputs"][node_id][output])) > 100: + logging.info("output : {} {} {}".format(node_id, prompt[node_id]["class_type"], + str(history_result["outputs"][node_id][output])[:100])) + else: + logging.info("output : {} {}".format(node_id, history_result["outputs"][node_id][output])) + + server.last_node_id = None + """if comfy.model_management.DISABLE_SMART_MEMORY: + comfy.model_management.unload_all_models()""" + return history_result + + +def recursive_delete(workflow, to_delete): + # workflow_copy = copy.deepcopy(workflow) + new_delete = [] + for node_id in to_delete: + for node_id2, node in workflow.items(): + for input_name, input_value in node["inputs"].items(): + if type(input_value) == list: + if len(input_value) > 0: + if input_value[0] == node_id: + new_delete.append(node_id2) + if node_id in workflow: + del workflow[node_id] + if len(new_delete) > 0: + workflow = recursive_delete(workflow, new_delete) + return workflow + + +class Workflow(SaveImage): + def __init__(self): + self.logger = Logger() + self.ws = None + + @classmethod + def INPUT_TYPES(cls): + return { + + "hidden": { + "workflows": ("STRING", {"default": ""}) + }} + + RETURN_TYPES = ( + AnyType("*"), AnyType("*"), AnyType("*"), AnyType("*"), AnyType("*"), AnyType("*"), AnyType("*"), AnyType("*"), + AnyType("*"), AnyType("*"), AnyType("*"), AnyType("*"), AnyType("*"), AnyType("*"), AnyType("*"), AnyType("*"), + ) + FUNCTION = "generate" + CATEGORY = "FlowChain ⛓️" + + OUTPUT_NODE = True + + @classmethod + def IS_CHANGED(s, workflows, **kworgs): + m = hashlib.sha256() + m.update(workflows.encode()) + return m.digest().hex() + + def generate(self, workflows, **kwargs): + # get current file path + + def get_workflow(workflow_name): + with urllib.request.urlopen( + "http://{}/flowchain/workflow?workflow_path={}".format(server_address, workflow_name)) as response: + workflow = json.loads(response.read()) + return workflow["workflow"] + + def populate_inputs(workflow, inputs, kwargs_values): + workflow_inputs = {k: v for k, v in workflow.items() if v["class_type"] == "WorkflowInput"} + for key, value in workflow_inputs.items(): + if value["inputs"]["Name"] in inputs: + if type(inputs[value["inputs"]["Name"]]) == list: + if value["inputs"]["Name"] in kwargs_values: + workflow[key]["inputs"]["default"] = kwargs_values[value["inputs"]["Name"]] + else: + workflow[key]["inputs"]["default"] = inputs[value["inputs"]["Name"]] + + workflow_inputs_images = {k: v for k, v in workflow.items() if + v["class_type"] == "WorkflowInput" and v["inputs"]["type"] == "IMAGE"} + for key, value in workflow_inputs_images.items(): + if "default" not in value["inputs"]: + workflow[key]["inputs"]["default"] = torch.tensor([]) + else: + if value["inputs"]["default"] == []: + workflow[key]["inputs"]["default"] = torch.tensor([]) + return workflow + + def treat_switch(workflow): + to_delete = [] + #do_net_delete = [] + switch_to_delete = [-1] + while len(switch_to_delete) > 0: + switch_nodes = {k: v for k, v in workflow.items() if + v["class_type"].startswith("Switch") and v["class_type"].endswith("[Crystools]")} + # order switch nodes by inputs.boolean value + switch_to_delete = [] + switch_nodes_copy = copy.deepcopy(switch_nodes) + for switch_id, switch_node in switch_nodes.items(): + # create list of inputs who have switch in their inputs + """inputs_from_switch = {node_id: node for node_id, node in workflow.items() if any( + input_value[0] == switch_id for input_value in node["inputs"].values() if type(input_value) == list)}""" + inputs_from_switch = [] + for node_ids, node in workflow.items(): + for input_name, input_value in node["inputs"].items(): + if type(input_value) == list: + if len(input_value) > 0: + if input_value[0] == switch_id: + inputs_from_switch.append({node_ids: input_name}) + # convert to dictionary + inputs_from_switch = {k: v for d in inputs_from_switch for k, v in d.items()} + switch = switch_nodes_copy[switch_id] + for node_id, input_name in inputs_from_switch.items(): + if type(switch["inputs"]["boolean"]) == list: + switch_boolean_value = workflow[switch["inputs"]["boolean"][0]]["inputs"] + + other_input_name = None + if "default" in switch_boolean_value: + other_input_name = "default" + elif "boolean" in switch_boolean_value: + other_input_name = "boolean" + + if other_input_name is not None: + if switch_boolean_value[other_input_name] == True: + if type(switch["inputs"]["on_true"]) == list: + workflow[node_id]["inputs"][input_name] = switch["inputs"]["on_true"] + if node_id in switch_nodes_copy: + switch_nodes_copy[node_id]["inputs"][input_name] = switch["inputs"]["on_true"] + else: + to_delete.append(node_id) + else: + if type(switch["inputs"]["on_false"]) == list: + workflow[node_id]["inputs"][input_name] = switch["inputs"]["on_false"] + if node_id in switch_nodes_copy: + switch_nodes_copy[node_id]["inputs"][input_name] = switch["inputs"]["on_false"] + else: + to_delete.append(node_id) + switch_to_delete.append(switch_id) + else: + if switch["inputs"]["boolean"] == True: + if type(switch["inputs"]["on_true"]) == list: + workflow[node_id]["inputs"][input_name] = switch["inputs"]["on_true"] + if node_id in switch_nodes_copy: + switch_nodes_copy[node_id]["inputs"][input_name] = switch["inputs"]["on_true"] + else: + to_delete.append(node_id) + else: + if type(switch["inputs"]["on_false"]) == list: + workflow[node_id]["inputs"][input_name] = switch["inputs"]["on_false"] + if node_id in switch_nodes_copy: + switch_nodes_copy[node_id]["inputs"][input_name] = switch["inputs"]["on_false"] + else: + to_delete.append(node_id) + switch_to_delete.append(switch_id) + print(switch_to_delete) + workflow = {k: v for k, v in workflow.items() if + not (v["class_type"].startswith("Switch") and v["class_type"].endswith( + "[Crystools]") and k in switch_to_delete)} + + return workflow, to_delete + + def treat_continue(workflow): + to_delete = [] + continue_nodes = {k: v for k, v in workflow.items() if + v["class_type"].startswith("WorkflowContinue")} + do_net_delete = [] + for continue_node_id, continue_node in continue_nodes.items(): + for node_id, node in workflow.items(): + for input_name, input_value in node["inputs"].items(): + if type(input_value) == list: + if len(input_value) > 0: + if input_value[0] == continue_node_id: + if type(continue_node["inputs"]["continue_workflow"]) == list: + input_other_node = \ + workflow[continue_node["inputs"]["continue_workflow"][0]][ + "inputs"] + other_input_name = None + if "default" in input_other_node: + other_input_name = "default" + elif "boolean" in input_other_node: + other_input_name = "boolean" + + if other_input_name is not None: + if input_other_node[other_input_name]: + workflow[node_id]["inputs"][input_name] = continue_node["inputs"]["input"] + else: + to_delete.append(node_id) + else: + do_net_delete.append(continue_node_id) + else: + if continue_node["inputs"]["continue_workflow"]: + workflow[node_id]["inputs"][input_name] = continue_node["inputs"]["input"] + else: + to_delete.append(node_id) + + workflow = {k: v for k, v in workflow.items() if + not (v["class_type"].startswith("WorkflowContinue") and k not in do_net_delete)} + return workflow, to_delete + + def redefine_id(subworkflow, max_id): + new_sub_workflow = {} + + for k, v in subworkflow.items(): + max_id += 1 + new_sub_workflow[str(max_id)] = v + # replace old id by new id items in inputs of workflow + for node_id, node in subworkflow.items(): + for input_name, input_value in node["inputs"].items(): + if type(input_value) == list: + if len(input_value) > 0: + if input_value[0] == k: + subworkflow[node_id]["inputs"][input_name][0] = str(max_id) + for node_id, node in new_sub_workflow.items(): + for input_name, input_value in node["inputs"].items(): + if type(input_value) == list: + if len(input_value) > 0: + if input_value[0] == k: + new_sub_workflow[node_id]["inputs"][input_name][0] = str(max_id) + return new_sub_workflow, max_id + + def change_subnode(subworkflow, node_id_to_find, value): + for node_id, node in subworkflow.items(): + for input_name, input_value in node["inputs"].items(): + if type(input_value) == list: + if len(input_value) > 0: + if input_value[0] == node_id_to_find: + subworkflow[node_id]["inputs"][input_name] = value + + return subworkflow + + def merge_inputs_outputs(workflow, workflow_name, subworkflow, workflow_outputs): + # get max workflow id + # coinvert workflow_outputs to list + workflow_outputs = list(workflow_outputs.values()) + workflow_node = {"node": {"id":k, **v} for k, v in workflow.items() if v["class_type"] == "Workflow" and v["inputs"]["workflows"] == workflow_name} + sub_input_nodes = {k: v for k, v in subworkflow.items() if v["class_type"] == "WorkflowInput"} + do_not_delete = [] + for sub_id, sub_node in sub_input_nodes.items(): + if sub_node["inputs"]["Name"] in workflow_node["node"]["inputs"]: + value = workflow_node["node"]["inputs"][sub_node["inputs"]["Name"]] + if type(value) == list: + subworkflow = change_subnode(subworkflow, sub_id, value) + else: + subworkflow[sub_id]["inputs"]["default"] = value + do_not_delete.append(sub_id) + + # remove input node + subworkflow = {k: v for k, v in subworkflow.items() if not (v["class_type"] == "WorkflowInput" and k not in do_not_delete)} + + sub_output_nodes = {k: v for k, v in subworkflow.items() if v["class_type"] == "WorkflowOutput"} + workflow_copy = copy.deepcopy(workflow) + for node_id, node in workflow_copy.items(): + for input_name, input_value in node["inputs"].items(): + if type(input_value) == list: + if len(input_value) > 0: + if input_value[0] == workflow_node["node"]["id"]: + for sub_output_id, sub_output_node in sub_output_nodes.items(): + if sub_output_node["inputs"]["Name"] == workflow_outputs[input_value[1]]["inputs"]["Name"]: + workflow[node_id]["inputs"][input_name] = sub_output_node["inputs"]["default"] + + # remove output node + subworkflow = {k: v for k, v in subworkflow.items() if not (v["class_type"] == "WorkflowOutput")} + + return workflow, subworkflow + + def clean_workflow(workflow, inputs=None, kwargs_values=None): + if kwargs_values is None: + kwargs_values = {} + if inputs is None: + inputs = {} + if inputs is not None: + workflow = populate_inputs(workflow, inputs, kwargs_values) + + workflow_outputs = {k: v for k, v in workflow.items() if v["class_type"] == "WorkflowOutput"} + + for output_id, output_node in workflow_outputs.items(): + workflow[output_id]["inputs"]["ui"] = False + + workflow, switch_to_delete = treat_switch(workflow) + workflow, continue_to_delete = treat_continue(workflow) + workflow = recursive_delete(workflow, switch_to_delete + continue_to_delete) + return workflow, workflow_outputs + + def get_recursive_workflow(workflows, max_id=0): + workflow = get_workflow(workflows) + workflow, max_id = redefine_id(workflow, max_id) + sub_workflows = {k: v for k, v in workflow.items() if v["class_type"] == "Workflow"} + for key, sub_workflow_node in sub_workflows.items(): + workflow_name = sub_workflow_node["inputs"]["workflows"] + subworkflow, max_id = get_recursive_workflow(workflow_name, max_id) + + #subworkflow = get_workflow(workflow_name) + #max_id = max([int(k) for k in workflow.keys() if k.isdigit()]) + + # change all id in subworkflow + #subworkflow = redefine_id(subworkflow["workflow"], max_id) + workflow_outputs_sub = {k: v for k, v in subworkflow.items() if v["class_type"] == "WorkflowOutput"} + workflow, subworkflow = merge_inputs_outputs(workflow, workflow_name, subworkflow, workflow_outputs_sub) + # sub_workflow, workflow_outputs_sub = treat_workflow(subworkflow) + workflow = {k: v for k, v in workflow.items() if + not (v["class_type"] == "Workflow" and v["inputs"]["workflows"] == workflow_name)} + # add subworkflow to workflow + workflow.update(subworkflow) + return workflow, max_id + + with urllib.request.urlopen("http://{}/queue".format(server_address)) as response: + queue_info = json.loads(response.read()) + + original_inputs = [v["inputs"] for k, v in queue_info["queue_running"][0][2].items() if + "workflows" in v["inputs"] and v["inputs"]["workflows"] == workflows][0] + + workflow, _ = get_recursive_workflow(workflows, 5000) + workflow, workflow_outputs = clean_workflow(workflow, original_inputs, kwargs) + workflow_outputs_id = [k for k, v in workflow.items() if v["class_type"] == "WorkflowOutput"] + + prompt_id = str(uuid.uuid4()) + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + servers = server.PromptServer(loop) + + servers.last_prompt_id = prompt_id + servers.client_id = client_id + execution_start_time = time.perf_counter() + logging.info("workflow : {}".format(workflows)) + history_result = execute(servers, workflow, prompt_id, {}, workflow_outputs_id) + current_time = time.perf_counter() + execution_time = current_time - execution_start_time + logging.info("Prompt executed in {:.2f} seconds".format(execution_time)) + comfy.model_management.unload_all_models() + del servers + gc.collect() + + output = [] + for id_node, node in workflow_outputs.items(): + if id_node in history_result["outputs"]: + mask = history_result["outputs"][id_node]["default"] + # create hash from mask + node name + """hash = hashlib.sha256(mask + hash = hash.update(node["inputs"]["Name"].encode()) + filename_prefix = node["inputs"]["Name"]+"/"+hash + if node["inputs"]["type"] == "IMAGE": + self.save_images(history_result["outputs"][id_node]["default"], filename_prefix) + elif node["inputs"]["type"] == "MASK": + preview = mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])).movedim(1, -1).expand(-1, -1, -1, 3) + self.save_images(preview, filename_prefix)""" + output.append(history_result["outputs"][id_node]["default"]) + else: + if node["inputs"]["type"] == "IMAGE" or node["inputs"]["type"] == "MASK": + black_image_np = np.zeros((255, 255, 3), dtype=np.uint8) + black_image_pil = Image.fromarray(black_image_np) + transform = transforms.ToTensor() + image_tensor = transform(black_image_pil) + image_tensor = image_tensor.permute(1, 2, 0) + image_tensor = image_tensor.unsqueeze(0) + output.append(image_tensor) + else: + output.append(None) + + return tuple(output) + # return tuple(queue[uid]["outputs"]) + + +NODE_CLASS_MAPPINGS_WORKFLOW = { + "Workflow": Workflow, +} + +NODE_DISPLAY_NAME_MAPPINGS_WORKFLOW = { + "Workflow": "Workflow (FlowChain ⛓️)", +} diff --git a/workflow_nodes.py b/workflow_nodes.py new file mode 100644 index 0000000..0b61d90 --- /dev/null +++ b/workflow_nodes.py @@ -0,0 +1,425 @@ +import torch +import numpy as np +from PIL import Image +import hashlib +from torchvision import transforms + + +class AnyType(str): + """A special class that is always equal in not equal comparisons. Credit to pythongosssss""" + + def __eq__(self, _) -> bool: + return True + + def __ne__(self, __value: object) -> bool: + return False + + +BOOLEAN = ("BOOLEAN", {"default": True}) +STRING = ("STRING", {"default": ""}) +any_input = AnyType("*") +node_type_list = ["none", "IMAGE", "MASK", "STRING", "INT", "FLOAT", "LATENT", "BOOLEAN", "CLIP", "CONDITIONING", "MODEL", "VAE"] +""" +class WorkflowOutputImage: + def __init__(self): + self.prompt_id = None + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "Name": STRING, + "default": ("IMAGE", {"default": []}) + }, + "hidden": { + "ui": BOOLEAN + } + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("images",) + FUNCTION = "execute" + OUTPUT_NODE = True + CATEGORY = "LipSync Studio 🎀" + + def execute(self, Name, default, ui=True): + if ui: + if default is None: + return (torch.tensor([]),) + return (default,) + else: + if default is None: + black_image_np = np.zeros((255, 255, 3), dtype=np.uint8) + black_image_pil = Image.fromarray(black_image_np) + transform = transforms.ToTensor() + image_tensor = transform(black_image_pil) + image_tensor = image_tensor.permute(1, 2, 0) + image_tensor = image_tensor.unsqueeze(0) + return {"ui": {"images": image_tensor}} + return {"ui": {"images": default}} + + +class WorkflowInputImage: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "Name": STRING, + "default": ("IMAGE", {"default": []}) + }, + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("images",) + FUNCTION = "execute" + CATEGORY = "LipSync Studio 🎀" + + def execute(self, Name, default): + # get current file path + return (default,) + + +class WorkflowInputString: + def __init__(self): + self.prompt_id = None + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "Name": STRING, + "default": STRING + } + } + + RETURN_TYPES = ("STRING",) + RETURN_NAMES = ("string",) + FUNCTION = "execute" + CATEGORY = "LipSync Studio 🎀" + + def execute(self, Name, default): + return (default,) + + +class WorkflowInputBoolean: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "Name": STRING, + "default": ("BOOLEAN", {"default": False}) + } + } + + RETURN_TYPES = ("BOOLEAN",) + RETURN_NAMES = ("boolean",) + FUNCTION = "execute" + CATEGORY = "LipSync Studio 🎀" + + def execute(self, Name, default): + return (default,) + + +class WorkflowInputInteger: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "Name": STRING, + "default": ("INT", {"default": 0}) + } + } + + RETURN_TYPES = ("INT",) + RETURN_NAMES = ("int",) + FUNCTION = "execute" + CATEGORY = "LipSync Studio 🎀" + + def execute(self, Name, default): + return (default,) + + +class WorkflowInputFloat: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "Name": STRING, + "default": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}) + } + } + + RETURN_TYPES = ("FLOAT",) + RETURN_NAMES = ("float",) + FUNCTION = "execute" + CATEGORY = "LipSync Studio 🎀" + + def execute(self, Name, default): + return (default,) + + +class WorkflowInputSwitch: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "Name": STRING, + "images": ("IMAGE", {"default": []}), + "default": BOOLEAN, + } + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("images",) + FUNCTION = "execute" + CATEGORY = "LipSync Studio 🎀" + + def execute(self, Name, images, default): + if default: + return (images,) + else: + return (images[0].unsqueeze(0),) + + +class WorkflowContinueImage: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "input": ("IMAGE", {"default": []}), + "continue_workflow": BOOLEAN, + } + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("output",) + FUNCTION = "execute" + CATEGORY = "LipSync Studio 🎀" + + @classmethod + def IS_CHANGED(s, input, continue_workflow): + m = hashlib.sha256() + if input is None: + return "0" + else: + m.update(input.encode()+str(continue_workflow).encode()) + return m.digest().hex() + + def execute(self, input, continue_workflow): + print("WorkflowContinue", continue_workflow) + if continue_workflow: + return (input,) + else: + return (input[0].unsqueeze(0),) + + +class WorkflowContinueLatent: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "input": ("LATENT", {"default": []}), + "continue_workflow": BOOLEAN, + } + } + + RETURN_TYPES = ("LATENT",) + RETURN_NAMES = ("output",) + FUNCTION = "execute" + CATEGORY = "LipSync Studio 🎀" + + @classmethod + def IS_CHANGED(s, input, continue_workflow): + m = hashlib.sha256() + m.update(input.encode()+str(continue_workflow).encode()) + return m.digest().hex() + + def execute(self, input, continue_workflow): + print("WorkflowContinue", continue_workflow) + if continue_workflow: + return (input,) + else: + ret = {"samples": input["samples"][0].unsqueeze(0)} + if "noise_mask" in input: + ret["noise_mask"] = input["noise_mask"][0].unsqueeze(0) + return (ret,) +""" + +class WorkflowContinue: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "input": ("IMAGE", {"default": []}), + "type": ( + ["none", "IMAGE", "LATENT"],), + "continue_workflow": BOOLEAN, + } + } + + RETURN_TYPES = (AnyType("*"),) + RETURN_NAMES = ("output",) + FUNCTION = "execute" + CATEGORY = "FlowChain ⛓️" + + @classmethod + def IS_CHANGED(s, input, type, continue_workflow): + m = hashlib.sha256() + if input is None: + return "0" + else: + m.update(input.encode()+str(continue_workflow).encode()) + return m.digest().hex() + + def execute(self, input, type, continue_workflow): + print("WorkflowContinue", continue_workflow) + if continue_workflow: + if type == "LATENT": + ret = {"samples": input["samples"][0].unsqueeze(0)} + if "noise_mask" in input: + ret["noise_mask"] = input["noise_mask"][0].unsqueeze(0) + return (ret,) + else: + return (input,) + else: + return (input[0].unsqueeze(0),) + + +class WorkflowInput: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return {"required": { + "Name": STRING, + "type": (node_type_list,), + "default": ("*",) + }} + + RETURN_TYPES = (AnyType("*"),) + RETURN_NAMES = ("output",) + FUNCTION = "execute" + CATEGORY = "FlowChain ⛓️" + #OUTPUT_NODE = True + + @classmethod + def IS_CHANGED(s, Name, type,default, **kwargs): + m = hashlib.sha256() + if default is not None: + m.update(str(default).encode()) + else: + m.update(Name.encode()+type.encode()) + return m.digest().hex() + + def execute(self, Name, type, default, **kwargs): + """if type == "SWITCH": + if "boolean" in kwargs: + if kwargs["boolean"]: + return (kwargs["default"],) + else: + return (kwargs["default"][0].unsqueeze(0),) + else: + return (kwargs["default"],) + else:""" + return (default,) + + +class WorkflowOutput: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return {"required": { + "Name": STRING, + "type": (node_type_list,) + }, + "hidden": { + "ui": BOOLEAN + }} + + RETURN_TYPES = (AnyType("*"),) + RETURN_NAMES = ("output",) + FUNCTION = "execute" + CATEGORY = "FlowChain ⛓️" + OUTPUT_NODE = True + + @classmethod + def IS_CHANGED(s, Name, type, ui=True, **kwargs): + m = hashlib.sha256() + m.update(Name.encode()+type.encode()) + return m.digest().hex() + + def execute(self, Name, type, ui=True, **kwargs): + if ui: + if kwargs["default"] is None: + return (torch.tensor([]),) + return (kwargs["default"],) + else: + if type in ["IMAGE", "MASK"]: + if kwargs["default"] is None: + black_image_np = np.zeros((255, 255, 3), dtype=np.uint8) + black_image_pil = Image.fromarray(black_image_np) + transform = transforms.ToTensor() + image_tensor = transform(black_image_pil) + image_tensor = image_tensor.permute(1, 2, 0) + image_tensor = image_tensor.unsqueeze(0) + return {"ui": {"default": image_tensor}} + return {"ui": {"default": kwargs["default"]}} + elif type == "LATENT": + if kwargs["default"] is None: + return {"ui": {"default": torch.tensor([])}} + return {"ui": {"default": kwargs["default"]}} + else: + ui = {"ui": {}} + ui["ui"]["default"] = kwargs["default"] + return ui + + + +NODE_CLASS_MAPPINGS_NODES = { + "WorkflowInput": WorkflowInput, + "WorkflowOutput": WorkflowOutput, + + #"WorkflowInputImage": WorkflowInputImage, + #"WorkflowInputString": WorkflowInputString, + #"WorkflowInputBoolean": WorkflowInputBoolean, + #"WorkflowInputInteger": WorkflowInputInteger, + #"WorkflowInputFloat": WorkflowInputFloat, + #"WorkflowOutputImage": WorkflowOutputImage, + #"WorkflowInputSwitch": WorkflowInputSwitch, + #"WorkflowContinueImage": WorkflowContinueImage, + #"WorkflowContinueLatent": WorkflowContinueLatent, + "WorkflowContinue": WorkflowContinue, + +} + +# A dictionary that contains the friendly/humanly readable titles for the nodes +NODE_DISPLAY_NAME_MAPPINGS_NODES = { + "WorkflowInput": "Workflow Input (FlowChain ⛓️)", + "WorkflowOutput": "Workflow Output (FlowChain ⛓️)", + #"WorkflowInputImage": "Workflow Input Image (Lipsync Studio)", + #"WorkflowInputString": "Workflow Input String (Lipsync Studio)", + #"WorkflowInputBoolean": "Workflow Input Boolean (Lipsync Studio)", + #"WorkflowInputInteger": "Workflow Input Integer (Lipsync Studio)", + #"WorkflowInputFloat": "Workflow Input Float (Lipsync Studio)", + #"WorkflowOutputImage": "Workflow Output Image (Lipsync Studio)", + #"WorkflowInputSwitch": "Workflow Input Switch (Lipsync Studio)", + #"WorkflowContinueImage": "Workflow Continue Image (Lipsync Studio)", + #"WorkflowContinueLatent": "Workflow Continue Latent (Lipsync Studio)", + "WorkflowContinue": "Workflow Continue (FlowChain ⛓️)", + # "VisualizeOpticalFlow": "Visualize optical flow", +}