7 Commits
Author SHA1 Message Date
City e413b3b72d Menu order 2024-01-04 04:41:52 +01:00
City d1d419b3f3 Update README.md 2024-01-04 04:36:44 +01:00
City a4eb4b09f5 Static workflow 2024-01-04 04:13:04 +01:00
City 9486d855b5 Restore chain arch 2024-01-04 03:42:45 +01:00
City 4c668c91f0 Rewrite part 1
Let's try this again
2024-01-04 00:28:15 +01:00
City f482f5ae8d Update readme.md 2023-09-03 21:06:21 +02:00
City 64fc4a5c46 Upload rudimentary distribution script 2023-09-03 21:03:52 +02:00
15 changed files with 813 additions and 394 deletions
+4
View File
@@ -1,3 +1,7 @@
mass-process/output
mass-process/image
mass-process/*.png
# Byte-compiled / optimized / DLL files
__pycache__/
*.py[cod]
+46 -20
View File
@@ -1,7 +1,7 @@
# ComfyUI_NetDist
Run ComfyUI workflows on multiple local GPUs/networked machines.
Also includes code to utilize in a render farm (save/load images to/from a server).
[NetDist_2xspeed.webm](https://github.com/city96/ComfyUI_NetDist/assets/125218114/b7ec2fcf-1e51-4b05-ad62-355da2a1bf6d)
## Install instructions:
There is currently a single external requirement, which is the `requests` library.
@@ -15,6 +15,47 @@ 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)
![NetDistSimple](https://github.com/city96/ComfyUI_NetDist/assets/125218114/dce5a155-2ffa-4979-b184-03de168beecb)
#### 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.
![NetDistMulti](https://github.com/city96/ComfyUI_NetDist/assets/125218114/2a0358aa-ab8e-47e2-82a2-7a27a17d0130)
#### 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)
![NetDistAdvanced](https://github.com/city96/ComfyUI_NetDist/assets/125218114/851c1ee6-edcf-4489-bab1-92ab9c5ef15e)
(This needs a fake image input to trigger, you can just give it a blank image).
![NetDistSaved](https://github.com/city96/ComfyUI_NetDist/assets/125218114/a39b5117-af1b-4f2c-a94e-5a330acc8ea4)
### Remote images
The `LoadImageUrl` ('Load Image (URL)') Node acts just like the normal 'Load Image' node.
@@ -24,27 +65,12 @@ 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:
- Have more `FetchRemote` nodes than `QueueRemote` ones.
- Queue a workflow on the same client multiple times.
- ~~Expect this to work smoothly.~~
## Roadmap
- Fix some edge cases, like linux controlling windows (`os.sep` mismatch).
- Switch to per-client batchsize.
- Upload rest of control software (external scheduler).
- Better workflow editing for static workflows.
- Handle multiple separate image output nodes.
+15 -1
View File
@@ -4,5 +4,19 @@ try:
except ImportError:
pass
else:
from .nodes.nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
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()}
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
+135
View File
@@ -0,0 +1,135 @@
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
+60
View File
@@ -0,0 +1,60 @@
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
+23
View File
@@ -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]
+28
View File
@@ -0,0 +1,28 @@
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"
+9
View File
@@ -0,0 +1,9 @@
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:
![job example](https://github.com/city96/ComfyUI_NetDist/assets/125218114/138ec97b-61a6-4631-a280-06b5c0e3c43d)
+159
View File
@@ -0,0 +1,159 @@
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()
+118
View File
@@ -0,0 +1,118 @@
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,
}
-358
View File
@@ -1,358 +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
# 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)
+10
View File
@@ -90,6 +90,9 @@ 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
@@ -115,3 +118,10 @@ 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,
}
-14
View File
@@ -1,14 +0,0 @@
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()}
+94
View File
@@ -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 = {
"RemoteQueueSimple" : RemoteQueueSimple,
"FetchRemote" : FetchRemote,
}
+111
View File
@@ -0,0 +1,111 @@
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,
}