WIP cache rewrite

A more correct implementation that instead extends the caching subsystem
of execution should be used, but is still work in progress
This commit is contained in:
Austin Mroz
2024-10-17 14:09:22 -05:00
parent dfaf02a946
commit 76d4d8e91e
2 changed files with 78 additions and 43 deletions
+2 -3
View File
@@ -1,6 +1,5 @@
from . import workflowcheckpointing, orchestration
from . import workflowcheckpointing
NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {}
WEB_DIRECTORY = "./web"
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS", "WEB_DIRECTORY"]
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+76 -40
View File
@@ -12,9 +12,9 @@ import hashlib
import time
import comfy.samplers
import execution
import server
import heapq
import execution
SAMPLER_NODES = ["SamplerCustom", "KSampler", "KSamplerAdvanced", "SamplerCustomAdvanced"]
@@ -31,6 +31,16 @@ async def get_header():
SALAD_TOKEN =(await r.json())['jwt']
return {'Authorization': SALAD_TOKEN}
def logError(func):
def wrapped(*args, **kwargs):
try:
func(*args, **kwargs)
except:
import traceback
traceback.print_exc()
raise
return func
class RequestLoop:
def __init__(self):
self.active_request = None
@@ -328,6 +338,9 @@ async def fetch_remote_files(remote_files, uid=None):
completion_futures = {}
def add_future(json_data):
if len(completion_futures.keys() == 0):
#For debugging, should be changed to assert later
return
index = max(completion_futures.keys())
json_data['extra_data']['completion_future'] = index
return json_data
@@ -377,14 +390,15 @@ async def post_prompt_remote(request):
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)
#object.__setattr__(prompt_route, 'handler', post_prompt_remote)
class CheckpointSampler(comfy.samplers.KSAMPLER):
def sample(self, *args, **kwargs):
args = list(args)
self.unique_id = server.PromptServer.instance.last_node_id
self.step = None
data, metadata = checkpoint.get(self.unique_id)
#data, metadata = checkpoint.get(self.unique_id)
data, metadata = None, None
if metadata is not None and 'step' in metadata:
data = data['x']
self.step = int(metadata['step'])
@@ -411,41 +425,57 @@ class CheckpointSampler(comfy.samplers.KSAMPLER):
if step == int(os.environ['FORCE_CRASH_AT']):
raise Exception("Simulated Crash")
original_recursive_execute = execution.execute
def recursive_execute_injection(*args):
unique_id = args[3]
class_type = args[1].get_node(unique_id)['class_type']
extra_data = args[4]
if class_type in SAMPLER_NODES:
data, metadata = checkpoint.get(unique_id)
if metadata is not None and 'step' in metadata:
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)):
if outputs[x] == 'tensor':
outputs[x] = list(data[str(x)])
elif outputs[x] == 'latent':
outputs[x] = [{'samples': l} for l in data[str(x)]]
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 args[2].outputs.get(unique_id) is not None:
data = {}
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])
outputs[x] = 'tensor'
elif isinstance(outputs[x][0], dict):
data[str(x)] = torch.stack([l['samples'] for l in outputs[x]])
outputs[x] = 'latent'
checkpoint.store(unique_id, data, {'completed': json.dumps(outputs)}, priority=1)
return res
class CheckpointCache(execution.HierarchicalCache):
def __init__(self, key_class):
self.injected = {}
super().__init__(key_class)
def get(self, node_id):
try:
return super().get(node_id)
except:
print("oppsie-woopsi")
if node_id in self.injected:
return self.injected[node_id]
#We don't want real ID. if any looping has occured, subresults must cache independently
node = self.dynprompt.get_node(node_id)
if node['class_type'] in SAMPLER_NODES and ('checkpointed'+node_id) not in self.injected:
data, metadata = checkpoint.get(node_id)
if metadata is not None and 'step' in metadata:
#TODO: redirect dynprompt and fill cache?
node = node.copy()
node['inputs'] = node['inputs'].copy()
node['inputs']['latent_image'] = ['checkpointed'+node_id, 0]
self.dynprompt.ephemeral_prompt[node_id] = node
self.injected['checkpointed'+node_id] = [[{'samples': data['x']}]]
elif metadata is not None and 'completed' in metadata:
outputs = json.loads(metadata['completed'])
for x in range(len(outputs)):
if outputs[x] == 'tensor':
outputs[x] = list(data[str(x)])
elif outputs[x] == 'latent':
outputs[x] = [{'samples': l} for l in data[str(x)]]
return outputs
return super().get(node_id)
def set(self, node_id, value):
try:
node = self.dynprompt.get_node(node_id)
if node['class_type'] in SAMPLER_NODES:
print(node_id, value)
return super().set(node_id, value)
breakpoint()
data = {}
for x in range(len(value)):
if isinstance(value[x][0], torch.Tensor):
data[str(x)] = torch.stack(value[x])
value[x] = 'tensor'
elif isinstance(value[x][0], dict):
data[str(x)] = torch.stack([l['samples'] for l in value[x]])
value[x] = 'latent'
checkpoint.store(node_id, data, {'completed': json.dumps(outputs)}, priority=1)
except:
breakpoint()
print("oopsie")
return super().set(node_id, value)
original_execute = execution.PromptExecutor.execute
def execute_injection(*args, **kwargs):
metadata = checkpoint.get('prompt')[1]
@@ -467,9 +497,15 @@ def execute_injection(*args, **kwargs):
if 'completion_future' in args[3]:
completion_futures[args[3]['completion_future']].set_result(outputs)
comfy.samplers.KSAMPLER = CheckpointSampler
execution.execute = recursive_execute_injection
orig_reset = execution.PromptExecutor.reset
def reset(self):
checkpoint.reset()
orig_reset(self)
execution.PromptExecutor.reset = reset
execution.PromptExecutor.execute = execute_injection
execution.HierarchicalCache = CheckpointCache
comfy.samplers.KSAMPLER = CheckpointSampler
NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {}