add ability to remember stack during batching for other nodes

This commit is contained in:
Daniel Martinek
2026-03-26 10:17:08 +01:00
parent 10deecb7b0
commit 943cb53ccb
6 changed files with 66 additions and 17 deletions
+1 -1
View File
@@ -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)
],
+12 -3
View File
@@ -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)
+13 -3
View File
@@ -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)
+13 -3
View File
@@ -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)
+13 -3
View File
@@ -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)
+14 -4
View File
@@ -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: