diff --git a/nodes/control.py b/nodes/control.py index 65523c5..0c6fbb3 100644 --- a/nodes/control.py +++ b/nodes/control.py @@ -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) diff --git a/nodes/images.py b/nodes/images.py index 52b5afc..c8f49d0 100644 --- a/nodes/images.py +++ b/nodes/images.py @@ -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,) diff --git a/nodes/misc.py b/nodes/misc.py deleted file mode 100644 index e490bac..0000000 --- a/nodes/misc.py +++ /dev/null @@ -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,) diff --git a/nodes/nodes.py b/nodes/nodes.py index 64e21b3..23b1aea 100644 --- a/nodes/nodes.py +++ b/nodes/nodes.py @@ -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,