From 83dde78f0c2d26b0f1151131815f7a16b741688c Mon Sep 17 00:00:00 2001 From: Austin Mroz Date: Fri, 14 Jun 2024 04:25:50 -0500 Subject: [PATCH] Improved isolation of saving checkpoints. In preperation for saving checkpoints remotely, all checkpoint transactions are now wrapped. Fixed step accumulation on consecutive failures Support all builtin samplers --- checkpointsampling.py | 122 +++++++++++++++++++++++++----------------- 1 file changed, 73 insertions(+), 49 deletions(-) diff --git a/checkpointsampling.py b/checkpointsampling.py index 6ece638..2842365 100644 --- a/checkpointsampling.py +++ b/checkpointsampling.py @@ -7,74 +7,98 @@ import comfy.samplers import execution import server +SAMPLER_NODES = ["SamplerCustom", "KSampler", "KSamplerAdvanced", "SamplerCustomAdvanced"] + +def store_checkpoint(unique_id, data, partial_progress_counter=-2, is_tensor=True): + """Swappable interface for saving checkpoints. + Implementation must be transactional: Either the whole thing completes, + or the prior checkpoint must be valid even if crash occurs mid execution""" + file = f"checkpoint/{unique_id}.{partial_progress_counter%2}.safetensor" + if is_tensor: + comfy.utils.save_torch_file({'x': data},file) + else: + with open(file, "w") as f: + json.dump(data,f) + with open(f"checkpoint/{unique_id}.txt", "w") as f: + f.write(str(partial_progress_counter)+"\n"+str(1*is_tensor)) +def get_checkpoint(unique_id): + """Returns the information previously saved""" + partial_progress = -2 + is_tensor = True + if os.path.exists(f"checkpoint/{unique_id}.txt"): + with open(f"checkpoint/{unique_id}.txt", "r") as f: + partial_progress = int(f.readline()) + is_tensor = f.readline() != '0' + file = f"checkpoint/{unique_id}.{partial_progress%2}.safetensor" + if not os.path.exists(file): + return None, None + if is_tensor: + res = safetensors.torch.load_file(file)['x'] + else: + with open(file, "r") as f: + res = json.load(f) + return partial_progress, res +def reset_checkpoints(unique_id=None): + """Clear all checkpoint information.""" + if unique_id is not None: + if os.path.exists(f"checkpoint/{unique_id}.0.safetensor"): + os.remove(f"checkpoint/{unique_id}.0.safetensor") + if os.path.exists(f"checkpoint/{unique_id}.1.safetensor"): + os.remove(f"checkpoint/{unique_id}.1.safetensor") + if os.path.exists(f"checkpoint/{unique_id}.txt"): + os.remove(f"checkpoint/{unique_id}.txt") + return + for file in os.listdir("checkpoint"): + os.remove(os.path.join("checkpoint", file)) + class CheckpointSampler(comfy.samplers.KSAMPLER): def sample(self, *args, **kwargs): args = list(args) self.unique_id = server.PromptServer.instance.last_node_id - step = None - with open(f"checkpoint/{self.unique_id}.json", "r") as f: - f.readline() - while line := f.readline(): - step = int(line) - if step is not None: + self.step, data = get_checkpoint(self.unique_id) + if self.step is not None: #checkpoint of execution exists - args[5] = safetensors.torch.load_file(f"checkpoint/{self.unique_id}.{step%2}.latent")['latent_tensor'].to(args[4].device) - args[1] = args[1][step:] + args[5] = data.to(args[4].device) + args[1] = args[1][self.step:] #disable added noise, as the checkpointed latent is already noised args[4][:] = 0 - original_callback = kwargs.get("callback", None) + original_callback = args[3] def callback(*args): - self.callback(args) + self.callback(*args) if original_callback is not None: return original_callback(*args) args[3] = callback res = super().sample(*args, **kwargs) - #cleanup checkpoints as execution has completed - #TODO: Have tracking for freshness and keep checkpoints post execution? - if os.path.exists(f"checkpoint/{self.unique_id}.0.latent"): - os.remove(f"checkpoint/{self.unique_id}.0.latent") - if os.path.exists(f"checkpoint/{self.unique_id}.1.latent"): - os.remove(f"checkpoint/{self.unique_id}.1.latent") - if os.path.exists(f"checkpoint/{self.unique_id}.json"): - os.remove(f"checkpoint/{self.unique_id}.json") + reset_checkpoints(self.unique_id) return res - def callback(self, args): - x = args[2] - print(x.sum()) - step = args[0] - comfy.utils.save_torch_file({"latent_tensor": x}, f"checkpoint/{self.unique_id}.{step%2}.latent") - with open(f"checkpoint/{self.unique_id}.json", "a") as f: - f.write("\n"+str(step)) - if step == 3: - if args[3] == 10: - raise Exception("test crash") - #manual exception raising + def callback(self, step, denoised, x, total_steps): + if self.step is not None: + step += self.step + store_checkpoint(self.unique_id, x, step) + original_recursive_execute = execution.recursive_execute def recursive_execute_injection(*args): + unique_id = args[3] - prev_prompt = None - if args[1][unique_id]['class_type'] == "SamplerCustom": - if os.path.exists(f"checkpoint/{unique_id}.json"): - with open(f"checkpoint/{unique_id}.json", "r") as f: - prev_prompt = json.loads(f.readline()) - if f.readline() == '': - #execution interrupted before checkpoint made - prev_prompt = None - #TODO: double check this is deep compare - if prev_prompt == args[1]: - #NOTE: to avoid precision lost on rescaling/ease of implementation, - #the tensor is loaded twice, this first load is just for dimensions, - #so the actual index doesn't matter - x = safetensors.torch.load_file(f"checkpoint/{unique_id}.0.latent")['latent_tensor'] + class_type = args[1][unique_id]['class_type'] + if len(args[5]) == 0: + _, prev_prompt = get_checkpoint('prompt') + if prev_prompt != args[1]: + reset_checkpoints() + store_checkpoint('prompt', args[1], is_tensor=False) + if class_type in SAMPLER_NODES: + step, x = get_checkpoint(unique_id) + if step is not None and step>0: args[1][unique_id]['inputs']['latent_image'] = ['checkpointed'+unique_id, 0] args[2]['checkpointed'+unique_id] = [[{'samples': x}]] - else: - with open(f"checkpoint/{unique_id}.json", "w") as f: - json.dump(args[1],f) - return original_recursive_execute(*args) - - + res = original_recursive_execute(*args) + #Conditionally save node output + #TODO: determine which non-sampler nodes are worth saving + if class_type in SAMPLER_NODES: + pass + #output = args[2][unique_id] + return res comfy.samplers.KSAMPLER = CheckpointSampler execution.recursive_execute = recursive_execute_injection