Rewrite OS separator logic

This commit is contained in:
City
2023-10-07 20:47:15 +02:00
parent 85a9cf5f1b
commit 69e2e2ed29
+26 -14
View File
@@ -1,3 +1,4 @@
import os
import time import time
import json import json
import torch import torch
@@ -165,7 +166,6 @@ class QueueRemote:
"multiline": False, "multiline": False,
"default": "http://127.0.0.1:8288/", "default": "http://127.0.0.1:8288/",
}), }),
"system": (["windows", "posix"],),
"batch_override": ("INT", {"default": 0, "min": 0, "max": 8}), "batch_override": ("INT", {"default": 0, "min": 0, "max": 8}),
"enabled": (["true", "false", "remote"],{"default": "true"}), "enabled": (["true", "false", "remote"],{"default": "true"}),
} }
@@ -177,7 +177,7 @@ class QueueRemote:
CATEGORY = "remote/advanced" CATEGORY = "remote/advanced"
TITLE = "Queue on remote (worker)" TITLE = "Queue on remote (worker)"
def queue_on_remote(self, remote_chain, remote_url, system, batch_override, enabled): def queue_on_remote(self, remote_chain, remote_url, batch_override, enabled):
batch = batch_override if batch_override > 0 else remote_chain["batch"] batch = batch_override if batch_override > 0 else remote_chain["batch"]
remote_chain["seed"] += batch remote_chain["seed"] += batch
remote_info = { # empty remote_info = { # empty
@@ -196,6 +196,7 @@ class QueueRemote:
remote_info["remote_url"] = remote_url remote_info["remote_url"] = remote_url
remote_info["job_id"] = remote_chain["job_id"] remote_info["job_id"] = remote_chain["job_id"]
### PROMPT LOGIC ###
prompt = deepcopy(remote_chain["prompt"]) prompt = deepcopy(remote_chain["prompt"])
to_del = [] to_del = []
def recursive_node_deletion(start_node): def recursive_node_deletion(start_node):
@@ -242,18 +243,31 @@ class QueueRemote:
if prompt[i]["class_type"] in ["SaveImage","PreviewImage"]: if prompt[i]["class_type"] in ["SaveImage","PreviewImage"]:
recursive_node_deletion(i) recursive_node_deletion(i)
prompt[str(max([int(x) for x in prompt.keys()])+1)] = output prompt[str(max([int(x) for x in prompt.keys()])+1)] = output
for i in to_del: del prompt[i]
if system == "posix": ### 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(): for i in prompt.keys():
if prompt[i]["class_type"] == "LoraLoader": if prompt[i]["class_type"] in sem_input_map.keys():
prompt[i]["inputs"]["lora_name"] = prompt[i]["inputs"]["lora_name"].replace("\\","/") key = sem_input_map[prompt[i]["class_type"]]
if prompt[i]["class_type"] == "VAELoader": prompt[i]["inputs"][key] = prompt[i]["inputs"][key].replace(sep_local, sep_remote)
prompt[i]["inputs"]["vae_name"] = prompt[i]["inputs"]["vae_name"].replace("\\","/")
if prompt[i]["class_type"] in ["CheckpointLoader","CheckpointLoaderSimple"]:
prompt[i]["inputs"]["ckpt_name"] = prompt[i]["inputs"]["ckpt_name"].replace("\\","/")
for i in to_del:
del prompt[i]
### REQ ###
data = { data = {
"prompt": prompt, "prompt": prompt,
"client_id": get_client_id(), "client_id": get_client_id(),
@@ -278,7 +292,6 @@ class QueueRemoteSingle():
"multiline": False, "multiline": False,
"default": "http://127.0.0.1:8288/", "default": "http://127.0.0.1:8288/",
}), }),
"system": (["windows", "posix"],),
"trigger": (["on_change", "always"],), "trigger": (["on_change", "always"],),
"batch_local": ("INT", {"default": 1, "min": 1, "max": 8}), "batch_local": ("INT", {"default": 1, "min": 1, "max": 8}),
"batch_remote": ("INT", {"default": 1, "min": 1, "max": 8}), "batch_remote": ("INT", {"default": 1, "min": 1, "max": 8}),
@@ -296,7 +309,7 @@ class QueueRemoteSingle():
CATEGORY = "remote" CATEGORY = "remote"
TITLE = "Queue on remote (single)" TITLE = "Queue on remote (single)"
def queue_on_remote(self, remote_url, system, trigger, batch_local, batch_remote, enabled, seed, prompt): def queue_on_remote(self, remote_url, trigger, batch_local, batch_remote, enabled, seed, prompt):
start = QueueRemoteChainStart() start = QueueRemoteChainStart()
remote_chain, = start.chain_start( remote_chain, = start.chain_start(
workflow = "current", workflow = "current",
@@ -309,7 +322,6 @@ class QueueRemoteSingle():
remote_chain, remote_info = queue.queue_on_remote( remote_chain, remote_info = queue.queue_on_remote(
remote_chain = remote_chain, remote_chain = remote_chain,
remote_url = remote_url, remote_url = remote_url,
system = system,
batch_override = batch_remote, batch_override = batch_remote,
enabled = enabled enabled = enabled
) )