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] 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,)