From 4c668c91f01ae0214a8841c9c5324fcbd5364c68 Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Thu, 4 Jan 2024 00:28:15 +0100 Subject: [PATCH 1/5] Rewrite part 1 Let's try this again --- __init__.py | 55 +++----- core/dispatch.py | 123 +++++++++++++++++ core/fetch.py | 56 ++++++++ core/utils.py | 23 ++++ nodes/{remote_images.py => images.py} | 38 +++++- nodes/misc.py | 28 ---- nodes/remote_control.py | 184 -------------------------- nodes/simple.py | 94 +++++++++++++ 8 files changed, 348 insertions(+), 253 deletions(-) create mode 100644 core/dispatch.py create mode 100644 core/fetch.py create mode 100644 core/utils.py rename nodes/{remote_images.py => images.py} (72%) delete mode 100644 nodes/misc.py delete mode 100644 nodes/remote_control.py create mode 100644 nodes/simple.py diff --git a/__init__.py b/__init__.py index 49d8dc3..62fffa3 100644 --- a/__init__.py +++ b/__init__.py @@ -1,44 +1,19 @@ -NODE_CLASS_MAPPINGS = {} -NODE_DISPLAY_NAME_MAPPINGS = {} +# only import if running as a custom node +try: + import comfy.utils +except ImportError: + pass +else: + NODE_CLASS_MAPPINGS = {} -def remote_control(): - global NODE_CLASS_MAPPINGS - global NODE_DISPLAY_NAME_MAPPINGS - from .nodes.remote_control import QueueRemote, FetchRemote - NODE_CLASS_MAPPINGS.update({ - "QueueRemote": QueueRemote, - "FetchRemote": FetchRemote, - }) - NODE_DISPLAY_NAME_MAPPINGS.update({ - "QueueRemote": "Queue on remote", - "FetchRemote": "Fetch from remote", - }) + from .nodes.images import NODE_CLASS_MAPPINGS as ImgNodes + NODE_CLASS_MAPPINGS.update(ImgNodes) -def remote_images(): - global NODE_CLASS_MAPPINGS - global NODE_DISPLAY_NAME_MAPPINGS - from .nodes.remote_images import LoadImageUrl, SaveImageUrl - NODE_CLASS_MAPPINGS.update({ - "LoadImageUrl": LoadImageUrl, - "SaveImageUrl": SaveImageUrl, - }) - NODE_DISPLAY_NAME_MAPPINGS.update({ - "LoadImageUrl": "Load Image (URL)", - "SaveImageUrl": "Save Image (URL)", - }) + from .nodes.simple import NODE_CLASS_MAPPINGS as NetNodes + NODE_CLASS_MAPPINGS.update(NetNodes) -def remote_misc(): - global NODE_CLASS_MAPPINGS - global NODE_DISPLAY_NAME_MAPPINGS - from .nodes.misc import CombineImageBatch - NODE_CLASS_MAPPINGS.update({ - "CombineImageBatch": CombineImageBatch, - }) - NODE_DISPLAY_NAME_MAPPINGS.update({ - "CombineImageBatch": "Combine images", - }) + from .nodes.advanced import NODE_CLASS_MAPPINGS as AdvNodes + NODE_CLASS_MAPPINGS.update(AdvNodes) -print("Loading network distribution node pack") -remote_control() -remote_images() -remote_misc() + 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/core/dispatch.py b/core/dispatch.py new file mode 100644 index 0000000..1eb1522 --- /dev/null +++ b/core/dispatch.py @@ -0,0 +1,123 @@ +import os +import time +import json +import torch +import random +import requests +import numpy as np +from PIL import Image +from copy import deepcopy + +from .utils import clean_url, get_client_id + +def clear_remote_queue(remote_url): + r = requests.get(f"{remote_url}/queue", timeout=4) + r.raise_for_status() + queue = r.json() + + to_cancel = [] + client_id = get_client_id() + for k in queue.get("queue_pending", []): + if k[3].get("client_id") == client_id: + to_cancel.append(k[1]) # job UUID + r = requests.post( + f"{remote_url}/queue", + json = {"delete" : to_cancel}, + timeout = 4, + ) + r.raise_for_status() + + for k in queue.get("queue_running", []): + if k[3].get("client_id") == client_id: + r = requests.post( + f"{remote_url}/interrupt", + json = {}, + timeout = 4, + ) + r.raise_for_status() + break + +def get_remote_os(remote_url): + url = f"{remote_url}/system_stats" + r = requests.get(url) + r.raise_for_status() + data = r.json() + return data["system"]["os"] + +def dispatch_to_remote(remote_url, prompt, job_id=f"{get_client_id()}-unknown"): + ### PROMPT LOGIC ### + prompt = deepcopy(prompt) + to_del = [] + def recursive_node_deletion(start_node): + target_nodes = [start_node] + if start_node not in to_del: + to_del.append(start_node) + while len(target_nodes) > 0: + new_targets = [] + for target in target_nodes: + for node in prompt.keys(): + inputs = prompt[node].get("inputs") + if not inputs: + continue + for i in inputs.values(): + if type(i) == list: + if len(i) > 0 and i[0] in to_del: + if node not in to_del: + to_del.append(node) + new_targets.append(node) + target_nodes += new_targets + target_nodes.remove(target) + + # find current node and disable all others + output_src = None + for i in prompt.keys(): + if prompt[i]["class_type"].startswith("RemoteQueue"): + if clean_url(prompt[i]["inputs"]["remote_url"]) == remote_url: + prompt[i]["inputs"]["enabled"] = "remote" + output_src = i + else: + prompt[i]["inputs"]["enabled"] = "false" + + output = None + for i in prompt.keys(): + # only leave current fetch but replace with PreviewImage + if prompt[i]["class_type"] == "FetchRemote": + if prompt[i]["inputs"]["remote_info"][0] == output_src: + output = { + "inputs": {"images": prompt[i]["inputs"]["final_image"]}, + "class_type": 'PreviewImage', + } + recursive_node_deletion(i) + # do not save output on remote + # todo: other output types + if prompt[i]["class_type"] in ["SaveImage", "PreviewImage"]: + recursive_node_deletion(i) + prompt[str(max([int(x) for x in prompt.keys()])+1)] = output + for i in to_del: del prompt[i] + + ### OS LOGIC ### + sep_remote = "\\" if get_remote_os(remote_url) == "nt" else "/" + sep_local = "\\" if os.name == "nt" else "/" + sem_input_map = { # class type : input to replace + "CheckpointLoaderSimple" : "ckpt_name", + "CheckpointLoader" : "ckpt_name", + "LoraLoader" : "lora_name", + "VAELoader" : "vae_name", + } + if sep_remote != sep_local: + for i in prompt.keys(): + if prompt[i]["class_type"] in sem_input_map.keys(): + key = sem_input_map[prompt[i]["class_type"]] + prompt[i]["inputs"][key] = prompt[i]["inputs"][key].replace(sep_local, sep_remote) + + ### SEND REQUEST ### + data = { + "prompt": prompt, + "client_id": get_client_id(), + "extra_data": { + "job_id": job_id, + } + } + ar = requests.post(f"{remote_url}/prompt", json=data, timeout=4) + ar.raise_for_status() + return diff --git a/core/fetch.py b/core/fetch.py new file mode 100644 index 0000000..cd2100c --- /dev/null +++ b/core/fetch.py @@ -0,0 +1,56 @@ +import time +import json +import torch +import requests +import numpy as np +from PIL import Image + +POLLING = 0.5 + +def wait_for_job(remote_url, job_id): + fail = 0 + while fail <= 3: + r = requests.get(f"{remote_url}/history", timeout=4) + try: + r.raise_for_status() + except Exception as e: + print("NetDist caught error while fetching output image:\n", e) + fail += 1 + continue + data = r.json() + if not data: + time.sleep(POLLING) + continue + for i,d in data.items(): + if d["prompt"][3].get("job_id") == job_id: + return d["outputs"][list(d["outputs"].keys())[-1]].get("images") + # todo: check if it's actually in the queue to avoid waiting forever + time.sleep(POLLING) + raise OSError("Failed to fetch image from remote client!") + +def fetch_from_remote(remote_url, job_id): + def img_to_torch(img): + image = img.convert("RGB") + image = np.array(image).astype(np.float32) / 255.0 + image = torch.from_numpy(image)[None,] + return image + + if not remote_url or not job_id: + return None + + images = [] + for i in wait_for_job(remote_url, job_id): + img_url = f"{remote_url}/view?filename={i['filename']}&subfolder={i['subfolder']}&type={i['type']}" + + ir = requests.get(img_url, stream=True, timeout=16) + ir.raise_for_status() + img = Image.open(ir.raw) + images.append(img_to_torch(img)) + + if len(images) == 0: + return None + + out = images[0] + for i in images[1:]: + out = torch.cat((out,i)) + return out diff --git a/core/utils.py b/core/utils.py new file mode 100644 index 0000000..8002a37 --- /dev/null +++ b/core/utils.py @@ -0,0 +1,23 @@ +import time +import random + +# set global ID once for entire session +try: GID +except NameError: + GID = ''.join(random.choice("abcdefghijklmnopqrstupvxyz") for x in range(5)) + print(f"NetDist: Set session ID to '{GID}'") + +def get_client_id(): + global GID + return(f"netdist-{GID}") + +def get_new_job_id(): + job_id = f"{get_client_id()}-{int(time.time()*1000)}" + time.sleep(0.1) # prevent ID mismatch, no matter how unlikely + return job_id + +def clean_url(raw, multi=False): + raw = raw.strip() + raw = raw.replace(' ', ',').replace('\n', ',').replace('\t', ',') + urls = [x.rstrip('/') for x in raw.split(',') if x.strip()] + return urls if multi else urls[0] diff --git a/nodes/remote_images.py b/nodes/images.py similarity index 72% rename from nodes/remote_images.py rename to nodes/images.py index ca8d32a..50ea73b 100644 --- a/nodes/remote_images.py +++ b/nodes/images.py @@ -23,6 +23,7 @@ class LoadImageUrl: RETURN_TYPES = ("IMAGE", "MASK") FUNCTION = "load_image_url" CATEGORY = "remote" + TITLE = "Load Image (URL)" def load_image_url(self, url): with requests.get(url, stream=True) as r: @@ -58,7 +59,8 @@ class SaveImageUrl: OUTPUT_NODE = True FUNCTION = "save_images" CATEGORY = "remote" - + TITLE = "Save Image (URL)" + def save_images(self, images, url, data_format, filename_prefix="ComfyUI", prompt=None, extra_pnginfo=None): filename = os.path.basename(os.path.normpath(filename_prefix)) @@ -86,3 +88,37 @@ class SaveImageUrl: with requests.post(url, json=data) as r: r.raise_for_status() return () + +class CombineImageBatch: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "images_a": ("IMAGE",), + "images_b": ("IMAGE",), + } + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("images",) + FUNCTION = "combine_images" + CATEGORY = "remote/misc" + TITLE = "Combine images" + + def combine_images(self,images_a,images_b): + try: + out = torch.cat((images_a,images_b), 0) + except RuntimeError: + print(f"Imagine size mismatch! {images_a.size()}, {images_b.size()}") + out = images_a + return (out,) + + +NODE_CLASS_MAPPINGS = { + "LoadImageUrl" : LoadImageUrl, + "SaveImageUrl" : SaveImageUrl, + "CombineImageBatch" : CombineImageBatch, +} diff --git a/nodes/misc.py b/nodes/misc.py deleted file mode 100644 index 6e74af7..0000000 --- a/nodes/misc.py +++ /dev/null @@ -1,28 +0,0 @@ -import torch -import torchvision - -class CombineImageBatch: - def __init__(self): - pass - - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "images_a": ("IMAGE",), - "images_b": ("IMAGE",), - } - } - - RETURN_TYPES = ("IMAGE",) - RETURN_NAMES = ("images",) - FUNCTION = "combine_images" - CATEGORY = "remote" - - def combine_images(self,images_a,images_b): - try: - out = torch.cat((images_a,images_b), 0) - except RuntimeError: - print(f"Imagine size mismatch! {images_a.size()}, {images_b.size()}") - out = images_a - return (out,) diff --git a/nodes/remote_control.py b/nodes/remote_control.py deleted file mode 100644 index bfd15de..0000000 --- a/nodes/remote_control.py +++ /dev/null @@ -1,184 +0,0 @@ -import time -import json -import torch -import random -import requests -import numpy as np -from PIL import Image -from copy import deepcopy - -def img_to_torch(img): - image = img.convert("RGB") - image = np.array(image).astype(np.float32) / 255.0 - image = torch.from_numpy(image)[None,] - return image - -class FetchRemote(): - def __init__(self): - pass - - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "final_image": ("IMAGE",), - "remote_info": ("REMOTE",), - }, - } - - RETURN_TYPES = ("IMAGE",) - FUNCTION = "get_remote_job" - CATEGORY = "remote" - - def wait_for_job(self,remote_url,job_id): - url = remote_url + "history" - - image_data = None - while not image_data: - r = requests.get(url) - r.raise_for_status() - data = r.json() - if not data: - time.sleep(0.5) - continue - for i,d in data.items(): - if d["prompt"][3].get("job_id") == job_id: - image_data = d["outputs"][list(d["outputs"].keys())[-1]].get("images") - time.sleep(0.5) - return image_data - - # remote_info can be none, but the node shouldn't exist at that point - def get_remote_job(self, final_image, remote_info): - images = [] - for i in self.wait_for_job(remote_info["remote_url"],remote_info["job_id"]): - img_url = f"{remote_info['remote_url']}view?filename={i['filename']}&subfolder={i['subfolder']}&type={i['type']}" - - ir = requests.get(img_url, stream=True) - ir.raise_for_status() - img = Image.open(ir.raw) - images.append(img_to_torch(img)) - - if len(images) == 0: - img = Image.new(mode="RGB", size=(768, 768)) - images.append(img_to_torch(img)) - - out = images[0] - for i in images[1:]: - out = torch.cat((out,i)) - - return (out,) - -class QueueRemote: - def __init__(self): - pass - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "remote_url": ("STRING", { - "multiline": False, - "default": "http://127.0.0.1:8188/", - }), - "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), - "batch_override": ("INT", {"default": 0, "min": 0, "max": 8}), - "system": (["windows", "posix"],), - "enabled": (["true", "false", "remote"],{"default": "true"}), - "node_id": ("INT", {"default": 1, "min": 1, "max": 64}), - }, - "hidden": { - "prompt": "PROMPT", - }, - } - - RETURN_TYPES = ("INT","REMOTE") - RETURN_NAMES = ("seed+batch","remote_info") - FUNCTION = "queue_on_remote" - CATEGORY = "remote" - - def queue_on_remote(self, remote_url, seed, batch_override, system, enabled, node_id, prompt): - def get_max_batch_size(prompt): - bs = 1 - for node in prompt.keys(): - if prompt[node]["class_type"] == "EmptyLatentImage": - for k,v in prompt[node]["inputs"].items(): - if k == "batch_size": - if batch_override > 0: - bs = batch_override - prompt[node]["inputs"]["batch_size"] = batch_override - else: - bs = max(bs,int(v)) - return bs - - if enabled == "false": - return (seed,{"remote_url":None,"job_id":None,}) - elif enabled == "remote": - return (seed+node_id*get_max_batch_size(prompt),{"remote_url":None,"job_id":None,}) - - new_prompt = deepcopy(prompt) - to_del = [] - def recursive_node_deletion(start_node): - target_nodes = [start_node] - if start_node not in to_del: - to_del.append(start_node) - while len(target_nodes) > 0: - new_targets = [] - for target in target_nodes: - for node in new_prompt.keys(): - inputs = new_prompt[node].get("inputs") - if not inputs: - continue - for i in inputs.values(): - if type(i) == list: - if len(i) > 0 and i[0] in to_del: - if node not in to_del: - to_del.append(node) - new_targets.append(node) - target_nodes += new_targets - target_nodes.remove(target) - - output_src = None - for i in new_prompt.keys(): - if new_prompt[i]["class_type"] == "QueueRemote": - if new_prompt[i]["inputs"]["remote_url"] == remote_url: - new_prompt[i]["inputs"]["enabled"] = "remote" - output_src = i - else: - new_prompt[i]["inputs"]["enabled"] = "false" - - output = None - for i in new_prompt.keys(): - # only leave current fetch but replace with PreviewImage - if new_prompt[i]["class_type"] == "FetchRemote": - if new_prompt[i]["inputs"]["remote_info"][0] == output_src: - output = { - 'inputs': {'images': new_prompt[i]["inputs"]["final_image"]}, - 'class_type': 'PreviewImage', - } - recursive_node_deletion(i) - # do not save output on remote - if new_prompt[i]["class_type"] in ["SaveImage","PreviewImage"]: - recursive_node_deletion(i) - new_prompt[str(max([int(x) for x in new_prompt.keys()])+1)] = output - - if system == "posix": - for i in new_prompt.keys(): - if new_prompt[i]["class_type"] == "LoraLoader": - new_prompt[i]["inputs"]["lora_name"] = new_prompt[i]["inputs"]["lora_name"].replace("\\","/") - if new_prompt[i]["class_type"] == "VAELoader": - new_prompt[i]["inputs"]["vae_name"] = new_prompt[i]["inputs"]["vae_name"].replace("\\","/") - if new_prompt[i]["class_type"] in ["CheckpointLoader","CheckpointLoaderSimple"]: - new_prompt[i]["inputs"]["ckpt_name"] = new_prompt[i]["inputs"]["ckpt_name"].replace("\\","/") - for i in to_del: - del new_prompt[i] - - job_id = f"netdist-{time.time()}" - data = { - "prompt": new_prompt, - "client_id": "netdist", - "extra_data": { - "job_id": job_id, - } - } - ar = requests.post(remote_url+"prompt", json=data) - ar.raise_for_status() - return (seed,{"remote_url":remote_url,"job_id":job_id}) diff --git a/nodes/simple.py b/nodes/simple.py new file mode 100644 index 0000000..ea626ea --- /dev/null +++ b/nodes/simple.py @@ -0,0 +1,94 @@ +from ..core.fetch import fetch_from_remote +from ..core.utils import clean_url, get_client_id, get_new_job_id +from ..core.dispatch import dispatch_to_remote, clear_remote_queue + +class FetchRemote(): + """ + Try to retrieve the final output image from the remote client. + On the remote client, this is replaced with a preview image node. + I.e. remote_info can be none, but the node shouldn't exist at that point + """ + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "final_image": ("IMAGE",), + "remote_info": ("REMINFO",), + }, + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "fetch" + CATEGORY = "remote" + TITLE = "Fetch from remote" + + def fetch(self, final_image, remote_info): + out = fetch_from_remote( + remote_url = remote_info.get("remote_url"), + job_id = remote_info.get("job_id"), + ) + if out is None: + out = final_image[:1] * 0.0 # black image + return (out,) + +class RemoteQueueSimple(): + """ + This is a "simplified" version without any extra controls. + """ + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "remote_url": ("STRING", { + "multiline": False, + "default": "http://127.0.0.1:8288/", + }), + "batch_local": ("INT", {"default": 1, "min": 1, "max": 8}), + "batch_remote": ("INT", {"default": 1, "min": 1, "max": 8}), + "trigger": (["on_change", "always"],), + "enabled": (["true", "false", "remote"],{"default": "true"}), + "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), + }, + "hidden": { + "prompt": "PROMPT", + }, + } + + RETURN_TYPES = ("INT", "INT", "REMINFO",) + RETURN_NAMES = ("seed", "batch", "remote_info",) + FUNCTION = "queue" + CATEGORY = "remote" + TITLE = "Queue on remote (single)" + + def queue(self, remote_url, batch_local, batch_remote, trigger, enabled, seed, prompt): + if enabled == "false": + return (seed, batch_local, {}) + if enabled == "remote": + return (seed+batch_local, batch_remote, {}) + + job_id = get_new_job_id() + remote_url = clean_url(remote_url) + clear_remote_queue(remote_url) + dispatch_to_remote(remote_url, prompt, job_id) + + remote_info = { + "remote_url" : remote_url, + "job_id" : job_id, + } + return (seed, batch_local, remote_info) + + @classmethod + def IS_CHANGED(self, remote_url, batch_local, batch_remote, trigger, enabled, seed, prompt): + uuid = f"W:{remote_url},B1:{batch_local},B2:{batch_remote},S:{seed},E:{enabled}" + return uuid if trigger == "on_change" else str(time.time()) + +NODE_CLASS_MAPPINGS = { + "FetchRemote" : FetchRemote, + "RemoteQueueSimple" : RemoteQueueSimple, +} From 9486d855b54afd682701565091ac1ebb5cdb8e06 Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Thu, 4 Jan 2024 03:42:45 +0100 Subject: [PATCH 2/5] Restore chain arch --- __init__.py | 9 ++-- nodes/advanced.py | 116 +++++++++++++++++++++++++++++++++++++++++++++ nodes/workflows.py | 111 +++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 233 insertions(+), 3 deletions(-) create mode 100644 nodes/advanced.py create mode 100644 nodes/workflows.py 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, +} From a4eb4b09f50b6d856669cdd4f5079275394db25b Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Thu, 4 Jan 2024 04:13:04 +0100 Subject: [PATCH 3/5] Static workflow --- core/dispatch.py | 20 ++++++++++++++++---- core/fetch.py | 6 +++++- nodes/advanced.py | 6 ++++-- nodes/workflows.py | 2 +- 4 files changed, 26 insertions(+), 8 deletions(-) diff --git a/core/dispatch.py b/core/dispatch.py index 1eb1522..661b001 100644 --- a/core/dispatch.py +++ b/core/dispatch.py @@ -39,12 +39,22 @@ def clear_remote_queue(remote_url): def get_remote_os(remote_url): url = f"{remote_url}/system_stats" - r = requests.get(url) + r = requests.get(url, timeout=4) r.raise_for_status() data = r.json() return data["system"]["os"] -def dispatch_to_remote(remote_url, prompt, job_id=f"{get_client_id()}-unknown"): +def get_output_nodes(remote_url): + # I'm 90% sure this could just use the + # list from the host but better safe than sorry + url = f"{remote_url}/object_info" + r = requests.get(url, timeout=4) + r.raise_for_status() + data = r.json() + out = [k for k, v in data.items() if v.get("output_node")] + return out + +def dispatch_to_remote(remote_url, prompt, job_id=f"{get_client_id()}-unknown", outputs="final_image"): ### PROMPT LOGIC ### prompt = deepcopy(prompt) to_del = [] @@ -78,6 +88,7 @@ def dispatch_to_remote(remote_url, prompt, job_id=f"{get_client_id()}-unknown"): else: prompt[i]["inputs"]["enabled"] = "false" + banned = [] if outputs == "any" else get_output_nodes(remote_url) output = None for i in prompt.keys(): # only leave current fetch but replace with PreviewImage @@ -90,9 +101,10 @@ def dispatch_to_remote(remote_url, prompt, job_id=f"{get_client_id()}-unknown"): recursive_node_deletion(i) # do not save output on remote # todo: other output types - if prompt[i]["class_type"] in ["SaveImage", "PreviewImage"]: + if prompt[i]["class_type"] in banned: recursive_node_deletion(i) - prompt[str(max([int(x) for x in prompt.keys()])+1)] = output + if output: + prompt[str(max([int(x) for x in prompt.keys()])+1)] = output for i in to_del: del prompt[i] ### OS LOGIC ### diff --git a/core/fetch.py b/core/fetch.py index cd2100c..96ed4ae 100644 --- a/core/fetch.py +++ b/core/fetch.py @@ -23,7 +23,11 @@ def wait_for_job(remote_url, job_id): continue for i,d in data.items(): if d["prompt"][3].get("job_id") == job_id: - return d["outputs"][list(d["outputs"].keys())[-1]].get("images") + # this needs to be less jank + if len(d["outputs"].keys()) > 0: + return d["outputs"][list(d["outputs"].keys())[-1]].get("images") + else: + return [] # todo: check if it's actually in the queue to avoid waiting forever time.sleep(POLLING) raise OSError("Failed to fetch image from remote client!") diff --git a/nodes/advanced.py b/nodes/advanced.py index 4f89dfe..18f258a 100644 --- a/nodes/advanced.py +++ b/nodes/advanced.py @@ -75,6 +75,7 @@ class RemoteQueueWorker: }), "batch_override": ("INT", {"default": 0, "min": 0, "max": 8}), "enabled": (["true", "false", "remote"],{"default": "true"}), + "outputs": (["final_image", "any"],{"default":"final_image"}), } } @@ -84,7 +85,7 @@ class RemoteQueueWorker: CATEGORY = "remote/advanced" TITLE = "Queue on remote (worker)" - def queue(self, remote_chain, remote_url, batch_override, enabled): + def queue(self, remote_chain, remote_url, batch_override, enabled, outputs): current_offset = remote_chain["seed_offset"] remote_chain["seed_offset"] += 1 if batch_override == 0 else batch_override if enabled == "false": @@ -101,7 +102,8 @@ class RemoteQueueWorker: dispatch_to_remote( remote_url, remote_chain["prompt"], - remote_chain["job_id"] + remote_chain["job_id"], + outputs, ) remote_info = { "remote_url" : remote_url, diff --git a/nodes/workflows.py b/nodes/workflows.py index 1195229..b0ed73d 100644 --- a/nodes/workflows.py +++ b/nodes/workflows.py @@ -53,7 +53,7 @@ class LoadDiskWorkflowJSON: TITLE = "Load workflow (disk)" def load_workflow(self, workflow): - json_path = folder_paths.get_annotated_filepath(cond) + json_path = folder_paths.get_annotated_filepath(workflow) with open(json_path) as f: data = json.loads(f.read()) return (data,) From d1d419b3f3f1381fddeb468d55a0367423d80bae Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Thu, 4 Jan 2024 04:36:44 +0100 Subject: [PATCH 4/5] Update README.md --- README.md | 66 ++++++++++++++++++++++++++++++++++++++----------------- 1 file changed, 46 insertions(+), 20 deletions(-) diff --git a/README.md b/README.md index 4b00f34..e8cc3cc 100644 --- a/README.md +++ b/README.md @@ -1,7 +1,7 @@ # ComfyUI_NetDist Run ComfyUI workflows on multiple local GPUs/networked machines. -Also includes code to utilize in a render farm (save/load images to/from a server). +[NetDist_2xspeed.webm](https://github.com/city96/ComfyUI_NetDist/assets/125218114/b7ec2fcf-1e51-4b05-ad62-355da2a1bf6d) ## Install instructions: There is currently a single external requirement, which is the `requests` library. @@ -15,6 +15,47 @@ git clone https://github.com/city96/ComfyUI_NetDist ComfyUI/custom_nodes/ComfyUI ``` ## Usage + +### Local Remote control +You will need at least two different ComfyUI instances. You can use two local GPUs by setting different `--port [port]` and `--cuda-device [number]` launch arguments. You'll most likely want `--port 8288 --cuda-device 1` + +#### Simple dual-GPU + +This is the simplest setup for people who have 2 GPUs or two separate PCs. It only requires two nodes to work. + +You can set the local/remote batch size, as well as when the node should trigger (set it to 'always' if it isn't getting executed - i.e. you changed a sampler setting but not the seed.) + +If you're running your second instance on a different PC, add `--listen` to your launch arguments and set the correct remote IP (open a terminal window and check with `ipconfig` on windows or `ip a` on linux). + +The `FetchRemote` ('Fetch from remote') node takes an image input. This should be your final image than you want to get back from your second instance (make sure not to route it back into itself). This node will wait for the second image to be generated (there's currently no preview/progress bar). + +Workflow JSON: [NetDistSimple.json](https://github.com/city96/ComfyUI_NetDist/files/13825326/NetDistSimple.json) + +![NetDistSimple](https://github.com/city96/ComfyUI_NetDist/assets/125218114/dce5a155-2ffa-4979-b184-03de168beecb) + +#### Simple multi-machine + +You can kind of scale the example above by connecting more of the simple queue nodes together, but the seed is a bit jank and you can get duplicate images if you try and reuse it. I guess just set the seed to randomized on both. + +![NetDistMulti](https://github.com/city96/ComfyUI_NetDist/assets/125218114/2a0358aa-ab8e-47e2-82a2-7a27a17d0130) + +#### Advanced + +This is mostly meant for more "advanced" setups with more than two GPUs. It allows easier per-batch overrides as well as setting a default batch size. + +It also allows using a workflow JSON as an input. To allow any workflow to run, the final image can be set to "any" instead of the default "final_image" (which would require the `FetchRemote` node to be in the workflow). + +I have nodes to save/load the workflows, but ideally there would be some nodes to also edit them - search and replace seed, etc. PRs welcome ;P + +Workflow JSON: [NetDistAdvanced.json](https://github.com/city96/ComfyUI_NetDist/files/13825337/NetDistAdvanced.json) + +![NetDistAdvanced](https://github.com/city96/ComfyUI_NetDist/assets/125218114/851c1ee6-edcf-4489-bab1-92ab9c5ef15e) + +(This needs a fake image input to trigger, you can just give it a blank image). + +![NetDistSaved](https://github.com/city96/ComfyUI_NetDist/assets/125218114/a39b5117-af1b-4f2c-a94e-5a330acc8ea4) + + ### Remote images The `LoadImageUrl` ('Load Image (URL)') Node acts just like the normal 'Load Image' node. @@ -24,27 +65,12 @@ The `SaveImageUrl` ('Save Image (URL)') Node sends a POST request to the target - The filenames are **not** guaranteed to be unique across batches since they aren't saved locally. You should handle this server-side. - No data is written to disk on the server. -### Local Remote control -You will need at least two different ComfyUI instances. You can use two local GPUs by setting different `--port [port]` and `--cuda-device [number]` launch arguments. - -The following video is an example of a multi-machine workflow. The `CombineImage` nodes aren't required, they just merge the output images into a single Preview. - -https://user-images.githubusercontent.com/125218114/234095447-85bd5111-d407-437a-a270-d159876b3a2a.mp4 - -**Chaining the seed is required**, as this allows each node to increment the seed (by `node_id*batch_size`). Simply connect the seed output of the first node to the seed input of the next one and eventually into the KSampler. - -The `FetchRemote` ('Fetch from remote') node takes an image input, this should be your final image (make sure not to route it back into itself) - -The `QueueRemote` ('Queue on remote') node will start the entire current workflow on the remote ComfyUI instance, with some changes: -- Disable all QueueRemote images (to stop recursion) -- Remove all SaveImage and PreviewImage nodes (not needed/makes it so there is only a single output) -- Replaces the `FetchRemote` ('Fetch from remote') node with a PreviewImage node, since this will be the only output -- The `FetchRemote` node (on the current workflow) will wait for the current job to finish on the remote machine. ### Things you probably shouldn't do: -- Have more `FetchRemote` nodes than `QueueRemote` ones. +- Queue a workflow on the same client multiple times. +- ~~Expect this to work smoothly.~~ ## Roadmap - Fix some edge cases, like linux controlling windows (`os.sep` mismatch). -- Switch to per-client batchsize. -- Upload rest of control software (external scheduler). +- Better workflow editing for static workflows. +- Handle multiple separate image output nodes. From e413b3b72d4fd0d921460dfa047deb78ba086956 Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Thu, 4 Jan 2024 04:41:52 +0100 Subject: [PATCH 5/5] Menu order --- nodes/images.py | 9 ++++++--- nodes/simple.py | 2 +- 2 files changed, 7 insertions(+), 4 deletions(-) diff --git a/nodes/images.py b/nodes/images.py index 50ea73b..c6e43a8 100644 --- a/nodes/images.py +++ b/nodes/images.py @@ -22,7 +22,7 @@ class LoadImageUrl: RETURN_TYPES = ("IMAGE", "MASK") FUNCTION = "load_image_url" - CATEGORY = "remote" + CATEGORY = "remote/image" TITLE = "Load Image (URL)" def load_image_url(self, url): @@ -58,7 +58,7 @@ class SaveImageUrl: RETURN_TYPES = () OUTPUT_NODE = True FUNCTION = "save_images" - CATEGORY = "remote" + CATEGORY = "remote/image" TITLE = "Save Image (URL)" def save_images(self, images, url, data_format, filename_prefix="ComfyUI", prompt=None, extra_pnginfo=None): @@ -90,6 +90,9 @@ class SaveImageUrl: return () class CombineImageBatch: + """ + This isn't needed anymore but I used it in too many places so I'm keeping it... + """ def __init__(self): pass @@ -105,7 +108,7 @@ class CombineImageBatch: RETURN_TYPES = ("IMAGE",) RETURN_NAMES = ("images",) FUNCTION = "combine_images" - CATEGORY = "remote/misc" + CATEGORY = "remote/image" TITLE = "Combine images" def combine_images(self,images_a,images_b): diff --git a/nodes/simple.py b/nodes/simple.py index ea626ea..872cca4 100644 --- a/nodes/simple.py +++ b/nodes/simple.py @@ -89,6 +89,6 @@ class RemoteQueueSimple(): return uuid if trigger == "on_change" else str(time.time()) NODE_CLASS_MAPPINGS = { - "FetchRemote" : FetchRemote, "RemoteQueueSimple" : RemoteQueueSimple, + "FetchRemote" : FetchRemote, }