fix always start with last output option
This commit is contained in:
+7
-11
@@ -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
@@ -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 = []
|
||||
|
||||
Reference in New Issue
Block a user