Add simple node for dual GPU

helps while testing as well
This commit is contained in:
City
2023-10-07 18:35:20 +02:00
parent 5dd1fa7410
commit 79a5622d32
4 changed files with 90 additions and 41 deletions
+57 -6
View File
@@ -94,7 +94,7 @@ class QueueRemoteChainStart:
RETURN_TYPES = ("REMCHAIN",)
RETURN_NAMES = ("remote_chain_start",)
FUNCTION = "chain_start"
CATEGORY = "remote"
CATEGORY = "remote/advanced"
TITLE = "Queue on remote (start of chain)"
def chain_start(self, workflow, trigger, batch, seed, prompt):
@@ -129,13 +129,12 @@ class QueueRemoteChainEnd:
RETURN_TYPES = ("INT", "INT")
RETURN_NAMES = ("seed", "batch")
FUNCTION = "chain_end"
CATEGORY = "remote"
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"]
print("###########REMQ",seed,batch)
return(seed,batch)
# @classmethod
@@ -165,8 +164,8 @@ class QueueRemote:
RETURN_TYPES = ("REMCHAIN", "REMINFO")
RETURN_NAMES = ("remote_chain", "remote_info")
FUNCTION = "queue_on_remote"
CATEGORY = "remote"
TITLE = "Queue on remote"
CATEGORY = "remote/advanced"
TITLE = "Queue on remote (worker)"
def queue_on_remote(self, remote_chain, remote_url, system, batch_override, enabled):
batch = batch_override if batch_override > 0 else remote_chain["batch"]
@@ -212,7 +211,7 @@ class QueueRemote:
# find current node and disable all others
output_src = None
for i in prompt.keys():
if prompt[i]["class_type"] == "QueueRemote":
if prompt[i]["class_type"] in ["QueueRemote", "QueueRemoteSingle"]:
if prompt[i]["inputs"]["remote_url"] == remote_url:
prompt[i]["inputs"]["enabled"] = "remote"
output_src = i
@@ -255,3 +254,55 @@ class QueueRemote:
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/",
}),
"system": (["windows", "posix"],),
"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, system, 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,
system = system,
batch_override = batch_remote,
enabled = enabled
)
end = QueueRemoteChainEnd()
out_seed, out_batch = end.chain_end(remote_chain)
return(out_seed, out_batch, remote_info)
+29 -2
View File
@@ -22,7 +22,7 @@ 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):
@@ -58,7 +58,7 @@ 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):
@@ -88,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,)
-29
View File
@@ -1,29 +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"
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,)
+4 -4
View File
@@ -1,11 +1,11 @@
from .control import QueueRemoteChainStart, QueueRemoteChainEnd, QueueRemote, FetchRemote
from .images import LoadImageUrl, SaveImageUrl
from .misc import CombineImageBatch
from .control import QueueRemoteChainStart, QueueRemoteChainEnd, QueueRemoteSingle, QueueRemote, FetchRemote
from .images import LoadImageUrl, SaveImageUrl, CombineImageBatch
NODE_CLASS_MAPPINGS = {
"QueueRemoteChainStart": QueueRemoteChainStart,
"QueueRemoteChainEnd": QueueRemoteChainEnd,
"QueueRemote": QueueRemote,
"QueueRemoteChainEnd": QueueRemoteChainEnd,
"QueueRemoteSingle" : QueueRemoteSingle,
"FetchRemote": FetchRemote,
"LoadImageUrl": LoadImageUrl,
"SaveImageUrl": SaveImageUrl,