add ability to remember stack during batching for other nodes
This commit is contained in:
@@ -32,7 +32,7 @@ class CLIPMathNode(io.ComfyNode):
|
||||
options=["do nothing","error","tile", "pad"],
|
||||
display_name="on size mismatch",
|
||||
default="error",
|
||||
tooltip="How to handle mismatched layer counts. For models, this usually defaults to broadcast (zero for missing layers)."
|
||||
tooltip="How to handle mismatched layer counts."
|
||||
),
|
||||
MrmthStack.Input(id="stack", tooltip="Access stack between nodes",optional=True)
|
||||
],
|
||||
|
||||
@@ -43,6 +43,14 @@ class ConditioningMathNode(io.ComfyNode):
|
||||
tooltip="How to handle mismatched image batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero."
|
||||
),
|
||||
io.Int.Input(id="batching", default=0),
|
||||
io.Bool.Input(
|
||||
id="remember_stack",
|
||||
default=False,
|
||||
display_name="Remember stack across batch",
|
||||
tooltip=(
|
||||
"If enabled, stack is copied at output leading to changes being remembered during batch operations (node runs multiple times in sucession). If disabled each batch gets it's own copy of the stack."
|
||||
),
|
||||
),
|
||||
MrmthStack.Input(id="stack",optional=True)
|
||||
],
|
||||
outputs=[
|
||||
@@ -52,19 +60,19 @@ class ConditioningMathNode(io.ComfyNode):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def check_lazy_status(cls, Expression,Expression_pi, V, F, length_mismatch="tile", batching=0, stack={}):
|
||||
def check_lazy_status(cls, Expression,Expression_pi, V, F, length_mismatch="tile", batching=0,remember_stack=False, stack={}):
|
||||
d = checkLazyNew(Expression,V,F)
|
||||
b = checkLazyNew(Expression_pi,V,F)
|
||||
return d|b
|
||||
|
||||
|
||||
@classmethod
|
||||
def execute(cls, V, F, Expression, Expression_pi, length_mismatch="tile", batching=0, stack={}):
|
||||
def execute(cls, V, F, Expression, Expression_pi, length_mismatch="tile", batching=0,remember_stack=False, stack={}):
|
||||
# Identify all present conditioning inputs
|
||||
tensor_keys = [k for k, v in V.items() if v is not None and isinstance(v, list) and len(v) > 0]
|
||||
if not tensor_keys:
|
||||
raise ValueError("At least one input is required.")
|
||||
stack = copy.deepcopy(stack) if stack is not None else {}
|
||||
stack = stack if remember_stack else (copy.deepcopy(stack) if stack is not None else {})
|
||||
|
||||
# Extract tensors and pooled outputs
|
||||
tensors = {}
|
||||
@@ -216,4 +224,5 @@ class ConditioningMathNode(io.ComfyNode):
|
||||
base[0] = (rtensor, new_dict)
|
||||
res_list = [base]
|
||||
|
||||
stack = stack if remember_stack else copy.deepcopy(stack)
|
||||
return (res_list,stack)
|
||||
|
||||
@@ -39,6 +39,14 @@ class ImageMathNode(io.ComfyNode):
|
||||
tooltip="How to handle mismatched image batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero."
|
||||
),
|
||||
io.Int.Input(id="batching", default=0),
|
||||
io.Bool.Input(
|
||||
id="remember_stack",
|
||||
default=False,
|
||||
display_name="Remember stack across batch",
|
||||
tooltip=(
|
||||
"If enabled, stack is copied at output leading to changes being remembered during batch operations (node runs multiple times in sucession). If disabled each batch gets it's own copy of the stack."
|
||||
),
|
||||
),
|
||||
MrmthStack.Input(id="stack", tooltip="Access stack between nodes",optional=True)
|
||||
],
|
||||
outputs=[
|
||||
@@ -48,11 +56,11 @@ class ImageMathNode(io.ComfyNode):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def check_lazy_status(cls, Expression, V, F, length_mismatch="tile",batching=0,stack={}):
|
||||
def check_lazy_status(cls, Expression, V, F, length_mismatch="tile",batching=0,remember_stack=False,stack={}):
|
||||
return checkLazyNew(Expression,V,F)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, V, F, Expression, length_mismatch="error",batching=0,stack={}):
|
||||
def execute(cls, V, F, Expression, length_mismatch="error",batching=0,remember_stack=False,stack={}):
|
||||
# I and F are Autogrow.Type which is dict[str, Any]
|
||||
|
||||
# Identify all present tensors and their keys
|
||||
@@ -61,7 +69,7 @@ class ImageMathNode(io.ComfyNode):
|
||||
raise ValueError("At least one input is required.")
|
||||
|
||||
tensors = [V[k] for k in tensor_keys]
|
||||
stack = copy.deepcopy(stack) if stack is not None else {}
|
||||
stack = stack if remember_stack else (copy.deepcopy(stack) if stack is not None else {})
|
||||
|
||||
# Normalize all tensors together to find the common target shape
|
||||
normalized_tensors = normalize_to_common_shape(*tensors, mode=length_mismatch)
|
||||
@@ -139,6 +147,8 @@ class ImageMathNode(io.ComfyNode):
|
||||
res_list = []
|
||||
for result_chunk in res:
|
||||
res_list.append(result_chunk)
|
||||
stack = stack if remember_stack else copy.deepcopy(stack)
|
||||
return (res_list, stack)
|
||||
else:
|
||||
stack = stack if remember_stack else copy.deepcopy(stack)
|
||||
return ([result], stack)
|
||||
|
||||
@@ -49,6 +49,14 @@ class LatentMathNode(io.ComfyNode):
|
||||
tooltip="How to handle mismatched latent batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero."
|
||||
),
|
||||
io.Int.Input(id="batching"),
|
||||
io.Bool.Input(
|
||||
id="remember_stack",
|
||||
default=False,
|
||||
display_name="Remember stack across batch",
|
||||
tooltip=(
|
||||
"If enabled, stack is copied at output leading to changes being remembered during batch operations (node runs multiple times in sucession). If disabled each batch gets it's own copy of the stack."
|
||||
),
|
||||
),
|
||||
MrmthStack.Input(id="stack", tooltip="Access stack between nodes",optional=True)
|
||||
],
|
||||
outputs=[
|
||||
@@ -61,11 +69,11 @@ class LatentMathNode(io.ComfyNode):
|
||||
tooltip = cleandoc(__doc__)
|
||||
|
||||
@classmethod
|
||||
def check_lazy_status(cls, Expression, V, F,batching, length_mismatch="tile",stack={}):
|
||||
def check_lazy_status(cls, Expression, V, F,batching, length_mismatch="tile",remember_stack=False,stack={}):
|
||||
return checkLazyNew(Expression,V,F)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, V, F, Expression,batching, length_mismatch="tile",stack={}) -> io.NodeOutput:
|
||||
def execute(cls, V, F, Expression,batching, length_mismatch="tile",remember_stack=False,stack={}) -> io.NodeOutput:
|
||||
# Determine reference latent
|
||||
ref_latent = None
|
||||
for lat in V.values():
|
||||
@@ -75,7 +83,7 @@ class LatentMathNode(io.ComfyNode):
|
||||
|
||||
if ref_latent is None:
|
||||
raise ValueError("At least one input is required.")
|
||||
stack = copy.deepcopy(stack) if stack is not None else {}
|
||||
stack = stack if remember_stack else (copy.deepcopy(stack) if stack is not None else {})
|
||||
|
||||
# Identify if any input is a NestedTensor and track original sizes for restoration
|
||||
stacked = False
|
||||
@@ -206,7 +214,9 @@ class LatentMathNode(io.ComfyNode):
|
||||
else:
|
||||
rl["samples"] = result_t
|
||||
results1.append(rl)
|
||||
stack = stack if remember_stack else copy.deepcopy(stack)
|
||||
return (results1,stack)
|
||||
rl = result_latent.copy()
|
||||
rl["samples"] = result_t
|
||||
stack = stack if remember_stack else copy.deepcopy(stack)
|
||||
return ([rl],stack)
|
||||
|
||||
@@ -39,6 +39,14 @@ class MaskMathNode(io.ComfyNode):
|
||||
tooltip="How to handle mismatched mask batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero."
|
||||
),
|
||||
io.Int.Input(id="batching", default=0),
|
||||
io.Bool.Input(
|
||||
id="remember_stack",
|
||||
default=False,
|
||||
display_name="Remember stack across batch",
|
||||
tooltip=(
|
||||
"If enabled, stack is copied at output leading to changes being remembered during batch operations (node runs multiple times in sucession). If disabled each batch gets it's own copy of the stack."
|
||||
),
|
||||
),
|
||||
MrmthStack.Input(id="stack", tooltip="Access stack between nodes",optional=True)
|
||||
],
|
||||
outputs=[
|
||||
@@ -48,18 +56,18 @@ class MaskMathNode(io.ComfyNode):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def check_lazy_status(cls, Expression, V, F, length_mismatch="tile",batching=0,stack={}):
|
||||
def check_lazy_status(cls, Expression, V, F, length_mismatch="tile",batching=0,remember_stack=False,stack={}):
|
||||
return checkLazyNew(Expression,V,F)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, V, F, Expression, length_mismatch="tile",batching=0,stack={}):
|
||||
def execute(cls, V, F, Expression, length_mismatch="tile",batching=0,remember_stack=False,stack={}):
|
||||
# Identify all present tensors and their keys
|
||||
tensor_keys = [k for k, v in V.items() if v is not None]
|
||||
if not tensor_keys:
|
||||
raise ValueError("At least one input is required.")
|
||||
|
||||
tensors = [V[k] for k in tensor_keys]
|
||||
stack = copy.deepcopy(stack) if stack is not None else {}
|
||||
stack = stack if remember_stack else (copy.deepcopy(stack) if stack is not None else {})
|
||||
# Normalize all tensors together
|
||||
normalized_tensors = normalize_to_common_shape(*tensors, mode=length_mismatch)
|
||||
V_norm = dict(zip(tensor_keys, normalized_tensors))
|
||||
@@ -132,6 +140,8 @@ class MaskMathNode(io.ComfyNode):
|
||||
res_list = []
|
||||
for result_chunk in res:
|
||||
res_list.append(result_chunk)
|
||||
stack = stack if remember_stack else copy.deepcopy(stack)
|
||||
return (res_list, stack)
|
||||
else:
|
||||
stack = stack if remember_stack else copy.deepcopy(stack)
|
||||
return ([result], stack)
|
||||
|
||||
@@ -36,6 +36,14 @@ class NoiseMathNode(io.ComfyNode):
|
||||
types=[io.String,MrmthParseTree],
|
||||
tooltip="Expression for noise",
|
||||
),
|
||||
io.Bool.Input(
|
||||
id="remember_stack",
|
||||
default=False,
|
||||
display_name="Remember stack across batch",
|
||||
tooltip=(
|
||||
"If enabled, stack is copied at output leading to changes being remembered during batch operations (node runs multiple times in sucession). If disabled each batch gets it's own copy of the stack."
|
||||
),
|
||||
),
|
||||
MrmthStack.Input(id="stack", tooltip="Access stack between nodes",optional=True)
|
||||
],
|
||||
outputs=[
|
||||
@@ -45,13 +53,15 @@ class NoiseMathNode(io.ComfyNode):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def check_lazy_status(cls, Noise, V, F,stack={}):
|
||||
def check_lazy_status(cls, Noise, V, F,remember_stack=False,stack={}):
|
||||
return checkLazyNew(Noise,V,F)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, Noise, V,F,stack={}):
|
||||
stack = copy.deepcopy(stack) if stack is not None else {}
|
||||
return (NoiseExecutor(V,F, Noise,stack),stack)
|
||||
def execute(cls, Noise, V,F,remember_stack=False,stack={}):
|
||||
stack = stack if remember_stack else (copy.deepcopy(stack) if stack is not None else {})
|
||||
executer = NoiseExecutor(V,F, Noise,stack)
|
||||
stack = stack if remember_stack else copy.deepcopy(stack)
|
||||
return (executer,stack)
|
||||
|
||||
|
||||
class NoiseExecutor:
|
||||
|
||||
Reference in New Issue
Block a user