21 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
Austin Mroz 6267d79a29 Swap license to AGPLV3
As the new code is more targeted to the interests of Banodoco and made
with the express purpose of being run as hosted server code, The license
has been changed to AGLPV3.
2024-06-25 21:05:03 -05:00
Austin Mroz 5d57303356 Fix mistaken index on priority update 2024-06-25 20:22:50 -05:00
Austin Mroz a7846f09c0 Re-enable net checkpointing under new fetch code 2024-06-25 19:39:47 -05:00
Austin Mroz 0b3f5fca9d Include orchestration, rework execution wrapping
Orchestration code (mostly endpoints for probes) was intended as a
seperate plugin, but given the tight tying of asynchronous dependency
code, it has been included in this history and will likely be branched
off as a dedicated repo not available for non client/server
architectures.

much of the execution wrapping code has been moved into a wrapper on
execute instead of recursive executed. This cleans up logic and fixes
issues with collecting all outputs when multiple output nodes exist
2024-06-25 15:49:21 -05:00
Austin Mroz ef3f67fcb3 Rework large file shunking
Rather than requesting as chunks, the stream is kept open and re-added
to the queue every 32MB. This removes the need for the server to report
ranges, or content length, or to recombine chunks, which greatly
simplifies logic.
2024-06-25 14:28:21 -05:00
Austin Mroz 9ada89ebe9 Fix chunk_index for big files, allow dump reqs
Added an environment variable to dump a prompt request and exit, for
easier testing.

Code for big files is executing, but I'm noticing but it produces
output files that are not correct and has incorrectly inserted bytes.

