|
|
@ -5,6 +5,7 @@ import json |
|
|
|
import threading |
|
|
|
import threading |
|
|
|
import heapq |
|
|
|
import heapq |
|
|
|
import traceback |
|
|
|
import traceback |
|
|
|
|
|
|
|
import asyncio |
|
|
|
|
|
|
|
|
|
|
|
if __name__ == "__main__": |
|
|
|
if __name__ == "__main__": |
|
|
|
if '--help' in sys.argv: |
|
|
|
if '--help' in sys.argv: |
|
|
@ -25,7 +26,6 @@ if '--dont-upcast-attention' in sys.argv: |
|
|
|
os.environ['ATTN_PRECISION'] = "fp16" |
|
|
|
os.environ['ATTN_PRECISION'] = "fp16" |
|
|
|
|
|
|
|
|
|
|
|
import torch |
|
|
|
import torch |
|
|
|
|
|
|
|
|
|
|
|
import nodes |
|
|
|
import nodes |
|
|
|
|
|
|
|
|
|
|
|
def get_input_data(inputs, class_def, outputs={}, prompt={}, extra_data={}): |
|
|
|
def get_input_data(inputs, class_def, outputs={}, prompt={}, extra_data={}): |
|
|
@ -286,16 +286,19 @@ def prompt_worker(q): |
|
|
|
q.task_done(item_id) |
|
|
|
q.task_done(item_id) |
|
|
|
|
|
|
|
|
|
|
|
class PromptQueue: |
|
|
|
class PromptQueue: |
|
|
|
def __init__(self): |
|
|
|
def __init__(self, socket_handler): |
|
|
|
|
|
|
|
self.socket_handler = socket_handler |
|
|
|
self.mutex = threading.RLock() |
|
|
|
self.mutex = threading.RLock() |
|
|
|
self.not_empty = threading.Condition(self.mutex) |
|
|
|
self.not_empty = threading.Condition(self.mutex) |
|
|
|
self.task_counter = 0 |
|
|
|
self.task_counter = 0 |
|
|
|
self.queue = [] |
|
|
|
self.queue = [] |
|
|
|
self.currently_running = {} |
|
|
|
self.currently_running = {} |
|
|
|
|
|
|
|
socket_handler.prompt_queue = self |
|
|
|
|
|
|
|
|
|
|
|
def put(self, item): |
|
|
|
def put(self, item): |
|
|
|
with self.mutex: |
|
|
|
with self.mutex: |
|
|
|
heapq.heappush(self.queue, item) |
|
|
|
heapq.heappush(self.queue, item) |
|
|
|
|
|
|
|
self.socket_handler.queue_updated(self) |
|
|
|
self.not_empty.notify() |
|
|
|
self.not_empty.notify() |
|
|
|
|
|
|
|
|
|
|
|
def get(self): |
|
|
|
def get(self): |
|
|
@ -306,11 +309,13 @@ class PromptQueue: |
|
|
|
i = self.task_counter |
|
|
|
i = self.task_counter |
|
|
|
self.currently_running[i] = copy.deepcopy(item) |
|
|
|
self.currently_running[i] = copy.deepcopy(item) |
|
|
|
self.task_counter += 1 |
|
|
|
self.task_counter += 1 |
|
|
|
|
|
|
|
self.socket_handler.queue_updated(self) |
|
|
|
return (item, i) |
|
|
|
return (item, i) |
|
|
|
|
|
|
|
|
|
|
|
def task_done(self, item_id): |
|
|
|
def task_done(self, item_id): |
|
|
|
with self.mutex: |
|
|
|
with self.mutex: |
|
|
|
self.currently_running.pop(item_id) |
|
|
|
self.currently_running.pop(item_id) |
|
|
|
|
|
|
|
self.socket_handler.queue_updated(self) |
|
|
|
|
|
|
|
|
|
|
|
def get_current_queue(self): |
|
|
|
def get_current_queue(self): |
|
|
|
with self.mutex: |
|
|
|
with self.mutex: |
|
|
@ -326,6 +331,7 @@ class PromptQueue: |
|
|
|
def wipe_queue(self): |
|
|
|
def wipe_queue(self): |
|
|
|
with self.mutex: |
|
|
|
with self.mutex: |
|
|
|
self.queue = [] |
|
|
|
self.queue = [] |
|
|
|
|
|
|
|
self.socket_handler.queue_updated(self) |
|
|
|
|
|
|
|
|
|
|
|
def delete_queue_item(self, function): |
|
|
|
def delete_queue_item(self, function): |
|
|
|
with self.mutex: |
|
|
|
with self.mutex: |
|
|
@ -336,35 +342,82 @@ class PromptQueue: |
|
|
|
else: |
|
|
|
else: |
|
|
|
self.queue.pop(x) |
|
|
|
self.queue.pop(x) |
|
|
|
heapq.heapify(self.queue) |
|
|
|
heapq.heapify(self.queue) |
|
|
|
|
|
|
|
self.socket_handler.queue_updated(self) |
|
|
|
return True |
|
|
|
return True |
|
|
|
return False |
|
|
|
return False |
|
|
|
|
|
|
|
|
|
|
|
from http.server import BaseHTTPRequestHandler, HTTPServer |
|
|
|
import aiohttp |
|
|
|
|
|
|
|
from aiohttp import web |
|
|
|
|
|
|
|
|
|
|
|
class PromptServer(BaseHTTPRequestHandler): |
|
|
|
def get_queue_info(prompt_queue): |
|
|
|
def _set_headers(self, code=200, ct='text/html'): |
|
|
|
|
|
|
|
self.send_response(code) |
|
|
|
|
|
|
|
self.send_header('Content-type', ct) |
|
|
|
|
|
|
|
self.end_headers() |
|
|
|
|
|
|
|
def log_message(self, format, *args): |
|
|
|
|
|
|
|
pass |
|
|
|
|
|
|
|
def do_GET(self): |
|
|
|
|
|
|
|
if self.path == "/prompt": |
|
|
|
|
|
|
|
self._set_headers(ct='application/json') |
|
|
|
|
|
|
|
prompt_info = {} |
|
|
|
prompt_info = {} |
|
|
|
exec_info = {} |
|
|
|
exec_info = {} |
|
|
|
exec_info['queue_remaining'] = self.server.prompt_queue.get_tasks_remaining() |
|
|
|
exec_info['queue_remaining'] = prompt_queue.get_tasks_remaining() |
|
|
|
prompt_info['exec_info'] = exec_info |
|
|
|
prompt_info['exec_info'] = exec_info |
|
|
|
self.wfile.write(json.dumps(prompt_info).encode('utf-8')) |
|
|
|
return prompt_info |
|
|
|
elif self.path == "/queue": |
|
|
|
|
|
|
|
self._set_headers(ct='application/json') |
|
|
|
class SocketHandler(): |
|
|
|
queue_info = {} |
|
|
|
def __init__(self, loop): |
|
|
|
current_queue = self.server.prompt_queue.get_current_queue() |
|
|
|
self.connected = set() |
|
|
|
queue_info['queue_running'] = current_queue[0] |
|
|
|
self.messages = asyncio.Queue() |
|
|
|
queue_info['queue_pending'] = current_queue[1] |
|
|
|
self.loop = loop |
|
|
|
self.wfile.write(json.dumps(queue_info).encode('utf-8')) |
|
|
|
|
|
|
|
elif self.path == "/object_info": |
|
|
|
async def publish_loop(self): |
|
|
|
self._set_headers(ct='application/json') |
|
|
|
while True: |
|
|
|
|
|
|
|
msg = await self.messages.get() |
|
|
|
|
|
|
|
await self.send(msg) |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def queue_updated(self, queue): |
|
|
|
|
|
|
|
# This is called by the queue processing thread so we need to make it thread safe |
|
|
|
|
|
|
|
loop.call_soon_threadsafe(self.messages.put_nowait, { 'type': 'status', 'status': get_queue_info(queue) }) |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
async def send(self, message, socket = None): |
|
|
|
|
|
|
|
if isinstance(message, str) == False: |
|
|
|
|
|
|
|
message = json.dumps(message) |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
if socket is None: |
|
|
|
|
|
|
|
for ws in self.connected: |
|
|
|
|
|
|
|
await ws.send_str(message) |
|
|
|
|
|
|
|
else: |
|
|
|
|
|
|
|
await socket.send_str(message) |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
async def process(self, request): |
|
|
|
|
|
|
|
ws = web.WebSocketResponse() |
|
|
|
|
|
|
|
await ws.prepare(request) |
|
|
|
|
|
|
|
self.connected.add(ws) |
|
|
|
|
|
|
|
try: |
|
|
|
|
|
|
|
# Send initial state to the new client |
|
|
|
|
|
|
|
await self.send({ 'type': 'status', 'status': get_queue_info(self.prompt_queue) }, ws) |
|
|
|
|
|
|
|
async for msg in ws: |
|
|
|
|
|
|
|
if msg.type == aiohttp.WSMsgType.ERROR: |
|
|
|
|
|
|
|
print('ws connection closed with exception %s' % ws.exception()) |
|
|
|
|
|
|
|
finally: |
|
|
|
|
|
|
|
self.connected.remove(ws) |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
return ws |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class PromptServer(): |
|
|
|
|
|
|
|
def __init__(self, prompt_queue, socket_handler): |
|
|
|
|
|
|
|
self.prompt_queue = prompt_queue |
|
|
|
|
|
|
|
self.socket_handler = socket_handler |
|
|
|
|
|
|
|
self.number = 0 |
|
|
|
|
|
|
|
self.app = web.Application() |
|
|
|
|
|
|
|
routes = web.RouteTableDef() |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@routes.get('/ws') |
|
|
|
|
|
|
|
async def websocket_handler(request): |
|
|
|
|
|
|
|
return await self.socket_handler.process(request) |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@routes.get("/") |
|
|
|
|
|
|
|
async def get_root(request): |
|
|
|
|
|
|
|
return aiohttp.web.HTTPFound('/index.html') |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@routes.get("/prompt") |
|
|
|
|
|
|
|
async def get_prompt(request): |
|
|
|
|
|
|
|
return web.json_response(get_queue_info(self.prompt_queue)) |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@routes.get("/object_info") |
|
|
|
|
|
|
|
async def get_object_info(request): |
|
|
|
out = {} |
|
|
|
out = {} |
|
|
|
for x in nodes.NODE_CLASS_MAPPINGS: |
|
|
|
for x in nodes.NODE_CLASS_MAPPINGS: |
|
|
|
obj_class = nodes.NODE_CLASS_MAPPINGS[x] |
|
|
|
obj_class = nodes.NODE_CLASS_MAPPINGS[x] |
|
|
@ -377,40 +430,32 @@ class PromptServer(BaseHTTPRequestHandler): |
|
|
|
if hasattr(obj_class, 'CATEGORY'): |
|
|
|
if hasattr(obj_class, 'CATEGORY'): |
|
|
|
info['category'] = obj_class.CATEGORY |
|
|
|
info['category'] = obj_class.CATEGORY |
|
|
|
out[x] = info |
|
|
|
out[x] = info |
|
|
|
self.wfile.write(json.dumps(out).encode('utf-8')) |
|
|
|
return web.json_response(out) |
|
|
|
elif self.path[1:] in os.listdir(self.server.server_dir): |
|
|
|
|
|
|
|
if self.path[1:].endswith('.css'): |
|
|
|
|
|
|
|
self._set_headers(ct='text/css') |
|
|
|
|
|
|
|
elif self.path[1:].endswith('.js'): |
|
|
|
|
|
|
|
self._set_headers(ct='text/javascript') |
|
|
|
|
|
|
|
else: |
|
|
|
|
|
|
|
self._set_headers() |
|
|
|
|
|
|
|
with open(os.path.join(self.server.server_dir, self.path[1:]), "rb") as f: |
|
|
|
|
|
|
|
self.wfile.write(f.read()) |
|
|
|
|
|
|
|
else: |
|
|
|
|
|
|
|
self._set_headers() |
|
|
|
|
|
|
|
with open(os.path.join(self.server.server_dir, "index.html"), "rb") as f: |
|
|
|
|
|
|
|
self.wfile.write(f.read()) |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def do_HEAD(self): |
|
|
|
@routes.get("/queue") |
|
|
|
self._set_headers() |
|
|
|
async def get_queue(request): |
|
|
|
|
|
|
|
queue_info = {} |
|
|
|
|
|
|
|
current_queue = self.prompt_queue.get_current_queue() |
|
|
|
|
|
|
|
queue_info['queue_running'] = current_queue[0] |
|
|
|
|
|
|
|
queue_info['queue_pending'] = current_queue[1] |
|
|
|
|
|
|
|
return web.json_response(queue_info) |
|
|
|
|
|
|
|
|
|
|
|
def do_POST(self): |
|
|
|
@routes.post("/prompt") |
|
|
|
|
|
|
|
async def post_prompt(request): |
|
|
|
|
|
|
|
print("got prompt") |
|
|
|
resp_code = 200 |
|
|
|
resp_code = 200 |
|
|
|
out_string = "" |
|
|
|
out_string = "" |
|
|
|
if self.path == "/prompt": |
|
|
|
json_data = await request.json() |
|
|
|
print("got prompt") |
|
|
|
|
|
|
|
data_string = self.rfile.read(int(self.headers['Content-Length'])) |
|
|
|
|
|
|
|
json_data = json.loads(data_string) |
|
|
|
|
|
|
|
if "number" in json_data: |
|
|
|
if "number" in json_data: |
|
|
|
number = float(json_data['number']) |
|
|
|
number = float(json_data['number']) |
|
|
|
else: |
|
|
|
else: |
|
|
|
number = self.server.number |
|
|
|
number = self.number |
|
|
|
if "front" in json_data: |
|
|
|
if "front" in json_data: |
|
|
|
if json_data['front']: |
|
|
|
if json_data['front']: |
|
|
|
number = -number |
|
|
|
number = -number |
|
|
|
|
|
|
|
|
|
|
|
self.server.number += 1 |
|
|
|
self.number += 1 |
|
|
|
if "prompt" in json_data: |
|
|
|
if "prompt" in json_data: |
|
|
|
prompt = json_data["prompt"] |
|
|
|
prompt = json_data["prompt"] |
|
|
|
valid = validate_prompt(prompt) |
|
|
|
valid = validate_prompt(prompt) |
|
|
@ -418,46 +463,54 @@ class PromptServer(BaseHTTPRequestHandler): |
|
|
|
if "extra_data" in json_data: |
|
|
|
if "extra_data" in json_data: |
|
|
|
extra_data = json_data["extra_data"] |
|
|
|
extra_data = json_data["extra_data"] |
|
|
|
if valid[0]: |
|
|
|
if valid[0]: |
|
|
|
self.server.prompt_queue.put((number, id(prompt), prompt, extra_data)) |
|
|
|
self.prompt_queue.put((number, id(prompt), prompt, extra_data)) |
|
|
|
else: |
|
|
|
else: |
|
|
|
resp_code = 400 |
|
|
|
resp_code = 400 |
|
|
|
out_string = valid[1] |
|
|
|
out_string = valid[1] |
|
|
|
print("invalid prompt:", valid[1]) |
|
|
|
print("invalid prompt:", valid[1]) |
|
|
|
elif self.path == "/queue": |
|
|
|
|
|
|
|
data_string = self.rfile.read(int(self.headers['Content-Length'])) |
|
|
|
return web.Response(body=out_string, status=resp_code) |
|
|
|
json_data = json.loads(data_string) |
|
|
|
|
|
|
|
|
|
|
|
@routes.post("/queue") |
|
|
|
|
|
|
|
async def post_queue(request): |
|
|
|
|
|
|
|
json_data = await request.json() |
|
|
|
if "clear" in json_data: |
|
|
|
if "clear" in json_data: |
|
|
|
if json_data["clear"]: |
|
|
|
if json_data["clear"]: |
|
|
|
self.server.prompt_queue.wipe_queue() |
|
|
|
self.prompt_queue.wipe_queue() |
|
|
|
if "delete" in json_data: |
|
|
|
if "delete" in json_data: |
|
|
|
to_delete = json_data['delete'] |
|
|
|
to_delete = json_data['delete'] |
|
|
|
for id_to_delete in to_delete: |
|
|
|
for id_to_delete in to_delete: |
|
|
|
delete_func = lambda a: a[1] == int(id_to_delete) |
|
|
|
delete_func = lambda a: a[1] == int(id_to_delete) |
|
|
|
self.server.prompt_queue.delete_queue_item(delete_func) |
|
|
|
self.prompt_queue.delete_queue_item(delete_func) |
|
|
|
|
|
|
|
|
|
|
|
self._set_headers(code=resp_code) |
|
|
|
return web.Response(status=200) |
|
|
|
self.end_headers() |
|
|
|
|
|
|
|
self.wfile.write(out_string.encode('utf8')) |
|
|
|
self.app.add_routes(routes) |
|
|
|
return |
|
|
|
self.app.add_routes([ |
|
|
|
|
|
|
|
web.static('/', os.path.join(os.path.dirname(os.path.realpath(__file__)), "webshit")), |
|
|
|
|
|
|
|
]) |
|
|
|
def run(prompt_queue, address='', port=8188): |
|
|
|
|
|
|
|
server_address = (address, port) |
|
|
|
async def start_server(server, address, port): |
|
|
|
httpd = HTTPServer(server_address, PromptServer) |
|
|
|
runner = web.AppRunner(server.app) |
|
|
|
httpd.server_dir = os.path.join(os.path.dirname(os.path.realpath(__file__)), "webshit") |
|
|
|
await runner.setup() |
|
|
|
httpd.prompt_queue = prompt_queue |
|
|
|
site = web.TCPSite(runner, address, port) |
|
|
|
httpd.number = 0 |
|
|
|
await site.start() |
|
|
|
if server_address[0] == '': |
|
|
|
|
|
|
|
addr = '0.0.0.0' |
|
|
|
if address == '': |
|
|
|
else: |
|
|
|
address = '0.0.0.0' |
|
|
|
addr = server_address[0] |
|
|
|
|
|
|
|
print("Starting server\n") |
|
|
|
print("Starting server\n") |
|
|
|
print("To see the GUI go to: http://{}:{}".format(addr, server_address[1])) |
|
|
|
print("To see the GUI go to: http://{}:{}".format(address, port)) |
|
|
|
httpd.serve_forever() |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
async def run(prompt_queue, socket_handler, address='', port=8188): |
|
|
|
|
|
|
|
server = PromptServer(prompt_queue, socket_handler) |
|
|
|
|
|
|
|
await asyncio.gather(start_server(server, address, port), socket_handler.publish_loop()) |
|
|
|
|
|
|
|
|
|
|
|
if __name__ == "__main__": |
|
|
|
if __name__ == "__main__": |
|
|
|
q = PromptQueue() |
|
|
|
loop = asyncio.new_event_loop() |
|
|
|
|
|
|
|
asyncio.set_event_loop(loop) |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
socket_handler = SocketHandler(loop) |
|
|
|
|
|
|
|
q = PromptQueue(socket_handler) |
|
|
|
threading.Thread(target=prompt_worker, daemon=True, args=(q,)).start() |
|
|
|
threading.Thread(target=prompt_worker, daemon=True, args=(q,)).start() |
|
|
|
if '--listen' in sys.argv: |
|
|
|
if '--listen' in sys.argv: |
|
|
|
address = '0.0.0.0' |
|
|
|
address = '0.0.0.0' |
|
|
@ -471,6 +524,9 @@ if __name__ == "__main__": |
|
|
|
except: |
|
|
|
except: |
|
|
|
pass |
|
|
|
pass |
|
|
|
|
|
|
|
|
|
|
|
run(q, address=address, port=port) |
|
|
|
try: |
|
|
|
|
|
|
|
loop.run_until_complete(run(q, socket_handler, address=address, port=port)) |
|
|
|
|
|
|
|
except KeyboardInterrupt: |
|
|
|
|
|
|
|
pass |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|