Restore chain arch

This commit is contained in:
City
2024-01-04 03:42:45 +01:00
parent 4c668c91f0
commit 9486d855b5
3 changed files with 233 additions and 3 deletions
+6 -3
View File
@@ -6,14 +6,17 @@ except ImportError:
else: else:
NODE_CLASS_MAPPINGS = {} NODE_CLASS_MAPPINGS = {}
from .nodes.images import NODE_CLASS_MAPPINGS as ImgNodes
NODE_CLASS_MAPPINGS.update(ImgNodes)
from .nodes.simple import NODE_CLASS_MAPPINGS as NetNodes from .nodes.simple import NODE_CLASS_MAPPINGS as NetNodes
NODE_CLASS_MAPPINGS.update(NetNodes) NODE_CLASS_MAPPINGS.update(NetNodes)
from .nodes.advanced import NODE_CLASS_MAPPINGS as AdvNodes from .nodes.advanced import NODE_CLASS_MAPPINGS as AdvNodes
NODE_CLASS_MAPPINGS.update(AdvNodes) NODE_CLASS_MAPPINGS.update(AdvNodes)
from .nodes.images import NODE_CLASS_MAPPINGS as ImgNodes
NODE_CLASS_MAPPINGS.update(ImgNodes)
from .nodes.workflows import NODE_CLASS_MAPPINGS as WrkNodes
NODE_CLASS_MAPPINGS.update(WrkNodes)
NODE_DISPLAY_NAME_MAPPINGS = {k:v.TITLE for k,v in NODE_CLASS_MAPPINGS.items()} NODE_DISPLAY_NAME_MAPPINGS = {k:v.TITLE for k,v in NODE_CLASS_MAPPINGS.items()}
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] __all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
+116
View File
@@ -0,0 +1,116 @@
from ..core.utils import clean_url, get_client_id, get_new_job_id
from ..core.dispatch import dispatch_to_remote, clear_remote_queue
class RemoteChainStart:
"""Merge required attributes into one [REMCHAIN]"""
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"workflow": ("JSON",),
"trigger": (["on_change", "always"],),
"batch": ("INT", {"default": 1, "min": 1, "max": 8}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
}
}
RETURN_TYPES = ("REMCHAIN",)
RETURN_NAMES = ("remote_chain",)
FUNCTION = "chain_start"
CATEGORY = "remote/advanced"
TITLE = "Queue on remote (start of chain)"
def chain_start(self, workflow, trigger, batch, seed):
remote_chain = {
"seed": seed,
"batch": batch,
"prompt": workflow,
"seed_offset": batch,
"job_id": get_new_job_id(),
}
return(remote_chain,)
@classmethod
def IS_CHANGED(self, workflow, trigger, batch, seed, prompt):
uuid = f"W:{workflow},B:{batch},S:{seed}"
return uuid if trigger == "on_change" else str(time.time())
class RemoteChainEnd:
"""Split [REMCHAIN] into local seed/batch"""
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"remote_chain": ("REMCHAIN",)
}
}
RETURN_TYPES = ("INT", "INT")
RETURN_NAMES = ("seed", "batch")
FUNCTION = "chain_end"
CATEGORY = "remote/advanced"
TITLE = "Queue on remote (end of chain)"
def chain_end(self, remote_chain):
seed = remote_chain["seed"]
batch = remote_chain["batch"]
return(seed,batch)
class RemoteQueueWorker:
"""Start job on remote worker"""
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"remote_chain": ("REMCHAIN",),
"remote_url": ("STRING", {
"multiline": False,
"default": "http://127.0.0.1:8288/",
}),
"batch_override": ("INT", {"default": 0, "min": 0, "max": 8}),
"enabled": (["true", "false", "remote"],{"default": "true"}),
}
}
RETURN_TYPES = ("REMCHAIN", "REMINFO")
RETURN_NAMES = ("remote_chain", "remote_info")
FUNCTION = "queue"
CATEGORY = "remote/advanced"
TITLE = "Queue on remote (worker)"
def queue(self, remote_chain, remote_url, batch_override, enabled):
current_offset = remote_chain["seed_offset"]
remote_chain["seed_offset"] += 1 if batch_override == 0 else batch_override
if enabled == "false":
return (remote_chain, {})
if enabled == "remote":
# apply offset from previous nodes in chain
remote_chain["seed"] += current_offset
if batch_override > 0:
remote_chain["batch"] = batch_override
return (remote_chain, {})
remote_url = clean_url(remote_url)
clear_remote_queue(remote_url)
dispatch_to_remote(
remote_url,
remote_chain["prompt"],
remote_chain["job_id"]
)
remote_info = {
"remote_url" : remote_url,
"job_id" : remote_chain["job_id"],
}
return (remote_chain, remote_info)
NODE_CLASS_MAPPINGS = {
"RemoteChainStart" : RemoteChainStart,
"RemoteQueueWorker" : RemoteQueueWorker,
"RemoteChainEnd" : RemoteChainEnd,
}
+111
View File
@@ -0,0 +1,111 @@
import os
import json
import hashlib
import folder_paths
class SaveDiskWorkflowJSON:
"""Save workflow to disk"""
def __init__(self):
self.output_dir = folder_paths.get_output_directory()
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"workflow": ("JSON", ),
"filename_prefix": ("STRING", {"default": "workflow/ComfyUI"}),
}
}
RETURN_TYPES = ()
FUNCTION = "save_workflow"
OUTPUT_NODE = True
CATEGORY = "remote/advanced"
TITLE = "Save workflow (disk)"
def save_workflow(self, workflow, filename_prefix):
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(filename_prefix, self.output_dir)
json_path = os.path.join(full_output_folder, f"{filename}_{counter:05}_.json")
with open(json_path, "w") as f:
f.write(json.dumps(workflow, indent=2))
return {}
class LoadDiskWorkflowJSON:
"""Load workflow JSON from disk"""
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
input_dir = folder_paths.get_input_directory()
files = [f for f in os.listdir(input_dir) if os.path.isfile(os.path.join(input_dir, f)) and f.endswith(".json")]
return {
"required": {
"workflow": [sorted(files),],
}
}
RETURN_TYPES = ("JSON",)
RETURN_NAMES = ("Workflow JSON",)
FUNCTION = "load_workflow"
CATEGORY = "remote/advanced"
TITLE = "Load workflow (disk)"
def load_workflow(self, workflow):
json_path = folder_paths.get_annotated_filepath(cond)
with open(json_path) as f:
data = json.loads(f.read())
return (data,)
@classmethod
def IS_CHANGED(s, workflow):
json_path = folder_paths.get_annotated_filepath(workflow)
m = hashlib.sha256()
with open(json_path, 'rb') as f:
m.update(f.read())
return m.digest().hex()
@classmethod
def VALIDATE_INPUTS(s, workflow):
if not folder_paths.exists_annotated_filepath(workflow):
return "Invalid JSON file: {}".format(workflow)
json_path = folder_paths.get_annotated_filepath(workflow)
with open(json_path) as f:
try: json.loads(f.read())
except:
return "Failed to read JSON file: {}".format(workflow)
return True
class LoadCurrentWorkflowJSON:
"""Fetch the current workflow/prompt as an API compatible JSON"""
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {},
"hidden": {
"prompt": "PROMPT",
},
}
RETURN_TYPES = ("JSON",)
RETURN_NAMES = ("Workflow JSON",)
FUNCTION = "load_workflow"
CATEGORY = "remote/advanced"
TITLE = "Load workflow (current)"
def load_workflow(self, prompt):
return (prompt,)
@classmethod
def IS_CHANGED(s, prompt):
return hashlib.sha256(json.dumps(prompt)).digest().hex()
NODE_CLASS_MAPPINGS = {
"SaveDiskWorkflowJSON": SaveDiskWorkflowJSON,
"LoadDiskWorkflowJSON": LoadDiskWorkflowJSON,
"LoadCurrentWorkflowJSON": LoadCurrentWorkflowJSON,
}