3 Commits
Author SHA1 Message Date
Austin Mroz 8a9e5d56cb Remove errant on_prompt handler
The orchestration code included functionality to track when execution
had completed which is not useful outside of orchestration. As the other
references to completion had been removed, this code was causing
exceptions during execution.
2024-10-17 14:57:44 -05:00
Austin Mroz 535ea76a2a Extract just the execution code for publishing
A minimal change set was made to allow the code to function after
changes to the execution model. The changes unrelated to workflow
orchestration have been pulled out so that the code is viable for
general use once again.
2024-10-17 14:45:25 -05:00
Austin Mroz e03d06d959 Fix license reference 2024-08-01 13:35:39 -05:00
3 changed files with 190 additions and 79 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"
+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.1.1"
license = "LICENSE" license = { file = "LICENSE" }
[project.urls] [project.urls]
Repository = "https://github.com/AustinMroz/ComfyUI-WorkflowCheckpointing" Repository = "https://github.com/AustinMroz/ComfyUI-WorkflowCheckpointing"
+187 -75
View File
@@ -13,6 +13,7 @@ import hashlib
import comfy.samplers import comfy.samplers
import execution import execution
import server import server
import heapq
SAMPLER_NODES = ["SamplerCustom", "KSampler", "KSamplerAdvanced", "SamplerCustomAdvanced"] SAMPLER_NODES = ["SamplerCustom", "KSampler", "KSamplerAdvanced", "SamplerCustomAdvanced"]
@@ -23,6 +24,7 @@ async def get_header():
return {'Salad-Api-Key': os.environ['SALAD_API_KEY']} return {'Salad-Api-Key': os.environ['SALAD_API_KEY']}
global SALAD_TOKEN global SALAD_TOKEN
if SALAD_TOKEN is None: if SALAD_TOKEN is None:
assert 'SALAD_MACHINE_ID' in os.environ, "SALAD_API_KEY must be provided if not deployed"
async with aiohttp.ClientSession() as session: async with aiohttp.ClientSession() as session:
async with session.get('http://169.254.169.254:80/v1/token') as r: async with session.get('http://169.254.169.254:80/v1/token') as r:
SALAD_TOKEN =(await r.json())['jwt'] SALAD_TOKEN =(await r.json())['jwt']
@@ -66,10 +68,8 @@ class RequestLoop:
async with s.delete(url, headers=await get_header()) as r: async with s.delete(url, headers=await get_header()) as r:
await r.text() await r.text()
async def _reset(self, s, uid): async def _reset(self, s, uid):
base_url = '/organizations/' + ORGANIZATION +'/files'
checkpoint_base = '/'.join([base_url, uid, 'checkpoint']) checkpoint_base = '/'.join([base_url, uid, 'checkpoint'])
checkpoint_base = 'https://storage-api.salad.com' + checkpoint_base async with s.get(base_url_path, headers=await get_header()) as r:
async with s.get(base_url, headers=await get_header()) as r:
js = await r.json() js = await r.json()
files = js['files'] files = js['files']
checkpoints = list(filter(lambda x: x['url'].startswith(checkpoint_base), files)) checkpoints = list(filter(lambda x: x['url'].startswith(checkpoint_base), files))
@@ -82,30 +82,147 @@ class RequestLoop:
async def process_requests(self): async def process_requests(self):
headers = await get_header() headers = await get_header()
async with aiohttp.ClientSession('https://storage-api.salad.com') as session: async with aiohttp.ClientSession('https://storage-api.salad.com') as session:
while True: try:
if self.do_reset != False: while True:
await self._reset(session, self.do_reset) if self.do_reset != False:
self.do_reset = False await self._reset(session, self.do_reset)
if self.active_request is None: self.do_reset = False
await asyncio.sleep(.1) if self.active_request is None:
else: await asyncio.sleep(.1)
req = self.active_request
fd = aiohttp.FormData({'file': req[1]})
async with session.put(req[0], headers=headers, data=fd) as r:
#We don't care about result, but must still await it
await r.text()
with self.mutex:
if not self.queue_high.empty():
self.active_request = self.queue_high.get()
else: else:
if self.low is not None: req = self.active_request
self.active_request = self.low fd = aiohttp.FormData({'file': req[1]})
self.low = None async with session.put(req[0], headers=headers, data=fd) as r:
#We don't care about result, but must still await it
await r.text()
with self.mutex:
if not self.queue_high.empty():
self.active_request = self.queue_high.get()
else: else:
self.active_request = None if self.low is not None:
self.active_request = self.low
self.low = None
else:
self.active_request = None
except:
#Exceptions from event loop get swallowed and kill the loop
import traceback
traceback.print_exc()
raise
class FetchQueue:
"""Modified priority queue implementation that tracks inflight and allows priority modification"""
def __init__(self):
self.lock = threading.RLock()
self.queue = []# queue contains priority, url, future
self.count = 0
self.consumed = {}
self.new_items = asyncio.Event()
def update_priority(self, i, priority):
#lock must already be acquired
future = self.queue[i][3]
if priority < self.queue[i][0]:
#priority is increased, invalidate old
self.queue[i] = (self.queue[i][0], self.queue[i][1], None, None)
heapq.heappush(self.queue, (priority, self.count, item, future))
self.count += 1
def requeue(self, future, item, dec_priority=1):
with self.lock:
priority = self.consumed[item][1] - dec_priority
heapq.heappush(self.queue, (priority, self.count, future, None))
self.count += 1
self.new_items.set()
def enqueue_checked(self, item, priority):
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)):
if self.queue[i][2] == item:
future = self.queue[i][3]
self.update_priority(i, priority)
return future
future = asyncio.Future()
heapq.heappush(self.queue, (priority, self.count, item, future))
self.count += 1
self.new_items.set()
return future
async def get(self):
while True:
await self.new_items.wait()
with self.lock:
priority, _, item, future = heapq.heappop(self.queue)
if len(self.queue) == 0:
self.new_items.clear()
if item is not None:
if isinstance(item, str):
self.consumed[item] = [future, priority]
return priority, item, future
else:
#item is future
item.set_result(True)
class FetchLoop:
def __init__(self):
self.queue = FetchQueue()
self.semaphore = asyncio.Semaphore(5)
self.cs = aiohttp.ClientSession()
event_loop = server.PromptServer.instance.loop
self.process_loop = event_loop.create_task(self.loop())
os.makedirs("fetches", exist_ok=True)
async def loop(self):
event_loop = server.PromptServer.instance.loop
while True:
await self.semaphore.acquire()
event_loop.create_task(self.fetch(*(await self.queue.get())))
def reset(self, url):
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):
chunk_size = 2**25 #32MB
headers = {}
if url.startswith(base_url):
headers.update(await get_header())
filename = os.path.join('fetches', string_hash(url))
try:
async with self.cs.get(url, headers=headers) as r:
with open(filename, 'wb') as f:
async for chunk in r.content.iter_chunked(chunk_size):
f.write(chunk)
if not r.content.is_eof():
awaken = asyncio.Future()
self.queue.requeue(awaken, url)
await awaken
future.set_result(filename)
except:
future.set_result(None)
raise
finally:
self.semaphore.release()
return
fetch_loop = FetchLoop()
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)
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)
ORGANIZATION = os.environ.get('SALAD_ORGANIZATION', None) ORGANIZATION = os.environ.get('SALAD_ORGANIZATION', None)
if ORGANIZATION is not None:
base_url_path = '/organizations/' + ORGANIZATION +'/files'
base_url = 'https://storage-api.salad.com' + base_url_path
class NetCheckpoint: class NetCheckpoint:
def __init__(self): def __init__(self):
self.requestloop = RequestLoop() self.requestloop = RequestLoop()
@@ -139,10 +256,12 @@ class NetCheckpoint:
if unique_id is not None: if unique_id is not None:
if os.path.exists(f"input/checkpoint/{unique_id}.checkpoint"): if os.path.exists(f"input/checkpoint/{unique_id}.checkpoint"):
os.remove(f"input/checkpoint/{unique_id}.checkpoint") os.remove(f"input/checkpoint/{unique_id}.checkpoint")
fetch_loop.reset('/'.join([base_url, self.uid, 'checkpoint', f'{unique_id}.checkpoint']))
return return
os.makedirs("input/checkpoint", exist_ok=True) os.makedirs("input/checkpoint", exist_ok=True)
for file in os.listdir("input/checkpoint"): for file in os.listdir("input/checkpoint"):
os.remove(os.path.join("input/checkpoint", file)) os.remove(os.path.join("input/checkpoint", file))
fetch_loop.reset('/'.join([base_url, self.uid, 'checkpoint', file]))
class FileCheckpoint: class FileCheckpoint:
def store(self, unique_id, tensors, metadata, priority=0): def store(self, unique_id, tensors, metadata, priority=0):
@@ -178,49 +297,46 @@ def file_hash(filename):
while n := f.readinto(b): while n := f.readinto(b):
h.update(b) h.update(b)
return h.hexdigest() return h.hexdigest()
def string_hash(s):
h = hashlib.sha256()
h.update(s.encode('utf-8'))
return h.hexdigest()
def fetch_remote_file(url, filepath, file_hash=None):
assert filepath.find("..") == -1, "Paths may not contain .."
return prepare_file(url, filepath, -1)
async def fetch_remote_file(session, file, semaphore):
filename = os.path.join("input", file['filepath'])
assert filename.find("..") == -1, "Paths may not contain .."
if os.path.exists(filename) and 'hash' in file and file_hash(filename) == file['hash']:
return
if file['url'].startswith('https://storage-api.salad.com/'):
headers = await get_header()
else:
headers = {}
async with semaphore:
async with session.get(file['url'], headers=headers) as r:
with open(filename, 'wb') as fd:
async for chunk in r.content.iter_chunked(2**16):
fd.write(chunk)
async def fetch_remote_files(remote_files, uid=None): async def fetch_remote_files(remote_files, uid=None):
#TODO: Add requested support for zip files? #TODO: Add requested support for zip files?
async with aiohttp.ClientSession() as s: if uid is not None:
base_url = 'https://storage-api.salad.com/organizations/' + ORGANIZATION +'/files' checkpoint_base = '/'.join([base_url_path, uid, 'checkpoint'])
if uid is not None: checkpoint_base = 'https://storage-api.salad.com'+ checkpoint_base
checkpoint_base = '/'.join([base_url, uid, 'checkpoint']) async with fetch_loop.cs.get(base_url, headers=await get_header()) as r:
async with s.get(base_url, headers=await get_header()) as r: js = await r.json()
js = await r.json() files = js['files']
files = js['files'] checkpoints = list(filter(lambda x: x['url'].startswith(checkpoint_base), files))
checkpoints = list(filter(lambda x: x['url'].startswith(checkpoint_base), files)) for cp in checkpoints:
for cp in checkpoints: cp['filepath'] = os.path.join('input/checkpoint',
cp['filepath'] = os.path.join('checkpoint', cp['url'][len(checkpoint_base)+1:])
cp['url'][len(checkpoint_base)+1:]) remote_files = itertools.chain(remote_files, checkpoints)
remote_files = itertools.chain(remote_files, checkpoints) fetches = []
semaphore = asyncio.Semaphore(5) for f in remote_files:
fetches = [asyncio.create_task(fetch_remote_file(s, f, semaphore)) for f in remote_files] fetches.append(fetch_remote_file(f['url'],f['filepath'], f.get('file_hash', None)))
if len(fetches) > 0: if len(fetches) > 0:
await asyncio.gather(*fetches) await asyncio.gather(*fetches)
prompt_route = next(filter(lambda x: x.path == '/prompt' and x.method == 'POST', prompt_route = next(filter(lambda x: x.path == '/prompt' and x.method == 'POST',
server.PromptServer.instance.routes)) server.PromptServer.instance.routes))
original_post_prompt = prompt_route.handler original_post_prompt = prompt_route.handler
async def post_prompt_remote(request): async def post_prompt_remote(request):
if 'dump_req' in os.environ:
with open('resp-dump.txt', 'wb') as f:
f.write(await request.read())
import sys
sys.exit()
json_data = await request.json() json_data = await request.json()
if "SALAD_ORGANIZATION" in os.environ: if "SALAD_ORGANIZATION" in os.environ:
extra_data = json_data.get("extra_data", {}) extra_data = json_data.get("extra_data", {})
#NOTE: Rendered obsolete by existing infrastructure, can be pruned
remote_files = extra_data.get("remote_files", []) remote_files = extra_data.get("remote_files", [])
uid = json_data.get("client_id", 'local') uid = json_data.get("client_id", 'local')
checkpoint.uid = uid checkpoint.uid = uid
@@ -261,26 +377,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 'checkpoints' in extra_data:
checkpoint.update(extra_data.pop('checkpoints'))
if 'prompt_checked' not in args[4]:
metadata = checkpoint.get('prompt')[1]
if metadata is None or json.loads(metadata['prompt']) != args[1]:
checkpoint.reset()
checkpoint.store('prompt', {'x': torch.ones(1)},
{'prompt': json.dumps(args[1])}, priority=2)
args[4]['prompt_checked'] = True
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)):
@@ -288,15 +394,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])
@@ -306,9 +412,15 @@ def recursive_execute_injection(*args):
outputs[x] = 'latent' outputs[x] = 'latent'
checkpoint.store(unique_id, data, {'completed': json.dumps(outputs)}, priority=1) checkpoint.store(unique_id, data, {'completed': json.dumps(outputs)}, priority=1)
return res return res
original_execute = execution.PromptExecutor.execute
def execute_injection(*args, **kwargs):
metadata = checkpoint.get('prompt')[1]
if metadata is None or json.loads(metadata['prompt']) != args[1]:
checkpoint.reset()
checkpoint.store('prompt', {'x': torch.ones(1)},
{'prompt': json.dumps(args[1])}, priority=2)
original_execute(*args, **kwargs)
comfy.samplers.KSAMPLER = CheckpointSampler comfy.samplers.KSAMPLER = CheckpointSampler
execution.recursive_execute = recursive_execute_injection execution.execute = recursive_execute_injection
execution.PromptExecutor.execute = execute_injection
NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {}