11 Commits
Author SHA1 Message Date
Austin Mroz dfaf02a946 Minimal compatibility with execution inversion
Bare bones changes to make individual node caching work post execution
inversion. This is likely non-viable moving forward information is
pushed into the output cache too late to block prior steps.

I plan a more correct implementation that instead wraps the cache. This
would mean that no changes to the actual execution pathing are required,
but prevents making things forward and backwards compatible
2024-08-22 13:38:37 -05:00
Austin Mroz fc677bb04c Fix license reference 2024-08-01 13:43:31 -05:00
Austin Mroz c286fbecd3 Support auth token in orchestration connection 2024-07-19 14:21:47 -05:00
Austin Mroz 1b8ed8a661 Add command for dumping logs 2024-07-03 16:55:07 -05:00
Austin Mroz 8e16554331 Unify Error responses, signing, caching
Error responses are now always part of a dict with lowercase error as
key.

Files are cached across restarts by url. Likely needs additional
testing for checkpoints.

Fixed an error where booleans could not be eserialized as part of the
change to sign file uploads.
2024-07-03 03:03:49 -05:00
Austin Mroz 9a42e29b49 Fruther orchestration implementation
Responses to orchestration are now wrapped and contain always contain a
message id to keep requests syncronized.

On connection failure, a reconnection is attempted.

Execution time is tracked and included in response data.

Response urls are now signed
2024-07-03 01:58:16 -05:00
Austin Mroz c3f11fa212 I attempt to keep socket open on exception
Exceptions while executing a command are now caught and sent back to the
socket server
2024-06-29 17:56:27 -05:00
Austin Mroz 1089d6e68f Orchestrate by outbound websocket
Responsiveness has been a point of major frustration when debugging.
In an attempt to isolate performance issues and move closer to the end
goal of sophisticated orchestration, an additional means of control has
been added where workers connect to a control server and open a
websocket for two-way communications
2024-06-29 15:26:37 -05:00
Austin Mroz 31e677ba25 Pass on_call args without parsing
With improved logging, I believe I misinterpreted the original error
message. The version of comfy used inside the current container is out
of date, and has fewer args for call_on_start.
2024-06-28 01:29:01 -05:00
Austin Mroz 0e0856ffcb Explicitly define args for wrapped server start 2024-06-27 14:01:41 -05:00
Austin Mroz 58d3d9e475 Fix ready probe, nested target folders
Fixed a minor mistake in reference used for readyness queue

