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, +}