diff --git a/__init__.py b/__init__.py index 62fffa3..a7e64d5 100644 --- a/__init__.py +++ b/__init__.py @@ -6,14 +6,17 @@ except ImportError: else: 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 NODE_CLASS_MAPPINGS.update(NetNodes) from .nodes.advanced import NODE_CLASS_MAPPINGS as 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()} __all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] diff --git a/nodes/advanced.py b/nodes/advanced.py new file mode 100644 index 0000000..4f89dfe --- /dev/null +++ b/nodes/advanced.py @@ -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, +} diff --git a/nodes/workflows.py b/nodes/workflows.py new file mode 100644 index 0000000..1195229 --- /dev/null +++ b/nodes/workflows.py @@ -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, +}