Implement remote storage of checkpoints of S4

Working implementation of external checkpoints using Salad Simple Storage
Service
This commit is contained in:
Austin Mroz
2024-06-17 22:06:10 -05:00
parent 2d61e214b3
commit 8f34bf21dd
+86 -12
View File
@@ -7,6 +7,8 @@ import asyncio
import queue
import threading
import logging
import itertools
import hashlib
import comfy.samplers
import execution
@@ -14,6 +16,18 @@ import server
SAMPLER_NODES = ["SamplerCustom", "KSampler", "KSamplerAdvanced", "SamplerCustomAdvanced"]
SALAD_TOKEN = None
async def get_header():
if 'SALAD_API_KEY' in os.environ:
#NOTE: Only for local testing. Do not add to container
return {'Salad-Api-Key': os.environ['SALAD_API_KEY']}
global SALAD_TOKEN
if SALAD_TOKEN is None:
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']
return {'Authorization': SALAD_TOKEN}
class RequestLoop:
def __init__(self):
self.active_request = None
@@ -42,10 +56,7 @@ class RequestLoop:
else:
self.low = req
async def process_requests(self):
async with aiohttp.ClientSession() as session:
async with session.get('http://169.254.169.254:80/v1/token') as r:
token = await r.json()['jwt']
headers = {'Authorization': token}
headers = await get_header()
async with aiohttp.ClientSession('https://storage-api.salad.com') as session:
while True:
if self.active_request is None:
@@ -53,7 +64,9 @@ class RequestLoop:
else:
try:
req = self.active_request
async with session.put(req[0], headers=headers, **req[1]) as r:
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()
except asyncio.CancelledError:
@@ -69,23 +82,22 @@ class RequestLoop:
else:
self.active_request = None
#Placeholder, would need to be pulled from salad
ORGANIZATION = "banodoco"
MACHINEID = "local"
ORGANIZATION = os.environ.get('SALAD_ORGANIZATION', None)
class NetCheckpoint:
def __init__(self):
self.requestloop = RequestLoop()
self.has_warned_size = False
assert ORGANIZATION is not None
def store(self, unique_id, tensors, metadata, priority=0):
file = "/".join(['', ORGANIZATION, MACHINEID, "checkpoint", f"{unique_id}.checkpoint"])
file = "/" + "/".join(['organizations', ORGANIZATION, 'files', self.uid,
"checkpoint", f"{unique_id}.checkpoint"])
data = safetensors.torch.save(tensors, metadata)
if len(data) > 10 ** 8:
if not self.has_warned_size:
logging.warning("Checkpoint is too large and has been skipped")
self.has_warned_size = True
return
self.requestloop.queue((file, {'data': data}), priority)
self.requestloop.queue((file, data), priority)
def get(self, unique_id):
"""Returns the information previously saved
If the request has checkpointed data, this information should
@@ -99,10 +111,12 @@ class NetCheckpoint:
return tensors, metadata
def reset(self, unique_id=None):
"""Clear all checkpoint information."""
#TODO: Find way to prune remote checkpoints
if unique_id is not None:
if os.path.exists(f"input/checkpoint/{unique_id}.checkpoint"):
os.remove(f"input/checkpoint/{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))
@@ -131,7 +145,61 @@ class FileCheckpoint:
for file in os.listdir("checkpoint"):
os.remove(os.path.join("checkpoint", file))
checkpoint = NetCheckpoint() if "USE_NET_CHECKPOINTING" in os.environ else FileCheckpoint()
checkpoint = NetCheckpoint() if "SALAD_ORGANIZATION" in os.environ else FileCheckpoint()
def file_hash(filename):
h = hashlib.sha256()
b = bytearray(10*1024*1024) # read 10 megabytes at a time
with open(filename, 'rb', buffering=0) as f:
while n := f.readinto(b):
h.update(b)
return h.hexdigest()
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)
@server.PromptServer.instance.routes.post("/prompt_remote")
async def post_prompt(request):
json_data = await request.json()
if "prompt" in json_data and "extra_data" in json_data and "SALAD_ORGANIZATION" in os.environ:
extra_data = json_data["extra_data"]
remote_files = extra_data.get("remote_files", [])
uid = extra_data.get("uid", 'local')
checkpoint.uid = uid
await fetch_remote_files(remote_files, uid=uid)
#probably a bad idea
return await server.PromptServer.instance.routes[15].handler(request)
class CheckpointSampler(comfy.samplers.KSAMPLER):
def sample(self, *args, **kwargs):
@@ -161,12 +229,18 @@ class CheckpointSampler(comfy.samplers.KSAMPLER):
step += self.step
data = safetensors.torch.save
checkpoint.store(self.unique_id, {'x':x}, {'step':str(step)})
if self.step is None and "FORCE_CRASH_AT" in os.environ:
if step == int(os.environ['FORCE_CRASH_AT']):
raise Exception("Simulated Crash")
original_recursive_execute = execution.recursive_execute
def recursive_execute_injection(*args):
unique_id = args[3]
class_type = args[1][unique_id]['class_type']
extra_data = args[4]
if 'checkpoints' in extra_data:
checkpoint.update(extra_data.pop('checkpoints'))
#Imperfect, is checked for each bubble down step
#Only applied once, but has unnecessary loads
if len(args[5]) == 0: