Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
dfaf02a946 | ||
|
|
fc677bb04c | ||
|
|
c286fbecd3 | ||
|
|
1b8ed8a661 | ||
|
|
8e16554331 | ||
|
|
9a42e29b49 | ||
|
|
c3f11fa212 | ||
|
|
1089d6e68f | ||
|
|
31e677ba25 | ||
|
|
0e0856ffcb | ||
|
|
58d3d9e475 | ||
|
|
6267d79a29 | ||
|
|
5d57303356 | ||
|
|
a7846f09c0 | ||
|
|
0b3f5fca9d | ||
|
|
ef3f67fcb3 | ||
|
|
9ada89ebe9 | ||
|
|
b02d301970 | ||
|
|
28a710ebf0 | ||
|
|
ec95a0c9b2 | ||
|
|
8debd2cfaf |
@@ -3,8 +3,7 @@ on:
|
||||
workflow_dispatch:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
- master
|
||||
- publish
|
||||
paths:
|
||||
- "pyproject.toml"
|
||||
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
GNU GENERAL PUBLIC LICENSE
|
||||
Version 3, 29 June 2007
|
||||
GNU AFFERO GENERAL PUBLIC LICENSE
|
||||
Version 3, 19 November 2007
|
||||
|
||||
Copyright (C) 2007 Free Software Foundation, Inc. <https://fsf.org/>
|
||||
Everyone is permitted to copy and distribute verbatim copies
|
||||
@@ -7,17 +7,15 @@
|
||||
|
||||
Preamble
|
||||
|
||||
The GNU General Public License is a free, copyleft license for
|
||||
software and other kinds of works.
|
||||
The GNU Affero General Public License is a free, copyleft license for
|
||||
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
|
||||
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
|
||||
software for all its users. We, the Free Software Foundation, use the
|
||||
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.
|
||||
software for all its users.
|
||||
|
||||
When we speak of free software, we are referring to freedom, not
|
||||
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
|
||||
free programs, and that you know you can do these things.
|
||||
|
||||
To protect your rights, we need to prevent others from denying you
|
||||
these rights or asking you to surrender the rights. Therefore, you have
|
||||
certain responsibilities if you distribute copies of the software, or if
|
||||
you modify it: responsibilities to respect the freedom of others.
|
||||
Developers that use our General Public Licenses protect your rights
|
||||
with two steps: (1) assert copyright on the software, and (2) offer
|
||||
you this License which gives you legal permission to copy, distribute
|
||||
and/or modify the software.
|
||||
|
||||
For example, if you distribute copies of such a program, whether
|
||||
gratis or for a fee, you must pass on to the recipients the same
|
||||
freedoms that you received. You must make sure that they, too, receive
|
||||
or can get the source code. And you must show them these terms so they
|
||||
know their rights.
|
||||
A secondary benefit of defending all users' freedom is that
|
||||
improvements made in alternate versions of the program, if they
|
||||
receive widespread use, become available for other developers to
|
||||
incorporate. Many developers of free software are heartened and
|
||||
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:
|
||||
(1) assert copyright on the software, and (2) offer you this License
|
||||
giving you legal permission to copy, distribute and/or modify it.
|
||||
The GNU Affero General Public License is designed specifically to
|
||||
ensure that, in such cases, the modified source code becomes available
|
||||
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
|
||||
that there is no warranty for this free software. For both users' and
|
||||
authors' sake, the GPL requires that modified versions be marked as
|
||||
changed, so that their problems will not be attributed erroneously to
|
||||
authors of previous versions.
|
||||
|
||||
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.
|
||||
An older license, called the Affero General Public License and
|
||||
published by Affero, was designed to accomplish similar goals. This is
|
||||
a different license, not a version of the Affero GPL, but Affero has
|
||||
released a new version of the Affero GPL which permits relicensing under
|
||||
this license.
|
||||
|
||||
The precise terms and conditions for copying, distribution and
|
||||
modification follow.
|
||||
@@ -72,7 +60,7 @@ modification follow.
|
||||
|
||||
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
|
||||
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
|
||||
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
|
||||
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
|
||||
License will continue to apply to the part which is the covered work,
|
||||
but the special requirements of the GNU Affero General Public License,
|
||||
section 13, concerning interaction through a network will apply to the
|
||||
combination as such.
|
||||
but the work with which it is combined will remain governed by version
|
||||
3 of the GNU General Public License.
|
||||
|
||||
14. Revised Versions of this License.
|
||||
|
||||
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
|
||||
be similar in spirit to the present version, but may differ in detail to
|
||||
the GNU Affero General Public License from time to time. Such new versions
|
||||
will be similar in spirit to the present version, but may differ in detail to
|
||||
address new problems or concerns.
|
||||
|
||||
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
|
||||
option of following the terms and conditions either of that numbered
|
||||
version or of any later version published by the Free Software
|
||||
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.
|
||||
|
||||
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
|
||||
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>
|
||||
|
||||
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
|
||||
the Free Software Foundation, either version 3 of the License, or
|
||||
it under the terms of the GNU Affero General Public License as published
|
||||
by the Free Software Foundation, either version 3 of the License, or
|
||||
(at your option) any later version.
|
||||
|
||||
This program is distributed in the hope that it will be useful,
|
||||
but WITHOUT ANY WARRANTY; without even the implied warranty of
|
||||
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/>.
|
||||
|
||||
Also add information on how to contact you by electronic and paper mail.
|
||||
|
||||
If the program does terminal interaction, make it output a short
|
||||
notice like this when it starts in an interactive mode:
|
||||
|
||||
<program> Copyright (C) <year> <name of author>
|
||||
This program comes with ABSOLUTELY NO WARRANTY; for details type `show w'.
|
||||
This is free software, and you are welcome to redistribute it
|
||||
under certain conditions; type `show c' for details.
|
||||
|
||||
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".
|
||||
If your software can interact with users remotely through a computer
|
||||
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
|
||||
interface could display a "Source" link that leads users to an archive
|
||||
of the code. There are many ways you could offer source, and different
|
||||
solutions will be better for different programs; see section 13 for the
|
||||
specific requirements.
|
||||
|
||||
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.
|
||||
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/>.
|
||||
|
||||
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
@@ -1,4 +1,4 @@
|
||||
from . import workflowcheckpointing
|
||||
from . import workflowcheckpointing, orchestration
|
||||
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
@@ -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
@@ -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"
|
||||
|
||||
+234
-73
@@ -9,10 +9,12 @@ import threading
|
||||
import logging
|
||||
import itertools
|
||||
import hashlib
|
||||
import time
|
||||
|
||||
import comfy.samplers
|
||||
import execution
|
||||
import server
|
||||
import heapq
|
||||
|
||||
SAMPLER_NODES = ["SamplerCustom", "KSampler", "KSamplerAdvanced", "SamplerCustomAdvanced"]
|
||||
|
||||
@@ -23,6 +25,7 @@ async def get_header():
|
||||
return {'Salad-Api-Key': os.environ['SALAD_API_KEY']}
|
||||
global SALAD_TOKEN
|
||||
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 session.get('http://169.254.169.254:80/v1/token') as r:
|
||||
SALAD_TOKEN =(await r.json())['jwt']
|
||||
@@ -66,10 +69,8 @@ class RequestLoop:
|
||||
async with s.delete(url, headers=await get_header()) as r:
|
||||
await r.text()
|
||||
async def _reset(self, s, uid):
|
||||
base_url = '/organizations/' + ORGANIZATION +'/files'
|
||||
checkpoint_base = '/'.join([base_url, uid, 'checkpoint'])
|
||||
checkpoint_base = 'https://storage-api.salad.com' + checkpoint_base
|
||||
async with s.get(base_url, headers=await get_header()) as r:
|
||||
async with s.get(base_url_path, headers=await get_header()) as r:
|
||||
js = await r.json()
|
||||
files = js['files']
|
||||
checkpoints = list(filter(lambda x: x['url'].startswith(checkpoint_base), files))
|
||||
@@ -82,30 +83,147 @@ class RequestLoop:
|
||||
async def process_requests(self):
|
||||
headers = await get_header()
|
||||
async with aiohttp.ClientSession('https://storage-api.salad.com') as session:
|
||||
while True:
|
||||
if self.do_reset != False:
|
||||
await self._reset(session, self.do_reset)
|
||||
self.do_reset = False
|
||||
if self.active_request is None:
|
||||
await asyncio.sleep(.1)
|
||||
else:
|
||||
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()
|
||||
try:
|
||||
while True:
|
||||
if self.do_reset != False:
|
||||
await self._reset(session, self.do_reset)
|
||||
self.do_reset = False
|
||||
if self.active_request is None:
|
||||
await asyncio.sleep(.1)
|
||||
else:
|
||||
if self.low is not None:
|
||||
self.active_request = self.low
|
||||
self.low = None
|
||||
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:
|
||||
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)
|
||||
if ORGANIZATION is not None:
|
||||
base_url_path = '/organizations/' + ORGANIZATION +'/files'
|
||||
base_url = 'https://storage-api.salad.com' + base_url_path
|
||||
class NetCheckpoint:
|
||||
def __init__(self):
|
||||
self.requestloop = RequestLoop()
|
||||
@@ -139,10 +257,12 @@ class NetCheckpoint:
|
||||
if unique_id is not None:
|
||||
if os.path.exists(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
|
||||
os.makedirs("input/checkpoint", exist_ok=True)
|
||||
for file in os.listdir("input/checkpoint"):
|
||||
os.remove(os.path.join("input/checkpoint", file))
|
||||
fetch_loop.reset('/'.join([base_url, self.uid, 'checkpoint', file]))
|
||||
|
||||
class FileCheckpoint:
|
||||
def store(self, unique_id, tensors, metadata, priority=0):
|
||||
@@ -178,54 +298,84 @@ def file_hash(filename):
|
||||
while n := f.readinto(b):
|
||||
h.update(b)
|
||||
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):
|
||||
#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:
|
||||
checkpoint_base = '/'.join([base_url, uid, 'checkpoint'])
|
||||
async with s.get(base_url, headers=await get_header()) as r:
|
||||
js = await r.json()
|
||||
files = js['files']
|
||||
checkpoints = list(filter(lambda x: x['url'].startswith(checkpoint_base), files))
|
||||
for cp in checkpoints:
|
||||
cp['filepath'] = os.path.join('checkpoint',
|
||||
cp['url'][len(checkpoint_base)+1:])
|
||||
remote_files = itertools.chain(remote_files, checkpoints)
|
||||
semaphore = asyncio.Semaphore(5)
|
||||
fetches = [asyncio.create_task(fetch_remote_file(s, f, semaphore)) for f in remote_files]
|
||||
if len(fetches) > 0:
|
||||
await asyncio.gather(*fetches)
|
||||
if uid is not None:
|
||||
checkpoint_base = '/'.join([base_url_path, uid, 'checkpoint'])
|
||||
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()
|
||||
files = js['files']
|
||||
checkpoints = list(filter(lambda x: x['url'].startswith(checkpoint_base), files))
|
||||
for cp in checkpoints:
|
||||
cp['filepath'] = os.path.join('input/checkpoint',
|
||||
cp['url'][len(checkpoint_base)+1:])
|
||||
remote_files = itertools.chain(remote_files, checkpoints)
|
||||
fetches = []
|
||||
for f in remote_files:
|
||||
fetches.append(fetch_remote_file(f['url'],f['filepath'], f.get('file_hash', None)))
|
||||
if len(fetches) > 0:
|
||||
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',
|
||||
server.PromptServer.instance.routes))
|
||||
original_post_prompt = prompt_route.handler
|
||||
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()
|
||||
if "SALAD_ORGANIZATION" in os.environ:
|
||||
extra_data = json_data.get("extra_data", {})
|
||||
#NOTE: Rendered obsolete by existing infrastructure, can be pruned
|
||||
remote_files = extra_data.get("remote_files", [])
|
||||
uid = json_data.get("client_id", 'local')
|
||||
checkpoint.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
|
||||
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']):
|
||||
raise Exception("Simulated Crash")
|
||||
|
||||
original_recursive_execute = execution.recursive_execute
|
||||
original_recursive_execute = execution.execute
|
||||
def recursive_execute_injection(*args):
|
||||
|
||||
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]
|
||||
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:
|
||||
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']}]]
|
||||
args[1].get_node(unique_id)['inputs']['latent_image'] = ['checkpointed'+unique_id, 0]
|
||||
args[2].outputs.set('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)):
|
||||
@@ -288,15 +428,15 @@ def recursive_execute_injection(*args):
|
||||
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
|
||||
args[2].outputs.set(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]:
|
||||
if class_type in SAMPLER_NODES and args[2].outputs.get(unique_id) is not None:
|
||||
data = {}
|
||||
outputs = args[2][unique_id].copy()
|
||||
outputs = args[2].outputs.get(unique_id).copy()
|
||||
for x in range(len(outputs)):
|
||||
if isinstance(outputs[x][0], torch.Tensor):
|
||||
data[str(x)] = torch.stack(outputs[x])
|
||||
@@ -306,9 +446,30 @@ def recursive_execute_injection(*args):
|
||||
outputs[x] = 'latent'
|
||||
checkpoint.store(unique_id, data, {'completed': json.dumps(outputs)}, priority=1)
|
||||
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
|
||||
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