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. diff --git a/__init__.py b/__init__.py index 49d8dc3..a7e64d5 100644 --- a/__init__.py +++ b/__init__.py @@ -1,44 +1,22 @@ -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.simple import NODE_CLASS_MAPPINGS as NetNodes + NODE_CLASS_MAPPINGS.update(NetNodes) -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.advanced import NODE_CLASS_MAPPINGS as AdvNodes + NODE_CLASS_MAPPINGS.update(AdvNodes) -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.images import NODE_CLASS_MAPPINGS as ImgNodes + NODE_CLASS_MAPPINGS.update(ImgNodes) -print("Loading network distribution node pack") -remote_control() -remote_images() -remote_misc() + 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/core/dispatch.py b/core/dispatch.py new file mode 100644 index 0000000..661b001 --- /dev/null +++ b/core/dispatch.py @@ -0,0 +1,135 @@ +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, timeout=4) + r.raise_for_status() + data = r.json() + return data["system"]["os"] + +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 = [] + 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" + + 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 + 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 banned: + recursive_node_deletion(i) + if output: + 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..96ed4ae --- /dev/null +++ b/core/fetch.py @@ -0,0 +1,60 @@ +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: + # 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!") + +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/advanced.py b/nodes/advanced.py new file mode 100644 index 0000000..18f258a --- /dev/null +++ b/nodes/advanced.py @@ -0,0 +1,118 @@ +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"}), + "outputs": (["final_image", "any"],{"default":"final_image"}), + } + } + + 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, outputs): + 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"], + outputs, + ) + 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/remote_images.py b/nodes/images.py similarity index 68% rename from nodes/remote_images.py rename to nodes/images.py index ca8d32a..c6e43a8 100644 --- a/nodes/remote_images.py +++ b/nodes/images.py @@ -22,7 +22,8 @@ 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): with requests.get(url, stream=True) as r: @@ -57,8 +58,9 @@ 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): filename = os.path.basename(os.path.normpath(filename_prefix)) @@ -86,3 +88,40 @@ class SaveImageUrl: with requests.post(url, json=data) as r: r.raise_for_status() 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 + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "images_a": ("IMAGE",), + "images_b": ("IMAGE",), + } + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("images",) + FUNCTION = "combine_images" + CATEGORY = "remote/image" + 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..872cca4 --- /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 = { + "RemoteQueueSimple" : RemoteQueueSimple, + "FetchRemote" : FetchRemote, +} diff --git a/nodes/workflows.py b/nodes/workflows.py new file mode 100644 index 0000000..b0ed73d --- /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(workflow) + 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, +}