From 943cb53ccb736be7db21b04f7a88f5c688f1664a Mon Sep 17 00:00:00 2001 From: Daniel Martinek Date: Thu, 26 Mar 2026 10:17:08 +0100 Subject: [PATCH] add ability to remember stack during batching for other nodes --- more_math/ClipMathNode.py | 2 +- more_math/ConditioningMathNode.py | 15 ++++++++++++--- more_math/ImageMathNode.py | 16 +++++++++++++--- more_math/LatentMathNode.py | 16 +++++++++++++--- more_math/MaskMathNode.py | 16 +++++++++++++--- more_math/NoiseMathNode.py | 18 ++++++++++++++---- 6 files changed, 66 insertions(+), 17 deletions(-) diff --git a/more_math/ClipMathNode.py b/more_math/ClipMathNode.py index d9a20a2..ac0383d 100644 --- a/more_math/ClipMathNode.py +++ b/more_math/ClipMathNode.py @@ -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) ], diff --git a/more_math/ConditioningMathNode.py b/more_math/ConditioningMathNode.py index 3d7f4a9..aa3fe1d 100644 --- a/more_math/ConditioningMathNode.py +++ b/more_math/ConditioningMathNode.py @@ -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) diff --git a/more_math/ImageMathNode.py b/more_math/ImageMathNode.py index 3cbb1ab..d53a6c9 100644 --- a/more_math/ImageMathNode.py +++ b/more_math/ImageMathNode.py @@ -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) diff --git a/more_math/LatentMathNode.py b/more_math/LatentMathNode.py index 50cfb3c..97b7e28 100644 --- a/more_math/LatentMathNode.py +++ b/more_math/LatentMathNode.py @@ -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) diff --git a/more_math/MaskMathNode.py b/more_math/MaskMathNode.py index 14527ec..49456b0 100644 --- a/more_math/MaskMathNode.py +++ b/more_math/MaskMathNode.py @@ -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) diff --git a/more_math/NoiseMathNode.py b/more_math/NoiseMathNode.py index db77460..3b0fedb 100644 --- a/more_math/NoiseMathNode.py +++ b/more_math/NoiseMathNode.py @@ -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: