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:
+73
-49
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user