When creating the linked file on disk, subfolder are created first if
needed
2024-06-26 18:40:04 -05:00
4 changed files with 104 additions and 28 deletions
+1 -2
View File
@@ -3,8 +3,7 @@ on:
workflow_dispatch: workflow_dispatch:
push: push:
branches: branches:
- main - publish
- master
paths: paths:
- "pyproject.toml" - "pyproject.toml"
+75 -10
View File
@@ -1,23 +1,27 @@
import json
import os
import server import server
import traceback
#Check for availability of workflow checkpointing #Check for availability of workflow checkpointing
import importlib import aiohttp
#wcp = importlib.import_module('custom_nodes.ComfyUI-WorkflowCheckpointing.workflowcheckpointing') import asyncio
from .workflowcheckpointing import post_prompt_remote
STATIC_AUTH_TOKEN = os.environ.get("STATIC_AUTH_TOKEN", None)
web = server.web web = server.web
ps = server.PromptServer.instance ps = server.PromptServer.instance
finished_startup = False finished_startup = False
original_server_start = ps.start original_server_start = ps.start
async def server_start(*args, **kwargs): async def server_start(address, port, verbose=True, call_on_start=None):
args = list(args) original_on_start= call_on_start
original_on_start= args[3] def on_start(*args, **kwargs):
def on_start(scheme, address, port):
if original_on_start is not None: if original_on_start is not None:
original_on_start(scheme, address, port) original_on_start(*args, **kwargs)
global finished_startup global finished_startup
finished_startup= True finished_startup= True
args[3] = on_start return await original_server_start(address, port, verbose, on_start)
return await original_server_start(*args, **kwargs)
ps.start = server_start ps.start = server_start
@@ -35,6 +39,67 @@ async def startup(request):
@ps.routes.get("/ready") @ps.routes.get("/ready")
async def ready(request): async def ready(request):
current_queue = ps.prompt_queue.get_current_queue() current_queue = ps.prompt_queue.get_current_queue()
if len(ps[0]) == 0 and len(ps[1]) == 0: if len(current_queue[0]) == 0 and len(current_queue[1]) == 0:
return web.json_response(current_queue) return web.json_response(current_queue)
return web.json_response(current_queue, status=503) return web.json_response(current_queue, status=503)
async def websocket_loop():
async with aiohttp.ClientSession() as session:
if 'ORCHESTRATION_SERVER' not in os.environ:
while True:
await asyncio.sleep(60)
if STATIC_AUTH_TOKEN:
headers = {"Authorization": f"Bearer {STATIC_AUTH_TOKEN}"}
else:
headers = None
async with session.ws_connect(os.environ["ORCHESTRATION_SERVER"],
headers=headers) as ws:
print("connected to server")
async for msg in ws:
try:
print("got command: " + str(msg))
if msg.type == aiohttp.WSMsgType.TEXT:
js = msg.json()
resp = {"message_id": js.get('message_id', 0)}
match js['command']:
case 'prompt':
#wrap as mock request
class MockRequest:
async def json(self):
return js['data']
out = await post_prompt_remote(MockRequest())
resp['data'] = json.loads(out.body._value)
case "queue":
resp['data'] = ps.prompt_queue.get_current_queue()
case "files":
resp['data'] = [f.name for f in os.scandir('fetches') if f.is_file()]
case "info":
resp['data'] = {}
if 'SALAD_MACHINE_ID' in os.environ:
resp['data']['machine_id'] = os.environ['SALAD_MACHINE_ID']
else:
resp['data']['machine_id'] = os.environ.get('HOSTNAME', 'local')
case "logs":
with open('comfyui.log', 'r') as f:
resp['data'] = f.read()
case _:
resp = {"error": "Unknown command"}
print(resp)
await ws.send_json(resp)
elif msg.type == aiohttp.WSMsgType.ERROR:
await ws.send_json({"error": "Received bad message"})
except Exception as e:
#NOTE: this will reraise if error was socket closing
await ws.send_json({"error": str(e)})
async def try_websocket():
while True:
try:
await websocket_loop()
except aiohttp.client_exceptions.ClientConnectorError:
print("disconnected")
except:
print(traceback.format_exc())
await asyncio.sleep(5)
print("Attempting re connection")
process_loop = ps.loop.create_task(try_websocket())
+2 -2
View File
@@ -1,8 +1,8 @@
[project] [project]
name = "comfyui-workflowcheckpointing" name = "comfyui-workflowcheckpointing"
description = "Automatically creates checkpoints during workflow execution. If If an workflow is canceled or ComfyUI crashes mid-execution, then these checkpoints are used when the workflow is re-queued to resume execution with minimal progress loss." description = "Automatically creates checkpoints during workflow execution. If If an workflow is canceled or ComfyUI crashes mid-execution, then these checkpoints are used when the workflow is re-queued to resume execution with minimal progress loss."
version = "1.0.0" version = "1.0.1"
license = "LICENSE" license = { file = "LICENSE" }
[project.urls] [project.urls]
Repository = "https://github.com/AustinMroz/ComfyUI-WorkflowCheckpointing" Repository = "https://github.com/AustinMroz/ComfyUI-WorkflowCheckpointing"
+25 -13
View File
@@ -9,6 +9,7 @@ import threading
import logging import logging
import itertools import itertools
import hashlib import hashlib
import time
import comfy.samplers import comfy.samplers
import execution import execution
@@ -136,6 +137,7 @@ class FetchQueue:
with self.lock: with self.lock:
if item in self.consumed: if item in self.consumed:
#TODO: Also update in queue #TODO: Also update in queue
#TODO: if complete check etag?
self.consumed[item][1] = min(self.consumed[item][1], priority) self.consumed[item][1] = min(self.consumed[item][1], priority)
return self.consumed[item][0] return self.consumed[item][0]
for i in range(len(self.queue)): for i in range(len(self.queue)):
@@ -180,6 +182,9 @@ class FetchLoop:
with self.queue.lock: with self.queue.lock:
if url in self.queue.consumed: if url in self.queue.consumed:
self.queue.consumed.pop(url) self.queue.consumed.pop(url)
hashloc = os.path.join('fetches', string_hash(url))
if os.path.exists(hashloc):
os.remove(hashloc)
def enqueue(self, url, priority=0): def enqueue(self, url, priority=0):
return self.queue.enqueue_checked(url, priority) return self.queue.enqueue_checked(url, priority)
async def fetch(self, priority, url, future): async def fetch(self, priority, url, future):
@@ -206,9 +211,12 @@ class FetchLoop:
return return
fetch_loop = FetchLoop() fetch_loop = FetchLoop()
async def prepare_file(url, path, priority): async def prepare_file(url, path, priority):
hashloc = os.path.join('fetches', string_hash(url))
if not os.path.exists(hashloc):
hashloc = await fetch_loop.enqueue(url, priority) hashloc = await fetch_loop.enqueue(url, priority)
if os.path.exists(path): if os.path.exists(path):
os.remove(path) os.remove(path)
os.makedirs(os.path.split(path)[0], exist_ok=True)
#TODO consider if symlinking would be better #TODO consider if symlinking would be better
os.link(hashloc, path) os.link(hashloc, path)
@@ -346,8 +354,10 @@ async def post_prompt_remote(request):
f = asyncio.Future() f = asyncio.Future()
index = max(completion_futures.keys(),default=0)+1 index = max(completion_futures.keys(),default=0)+1
completion_futures[index] = f completion_futures[index] = f
start_time = time.perf_counter()
base_res = await original_post_prompt(request) base_res = await original_post_prompt(request)
outputs = await f outputs = await f
execution_time = time.perf_counter() - start_time
completion_futures.pop(index) completion_futures.pop(index)
if "SALAD_ORGANIZATION" in os.environ: if "SALAD_ORGANIZATION" in os.environ:
async with aiohttp.ClientSession('https://storage-api.salad.com') as s: async with aiohttp.ClientSession('https://storage-api.salad.com') as s:
@@ -356,13 +366,16 @@ async def post_prompt_remote(request):
with open(outputs[i], 'rb') as f: with open(outputs[i], 'rb') as f:
data = f.read() data = f.read()
#TODO support uploads > 100MB/ memory optimizations #TODO support uploads > 100MB/ memory optimizations
fd = {'file': data} fd = {'file': data, 'sign': 'true'}
url = '/'.join([base_url_path, uid, 'outputs', outputs[i]]) url = '/'.join([base_url_path, uid, 'outputs', outputs[i]])
async with s.put(url, headers=headers, data=fd) as r: async with s.put(url, headers=headers, data=fd) as r:
await r.text() url = (await r.json())['url']
outputs[i] = url outputs[i] = url
res = base_res.text[:-1] + ', "outputs": ' + json.dumps(outputs) + '}' json_output = json.loads(base_res.text)
return server.web.Response(body=res) json_output['outputs'] = outputs
json_output['execution_time'] = execution_time
json_output['machineid'] = os.environ.get('SALAD_MACHINE_ID', "local")
return server.web.Response(body=json.dumps(json_output))
#Dangerous #Dangerous
object.__setattr__(prompt_route, 'handler', post_prompt_remote) object.__setattr__(prompt_route, 'handler', post_prompt_remote)
@@ -398,17 +411,16 @@ class CheckpointSampler(comfy.samplers.KSAMPLER):
if step == int(os.environ['FORCE_CRASH_AT']): if step == int(os.environ['FORCE_CRASH_AT']):
raise Exception("Simulated Crash") raise Exception("Simulated Crash")
original_recursive_execute = execution.recursive_execute original_recursive_execute = execution.execute
def recursive_execute_injection(*args): def recursive_execute_injection(*args):
unique_id = args[3] unique_id = args[3]
class_type = args[1][unique_id]['class_type'] class_type = args[1].get_node(unique_id)['class_type']
extra_data = args[4] extra_data = args[4]
if class_type in SAMPLER_NODES: if class_type in SAMPLER_NODES:
data, metadata = checkpoint.get(unique_id) data, metadata = checkpoint.get(unique_id)
if metadata is not None and 'step' in metadata: if metadata is not None and 'step' in metadata:
args[1][unique_id]['inputs']['latent_image'] = ['checkpointed'+unique_id, 0] args[1].get_node(unique_id)['inputs']['latent_image'] = ['checkpointed'+unique_id, 0]
args[2]['checkpointed'+unique_id] = [[{'samples': data['x']}]] args[2].outputs.set('checkpointed'+unique_id, [[{'samples': data['x']}]])
elif metadata is not None and 'completed' in metadata: elif metadata is not None and 'completed' in metadata:
outputs = json.loads(metadata['completed']) outputs = json.loads(metadata['completed'])
for x in range(len(outputs)): for x in range(len(outputs)):
@@ -416,15 +428,15 @@ def recursive_execute_injection(*args):
outputs[x] = list(data[str(x)]) outputs[x] = list(data[str(x)])
elif outputs[x] == 'latent': elif outputs[x] == 'latent':
outputs[x] = [{'samples': l} for l in data[str(x)]] outputs[x] = [{'samples': l} for l in data[str(x)]]
args[2][unique_id] = outputs args[2].outputs.set(unique_id, outputs)
return True, None, None return True, None, None
res = original_recursive_execute(*args) res = original_recursive_execute(*args)
#Conditionally save node output #Conditionally save node output
#TODO: determine which non-sampler nodes are worth saving #TODO: determine which non-sampler nodes are worth saving
if class_type in SAMPLER_NODES and unique_id in args[2]: if class_type in SAMPLER_NODES and args[2].outputs.get(unique_id) is not None:
data = {} data = {}
outputs = args[2][unique_id].copy() outputs = args[2].outputs.get(unique_id).copy()
for x in range(len(outputs)): for x in range(len(outputs)):
if isinstance(outputs[x][0], torch.Tensor): if isinstance(outputs[x][0], torch.Tensor):
data[str(x)] = torch.stack(outputs[x]) data[str(x)] = torch.stack(outputs[x])
@@ -456,7 +468,7 @@ def execute_injection(*args, **kwargs):
completion_futures[args[3]['completion_future']].set_result(outputs) completion_futures[args[3]['completion_future']].set_result(outputs)
comfy.samplers.KSAMPLER = CheckpointSampler comfy.samplers.KSAMPLER = CheckpointSampler
execution.recursive_execute = recursive_execute_injection execution.execute = recursive_execute_injection
execution.PromptExecutor.execute = execute_injection execution.PromptExecutor.execute = execute_injection
NODE_CLASS_MAPPINGS = {} NODE_CLASS_MAPPINGS = {}