Will either disable or swap to different method of keeping the request
but re-acquiring the connection semaphore
2024-06-25 12:29:52 -05:00
Austin Mroz b02d301970 Implement standalone priority queue for fetches 2024-06-25 03:24:00 -05:00
Austin Mroz 28a710ebf0 Include output files in response
The list of output files needs to bubble back up to the async request
code so that it can be included in the response.
2024-06-24 18:25:37 -05:00
Austin Mroz ec95a0c9b2 Add output tracking.
Will need additional pass to figure out sending the result
2024-06-24 13:22:18 -05:00
Austin Mroz 8debd2cfaf Implement priority system for fetches. 2024-06-24 00:41:22 -05:00
6 changed files with 408 additions and 156 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"
+65 -78
View File
@@ -1,5 +1,5 @@
GNU GENERAL PUBLIC LICENSE GNU AFFERO GENERAL PUBLIC LICENSE
Version 3, 29 June 2007 Version 3, 19 November 2007
Copyright (C) 2007 Free Software Foundation, Inc. <https://fsf.org/> Copyright (C) 2007 Free Software Foundation, Inc. <https://fsf.org/>
Everyone is permitted to copy and distribute verbatim copies Everyone is permitted to copy and distribute verbatim copies
@@ -7,17 +7,15 @@
Preamble Preamble
The GNU General Public License is a free, copyleft license for The GNU Affero General Public License is a free, copyleft license for
software and other kinds of works. software and other kinds of works, specifically designed to ensure
cooperation with the community in the case of network server software.
The licenses for most software and other practical works are designed The licenses for most software and other practical works are designed
to take away your freedom to share and change the works. By contrast, to take away your freedom to share and change the works. By contrast,
the GNU General Public License is intended to guarantee your freedom to our General Public Licenses are intended to guarantee your freedom to
share and change all versions of a program--to make sure it remains free share and change all versions of a program--to make sure it remains free
software for all its users. We, the Free Software Foundation, use the software for all its users.
GNU General Public License for most of our software; it applies also to
any other work released this way by its authors. You can apply it to
your programs, too.
When we speak of free software, we are referring to freedom, not When we speak of free software, we are referring to freedom, not
price. Our General Public Licenses are designed to make sure that you price. Our General Public Licenses are designed to make sure that you
@@ -26,44 +24,34 @@ them if you wish), that you receive source code or can get it if you
want it, that you can change the software or use pieces of it in new want it, that you can change the software or use pieces of it in new
free programs, and that you know you can do these things. free programs, and that you know you can do these things.
To protect your rights, we need to prevent others from denying you Developers that use our General Public Licenses protect your rights
these rights or asking you to surrender the rights. Therefore, you have with two steps: (1) assert copyright on the software, and (2) offer
certain responsibilities if you distribute copies of the software, or if you this License which gives you legal permission to copy, distribute
you modify it: responsibilities to respect the freedom of others. and/or modify the software.
For example, if you distribute copies of such a program, whether A secondary benefit of defending all users' freedom is that
gratis or for a fee, you must pass on to the recipients the same improvements made in alternate versions of the program, if they
freedoms that you received. You must make sure that they, too, receive receive widespread use, become available for other developers to
or can get the source code. And you must show them these terms so they incorporate. Many developers of free software are heartened and
know their rights. encouraged by the resulting cooperation. However, in the case of
software used on network servers, this result may fail to come about.
The GNU General Public License permits making a modified version and
letting the public access it on a server without ever releasing its
source code to the public.
Developers that use the GNU GPL protect your rights with two steps: The GNU Affero General Public License is designed specifically to
(1) assert copyright on the software, and (2) offer you this License ensure that, in such cases, the modified source code becomes available
giving you legal permission to copy, distribute and/or modify it. to the community. It requires the operator of a network server to
provide the source code of the modified version running there to the
users of that server. Therefore, public use of a modified version, on
a publicly accessible server, gives the public access to the source
code of the modified version.
For the developers' and authors' protection, the GPL clearly explains An older license, called the Affero General Public License and
that there is no warranty for this free software. For both users' and published by Affero, was designed to accomplish similar goals. This is
authors' sake, the GPL requires that modified versions be marked as a different license, not a version of the Affero GPL, but Affero has
changed, so that their problems will not be attributed erroneously to released a new version of the Affero GPL which permits relicensing under
authors of previous versions. this license.
Some devices are designed to deny users access to install or run
modified versions of the software inside them, although the manufacturer
can do so. This is fundamentally incompatible with the aim of
protecting users' freedom to change the software. The systematic
pattern of such abuse occurs in the area of products for individuals to
use, which is precisely where it is most unacceptable. Therefore, we
have designed this version of the GPL to prohibit the practice for those
products. If such problems arise substantially in other domains, we
stand ready to extend this provision to those domains in future versions
of the GPL, as needed to protect the freedom of users.
Finally, every program is threatened constantly by software patents.
States should not allow patents to restrict development and use of
software on general-purpose computers, but in those that do, we wish to
avoid the special danger that patents applied to a free program could
make it effectively proprietary. To prevent this, the GPL assures that
patents cannot be used to render the program non-free.
The precise terms and conditions for copying, distribution and The precise terms and conditions for copying, distribution and
modification follow. modification follow.
@@ -72,7 +60,7 @@ modification follow.
0. Definitions. 0. Definitions.
"This License" refers to version 3 of the GNU General Public License. "This License" refers to version 3 of the GNU Affero General Public License.
"Copyright" also means copyright-like laws that apply to other kinds of "Copyright" also means copyright-like laws that apply to other kinds of
works, such as semiconductor masks. works, such as semiconductor masks.
@@ -549,35 +537,45 @@ to collect a royalty for further conveying from those to whom you convey
the Program, the only way you could satisfy both those terms and this the Program, the only way you could satisfy both those terms and this
License would be to refrain entirely from conveying the Program. License would be to refrain entirely from conveying the Program.
13. Use with the GNU Affero General Public License. 13. Remote Network Interaction; Use with the GNU General Public License.
Notwithstanding any other provision of this License, if you modify the
Program, your modified version must prominently offer all users
interacting with it remotely through a computer network (if your version
supports such interaction) an opportunity to receive the Corresponding
Source of your version by providing access to the Corresponding Source
from a network server at no charge, through some standard or customary
means of facilitating copying of software. This Corresponding Source
shall include the Corresponding Source for any work covered by version 3
of the GNU General Public License that is incorporated pursuant to the
following paragraph.
Notwithstanding any other provision of this License, you have Notwithstanding any other provision of this License, you have
permission to link or combine any covered work with a work licensed permission to link or combine any covered work with a work licensed
under version 3 of the GNU Affero General Public License into a single under version 3 of the GNU General Public License into a single
combined work, and to convey the resulting work. The terms of this combined work, and to convey the resulting work. The terms of this
License will continue to apply to the part which is the covered work, License will continue to apply to the part which is the covered work,
but the special requirements of the GNU Affero General Public License, but the work with which it is combined will remain governed by version
section 13, concerning interaction through a network will apply to the 3 of the GNU General Public License.
combination as such.
14. Revised Versions of this License. 14. Revised Versions of this License.
The Free Software Foundation may publish revised and/or new versions of The Free Software Foundation may publish revised and/or new versions of
the GNU General Public License from time to time. Such new versions will the GNU Affero General Public License from time to time. Such new versions
be similar in spirit to the present version, but may differ in detail to will be similar in spirit to the present version, but may differ in detail to
address new problems or concerns. address new problems or concerns.
Each version is given a distinguishing version number. If the Each version is given a distinguishing version number. If the
Program specifies that a certain numbered version of the GNU General Program specifies that a certain numbered version of the GNU Affero General
Public License "or any later version" applies to it, you have the Public License "or any later version" applies to it, you have the
option of following the terms and conditions either of that numbered option of following the terms and conditions either of that numbered
version or of any later version published by the Free Software version or of any later version published by the Free Software
Foundation. If the Program does not specify a version number of the Foundation. If the Program does not specify a version number of the
GNU General Public License, you may choose any version ever published GNU Affero General Public License, you may choose any version ever published
by the Free Software Foundation. by the Free Software Foundation.
If the Program specifies that a proxy can decide which future If the Program specifies that a proxy can decide which future
versions of the GNU General Public License can be used, that proxy's versions of the GNU Affero General Public License can be used, that proxy's
public statement of acceptance of a version permanently authorizes you public statement of acceptance of a version permanently authorizes you
to choose that version for the Program. to choose that version for the Program.
@@ -635,40 +633,29 @@ the "copyright" line and a pointer to where the full notice is found.
Copyright (C) <year> <name of author> Copyright (C) <year> <name of author>
This program is free software: you can redistribute it and/or modify This program is free software: you can redistribute it and/or modify
it under the terms of the GNU General Public License as published by it under the terms of the GNU Affero General Public License as published
the Free Software Foundation, either version 3 of the License, or by the Free Software Foundation, either version 3 of the License, or
(at your option) any later version. (at your option) any later version.
This program is distributed in the hope that it will be useful, This program is distributed in the hope that it will be useful,
but WITHOUT ANY WARRANTY; without even the implied warranty of but WITHOUT ANY WARRANTY; without even the implied warranty of
MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
GNU General Public License for more details. GNU Affero General Public License for more details.
You should have received a copy of the GNU General Public License You should have received a copy of the GNU Affero General Public License
along with this program. If not, see <https://www.gnu.org/licenses/>. along with this program. If not, see <https://www.gnu.org/licenses/>.
Also add information on how to contact you by electronic and paper mail. Also add information on how to contact you by electronic and paper mail.
If the program does terminal interaction, make it output a short If your software can interact with users remotely through a computer
notice like this when it starts in an interactive mode: network, you should also make sure that it provides a way for users to
get its source. For example, if your program is a web application, its
<program> Copyright (C) <year> <name of author> interface could display a "Source" link that leads users to an archive
This program comes with ABSOLUTELY NO WARRANTY; for details type `show w'. of the code. There are many ways you could offer source, and different
This is free software, and you are welcome to redistribute it solutions will be better for different programs; see section 13 for the
under certain conditions; type `show c' for details. specific requirements.
The hypothetical commands `show w' and `show c' should show the appropriate
parts of the General Public License. Of course, your program's commands
might be different; for a GUI interface, you would use an "about box".
You should also get your employer (if you work as a programmer) or school, You should also get your employer (if you work as a programmer) or school,
if any, to sign a "copyright disclaimer" for the program, if necessary. if any, to sign a "copyright disclaimer" for the program, if necessary.
For more information on this, and how to apply and follow the GNU GPL, see For more information on this, and how to apply and follow the GNU AGPL, see
<https://www.gnu.org/licenses/>. <https://www.gnu.org/licenses/>.
The GNU General Public License does not permit incorporating your program
into proprietary programs. If your program is a subroutine library, you
may consider it more useful to permit linking proprietary applications with
the library. If this is what you want to do, use the GNU Lesser General
Public License instead of this License. But first, please read
<https://www.gnu.org/licenses/why-not-lgpl.html>.
+1 -1
View File
@@ -1,4 +1,4 @@
from . import workflowcheckpointing from . import workflowcheckpointing, orchestration
NODE_CLASS_MAPPINGS = {} NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {} NODE_DISPLAY_NAME_MAPPINGS = {}
+105
View File
@@ -0,0 +1,105 @@
import json
import os
import server
import traceback
#Check for availability of workflow checkpointing
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(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(*args, **kwargs)
global finished_startup
finished_startup= True
return await original_server_start(address, port, verbose, on_start)
ps.start = server_start
@ps.routes.get("/health")
async def heath(request):
#while any of the server endpoints could likely be used
return web.json_response([])
@ps.routes.get("/startup")
async def startup(request):
if finished_startup:
return web.json_response([])
return web.Response(status=503)
@ps.routes.get("/ready")
async def ready(request):
current_queue = ps.prompt_queue.get_current_queue()
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
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"
+205 -44
View File
@@ -9,10 +9,12 @@ 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
import server import server
import heapq
SAMPLER_NODES = ["SamplerCustom", "KSampler", "KSamplerAdvanced", "SamplerCustomAdvanced"] SAMPLER_NODES = ["SamplerCustom", "KSampler", "KSamplerAdvanced", "SamplerCustomAdvanced"]
@@ -23,6 +25,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 +69,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,6 +83,7 @@ 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:
try:
while True: while True:
if self.do_reset != False: if self.do_reset != False:
await self._reset(session, self.do_reset) await self._reset(session, self.do_reset)
@@ -104,8 +106,124 @@ class RequestLoop:
self.low = None self.low = None
else: else:
self.active_request = None 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 +257,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,54 +298,84 @@ 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:
base_url = 'https://storage-api.salad.com/organizations/' + ORGANIZATION +'/files'
if uid is not None: if uid is not None:
checkpoint_base = '/'.join([base_url, uid, 'checkpoint']) checkpoint_base = '/'.join([base_url_path, uid, 'checkpoint'])
async with s.get(base_url, headers=await get_header()) as r: checkpoint_base = 'https://storage-api.salad.com'+ checkpoint_base
async with fetch_loop.cs.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('checkpoint', cp['filepath'] = os.path.join('input/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)
semaphore = asyncio.Semaphore(5) fetches = []
fetches = [asyncio.create_task(fetch_remote_file(s, f, semaphore)) for f in remote_files] 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)
completion_futures = {}
def add_future(json_data):
index = max(completion_futures.keys())
json_data['extra_data']['completion_future'] = index
return json_data
server.PromptServer.instance.add_on_prompt_handler(add_future)
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
await fetch_remote_files(remote_files, uid=uid) await fetch_remote_files(remote_files, uid=uid)
return await original_post_prompt(request) if 'prompt' not in json_data:
return server.web.json_response("PreLoad Complete")
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:
headers = await get_header()
for i in range(len(outputs)):
with open(outputs[i], 'rb') as f:
data = f.read()
#TODO support uploads > 100MB/ memory optimizations
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:
url = (await r.json())['url']
outputs[i] = url
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 #Dangerous
object.__setattr__(prompt_route, 'handler', post_prompt_remote) object.__setattr__(prompt_route, 'handler', post_prompt_remote)
@@ -261,26 +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 '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 +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])
@@ -306,9 +446,30 @@ 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)
prev_outputs = {}
os.makedirs("temp", exist_ok=True)
#TODO: Consider subdir recursing?
for item in itertools.chain(os.scandir("output"), os.scandir("temp")):
if item.is_file():
prev_outputs[item.path] = item.stat().st_mtime
original_execute(*args, **kwargs)
outputs = []
for item in itertools.chain(os.scandir("output"), os.scandir("temp")):
if item.is_file() and prev_outputs.get(item.path, 0) < item.stat().st_mtime:
outputs.append(item.path)
if 'completion_future' in args[3]:
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
NODE_CLASS_MAPPINGS = {} NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {} NODE_DISPLAY_NAME_MAPPINGS = {}