Initial commit for Comfyui-FlowChain
@@ -0,0 +1,3 @@
|
||||
docs/assets/demo.gif
|
||||
.git/*
|
||||
**/__pycache__/
|
||||
@@ -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!
|
||||
@@ -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.
|
||||
@@ -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)
|
||||
|
||||
<img src="docs/assets/demo.gif" width="100%">
|
||||
|
||||
## 🚀 All Nodes
|
||||
|
||||
<img src="docs/assets/allnodes.png" width="100%">
|
||||
|
||||
## 📖 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 |
|
||||
|:-------------------------------------------------:|:--------------------|:------------------------------------------------------------------------------------------------------------:|:----------------:|
|
||||
| <img src="docs/assets/workflow.png" width="100%"> | _Workflow_ | Node that allows loading workflows in API format. It will show Inputs and Outputs into the loaded Workflows | FlowChain ⛓️ |
|
||||
| <img src="docs/assets/Input.png" width="100%"> | _Workflow Input_ | Node used to declare the inputs of your workflows. | FlowChain ⛓️ |
|
||||
| <img src="docs/assets/output.png" width="100%"> | _Workflow Output_ | Node used to declare the outputs of your workflows. | FlowChain ⛓️ |
|
||||
| <img src="docs/assets/Continue.png" width="100%"> | _Workflow Continue_ | Node to stop/Continue the workflow process. | FlowChain ⛓️ |
|
||||
| <img src="docs/assets/lipsync.png" width="100%"> | _Workflow Lipsync_ | Extra Node to use LipSync Studio via API | FlowChain ⛓️ |
|
||||
|
||||
|
||||
## 📺 Tutorial
|
||||
- [Here](https://youtu.be/B84A5alpPDc)
|
||||
|
||||
# 🐍 Usage
|
||||
|
||||
## ⛓️ Workflow Node
|
||||

|
||||
|
||||
- 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.
|
||||
|
||||
<img src="docs/assets/save_as_api.png">
|
||||
|
||||
If you don't see **"Export (API format)"** options in Comfyui do this :
|
||||
- go to Settings
|
||||
- Activate the **Dev Mode** options
|
||||
|
||||

|
||||
|
||||
-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
|
||||

|
||||
|
||||
- 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.
|
||||
|
||||
- 
|
||||
|
||||
## ⛓️ Output Node
|
||||

|
||||
|
||||
- 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.
|
||||
|
||||

|
||||
|
||||
## ⛓️ Continue Node
|
||||

|
||||
|
||||
- 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.
|
||||
|
||||

|
||||
|
||||
- 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.
|
||||
|
||||

|
||||
|
||||
## 🔉👄 Workflow LipSync Node
|
||||

|
||||
|
||||
- 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.
|
||||
|
||||

|
||||
|
||||
## 💪 Special things to know
|
||||
|
||||
the **"🪛 Switch"** nodes from [Crystools](https://github.com/crystian/ComfyUI-Crystools) have a particular place in **workflow Node**
|
||||
|
||||

|
||||
|
||||
Let's illustrate this with an example:
|
||||
|
||||

|
||||
|
||||
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.
|
||||
|
||||

|
||||
|
||||
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).
|
||||
@@ -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"
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
from .logger import *
|
||||
from .keys import *
|
||||
from .types import *
|
||||
from .config import *
|
||||
from .common import *
|
||||
from .version import *
|
||||
@@ -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"
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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"
|
||||
@@ -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)
|
||||
@@ -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("*")
|
||||
@@ -0,0 +1 @@
|
||||
version = "1.15.0"
|
||||
|
After Width: | Height: | Size: 13 KiB |
|
After Width: | Height: | Size: 7.5 KiB |
|
After Width: | Height: | Size: 109 KiB |
|
After Width: | Height: | Size: 22 KiB |
|
After Width: | Height: | Size: 71 KiB |
|
After Width: | Height: | Size: 230 KiB |
|
After Width: | Height: | Size: 26 KiB |
|
After Width: | Height: | Size: 144 KiB |
|
After Width: | Height: | Size: 20 KiB |
|
After Width: | Height: | Size: 6.5 KiB |
|
After Width: | Height: | Size: 22 KiB |
|
After Width: | Height: | Size: 30 KiB |
|
After Width: | Height: | Size: 25 KiB |
|
After Width: | Height: | Size: 517 KiB |
|
After Width: | Height: | Size: 578 KiB |
|
After Width: | Height: | Size: 12 KiB |
|
After Width: | Height: | Size: 22 KiB |
|
After Width: | Height: | Size: 12 KiB |
|
After Width: | Height: | Size: 1.2 KiB |
|
After Width: | Height: | Size: 267 KiB |
|
After Width: | Height: | Size: 272 KiB |
|
After Width: | Height: | Size: 39 KiB |
|
After Width: | Height: | Size: 22 KiB |
|
After Width: | Height: | Size: 21 KiB |
|
After Width: | Height: | Size: 42 KiB |
|
After Width: | Height: | Size: 45 KiB |
@@ -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 ⛓️)",
|
||||
}
|
||||
@@ -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
|
||||
"""
|
||||
@@ -0,0 +1 @@
|
||||
gradio_client==0.8.0
|
||||
@@ -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
|
||||
@@ -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")
|
||||
@@ -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)
|
||||
@@ -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 ⛓️)",
|
||||
}
|
||||
@@ -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",
|
||||
}
|
||||