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
This commit is contained in:
Austin Mroz
2024-06-14 04:25:50 -05:00
parent 531fa3f9f2
commit 83dde78f0c
+73 -49
View File
@@ -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