Compare commits
12
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
76d4d8e91e | ||
|
|
dfaf02a946 | ||
|
|
fc677bb04c | ||
|
|
c286fbecd3 | ||
|
|
1b8ed8a661 | ||
|
|
8e16554331 | ||
|
|
9a42e29b49 | ||
|
|
c3f11fa212 | ||
|
|
1089d6e68f | ||
|
|
31e677ba25 | ||
|
|
0e0856ffcb | ||
|
|
58d3d9e475 |
@@ -3,8 +3,7 @@ on:
|
||||
workflow_dispatch:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
- master
|
||||
- publish
|
||||
paths:
|
||||
- "pyproject.toml"
|
||||
|
||||
|
||||
+2
-3
@@ -1,6 +1,5 @@
|
||||
from . import workflowcheckpointing, orchestration
|
||||
from . import workflowcheckpointing
|
||||
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
WEB_DIRECTORY = "./web"
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS", "WEB_DIRECTORY"]
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
|
||||
+75
-10
@@ -1,23 +1,27 @@
|
||||
import json
|
||||
import os
|
||||
import server
|
||||
import traceback
|
||||
#Check for availability of workflow checkpointing
|
||||
import importlib
|
||||
#wcp = importlib.import_module('custom_nodes.ComfyUI-WorkflowCheckpointing.workflowcheckpointing')
|
||||
import aiohttp
|
||||
import asyncio
|
||||
from .workflowcheckpointing import post_prompt_remote
|
||||
|
||||
STATIC_AUTH_TOKEN = os.environ.get("STATIC_AUTH_TOKEN", None)
|
||||
|
||||
web = server.web
|
||||
ps = server.PromptServer.instance
|
||||
|
||||
finished_startup = False
|
||||
original_server_start = ps.start
|
||||
async def server_start(*args, **kwargs):
|
||||
args = list(args)
|
||||
original_on_start= args[3]
|
||||
def on_start(scheme, address, port):
|
||||
async def server_start(address, port, verbose=True, call_on_start=None):
|
||||
original_on_start= call_on_start
|
||||
def on_start(*args, **kwargs):
|
||||
if original_on_start is not None:
|
||||
original_on_start(scheme, address, port)
|
||||
original_on_start(*args, **kwargs)
|
||||
global finished_startup
|
||||
finished_startup= True
|
||||
args[3] = on_start
|
||||
return await original_server_start(*args, **kwargs)
|
||||
return await original_server_start(address, port, verbose, on_start)
|
||||
ps.start = server_start
|
||||
|
||||
|
||||
@@ -35,6 +39,67 @@ async def startup(request):
|
||||
@ps.routes.get("/ready")
|
||||
async def ready(request):
|
||||
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, 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
@@ -1,8 +1,8 @@
|
||||
[project]
|
||||
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."
|
||||
version = "1.0.0"
|
||||
license = "LICENSE"
|
||||
version = "1.0.1"
|
||||
license = { file = "LICENSE" }
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/AustinMroz/ComfyUI-WorkflowCheckpointing"
|
||||
|
||||
+94
-46
@@ -9,11 +9,12 @@ import threading
|
||||
import logging
|
||||
import itertools
|
||||
import hashlib
|
||||
import time
|
||||
|
||||
import comfy.samplers
|
||||
import execution
|
||||
import server
|
||||
import heapq
|
||||
import execution
|
||||
|
||||
SAMPLER_NODES = ["SamplerCustom", "KSampler", "KSamplerAdvanced", "SamplerCustomAdvanced"]
|
||||
|
||||
@@ -30,6 +31,16 @@ async def get_header():
|
||||
SALAD_TOKEN =(await r.json())['jwt']
|
||||
return {'Authorization': SALAD_TOKEN}
|
||||
|
||||
def logError(func):
|
||||
def wrapped(*args, **kwargs):
|
||||
try:
|
||||
func(*args, **kwargs)
|
||||
except:
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
raise
|
||||
return func
|
||||
|
||||
class RequestLoop:
|
||||
def __init__(self):
|
||||
self.active_request = None
|
||||
@@ -136,6 +147,7 @@ class FetchQueue:
|
||||
with self.lock:
|
||||
if item in self.consumed:
|
||||
#TODO: Also update in queue
|
||||
#TODO: if complete check etag?
|
||||
self.consumed[item][1] = min(self.consumed[item][1], priority)
|
||||
return self.consumed[item][0]
|
||||
for i in range(len(self.queue)):
|
||||
@@ -180,6 +192,9 @@ class FetchLoop:
|
||||
with self.queue.lock:
|
||||
if url in self.queue.consumed:
|
||||
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):
|
||||
return self.queue.enqueue_checked(url, priority)
|
||||
async def fetch(self, priority, url, future):
|
||||
@@ -206,9 +221,12 @@ class FetchLoop:
|
||||
return
|
||||
fetch_loop = FetchLoop()
|
||||
async def prepare_file(url, path, priority):
|
||||
hashloc = await fetch_loop.enqueue(url, priority)
|
||||
hashloc = os.path.join('fetches', string_hash(url))
|
||||
if not os.path.exists(hashloc):
|
||||
hashloc = await fetch_loop.enqueue(url, priority)
|
||||
if os.path.exists(path):
|
||||
os.remove(path)
|
||||
os.makedirs(os.path.split(path)[0], exist_ok=True)
|
||||
#TODO consider if symlinking would be better
|
||||
os.link(hashloc, path)
|
||||
|
||||
@@ -320,6 +338,9 @@ async def fetch_remote_files(remote_files, uid=None):
|
||||
|
||||
completion_futures = {}
|
||||
def add_future(json_data):
|
||||
if len(completion_futures.keys() == 0):
|
||||
#For debugging, should be changed to assert later
|
||||
return
|
||||
index = max(completion_futures.keys())
|
||||
json_data['extra_data']['completion_future'] = index
|
||||
return json_data
|
||||
@@ -346,8 +367,10 @@ async def post_prompt_remote(request):
|
||||
f = asyncio.Future()
|
||||
index = max(completion_futures.keys(),default=0)+1
|
||||
completion_futures[index] = f
|
||||
start_time = time.perf_counter()
|
||||
base_res = await original_post_prompt(request)
|
||||
outputs = await f
|
||||
execution_time = time.perf_counter() - start_time
|
||||
completion_futures.pop(index)
|
||||
if "SALAD_ORGANIZATION" in os.environ:
|
||||
async with aiohttp.ClientSession('https://storage-api.salad.com') as s:
|
||||
@@ -356,22 +379,26 @@ async def post_prompt_remote(request):
|
||||
with open(outputs[i], 'rb') as f:
|
||||
data = f.read()
|
||||
#TODO support uploads > 100MB/ memory optimizations
|
||||
fd = {'file': data}
|
||||
fd = {'file': data, 'sign': 'true'}
|
||||
url = '/'.join([base_url_path, uid, 'outputs', outputs[i]])
|
||||
async with s.put(url, headers=headers, data=fd) as r:
|
||||
await r.text()
|
||||
url = (await r.json())['url']
|
||||
outputs[i] = url
|
||||
res = base_res.text[:-1] + ', "outputs": ' + json.dumps(outputs) + '}'
|
||||
return server.web.Response(body=res)
|
||||
json_output = json.loads(base_res.text)
|
||||
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
|
||||
object.__setattr__(prompt_route, 'handler', post_prompt_remote)
|
||||
#object.__setattr__(prompt_route, 'handler', post_prompt_remote)
|
||||
|
||||
class CheckpointSampler(comfy.samplers.KSAMPLER):
|
||||
def sample(self, *args, **kwargs):
|
||||
args = list(args)
|
||||
self.unique_id = server.PromptServer.instance.last_node_id
|
||||
self.step = None
|
||||
data, metadata = checkpoint.get(self.unique_id)
|
||||
#data, metadata = checkpoint.get(self.unique_id)
|
||||
data, metadata = None, None
|
||||
if metadata is not None and 'step' in metadata:
|
||||
data = data['x']
|
||||
self.step = int(metadata['step'])
|
||||
@@ -398,42 +425,57 @@ class CheckpointSampler(comfy.samplers.KSAMPLER):
|
||||
if step == int(os.environ['FORCE_CRASH_AT']):
|
||||
raise Exception("Simulated Crash")
|
||||
|
||||
original_recursive_execute = execution.recursive_execute
|
||||
def recursive_execute_injection(*args):
|
||||
|
||||
unique_id = args[3]
|
||||
class_type = args[1][unique_id]['class_type']
|
||||
extra_data = args[4]
|
||||
if class_type in SAMPLER_NODES:
|
||||
data, metadata = checkpoint.get(unique_id)
|
||||
if metadata is not None and 'step' in metadata:
|
||||
args[1][unique_id]['inputs']['latent_image'] = ['checkpointed'+unique_id, 0]
|
||||
args[2]['checkpointed'+unique_id] = [[{'samples': data['x']}]]
|
||||
elif metadata is not None and 'completed' in metadata:
|
||||
outputs = json.loads(metadata['completed'])
|
||||
for x in range(len(outputs)):
|
||||
if outputs[x] == 'tensor':
|
||||
outputs[x] = list(data[str(x)])
|
||||
elif outputs[x] == 'latent':
|
||||
outputs[x] = [{'samples': l} for l in data[str(x)]]
|
||||
args[2][unique_id] = outputs
|
||||
return True, None, None
|
||||
|
||||
res = original_recursive_execute(*args)
|
||||
#Conditionally save node output
|
||||
#TODO: determine which non-sampler nodes are worth saving
|
||||
if class_type in SAMPLER_NODES and unique_id in args[2]:
|
||||
data = {}
|
||||
outputs = args[2][unique_id].copy()
|
||||
for x in range(len(outputs)):
|
||||
if isinstance(outputs[x][0], torch.Tensor):
|
||||
data[str(x)] = torch.stack(outputs[x])
|
||||
outputs[x] = 'tensor'
|
||||
elif isinstance(outputs[x][0], dict):
|
||||
data[str(x)] = torch.stack([l['samples'] for l in outputs[x]])
|
||||
outputs[x] = 'latent'
|
||||
checkpoint.store(unique_id, data, {'completed': json.dumps(outputs)}, priority=1)
|
||||
return res
|
||||
class CheckpointCache(execution.HierarchicalCache):
|
||||
def __init__(self, key_class):
|
||||
self.injected = {}
|
||||
super().__init__(key_class)
|
||||
def get(self, node_id):
|
||||
try:
|
||||
return super().get(node_id)
|
||||
except:
|
||||
print("oppsie-woopsi")
|
||||
if node_id in self.injected:
|
||||
return self.injected[node_id]
|
||||
#We don't want real ID. if any looping has occured, subresults must cache independently
|
||||
node = self.dynprompt.get_node(node_id)
|
||||
if node['class_type'] in SAMPLER_NODES and ('checkpointed'+node_id) not in self.injected:
|
||||
data, metadata = checkpoint.get(node_id)
|
||||
if metadata is not None and 'step' in metadata:
|
||||
#TODO: redirect dynprompt and fill cache?
|
||||
node = node.copy()
|
||||
node['inputs'] = node['inputs'].copy()
|
||||
node['inputs']['latent_image'] = ['checkpointed'+node_id, 0]
|
||||
self.dynprompt.ephemeral_prompt[node_id] = node
|
||||
self.injected['checkpointed'+node_id] = [[{'samples': data['x']}]]
|
||||
elif metadata is not None and 'completed' in metadata:
|
||||
outputs = json.loads(metadata['completed'])
|
||||
for x in range(len(outputs)):
|
||||
if outputs[x] == 'tensor':
|
||||
outputs[x] = list(data[str(x)])
|
||||
elif outputs[x] == 'latent':
|
||||
outputs[x] = [{'samples': l} for l in data[str(x)]]
|
||||
return outputs
|
||||
return super().get(node_id)
|
||||
def set(self, node_id, value):
|
||||
try:
|
||||
node = self.dynprompt.get_node(node_id)
|
||||
if node['class_type'] in SAMPLER_NODES:
|
||||
print(node_id, value)
|
||||
return super().set(node_id, value)
|
||||
breakpoint()
|
||||
data = {}
|
||||
for x in range(len(value)):
|
||||
if isinstance(value[x][0], torch.Tensor):
|
||||
data[str(x)] = torch.stack(value[x])
|
||||
value[x] = 'tensor'
|
||||
elif isinstance(value[x][0], dict):
|
||||
data[str(x)] = torch.stack([l['samples'] for l in value[x]])
|
||||
value[x] = 'latent'
|
||||
checkpoint.store(node_id, data, {'completed': json.dumps(outputs)}, priority=1)
|
||||
except:
|
||||
breakpoint()
|
||||
print("oopsie")
|
||||
return super().set(node_id, value)
|
||||
original_execute = execution.PromptExecutor.execute
|
||||
def execute_injection(*args, **kwargs):
|
||||
metadata = checkpoint.get('prompt')[1]
|
||||
@@ -455,9 +497,15 @@ def execute_injection(*args, **kwargs):
|
||||
if 'completion_future' in args[3]:
|
||||
completion_futures[args[3]['completion_future']].set_result(outputs)
|
||||
|
||||
comfy.samplers.KSAMPLER = CheckpointSampler
|
||||
execution.recursive_execute = recursive_execute_injection
|
||||
orig_reset = execution.PromptExecutor.reset
|
||||
def reset(self):
|
||||
checkpoint.reset()
|
||||
orig_reset(self)
|
||||
|
||||
execution.PromptExecutor.reset = reset
|
||||
execution.PromptExecutor.execute = execute_injection
|
||||
execution.HierarchicalCache = CheckpointCache
|
||||
comfy.samplers.KSAMPLER = CheckpointSampler
|
||||
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
Reference in New Issue
Block a user