Initial refractor
This commit is contained in:
+5
-1
@@ -4,12 +4,16 @@ NODE_DISPLAY_NAME_MAPPINGS = {}
|
|||||||
def remote_control():
|
def remote_control():
|
||||||
global NODE_CLASS_MAPPINGS
|
global NODE_CLASS_MAPPINGS
|
||||||
global NODE_DISPLAY_NAME_MAPPINGS
|
global NODE_DISPLAY_NAME_MAPPINGS
|
||||||
from .nodes.remote_control import QueueRemote, FetchRemote
|
from .nodes.remote_control import QueueRemoteChainStart, QueueRemoteChainEnd, QueueRemote, FetchRemote
|
||||||
NODE_CLASS_MAPPINGS.update({
|
NODE_CLASS_MAPPINGS.update({
|
||||||
|
"QueueRemoteChainStart": QueueRemoteChainStart,
|
||||||
|
"QueueRemoteChainEnd": QueueRemoteChainEnd,
|
||||||
"QueueRemote": QueueRemote,
|
"QueueRemote": QueueRemote,
|
||||||
"FetchRemote": FetchRemote,
|
"FetchRemote": FetchRemote,
|
||||||
})
|
})
|
||||||
NODE_DISPLAY_NAME_MAPPINGS.update({
|
NODE_DISPLAY_NAME_MAPPINGS.update({
|
||||||
|
"QueueRemoteChainStart": "Queue on remote (start of chain)",
|
||||||
|
"QueueRemoteChainEnd": "Queue on remote (end of chain)",
|
||||||
"QueueRemote": "Queue on remote",
|
"QueueRemote": "Queue on remote",
|
||||||
"FetchRemote": "Fetch from remote",
|
"FetchRemote": "Fetch from remote",
|
||||||
})
|
})
|
||||||
|
|||||||
+132
-63
@@ -7,22 +7,17 @@ import numpy as np
|
|||||||
from PIL import Image
|
from PIL import Image
|
||||||
from copy import deepcopy
|
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():
|
class FetchRemote():
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s):
|
||||||
return {
|
return {
|
||||||
"required": {
|
"required": {
|
||||||
"final_image": ("IMAGE",),
|
"final_image": ("IMAGE",),
|
||||||
"remote_info": ("REMOTE",),
|
"remote_info": ("REMINFO",),
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -32,7 +27,7 @@ class FetchRemote():
|
|||||||
|
|
||||||
def wait_for_job(self,remote_url,job_id):
|
def wait_for_job(self,remote_url,job_id):
|
||||||
url = remote_url + "history"
|
url = remote_url + "history"
|
||||||
|
|
||||||
image_data = None
|
image_data = None
|
||||||
while not image_data:
|
while not image_data:
|
||||||
r = requests.get(url)
|
r = requests.get(url)
|
||||||
@@ -49,6 +44,15 @@ class FetchRemote():
|
|||||||
|
|
||||||
# remote_info can be none, but the node shouldn't exist at that point
|
# remote_info can be none, but the node shouldn't exist at that point
|
||||||
def get_remote_job(self, final_image, remote_info):
|
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 = []
|
images = []
|
||||||
for i in self.wait_for_job(remote_info["remote_url"],remote_info["job_id"]):
|
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']}"
|
img_url = f"{remote_info['remote_url']}view?filename={i['filename']}&subfolder={i['subfolder']}&type={i['type']}"
|
||||||
@@ -68,6 +72,75 @@ class FetchRemote():
|
|||||||
|
|
||||||
return (out,)
|
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"
|
||||||
|
|
||||||
|
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"netdist-{time.time()}"
|
||||||
|
}
|
||||||
|
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"
|
||||||
|
|
||||||
|
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
|
||||||
|
# def IS_CHANGED(self, remote_chain_end):
|
||||||
|
# uid = f"S:{remote_chain_end['seed']}-B:{remote_chain_end['batch']}"
|
||||||
|
# return uid
|
||||||
|
|
||||||
|
|
||||||
class QueueRemote:
|
class QueueRemote:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
pass
|
pass
|
||||||
@@ -75,46 +148,42 @@ class QueueRemote:
|
|||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s):
|
||||||
return {
|
return {
|
||||||
"required": {
|
"required": {
|
||||||
|
"remote_chain": ("REMCHAIN",),
|
||||||
"remote_url": ("STRING", {
|
"remote_url": ("STRING", {
|
||||||
"multiline": False,
|
"multiline": False,
|
||||||
"default": "http://127.0.0.1:8188/",
|
"default": "http://127.0.0.1:8288/",
|
||||||
}),
|
}),
|
||||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
|
||||||
"batch_override": ("INT", {"default": 0, "min": 0, "max": 8}),
|
|
||||||
"system": (["windows", "posix"],),
|
"system": (["windows", "posix"],),
|
||||||
|
"batch_override": ("INT", {"default": 0, "min": 0, "max": 8}),
|
||||||
"enabled": (["true", "false", "remote"],{"default": "true"}),
|
"enabled": (["true", "false", "remote"],{"default": "true"}),
|
||||||
"node_id": ("INT", {"default": 1, "min": 1, "max": 64}),
|
}
|
||||||
},
|
|
||||||
"hidden": {
|
|
||||||
"prompt": "PROMPT",
|
|
||||||
},
|
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ("INT","REMOTE")
|
RETURN_TYPES = ("REMCHAIN", "REMINFO")
|
||||||
RETURN_NAMES = ("seed+batch","remote_info")
|
RETURN_NAMES = ("remote_chain", "remote_info")
|
||||||
FUNCTION = "queue_on_remote"
|
FUNCTION = "queue_on_remote"
|
||||||
CATEGORY = "remote"
|
CATEGORY = "remote"
|
||||||
|
|
||||||
def queue_on_remote(self, remote_url, seed, batch_override, system, enabled, node_id, prompt):
|
def queue_on_remote(self, remote_chain, remote_url, system, batch_override, enabled):
|
||||||
def get_max_batch_size(prompt):
|
batch = batch_override if batch_override > 0 else remote_chain["batch"]
|
||||||
bs = 1
|
remote_chain["seed"] += batch
|
||||||
for node in prompt.keys():
|
remote_info = { # empty
|
||||||
if prompt[node]["class_type"] == "EmptyLatentImage":
|
"remote_url": None,
|
||||||
for k,v in prompt[node]["inputs"].items():
|
"job_id": None,
|
||||||
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":
|
if enabled == "false":
|
||||||
return (seed,{"remote_url":None,"job_id":None,})
|
return(remote_chain, remote_info)
|
||||||
elif enabled == "remote":
|
elif enabled == "remote":
|
||||||
return (seed+node_id*get_max_batch_size(prompt),{"remote_url":None,"job_id":None,})
|
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"]
|
||||||
|
|
||||||
new_prompt = deepcopy(prompt)
|
prompt = deepcopy(remote_chain["prompt"])
|
||||||
to_del = []
|
to_del = []
|
||||||
def recursive_node_deletion(start_node):
|
def recursive_node_deletion(start_node):
|
||||||
target_nodes = [start_node]
|
target_nodes = [start_node]
|
||||||
@@ -123,8 +192,8 @@ class QueueRemote:
|
|||||||
while len(target_nodes) > 0:
|
while len(target_nodes) > 0:
|
||||||
new_targets = []
|
new_targets = []
|
||||||
for target in target_nodes:
|
for target in target_nodes:
|
||||||
for node in new_prompt.keys():
|
for node in prompt.keys():
|
||||||
inputs = new_prompt[node].get("inputs")
|
inputs = prompt[node].get("inputs")
|
||||||
if not inputs:
|
if not inputs:
|
||||||
continue
|
continue
|
||||||
for i in inputs.values():
|
for i in inputs.values():
|
||||||
@@ -135,50 +204,50 @@ class QueueRemote:
|
|||||||
new_targets.append(node)
|
new_targets.append(node)
|
||||||
target_nodes += new_targets
|
target_nodes += new_targets
|
||||||
target_nodes.remove(target)
|
target_nodes.remove(target)
|
||||||
|
|
||||||
|
# find current node and disable all others
|
||||||
output_src = None
|
output_src = None
|
||||||
for i in new_prompt.keys():
|
for i in prompt.keys():
|
||||||
if new_prompt[i]["class_type"] == "QueueRemote":
|
if prompt[i]["class_type"] == "QueueRemote":
|
||||||
if new_prompt[i]["inputs"]["remote_url"] == remote_url:
|
if prompt[i]["inputs"]["remote_url"] == remote_url:
|
||||||
new_prompt[i]["inputs"]["enabled"] = "remote"
|
prompt[i]["inputs"]["enabled"] = "remote"
|
||||||
output_src = i
|
output_src = i
|
||||||
else:
|
else:
|
||||||
new_prompt[i]["inputs"]["enabled"] = "false"
|
prompt[i]["inputs"]["enabled"] = "false"
|
||||||
|
|
||||||
output = None
|
output = None
|
||||||
for i in new_prompt.keys():
|
for i in prompt.keys():
|
||||||
# only leave current fetch but replace with PreviewImage
|
# only leave current fetch but replace with PreviewImage
|
||||||
if new_prompt[i]["class_type"] == "FetchRemote":
|
if prompt[i]["class_type"] == "FetchRemote":
|
||||||
if new_prompt[i]["inputs"]["remote_info"][0] == output_src:
|
if prompt[i]["inputs"]["remote_info"][0] == output_src:
|
||||||
output = {
|
output = {
|
||||||
'inputs': {'images': new_prompt[i]["inputs"]["final_image"]},
|
'inputs': {'images': prompt[i]["inputs"]["final_image"]},
|
||||||
'class_type': 'PreviewImage',
|
'class_type': 'PreviewImage',
|
||||||
}
|
}
|
||||||
recursive_node_deletion(i)
|
recursive_node_deletion(i)
|
||||||
# do not save output on remote
|
# do not save output on remote
|
||||||
if new_prompt[i]["class_type"] in ["SaveImage","PreviewImage"]:
|
if prompt[i]["class_type"] in ["SaveImage","PreviewImage"]:
|
||||||
recursive_node_deletion(i)
|
recursive_node_deletion(i)
|
||||||
new_prompt[str(max([int(x) for x in new_prompt.keys()])+1)] = output
|
prompt[str(max([int(x) for x in prompt.keys()])+1)] = output
|
||||||
|
|
||||||
if system == "posix":
|
if system == "posix":
|
||||||
for i in new_prompt.keys():
|
for i in prompt.keys():
|
||||||
if new_prompt[i]["class_type"] == "LoraLoader":
|
if prompt[i]["class_type"] == "LoraLoader":
|
||||||
new_prompt[i]["inputs"]["lora_name"] = new_prompt[i]["inputs"]["lora_name"].replace("\\","/")
|
prompt[i]["inputs"]["lora_name"] = prompt[i]["inputs"]["lora_name"].replace("\\","/")
|
||||||
if new_prompt[i]["class_type"] == "VAELoader":
|
if prompt[i]["class_type"] == "VAELoader":
|
||||||
new_prompt[i]["inputs"]["vae_name"] = new_prompt[i]["inputs"]["vae_name"].replace("\\","/")
|
prompt[i]["inputs"]["vae_name"] = prompt[i]["inputs"]["vae_name"].replace("\\","/")
|
||||||
if new_prompt[i]["class_type"] in ["CheckpointLoader","CheckpointLoaderSimple"]:
|
if prompt[i]["class_type"] in ["CheckpointLoader","CheckpointLoaderSimple"]:
|
||||||
new_prompt[i]["inputs"]["ckpt_name"] = new_prompt[i]["inputs"]["ckpt_name"].replace("\\","/")
|
prompt[i]["inputs"]["ckpt_name"] = prompt[i]["inputs"]["ckpt_name"].replace("\\","/")
|
||||||
for i in to_del:
|
for i in to_del:
|
||||||
del new_prompt[i]
|
del prompt[i]
|
||||||
|
|
||||||
job_id = f"netdist-{time.time()}"
|
|
||||||
data = {
|
data = {
|
||||||
"prompt": new_prompt,
|
"prompt": prompt,
|
||||||
"client_id": "netdist",
|
"client_id": "netdist",
|
||||||
"extra_data": {
|
"extra_data": {
|
||||||
"job_id": job_id,
|
"job_id": remote_info["job_id"],
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
ar = requests.post(remote_url+"prompt", json=data)
|
ar = requests.post(remote_url+"prompt", json=data)
|
||||||
ar.raise_for_status()
|
ar.raise_for_status()
|
||||||
return (seed,{"remote_url":remote_url,"job_id":job_id})
|
return(remote_chain, remote_info)
|
||||||
|
|||||||
Reference in New Issue
Block a user