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:
|
workflow_dispatch:
|
||||||
push:
|
push:
|
||||||
branches:
|
branches:
|
||||||
- main
|
- publish
|
||||||
- master
|
|
||||||
paths:
|
paths:
|
||||||
- "pyproject.toml"
|
- "pyproject.toml"
|
||||||
|
|
||||||
|
|||||||
@@ -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
@@ -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 = {}
|
||||||
|
|||||||
@@ -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]
|
[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
@@ -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 = {}
|
||||||
|
|||||||
Reference in New Issue
Block a user