Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e7dce371ce | ||
|
|
69e2e2ed29 | ||
|
|
85a9cf5f1b | ||
|
|
79a5622d32 | ||
|
|
5dd1fa7410 | ||
|
|
91bb2fb78b |
@@ -1,7 +1,3 @@
|
||||
mass-process/output
|
||||
mass-process/image
|
||||
mass-process/*.png
|
||||
|
||||
# Byte-compiled / optimized / DLL files
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# ComfyUI_NetDist
|
||||
Run ComfyUI workflows on multiple local GPUs/networked machines.
|
||||
|
||||
[NetDist_2xspeed.webm](https://github.com/city96/ComfyUI_NetDist/assets/125218114/b7ec2fcf-1e51-4b05-ad62-355da2a1bf6d)
|
||||
Also includes code to utilize in a render farm (save/load images to/from a server).
|
||||
|
||||
## Install instructions:
|
||||
There is currently a single external requirement, which is the `requests` library.
|
||||
@@ -15,47 +15,6 @@ 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)
|
||||
|
||||

|
||||
|
||||
#### 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.
|
||||
|
||||

|
||||
|
||||
#### 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)
|
||||
|
||||

|
||||
|
||||
(This needs a fake image input to trigger, you can just give it a blank image).
|
||||
|
||||

|
||||
|
||||
|
||||
### Remote images
|
||||
The `LoadImageUrl` ('Load Image (URL)') Node acts just like the normal 'Load Image' node.
|
||||
|
||||
@@ -65,12 +24,27 @@ 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:
|
||||
- Queue a workflow on the same client multiple times.
|
||||
- ~~Expect this to work smoothly.~~
|
||||
- Have more `FetchRemote` nodes than `QueueRemote` ones.
|
||||
|
||||
## Roadmap
|
||||
- Fix some edge cases, like linux controlling windows (`os.sep` mismatch).
|
||||
- Better workflow editing for static workflows.
|
||||
- Handle multiple separate image output nodes.
|
||||
- Switch to per-client batchsize.
|
||||
- Upload rest of control software (external scheduler).
|
||||
|
||||
+1
-15
@@ -4,19 +4,5 @@ try:
|
||||
except ImportError:
|
||||
pass
|
||||
else:
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
|
||||
from .nodes.simple import NODE_CLASS_MAPPINGS as NetNodes
|
||||
NODE_CLASS_MAPPINGS.update(NetNodes)
|
||||
|
||||
from .nodes.advanced import NODE_CLASS_MAPPINGS as AdvNodes
|
||||
NODE_CLASS_MAPPINGS.update(AdvNodes)
|
||||
|
||||
from .nodes.images import NODE_CLASS_MAPPINGS as ImgNodes
|
||||
NODE_CLASS_MAPPINGS.update(ImgNodes)
|
||||
|
||||
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()}
|
||||
from .nodes.nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
|
||||
|
||||
@@ -1,135 +0,0 @@
|
||||
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
|
||||
@@ -1,60 +0,0 @@
|
||||
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
|
||||
@@ -1,23 +0,0 @@
|
||||
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]
|
||||
@@ -1,28 +0,0 @@
|
||||
workflow: "job.example.png" # Your actual workflow to distribute.
|
||||
job_start: 1 # Start of job_num, will be incremented
|
||||
job_end: 100 # until it reaches job_end.
|
||||
|
||||
workers: # List of workers the server will connect to.
|
||||
"RTX3080@LOC": # Client nickname. Can be anything.
|
||||
url: "http://127.0.0.1:8188/" # ComfyUI URL (will be used for API).
|
||||
system: "windows" # 'windows' or 'posix'.
|
||||
"P40_0@NET":
|
||||
url: "http://192.168.4.6:8188/"
|
||||
system: "posix"
|
||||
"P40_1@NET":
|
||||
url: "http://192.168.4.6:8288/"
|
||||
system: "posix"
|
||||
|
||||
# Replace specific strings in the workflow inputs.
|
||||
# Source string should be the one present in your saved workflow.
|
||||
# Destination string will be the one the clients see.
|
||||
replacement:
|
||||
-
|
||||
src: "http://127.0.0.1:8080/image/openpose/0000.png"
|
||||
dst: "http://127.0.0.1:8080/image/openpose/{job_num:04}.png"
|
||||
-
|
||||
src: "http://192.168.4.4:8080/0000.png"
|
||||
dst: "http://192.168.4.4:8080/{job_num:04}.png"
|
||||
-
|
||||
src: "http://example.lan/upload/output/0000.png"
|
||||
dst: "http://example.lan/upload/output/{job_num:04}.png"
|
||||
@@ -1,9 +0,0 @@
|
||||
This is the script I used when I had to mass-process some controlnet inputs for an animation. It was requested that I publish this in [issue#2](https://github.com/city96/ComfyUI_NetDist/issues/2#issuecomment-1696342450)
|
||||
|
||||
Note that you'll have to make all images accessible to the clients, otherwise they fail. I was using the load from URL nodes for this. For testing, I simply used the built-in python web server with `python -m http.server 8080`, but you can use your own web server to host them.
|
||||
|
||||
I don't currently have seed randomization or proper output handling. For the later, just disable all the save/preview image nodes other than the one you'll be using as the final output.
|
||||
|
||||
Here's a sample workflow:
|
||||
|
||||

|
||||
@@ -1,159 +0,0 @@
|
||||
import os
|
||||
import time
|
||||
import yaml
|
||||
import json
|
||||
import requests
|
||||
import argparse
|
||||
from PIL import Image
|
||||
from tqdm import tqdm
|
||||
from queue import Queue
|
||||
from copy import deepcopy
|
||||
from threading import Thread
|
||||
|
||||
class JobShard:
|
||||
def __init__(self, workflow, job_num):
|
||||
self.workflow = workflow # raw workflow
|
||||
self.job_num = job_num # numerical ID of job
|
||||
self.prompt = None # created when assigned to worker
|
||||
self.job_id = None # ^
|
||||
|
||||
def format_workflow(self, rep, system, job_num):
|
||||
w = deepcopy(self.workflow)
|
||||
for i in w.keys():
|
||||
# Fix path mismatch
|
||||
ct = w[i]["class_type"]
|
||||
pr = ("\\","/") if system == "posix" else ("/","\\")
|
||||
if ct == "LoraLoader":
|
||||
w[i]["inputs"]["lora_name"] = w[i]["inputs"]["lora_name"].replace(*pr)
|
||||
elif ct == "VAELoader":
|
||||
w[i]["inputs"]["vae_name"] = w[i]["inputs"]["vae_name"].replace(*pr)
|
||||
elif ct in ["CheckpointLoader","CheckpointLoaderSimple"]:
|
||||
w[i]["inputs"]["ckpt_name"] = w[i]["inputs"]["ckpt_name"].replace(*pr)
|
||||
# replace strings
|
||||
for k in w[i].get("inputs",{}).keys():
|
||||
src = w[i]["inputs"][k]
|
||||
dst = [x["dst"] for x in rep if x["src"] == src]
|
||||
if dst:
|
||||
w[i]["inputs"][k] = dst[0].format(job_num=job_num)
|
||||
self.prompt = w
|
||||
|
||||
def assign(self, worker):
|
||||
self.format_workflow(worker.conf["replacement"], worker.system, self.job_num)
|
||||
self.job_id = f"{worker.name}-{self.job_num}@{int(time.time())}"
|
||||
|
||||
class Worker:
|
||||
def __init__(self, name, system, url, conf, jobs, prog):
|
||||
self.name = name
|
||||
self.url = url.rstrip("/") if url.endswith("/") else url
|
||||
self.system = system.lower().strip()
|
||||
self.conf = conf # global config
|
||||
self.jobs = jobs # queue of all jobs
|
||||
self.prog = prog # progress bar
|
||||
self.job = None
|
||||
|
||||
def is_busy(self):
|
||||
busy = True if self.job else False
|
||||
return busy
|
||||
|
||||
def run(self):
|
||||
while not self.jobs.empty():
|
||||
self.job = self.jobs.get()
|
||||
self.job.assign(self)
|
||||
self.start_job()
|
||||
self.fetch_job()
|
||||
self.job = None
|
||||
self.jobs.task_done()
|
||||
self.prog.update()
|
||||
|
||||
def start_job(self):
|
||||
url = f"{self.url}/prompt"
|
||||
data = {
|
||||
"prompt": self.job.prompt,
|
||||
"client_id": "netdist-mass",
|
||||
"extra_data": {
|
||||
"job_id": self.job.job_id,
|
||||
}
|
||||
}
|
||||
r = requests.post(url, json=data)
|
||||
r.raise_for_status()
|
||||
|
||||
def wait_for_job(self):
|
||||
url = self.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") == self.job.job_id:
|
||||
image_data = d["outputs"][list(d["outputs"].keys())[-1]].get("images")
|
||||
break
|
||||
time.sleep(0.5)
|
||||
return image_data
|
||||
|
||||
def fetch_job(self):
|
||||
images = []
|
||||
for i in self.wait_for_job():
|
||||
img_url = f"{self.url}/view?filename={i['filename']}&subfolder={i['subfolder']}&type={i['type']}"
|
||||
ir = requests.get(img_url, stream=True)
|
||||
ir.raise_for_status()
|
||||
images.append(Image.open(ir.raw))
|
||||
|
||||
if len(images) == 0:
|
||||
print(f"{self.name}@{self.url} job failed")
|
||||
elif len(images) == 1:
|
||||
images[0].save(f"output/{self.job.job_num}.png")
|
||||
else:
|
||||
for i in range(len(images)):
|
||||
images[i].save(f"output/{self.job.job_num}.{i}.png")
|
||||
|
||||
def get_workflow(path):
|
||||
if path.endswith(".png"):
|
||||
img = Image.open(path)
|
||||
data = json.loads(
|
||||
img.text.get("prompt")
|
||||
)
|
||||
else:
|
||||
print("invalid input workflow")
|
||||
exit(1)
|
||||
return data
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument('--conf', required=True, help="Config file describing job.")
|
||||
args = parser.parse_args()
|
||||
|
||||
with open(args.conf) as f:
|
||||
conf = yaml.safe_load(f.read())
|
||||
|
||||
if not os.path.isdir("output"):
|
||||
os.mkdir("output")
|
||||
|
||||
# creqte queue with jobs
|
||||
jobs = Queue()
|
||||
wf = get_workflow(conf["workflow"])
|
||||
for job_num in range(conf["job_start"],conf["job_end"]):
|
||||
jobs.put(
|
||||
JobShard(wf, job_num)
|
||||
)
|
||||
prog = tqdm(total=jobs.qsize())
|
||||
|
||||
# initialize workers
|
||||
workers = []
|
||||
for name, k in conf["workers"].items():
|
||||
workers.append(Worker(
|
||||
name=name,
|
||||
system=k["system"],
|
||||
url=k["url"],
|
||||
jobs=jobs,
|
||||
prog=prog,
|
||||
conf=conf)
|
||||
)
|
||||
|
||||
|
||||
# execute all
|
||||
[Thread(target=w.run, daemon=True).start() for w in workers]
|
||||
jobs.join()
|
||||
@@ -1,118 +0,0 @@
|
||||
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,
|
||||
}
|
||||
@@ -0,0 +1,358 @@
|
||||
import os
|
||||
import time
|
||||
import json
|
||||
import torch
|
||||
import random
|
||||
import requests
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
from copy import deepcopy
|
||||
|
||||
|
||||
# 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}")
|
||||
|
||||
|
||||
class FetchRemote():
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"final_image": ("IMAGE",),
|
||||
"remote_info": ("REMINFO",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "get_remote_job"
|
||||
CATEGORY = "remote"
|
||||
TITLE = "Fetch from 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):
|
||||
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']}"
|
||||
|
||||
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 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/advanced"
|
||||
TITLE = "Queue on remote (start of chain)"
|
||||
|
||||
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"{get_client_id()}-{int(time.time()*1000*1000)}"
|
||||
}
|
||||
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/advanced"
|
||||
TITLE = "Queue on remote (end of chain)"
|
||||
|
||||
def chain_end(self, remote_chain_end):
|
||||
seed = remote_chain_end["current_seed"]
|
||||
batch = remote_chain_end["current_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
|
||||
@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"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("REMCHAIN", "REMINFO")
|
||||
RETURN_NAMES = ("remote_chain", "remote_info")
|
||||
FUNCTION = "queue_on_remote"
|
||||
CATEGORY = "remote/advanced"
|
||||
TITLE = "Queue on remote (worker)"
|
||||
|
||||
def clear_remote_queue(self, remote_url):
|
||||
def parse_job(data):
|
||||
gid = data[1]
|
||||
client_id = data[3].get("client_id")
|
||||
job_id = data[3].get("job_id")
|
||||
return gid, client_id, job_id
|
||||
|
||||
r = requests.get(remote_url + "queue")
|
||||
r.raise_for_status()
|
||||
queue = r.json()
|
||||
|
||||
to_cancel = []
|
||||
for k in queue.get("queue_pending", []):
|
||||
if k[3].get("client_id") == get_client_id():
|
||||
to_cancel.append(k[1]) # job UUID
|
||||
r = requests.post(
|
||||
remote_url+"queue",
|
||||
json={ "delete" : to_cancel }
|
||||
)
|
||||
r.raise_for_status()
|
||||
|
||||
for k in queue.get("queue_running", []):
|
||||
if k[3].get("client_id") == get_client_id():
|
||||
r = requests.post(remote_url+"interrupt", json={})
|
||||
r.raise_for_status()
|
||||
break
|
||||
|
||||
def queue_on_remote(self, remote_chain, remote_url, 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(remote_chain, remote_info)
|
||||
elif enabled == "remote":
|
||||
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"]
|
||||
|
||||
### PROMPT LOGIC ###
|
||||
prompt = deepcopy(remote_chain["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"] in ["QueueRemote", "QueueRemoteSingle"]:
|
||||
if 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
|
||||
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 ###
|
||||
def get_remote_os(remote_url):
|
||||
url = remote_url + "system_stats"
|
||||
r = requests.get(url)
|
||||
r.raise_for_status()
|
||||
data = r.json()
|
||||
return data["system"]["os"]
|
||||
|
||||
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)
|
||||
|
||||
### REQ ###
|
||||
data = {
|
||||
"prompt": prompt,
|
||||
"client_id": get_client_id(),
|
||||
"extra_data": {
|
||||
"job_id": remote_info["job_id"],
|
||||
}
|
||||
}
|
||||
self.clear_remote_queue(remote_url)
|
||||
ar = requests.post(remote_url+"prompt", json=data)
|
||||
ar.raise_for_status()
|
||||
return(remote_chain, remote_info)
|
||||
|
||||
|
||||
class QueueRemoteSingle():
|
||||
"""This just abstracts most of the code when only using two GPUs."""
|
||||
def __init__(self):
|
||||
pass
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"remote_url": ("STRING", {
|
||||
"multiline": False,
|
||||
"default": "http://127.0.0.1:8288/",
|
||||
}),
|
||||
"trigger": (["on_change", "always"],),
|
||||
"batch_local": ("INT", {"default": 1, "min": 1, "max": 8}),
|
||||
"batch_remote": ("INT", {"default": 1, "min": 1, "max": 8}),
|
||||
"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_on_remote"
|
||||
CATEGORY = "remote"
|
||||
TITLE = "Queue on remote (single)"
|
||||
|
||||
def queue_on_remote(self, remote_url, trigger, batch_local, batch_remote, enabled, seed, prompt):
|
||||
start = QueueRemoteChainStart()
|
||||
remote_chain, = start.chain_start(
|
||||
workflow = "current",
|
||||
trigger = trigger,
|
||||
batch = batch_local,
|
||||
seed = seed,
|
||||
prompt = prompt
|
||||
)
|
||||
queue = QueueRemote()
|
||||
remote_chain, remote_info = queue.queue_on_remote(
|
||||
remote_chain = remote_chain,
|
||||
remote_url = remote_url,
|
||||
batch_override = batch_remote,
|
||||
enabled = enabled
|
||||
)
|
||||
end = QueueRemoteChainEnd()
|
||||
out_seed, out_batch = end.chain_end(remote_chain)
|
||||
return(out_seed, out_batch, remote_info)
|
||||
+1
-11
@@ -60,7 +60,7 @@ class SaveImageUrl:
|
||||
FUNCTION = "save_images"
|
||||
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))
|
||||
|
||||
@@ -90,9 +90,6 @@ class SaveImageUrl:
|
||||
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
|
||||
|
||||
@@ -118,10 +115,3 @@ class CombineImageBatch:
|
||||
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,
|
||||
}
|
||||
|
||||
@@ -0,0 +1,14 @@
|
||||
from .control import QueueRemoteChainStart, QueueRemoteChainEnd, QueueRemoteSingle, QueueRemote, FetchRemote
|
||||
from .images import LoadImageUrl, SaveImageUrl, CombineImageBatch
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"QueueRemoteChainStart": QueueRemoteChainStart,
|
||||
"QueueRemote": QueueRemote,
|
||||
"QueueRemoteChainEnd": QueueRemoteChainEnd,
|
||||
"QueueRemoteSingle" : QueueRemoteSingle,
|
||||
"FetchRemote": FetchRemote,
|
||||
"LoadImageUrl": LoadImageUrl,
|
||||
"SaveImageUrl": SaveImageUrl,
|
||||
"CombineImageBatch": CombineImageBatch,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {k:v.TITLE for k,v in NODE_CLASS_MAPPINGS.items()}
|
||||
@@ -1,94 +0,0 @@
|
||||
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,
|
||||
}
|
||||
@@ -1,111 +0,0 @@
|
||||
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,
|
||||
}
|
||||
Reference in New Issue
Block a user