Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8a9e5d56cb | ||
|
|
535ea76a2a | ||
|
|
e03d06d959 |
@@ -3,8 +3,7 @@ on:
|
|||||||
workflow_dispatch:
|
workflow_dispatch:
|
||||||
push:
|
push:
|
||||||
branches:
|
branches:
|
||||||
- main
|
- publish
|
||||||
- master
|
|
||||||
paths:
|
paths:
|
||||||
- "pyproject.toml"
|
- "pyproject.toml"
|
||||||
|
|
||||||
|
|||||||
+2
-2
@@ -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
@@ -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 = {}
|
|
||||||
|
|||||||
Reference in New Issue
Block a user