142
This commit is contained in:
+29
-13
@@ -207,22 +207,33 @@ def mask_from_data(data) -> torch.Tensor:
|
||||
return mask.unsqueeze(0)
|
||||
|
||||
class InOutStore:
|
||||
previous_inputs:list[Any] = []
|
||||
last_output:tuple[torch.Tensor, torch.Tensor|None, str, str, str]|None = None
|
||||
|
||||
stores:dict[str, "InOutStore"] = {}
|
||||
@classmethod
|
||||
def check_input_unchanged(cls, *args) -> bool:
|
||||
def get_store(cls, graph_id:str) -> "InOutStore":
|
||||
if graph_id not in cls.stores:
|
||||
cls.stores[graph_id] = InOutStore()
|
||||
return cls.stores[graph_id]
|
||||
|
||||
def __init__(self):
|
||||
self.previous_inputs:list[Any] = []
|
||||
self.last_output:tuple[torch.Tensor, torch.Tensor|None, str, str, str]|None = None
|
||||
|
||||
def get_last(self) -> tuple[torch.Tensor, torch.Tensor|None, str, str, str]:
|
||||
assert self.last_output is not None, "No last output stored"
|
||||
return self.last_output
|
||||
|
||||
def check_input_unchanged(self, *args) -> bool:
|
||||
def make_copy(x): return x.clone() if isinstance(x, torch.Tensor) else x
|
||||
try:
|
||||
if len(cls.previous_inputs)!=len(args): return False
|
||||
for prev, new in zip(cls.previous_inputs, args):
|
||||
if len(self.previous_inputs)!=len(args): return False
|
||||
for prev, new in zip(self.previous_inputs, args):
|
||||
if isinstance(prev, torch.Tensor) and isinstance(new, torch.Tensor):
|
||||
if not torch.equal(prev, new): return False
|
||||
else:
|
||||
if prev != new: return False
|
||||
return True
|
||||
finally:
|
||||
cls.previous_inputs = [ make_copy(x) for x in args ]
|
||||
self.previous_inputs = [ make_copy(x) for x in args ]
|
||||
|
||||
class MaskImageFilter(io.ComfyNode, FilterNodeBase):
|
||||
@classmethod
|
||||
@@ -240,7 +251,10 @@ class MaskImageFilter(io.ComfyNode, FilterNodeBase):
|
||||
io.String.Input("extra1", default="", optional=True),
|
||||
io.String.Input("extra2", default="", optional=True),
|
||||
io.String.Input("extra3", default="", optional=True),
|
||||
io.String.Input("graph_id", default="")
|
||||
io.String.Input("graph_id", default=""),
|
||||
],
|
||||
hidden=[
|
||||
io.Hidden.unique_id,
|
||||
],
|
||||
outputs = [
|
||||
io.Image.Output("image", display_name="image"),
|
||||
@@ -254,13 +268,15 @@ class MaskImageFilter(io.ComfyNode, FilterNodeBase):
|
||||
|
||||
@classmethod
|
||||
def execute(cls, image, timeout, if_no_mask, graph_id, if_inputs_unchanged="Run normally", mask=None, extra1="", extra2="", extra3="", tip="", **kwargs): # type: ignore
|
||||
iostore = InOutStore.get_store(f"{graph_id}_{cls.hidden.unique_id}")
|
||||
|
||||
# check if everything is unchanged (and store these inputs for next check)
|
||||
if InOutStore.check_input_unchanged(image, timeout, if_no_mask, graph_id, mask, extra1, extra2, extra3, tip) and InOutStore.last_output is not None:
|
||||
if iostore.check_input_unchanged(image, timeout, if_no_mask, graph_id, mask, extra1, extra2, extra3, tip) and iostore.last_output is not None:
|
||||
if if_inputs_unchanged == "Start with last output":
|
||||
image, mask, extra1, extra2, extra3 = InOutStore.last_output
|
||||
image, mask, extra1, extra2, extra3 = iostore.get_last()
|
||||
mask = 1.0 - mask if mask is not None else None # The mask editor works in inverse
|
||||
elif if_inputs_unchanged == "Resend last output":
|
||||
return io.NodeOutput( *InOutStore.last_output )
|
||||
return io.NodeOutput( *iostore.get_last() )
|
||||
|
||||
if mask is not None and mask.shape[:3] == image.shape[:3] and not torch.all(mask==0):
|
||||
input_to_send = torch.cat((image, mask.unsqueeze(-1)), dim=-1)
|
||||
@@ -284,10 +300,10 @@ class MaskImageFilter(io.ComfyNode, FilterNodeBase):
|
||||
if mask is None: mask = torch.zeros_like(image[...,0])
|
||||
if if_no_mask == 'cancel' and torch.all(mask==0): raise InterruptProcessingException()
|
||||
|
||||
InOutStore.last_output = ( image.clone(), mask.clone(), *response.get_extras((extra1, extra2, extra3)) )
|
||||
iostore.last_output = ( image.clone(), mask.clone(), *response.get_extras((extra1, extra2, extra3)) )
|
||||
if (image.shape[0:3] != mask.shape[0:3]):
|
||||
print(f"Mask shape {mask.shape} does not match image shape {image.shape}")
|
||||
return io.NodeOutput( *InOutStore.last_output )
|
||||
return io.NodeOutput( *iostore.get_last() )
|
||||
|
||||
@classmethod
|
||||
def fingerprint_inputs(cls, **kwargs) -> Any:
|
||||
|
||||
Reference in New Issue
Block a user