6 Commits
Author SHA1 Message Date
City e7dce371ce Clear queue before new job 2023-10-07 21:11:41 +02:00
City 69e2e2ed29 Rewrite OS separator logic 2023-10-07 20:47:15 +02:00
City 85a9cf5f1b Random global ID 2023-10-07 20:27:23 +02:00
City 79a5622d32 Add simple node for dual GPU
helps while testing as well
2023-10-07 18:35:20 +02:00
City 5dd1fa7410 Cleanup 2023-10-07 16:56:12 +02:00
City 91bb2fb78b Initial refractor 2023-08-10 22:21:49 +02:00
6 changed files with 411 additions and 258 deletions
+8 -44
View File
@@ -1,44 +1,8 @@
NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {}
def remote_control():
global NODE_CLASS_MAPPINGS
global NODE_DISPLAY_NAME_MAPPINGS
from .nodes.remote_control import QueueRemote, FetchRemote
NODE_CLASS_MAPPINGS.update({
"QueueRemote": QueueRemote,
"FetchRemote": FetchRemote,
})
NODE_DISPLAY_NAME_MAPPINGS.update({
"QueueRemote": "Queue on remote",
"FetchRemote": "Fetch from remote",
})
def remote_images():
global NODE_CLASS_MAPPINGS
global NODE_DISPLAY_NAME_MAPPINGS
from .nodes.remote_images import LoadImageUrl, SaveImageUrl
NODE_CLASS_MAPPINGS.update({
"LoadImageUrl": LoadImageUrl,
"SaveImageUrl": SaveImageUrl,
})
NODE_DISPLAY_NAME_MAPPINGS.update({
"LoadImageUrl": "Load Image (URL)",
"SaveImageUrl": "Save Image (URL)",
})
def remote_misc():
global NODE_CLASS_MAPPINGS
global NODE_DISPLAY_NAME_MAPPINGS
from .nodes.misc import CombineImageBatch
NODE_CLASS_MAPPINGS.update({
"CombineImageBatch": CombineImageBatch,
})
NODE_DISPLAY_NAME_MAPPINGS.update({
"CombineImageBatch": "Combine images",
})
print("Loading network distribution node pack")
remote_control()
remote_images()
remote_misc()
# only import if running as a custom node
try:
import comfy.utils
except ImportError:
pass
else:
from .nodes.nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
+358
View File
@@ -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)
+31 -2
View File
@@ -22,7 +22,8 @@ class LoadImageUrl:
RETURN_TYPES = ("IMAGE", "MASK")
FUNCTION = "load_image_url"
CATEGORY = "remote"
CATEGORY = "remote/image"
TITLE = "Load Image (URL)"
def load_image_url(self, url):
with requests.get(url, stream=True) as r:
@@ -57,7 +58,8 @@ class SaveImageUrl:
RETURN_TYPES = ()
OUTPUT_NODE = True
FUNCTION = "save_images"
CATEGORY = "remote"
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))
@@ -86,3 +88,30 @@ class SaveImageUrl:
with requests.post(url, json=data) as r:
r.raise_for_status()
return ()
class CombineImageBatch:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"images_a": ("IMAGE",),
"images_b": ("IMAGE",),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("images",)
FUNCTION = "combine_images"
CATEGORY = "remote/image"
TITLE = "Combine images"
def combine_images(self,images_a,images_b):
try:
out = torch.cat((images_a,images_b), 0)
except RuntimeError:
print(f"Imagine size mismatch! {images_a.size()}, {images_b.size()}")
out = images_a
return (out,)
-28
View File
@@ -1,28 +0,0 @@
import torch
import torchvision
class CombineImageBatch:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"images_a": ("IMAGE",),
"images_b": ("IMAGE",),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("images",)
FUNCTION = "combine_images"
CATEGORY = "remote"
def combine_images(self,images_a,images_b):
try:
out = torch.cat((images_a,images_b), 0)
except RuntimeError:
print(f"Imagine size mismatch! {images_a.size()}, {images_b.size()}")
out = images_a
return (out,)
+14
View File
@@ -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()}
-184
View File
@@ -1,184 +0,0 @@
import time
import json
import torch
import random
import requests
import numpy as np
from PIL import Image
from copy import deepcopy
def img_to_torch(img):
image = img.convert("RGB")
image = np.array(image).astype(np.float32) / 255.0
image = torch.from_numpy(image)[None,]
return image
class FetchRemote():
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"final_image": ("IMAGE",),
"remote_info": ("REMOTE",),
},
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "get_remote_job"
CATEGORY = "remote"
def wait_for_job(self,remote_url,job_id):
url = remote_url + "history"
image_data = None
while not image_data:
r = requests.get(url)
r.raise_for_status()
data = r.json()
if not data:
time.sleep(0.5)
continue
for i,d in data.items():
if d["prompt"][3].get("job_id") == job_id:
image_data = d["outputs"][list(d["outputs"].keys())[-1]].get("images")
time.sleep(0.5)
return image_data
# remote_info can be none, but the node shouldn't exist at that point
def get_remote_job(self, final_image, remote_info):
images = []
for i in self.wait_for_job(remote_info["remote_url"],remote_info["job_id"]):
img_url = f"{remote_info['remote_url']}view?filename={i['filename']}&subfolder={i['subfolder']}&type={i['type']}"
ir = requests.get(img_url, stream=True)
ir.raise_for_status()
img = Image.open(ir.raw)
images.append(img_to_torch(img))
if len(images) == 0:
img = Image.new(mode="RGB", size=(768, 768))
images.append(img_to_torch(img))
out = images[0]
for i in images[1:]:
out = torch.cat((out,i))
return (out,)
class QueueRemote:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"remote_url": ("STRING", {
"multiline": False,
"default": "http://127.0.0.1:8188/",
}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"batch_override": ("INT", {"default": 0, "min": 0, "max": 8}),
"system": (["windows", "posix"],),
"enabled": (["true", "false", "remote"],{"default": "true"}),
"node_id": ("INT", {"default": 1, "min": 1, "max": 64}),
},
"hidden": {
"prompt": "PROMPT",
},
}
RETURN_TYPES = ("INT","REMOTE")
RETURN_NAMES = ("seed+batch","remote_info")
FUNCTION = "queue_on_remote"
CATEGORY = "remote"
def queue_on_remote(self, remote_url, seed, 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
if enabled == "false":
return (seed,{"remote_url":None,"job_id":None,})
elif enabled == "remote":
return (seed+node_id*get_max_batch_size(prompt),{"remote_url":None,"job_id":None,})
new_prompt = deepcopy(prompt)
to_del = []
def recursive_node_deletion(start_node):
target_nodes = [start_node]
if start_node not in to_del:
to_del.append(start_node)
while len(target_nodes) > 0:
new_targets = []
for target in target_nodes:
for node in new_prompt.keys():
inputs = new_prompt[node].get("inputs")
if not inputs:
continue
for i in inputs.values():
if type(i) == list:
if len(i) > 0 and i[0] in to_del:
if node not in to_del:
to_del.append(node)
new_targets.append(node)
target_nodes += new_targets
target_nodes.remove(target)
output_src = None
for i in new_prompt.keys():
if new_prompt[i]["class_type"] == "QueueRemote":
if new_prompt[i]["inputs"]["remote_url"] == remote_url:
new_prompt[i]["inputs"]["enabled"] = "remote"
output_src = i
else:
new_prompt[i]["inputs"]["enabled"] = "false"
output = None
for i in new_prompt.keys():
# only leave current fetch but replace with PreviewImage
if new_prompt[i]["class_type"] == "FetchRemote":
if new_prompt[i]["inputs"]["remote_info"][0] == output_src:
output = {
'inputs': {'images': new_prompt[i]["inputs"]["final_image"]},
'class_type': 'PreviewImage',
}
recursive_node_deletion(i)
# do not save output on remote
if new_prompt[i]["class_type"] in ["SaveImage","PreviewImage"]:
recursive_node_deletion(i)
new_prompt[str(max([int(x) for x in new_prompt.keys()])+1)] = output
if system == "posix":
for i in new_prompt.keys():
if new_prompt[i]["class_type"] == "LoraLoader":
new_prompt[i]["inputs"]["lora_name"] = new_prompt[i]["inputs"]["lora_name"].replace("\\","/")
if new_prompt[i]["class_type"] == "VAELoader":
new_prompt[i]["inputs"]["vae_name"] = new_prompt[i]["inputs"]["vae_name"].replace("\\","/")
if new_prompt[i]["class_type"] in ["CheckpointLoader","CheckpointLoaderSimple"]:
new_prompt[i]["inputs"]["ckpt_name"] = new_prompt[i]["inputs"]["ckpt_name"].replace("\\","/")
for i in to_del:
del new_prompt[i]
job_id = f"netdist-{time.time()}"
data = {
"prompt": new_prompt,
"client_id": "netdist",
"extra_data": {
"job_id": job_id,
}
}
ar = requests.post(remote_url+"prompt", json=data)
ar.raise_for_status()
return (seed,{"remote_url":remote_url,"job_id":job_id})