fix always start with last output option

This commit is contained in:
Chris
2026-06-03 17:00:46 +10:00
parent 0da6b06464
commit 864429ec2a
2 changed files with 15 additions and 28 deletions
+7 -11
View File
@@ -25,17 +25,15 @@ class FilterNodeBase:
return cls._preview_image.save_images(images, **kwargs)['ui']['images']
@classmethod
def load_mask(cls, file:str, type:str="clipspace", append=" [input]") -> torch.Tensor:
def load_mask(cls, file:str|Path, type:str="clipspace", append=" [input]") -> torch.Tensor:
f = os.path.join(type, file)+append
path = folder_paths.get_annotated_filepath(f)
print(f"Loading mask from {f}")
return cls._load_image.load_image(f)[1]
@classmethod
def newest_mask_file(cls) -> Path:
def newest_mask_file(cls) -> Path|None:
dr = Path(folder_paths.get_input_directory()) / 'clipspace'
masked_files = list(dr.glob("*masked*"))
return max([f for f in masked_files], key=lambda item: item.stat().st_ctime) if masked_files else None
return max([f for f in masked_files], key=lambda item: item.stat().st_birthtime) if masked_files else None
@classmethod
def fingerprint_inputs(cls, **kwargs): # type: ignore
@@ -279,8 +277,9 @@ class MaskImageFilter(FilterNodeBase, io.ComfyNode):
mask=None, audiofile="", extra1="", extra2="", extra3="", tip="", **kwargs):
iostore = InOutStore.get_store(f"{graph_id}_{cls.hidden.unique_id}")
if if_inputs_unchanged == "Always start with last output" and iostore.have_last_output:
if iostore.check_input_tensors_congruent(image):
if (if_inputs_unchanged == "Always start with last output" and
iostore.have_last_output and
iostore.check_input_image_congruent(image)):
image, mask, extra1, extra2, extra3 = iostore.get_last_outputs()
mask = 1.0 - mask if mask is not None else None # The mask editor works in inverse
@@ -296,7 +295,6 @@ class MaskImageFilter(FilterNodeBase, io.ComfyNode):
if unchanged_in:
if if_inputs_unchanged == "Start with last output":
print("\n\nStarting with last output\n\n")
image, mask, extra1, extra2, extra3 = iostore.get_last_outputs()
mask = 1.0 - mask if mask is not None else None # The mask editor works in inverse
elif if_inputs_unchanged == "Resend last output":
@@ -327,13 +325,11 @@ class MaskImageFilter(FilterNodeBase, io.ComfyNode):
(time.monotonic()-started_waiting_at < 5)): time.sleep(1)
if (mask_file==last_mask_file):
print("Didn't get a new mask file - using input mask or image")
mask = mask if mask is not None else cls.load_mask(urls[0]['filename']+" [temp]")
else:
elif (mask_file is not None):
mask = cls.load_mask(mask_file)
if mask is None:
print("No mask file - setting blank")
mask = torch.zeros_like(image[...,0])
if if_no_mask == 'cancel' and torch.all(mask==0): raise InterruptProcessingException()
+8 -17
View File
@@ -3,8 +3,6 @@ from typing import Any
outputs_type = tuple[torch.Tensor, torch.Tensor|None, str, str, str]
def types(l:list[Any]) -> str: return ":".join(str(x.__class__) for x in l)
def make_copy(x): return x.clone() if isinstance(x, torch.Tensor) else x
class InOutStore:
@@ -32,40 +30,33 @@ class InOutStore:
return self.last_output
def update_last_outputs(self, outputs:outputs_type):
self.last_output = [ make_copy(x) for x in outputs ]
self.last_output = tuple( make_copy(x) for x in outputs ) # type: ignore
def update_last_inputs(self, *args):
self.last_inputs = [ make_copy(x) for x in args ]
def compare_with_last_inputs(self, *args) -> bool:
if self.last_inputs is None: return False # first time we've been called
if not types(self.last_inputs) == types(args):
assert False, "Called compare_with_last_inputs with different classes"
# compare the tensors
for prev, new in zip(self.last_inputs, args):
if isinstance(prev, torch.Tensor) and not torch.equal(prev, new) and not torch.equal(1-prev, new):
print("Tensor input changed")
if isinstance(prev, torch.Tensor) and isinstance(new, torch.Tensor) and not torch.equal(prev, new) and not torch.equal(1-prev, new):
return False
if (isinstance(prev, torch.Tensor) and new is None) or (isinstance(new, torch.Tensor) and prev is None):
return False
# compare the non-tensors
if self.tensor_free_hash(*args) != self.tensor_free_hash(self.last_inputs):
print("Non tensor input changed")
return False
return True
def check_input_tensors_congruent(self, *args) -> bool:
def check_input_image_congruent(self, image:torch.Tensor) -> bool:
if self.last_inputs is None: return False
for a,i in enumerate(args):
assert isinstance(a,torch.Tensor)
if self.last_input_tensors[i].shape != a.shape: return False
return True
return (image.shape == self.last_input_tensors[0].shape)
def tensor_free_hash(self, *args):
tf = ",".join( str(v) for v in flatten(args) if not isinstance(v, torch.Tensor) and v )
print(f"Hashing {tf} ... ")
return hash( tf )
return hash( ",".join( str(v) for v in flatten(args) if not isinstance(v, torch.Tensor) ))
def flatten(args) -> list[Any]:
f = []