From d814a30ba799ef5815b3cdef44542ee22d5905a9 Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Mon, 24 Apr 2023 21:02:08 +0200 Subject: [PATCH] Initial version --- __init__.py | 44 ++++++++++ nodes/misc.py | 28 +++++++ nodes/remote_control.py | 180 ++++++++++++++++++++++++++++++++++++++++ nodes/remote_images.py | 88 ++++++++++++++++++++ 4 files changed, 340 insertions(+) create mode 100644 __init__.py create mode 100644 nodes/misc.py create mode 100644 nodes/remote_control.py create mode 100644 nodes/remote_images.py diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..49d8dc3 --- /dev/null +++ b/__init__.py @@ -0,0 +1,44 @@ +NODE_CLASS_MAPPINGS = {} +NODE_DISPLAY_NAME_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", + }) + +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)", + }) + +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", + }) + +print("Loading network distribution node pack") +remote_control() +remote_images() +remote_misc() diff --git a/nodes/misc.py b/nodes/misc.py new file mode 100644 index 0000000..6cd94f3 --- /dev/null +++ b/nodes/misc.py @@ -0,0 +1,28 @@ +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)) + 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 new file mode 100644 index 0000000..073f416 --- /dev/null +++ b/nodes/remote_control.py @@ -0,0 +1,180 @@ +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}), + "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, 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": + 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"]: + print("!!!!!!!!!!") + 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/remote_images.py b/nodes/remote_images.py new file mode 100644 index 0000000..ca8d32a --- /dev/null +++ b/nodes/remote_images.py @@ -0,0 +1,88 @@ +import os +import json +import torch +import requests +import numpy as np +from PIL import Image +from PIL.PngImagePlugin import PngInfo +from base64 import b64encode +from io import BytesIO + +class LoadImageUrl: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "url": ("STRING", { "multiline": False, }) + } + } + + RETURN_TYPES = ("IMAGE", "MASK") + FUNCTION = "load_image_url" + CATEGORY = "remote" + + def load_image_url(self, url): + with requests.get(url, stream=True) as r: + r.raise_for_status() + i = Image.open(r.raw) + image = i.convert("RGB") + image = np.array(image).astype(np.float32) / 255.0 + image = torch.from_numpy(image)[None,] + if 'A' in i.getbands(): + mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0 + mask = 1. - torch.from_numpy(mask) + else: + mask = torch.zeros((64,64), dtype=torch.float32, device="cpu") + return (image, mask) + +class SaveImageUrl: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "images": ("IMAGE", ), + "url": ("STRING", { "multiline": False, }), + "filename_prefix": ("STRING", {"default": "ComfyUI"}), + "data_format": (["HTML_image", "Raw_data"],) + }, + "hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"}, + } + + RETURN_TYPES = () + OUTPUT_NODE = True + FUNCTION = "save_images" + CATEGORY = "remote" + + 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)) + + counter = 1 + data = {} + for image in images: + i = 255. * image.cpu().numpy() + img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8)) + meta = PngInfo() + if prompt is not None: + meta.add_text("prompt", json.dumps(prompt)) + if extra_pnginfo is not None: + for x in extra_pnginfo: + meta.add_text(x, json.dumps(extra_pnginfo[x])) + + file = f"{filename}_{counter:05}.png" + + buffer = BytesIO() + img.save(buffer, "png", pnginfo=meta, compress_level=4) + buffer.seek(0) + encoded = b64encode(buffer.read()).decode('utf-8') + data[file] = f"data:image/png;base64,{encoded}" if data_format == "HTML_image" else encoded + counter += 1 + + with requests.post(url, json=data) as r: + r.raise_for_status() + return ()