Files
city96-ComfyUI_NetDist/mass-process/server.py
T

160 lines
4.2 KiB
Python

import os
import time
import yaml
import json
import requests
import argparse
from PIL import Image
from tqdm import tqdm
from queue import Queue
from copy import deepcopy
from threading import Thread
class JobShard:
def __init__(self, workflow, job_num):
self.workflow = workflow # raw workflow
self.job_num = job_num # numerical ID of job
self.prompt = None # created when assigned to worker
self.job_id = None # ^
def format_workflow(self, rep, system, job_num):
w = deepcopy(self.workflow)
for i in w.keys():
# Fix path mismatch
ct = w[i]["class_type"]
pr = ("\\","/") if system == "posix" else ("/","\\")
if ct == "LoraLoader":
w[i]["inputs"]["lora_name"] = w[i]["inputs"]["lora_name"].replace(*pr)
elif ct == "VAELoader":
w[i]["inputs"]["vae_name"] = w[i]["inputs"]["vae_name"].replace(*pr)
elif ct in ["CheckpointLoader","CheckpointLoaderSimple"]:
w[i]["inputs"]["ckpt_name"] = w[i]["inputs"]["ckpt_name"].replace(*pr)
# replace strings
for k in w[i].get("inputs",{}).keys():
src = w[i]["inputs"][k]
dst = [x["dst"] for x in rep if x["src"] == src]
if dst:
w[i]["inputs"][k] = dst[0].format(job_num=job_num)
self.prompt = w
def assign(self, worker):
self.format_workflow(worker.conf["replacement"], worker.system, self.job_num)
self.job_id = f"{worker.name}-{self.job_num}@{int(time.time())}"
class Worker:
def __init__(self, name, system, url, conf, jobs, prog):
self.name = name
self.url = url.rstrip("/") if url.endswith("/") else url
self.system = system.lower().strip()
self.conf = conf # global config
self.jobs = jobs # queue of all jobs
self.prog = prog # progress bar
self.job = None
def is_busy(self):
busy = True if self.job else False
return busy
def run(self):
while not self.jobs.empty():
self.job = self.jobs.get()
self.job.assign(self)
self.start_job()
self.fetch_job()
self.job = None
self.jobs.task_done()
self.prog.update()
def start_job(self):
url = f"{self.url}/prompt"
data = {
"prompt": self.job.prompt,
"client_id": "netdist-mass",
"extra_data": {
"job_id": self.job.job_id,
}
}
r = requests.post(url, json=data)
r.raise_for_status()
def wait_for_job(self):
url = self.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") == self.job.job_id:
image_data = d["outputs"][list(d["outputs"].keys())[-1]].get("images")
break
time.sleep(0.5)
return image_data
def fetch_job(self):
images = []
for i in self.wait_for_job():
img_url = f"{self.url}/view?filename={i['filename']}&subfolder={i['subfolder']}&type={i['type']}"
ir = requests.get(img_url, stream=True)
ir.raise_for_status()
images.append(Image.open(ir.raw))
if len(images) == 0:
print(f"{self.name}@{self.url} job failed")
elif len(images) == 1:
images[0].save(f"output/{self.job.job_num}.png")
else:
for i in range(len(images)):
images[i].save(f"output/{self.job.job_num}.{i}.png")
def get_workflow(path):
if path.endswith(".png"):
img = Image.open(path)
data = json.loads(
img.text.get("prompt")
)
else:
print("invalid input workflow")
exit(1)
return data
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument('--conf', required=True, help="Config file describing job.")
args = parser.parse_args()
with open(args.conf) as f:
conf = yaml.safe_load(f.read())
if not os.path.isdir("output"):
os.mkdir("output")
# creqte queue with jobs
jobs = Queue()
wf = get_workflow(conf["workflow"])
for job_num in range(conf["job_start"],conf["job_end"]):
jobs.put(
JobShard(wf, job_num)
)
prog = tqdm(total=jobs.qsize())
# initialize workers
workers = []
for name, k in conf["workers"].items():
workers.append(Worker(
name=name,
system=k["system"],
url=k["url"],
jobs=jobs,
prog=prog,
conf=conf)
)
# execute all
[Thread(target=w.run, daemon=True).start() for w in workers]
jobs.join()