diff --git a/__init__.py b/__init__.py index 49d8dc3..ffdd15c 100644 --- a/__init__.py +++ b/__init__.py @@ -4,12 +4,16 @@ NODE_DISPLAY_NAME_MAPPINGS = {} def remote_control(): global NODE_CLASS_MAPPINGS global NODE_DISPLAY_NAME_MAPPINGS - from .nodes.remote_control import QueueRemote, FetchRemote + from .nodes.remote_control import QueueRemoteChainStart, QueueRemoteChainEnd, QueueRemote, FetchRemote NODE_CLASS_MAPPINGS.update({ + "QueueRemoteChainStart": QueueRemoteChainStart, + "QueueRemoteChainEnd": QueueRemoteChainEnd, "QueueRemote": QueueRemote, "FetchRemote": FetchRemote, }) NODE_DISPLAY_NAME_MAPPINGS.update({ + "QueueRemoteChainStart": "Queue on remote (start of chain)", + "QueueRemoteChainEnd": "Queue on remote (end of chain)", "QueueRemote": "Queue on remote", "FetchRemote": "Fetch from remote", }) diff --git a/nodes/remote_control.py b/nodes/remote_control.py index bfd15de..b72a380 100644 --- a/nodes/remote_control.py +++ b/nodes/remote_control.py @@ -7,22 +7,17 @@ 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",), + "remote_info": ("REMINFO",), }, } @@ -32,7 +27,7 @@ class FetchRemote(): def wait_for_job(self,remote_url,job_id): url = remote_url + "history" - + image_data = None while not image_data: r = requests.get(url) @@ -49,6 +44,15 @@ class FetchRemote(): # remote_info can be none, but the node shouldn't exist at that point def get_remote_job(self, final_image, remote_info): + 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_info["remote_url"] or not remote_info["job_id"]: + return (torch.empty(0,0,0,0),) + 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']}" @@ -68,6 +72,75 @@ class FetchRemote(): return (out,) + +class QueueRemoteChainStart: + def __init__(self): + pass + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "workflow": (["current"],), + "trigger": (["on_change", "always"],), + "batch": ("INT", {"default": 1, "min": 1, "max": 8}), + "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), + }, + "hidden": { + "prompt": "PROMPT", + }, + } + + RETURN_TYPES = ("REMCHAIN",) + RETURN_NAMES = ("remote_chain_start",) + FUNCTION = "chain_start" + CATEGORY = "remote" + + def chain_start(self, workflow, trigger, batch, seed, prompt): + remote_chain = { + "seed": seed+batch, + "batch": batch, + "prompt": prompt, + "current_seed": seed+batch, + "current_batch": batch, + "job_id": f"netdist-{time.time()}" + } + return(remote_chain,) + + @classmethod + def IS_CHANGED(self, workflow, trigger, batch, seed, prompt): + # don't trigger on workflow change, only input change + uuid = f"W:{workflow},B:{batch},S:{seed}" + return uuid if trigger == "on_change" else str(time.time()) + + +class QueueRemoteChainEnd: + def __init__(self): + pass + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "remote_chain_end": ("REMCHAIN",) + } + } + + RETURN_TYPES = ("INT", "INT") + RETURN_NAMES = ("seed", "batch") + FUNCTION = "chain_end" + CATEGORY = "remote" + + def chain_end(self, remote_chain_end): + seed = remote_chain_end["current_seed"] + batch = remote_chain_end["current_batch"] + print("###########REMQ",seed,batch) + return(seed,batch) + + # @classmethod + # def IS_CHANGED(self, remote_chain_end): + # uid = f"S:{remote_chain_end['seed']}-B:{remote_chain_end['batch']}" + # return uid + + class QueueRemote: def __init__(self): pass @@ -75,46 +148,42 @@ class QueueRemote: def INPUT_TYPES(s): return { "required": { + "remote_chain": ("REMCHAIN",), "remote_url": ("STRING", { "multiline": False, - "default": "http://127.0.0.1:8188/", + "default": "http://127.0.0.1:8288/", }), - "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), - "batch_override": ("INT", {"default": 0, "min": 0, "max": 8}), "system": (["windows", "posix"],), + "batch_override": ("INT", {"default": 0, "min": 0, "max": 8}), "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") + + RETURN_TYPES = ("REMCHAIN", "REMINFO") + RETURN_NAMES = ("remote_chain", "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 + + def queue_on_remote(self, remote_chain, remote_url, system, batch_override, enabled): + batch = batch_override if batch_override > 0 else remote_chain["batch"] + remote_chain["seed"] += batch + remote_info = { # empty + "remote_url": None, + "job_id": None, + } if enabled == "false": - return (seed,{"remote_url":None,"job_id":None,}) + return(remote_chain, remote_info) elif enabled == "remote": - return (seed+node_id*get_max_batch_size(prompt),{"remote_url":None,"job_id":None,}) + remote_chain["current_seed"] = remote_chain["seed"] # hasn't run yet + remote_chain["current_batch"] = batch + # print(remote_chain) + return(remote_chain, remote_info) # + else: + remote_info["remote_url"] = remote_url + remote_info["job_id"] = remote_chain["job_id"] - new_prompt = deepcopy(prompt) + prompt = deepcopy(remote_chain["prompt"]) to_del = [] def recursive_node_deletion(start_node): target_nodes = [start_node] @@ -123,8 +192,8 @@ class QueueRemote: while len(target_nodes) > 0: new_targets = [] for target in target_nodes: - for node in new_prompt.keys(): - inputs = new_prompt[node].get("inputs") + for node in prompt.keys(): + inputs = prompt[node].get("inputs") if not inputs: continue for i in inputs.values(): @@ -135,50 +204,50 @@ class QueueRemote: 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 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" + for i in prompt.keys(): + if prompt[i]["class_type"] == "QueueRemote": + if prompt[i]["inputs"]["remote_url"] == remote_url: + prompt[i]["inputs"]["enabled"] = "remote" output_src = i else: - new_prompt[i]["inputs"]["enabled"] = "false" - + prompt[i]["inputs"]["enabled"] = "false" + output = None - for i in new_prompt.keys(): + for i in 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: + if prompt[i]["class_type"] == "FetchRemote": + if prompt[i]["inputs"]["remote_info"][0] == output_src: output = { - 'inputs': {'images': new_prompt[i]["inputs"]["final_image"]}, + 'inputs': {'images': 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"]: + if 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 + prompt[str(max([int(x) for x in 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 prompt.keys(): + if prompt[i]["class_type"] == "LoraLoader": + prompt[i]["inputs"]["lora_name"] = prompt[i]["inputs"]["lora_name"].replace("\\","/") + if prompt[i]["class_type"] == "VAELoader": + prompt[i]["inputs"]["vae_name"] = prompt[i]["inputs"]["vae_name"].replace("\\","/") + if prompt[i]["class_type"] in ["CheckpointLoader","CheckpointLoaderSimple"]: + prompt[i]["inputs"]["ckpt_name"] = prompt[i]["inputs"]["ckpt_name"].replace("\\","/") for i in to_del: - del new_prompt[i] + del prompt[i] - job_id = f"netdist-{time.time()}" data = { - "prompt": new_prompt, + "prompt": prompt, "client_id": "netdist", "extra_data": { - "job_id": job_id, + "job_id": remote_info["job_id"], } } ar = requests.post(remote_url+"prompt", json=data) ar.raise_for_status() - return (seed,{"remote_url":remote_url,"job_id":job_id}) + return(remote_chain, remote_info)