From eae6b03f25bfc15005bc418aae50dd3e59465dd8 Mon Sep 17 00:00:00 2001 From: mcDandy Date: Mon, 2 Feb 2026 12:48:13 +0100 Subject: [PATCH] Add stack passing --- more_math/AudioMathNode.py | 13 ++++++++----- more_math/ClipMathNode.py | 13 ++++++++----- more_math/ConditioningMathNode.py | 14 ++++++++------ more_math/FloatMathNode.py | 11 +++++++---- more_math/GuiderMathNode.py | 16 ++++++++++------ more_math/ImageMathNode.py | 13 ++++++++----- more_math/LatentMathNode.py | 16 ++++++++++------ more_math/MaskMathNode.py | 13 ++++++++----- more_math/ModelMathNode.py | 15 +++++++++------ more_math/NoiseMathNode.py | 15 +++++++++------ more_math/SigmasMathNode.py | 13 ++++++++----- more_math/Stack.py | 13 +++++++++++++ more_math/VaeMathNode.py | 14 +++++++++----- more_math/VideoMathNode.py | 16 +++++++++------- more_math/modelLikeCommon.py | 5 ++--- 15 files changed, 126 insertions(+), 74 deletions(-) create mode 100644 more_math/Stack.py diff --git a/more_math/AudioMathNode.py b/more_math/AudioMathNode.py index 741f76c..4d37bb9 100644 --- a/more_math/AudioMathNode.py +++ b/more_math/AudioMathNode.py @@ -14,6 +14,7 @@ from antlr4 import InputStream, CommonTokenStream from .Parser.MathExprLexer import MathExprLexer from .Parser.MathExprParser import MathExprParser import re +from .Stack import MrmthStack class AudioMathNode(io.ComfyNode): """ @@ -40,15 +41,17 @@ class AudioMathNode(io.ComfyNode): options=["tile", "error", "pad"], default="error", tooltip="How to handle mismatched image batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero." - ) + ), + MrmthStack.Input(id="stack", tooltip="Access stack between nodes",optional=True) ], outputs=[ io.Audio.Output(), + MrmthStack.Output(), ], ) @classmethod - def check_lazy_status(cls, Expression, V, F, length_mismatch="tile"): + def check_lazy_status(cls, Expression, V, F, length_mismatch="tile",stack=[]): input_stream = InputStream(Expression) lexer = MathExprLexer(input_stream) @@ -80,7 +83,7 @@ class AudioMathNode(io.ComfyNode): return needed1 @classmethod - def execute(cls, V, F, Expression, length_mismatch="tile"): + def execute(cls, V, F, Expression, length_mismatch="tile",stack=[]): # Identify all present audio inputs and their keys tensor_keys = [k for k, v in V.items() if v is not None and isinstance(v, dict) and "waveform" in v] if not tensor_keys: @@ -145,7 +148,7 @@ class AudioMathNode(io.ComfyNode): variables[k] = val if val is not None else 0.0 tree = parse_expr(Expression); - visitor = UnifiedMathVisitor(variables, a_w.shape,a_w.device) + visitor = UnifiedMathVisitor(variables, a_w.shape,a_w.device,state_storage=stack) result = visitor.visit(tree) result = as_tensor(result, a_w.shape) - return ({"waveform":result,"sample_rate":sample_rate},) + return ({"waveform":result,"sample_rate":sample_rate},stack) diff --git a/more_math/ClipMathNode.py b/more_math/ClipMathNode.py index 9de1a5c..0c5d7ca 100644 --- a/more_math/ClipMathNode.py +++ b/more_math/ClipMathNode.py @@ -4,6 +4,7 @@ from antlr4 import InputStream, CommonTokenStream from .Parser.MathExprLexer import MathExprLexer from .Parser.MathExprParser import MathExprParser import re +from .Stack import MrmthStack class CLIPMathNode(io.ComfyNode): @@ -26,17 +27,19 @@ class CLIPMathNode(io.ComfyNode): options=["tile", "error", "pad"], default="error", tooltip="How to handle mismatched layer counts. For models, this usually defaults to broadcast (zero for missing layers)." - ) + ), + MrmthStack.Input(id="stack", tooltip="Access stack between nodes",optional=True) ], outputs=[ io.Clip.Output(), + MrmthStack.Output(), ], ) tooltip = cleandoc(__doc__) @classmethod - def check_lazy_status(cls, Expression, V, F, length_mismatch="tile"): + def check_lazy_status(cls, Expression, V, F, length_mismatch="tile",stack=[]): input_stream = InputStream(Expression) lexer = MathExprLexer(input_stream) @@ -68,7 +71,7 @@ class CLIPMathNode(io.ComfyNode): return needed1 @classmethod - def execute(cls, V, F, Expression, length_mismatch="tile") -> io.NodeOutput: + def execute(cls, V, F, Expression, length_mismatch="tile",stack=[]) -> io.NodeOutput: # Determine reference CLIP a = V.get("V0") if a is None: @@ -93,9 +96,9 @@ class CLIPMathNode(io.ComfyNode): # The prompt says aliases are supported in check_lazy_status. Variables map in helper handles logic. aliases = {"a": "V0", "b": "V1", "c": "V2", "d": "V3", "w": "F0", "x": "F1", "y": "F2", "z": "F3"} - patches = calculate_patches_autogrow(Expression, V=patchers_V, F=F, mapping=aliases) + patches = calculate_patches_autogrow(Expression, V=patchers_V, F=F, mapping=aliases,stack=stack) out_clip = a.clone() if patches: out_clip.add_patches(patches, 1.0, 1.0) - return (out_clip,) + return (out_clip,stack) diff --git a/more_math/ConditioningMathNode.py b/more_math/ConditioningMathNode.py index c25cfcd..753d359 100644 --- a/more_math/ConditioningMathNode.py +++ b/more_math/ConditioningMathNode.py @@ -8,6 +8,7 @@ from .Parser.MathExprLexer import MathExprLexer from .Parser.MathExprParser import MathExprParser import re import copy +from .Stack import MrmthStack class ConditioningMathNode(io.ComfyNode): """ @@ -36,15 +37,17 @@ class ConditioningMathNode(io.ComfyNode): default="error", 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") + io.Int.Input(id="batching"), + MrmthStack.Input(id="stack",optional=True) ], outputs=[ io.Conditioning.Output(is_output_list=True), + MrmthStack.Output() ], ) @classmethod - def check_lazy_status(cls, Expression,Expression_pi, V, F,batching, length_mismatch="tile"): + def check_lazy_status(cls, Expression,Expression_pi, V, F,batching, length_mismatch="tile",stack=[]): input_stream = InputStream(Expression) lexer = MathExprLexer(input_stream) @@ -81,7 +84,7 @@ class ConditioningMathNode(io.ComfyNode): return needed1 @classmethod - def execute(cls, V, F, Expression, Expression_pi,batching, length_mismatch="tile"): + def execute(cls, V, F, Expression, Expression_pi,batching, length_mismatch="tile",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: @@ -90,7 +93,6 @@ class ConditioningMathNode(io.ComfyNode): # Extract tensors and pooled outputs tensors = {} pooled_outputs = {} - ss = dict() for key in tensor_keys: conditioning = V[key] tensors[key] = conditioning[0][0] @@ -152,7 +154,7 @@ class ConditioningMathNode(io.ComfyNode): # Execute Expression (Main Tensor) tree = parse_expr(Expression) - visitor = UnifiedMathVisitor(variables, a.shape,a.device, state_storage=ss) + visitor = UnifiedMathVisitor(variables, a.shape,a.device, state_storage=stack) rtensor = visitor.visit(tree) rtensor = as_tensor(rtensor, a.shape) @@ -193,7 +195,7 @@ class ConditioningMathNode(io.ComfyNode): # Execute Expression_pi (Pooled Output) tree_pi = parse_expr(Expression_pi) - visitor_pi = UnifiedMathVisitor(variables_pi, a_p.shape,a_p.device, state_storage=ss) + visitor_pi = UnifiedMathVisitor(variables_pi, a_p.shape,a_p.device, state_storage=stack) rpooled_raw = visitor_pi.visit(tree_pi) rpooled = as_tensor(rpooled_raw, a_p.shape) diff --git a/more_math/FloatMathNode.py b/more_math/FloatMathNode.py index 9a8b15e..c1fd9d8 100644 --- a/more_math/FloatMathNode.py +++ b/more_math/FloatMathNode.py @@ -9,6 +9,7 @@ from antlr4 import InputStream, CommonTokenStream from .Parser.MathExprLexer import MathExprLexer from .Parser.MathExprParser import MathExprParser import re +from .Stack import MrmthStack class FloatMathNode(io.ComfyNode): @@ -29,16 +30,18 @@ class FloatMathNode(io.ComfyNode): inputs=[ io.Autogrow.Input(id="V",template=io.Autogrow.TemplatePrefix(io.Float.Input("values"), prefix="V", min=1, max=50)), io.String.Input(id="FloatFunc", default="a*(1-w)+b*w", tooltip="Expression to use on inputs"), + MrmthStack.Input(id="stack", tooltip="Access stack between nodes",optional=True) ], outputs=[ io.Float.Output(), + MrmthStack.Output(), ], ) tooltip = cleandoc(__doc__) @classmethod - def check_lazy_status(cls, FloatFunc, V): + def check_lazy_status(cls, FloatFunc, V,stack=[]): input_stream = InputStream(FloatFunc) lexer = MathExprLexer(input_stream) stream = CommonTokenStream(lexer) @@ -69,7 +72,7 @@ class FloatMathNode(io.ComfyNode): return needed1 @classmethod - def execute(cls, FloatFunc, V): + def execute(cls, FloatFunc, V,stack=[]): variables = {} # Populate aliases @@ -95,9 +98,9 @@ class FloatMathNode(io.ComfyNode): tree = parse_expr(FloatFunc); # scalar execution # UnifiedMathVisitor expects variables and a shape. Shape [1] for scalar? - visitor = UnifiedMathVisitor(variables, [1]) + visitor = UnifiedMathVisitor(variables, [1],state_storage=stack) result = visitor.visit(tree) # Result might be float or tensor(scalar) if torch.is_tensor(result): result = result[0].item() - return (float(result),) + return (float(result),stack) diff --git a/more_math/GuiderMathNode.py b/more_math/GuiderMathNode.py index 284eaf5..1d67f00 100644 --- a/more_math/GuiderMathNode.py +++ b/more_math/GuiderMathNode.py @@ -1,3 +1,4 @@ +from numpy import stack import torch import re from antlr4 import InputStream, CommonTokenStream @@ -19,6 +20,7 @@ import comfy.model_patcher import comfy.utils import comfy.hooks import comfy.samplers +from .Stack import MrmthStack class GuiderMathNode(io.ComfyNode): @@ -37,14 +39,16 @@ class GuiderMathNode(io.ComfyNode): io.Autogrow.Input(id="F", template=io.Autogrow.TemplatePrefix(io.Float.Input("float", default=0.0, optional=True, lazy=True, force_input=True), prefix="F", min=1, max=50)), io.String.Input(id="Expression", default="G0*(1-F0)+G1*F0", tooltip="Expression to apply on input guiders. Aliases: a=G0, b=G1, c=G2, d=G3, w=F0, x=F1, y=F2, z=F3. Context: steps, current_step"), io.String.Input(id="Expression1", default="G0*(1-F0)+G1*F0", tooltip="Expression to apply after generation finishes."), + MrmthStack.Input(id="stack", tooltip="Access stack between nodes",optional=True) ], outputs=[ io.Guider.Output(), + MrmthStack.Output() ], ) @classmethod - def check_lazy_status(cls, Expression,Expression1, V, F): + def check_lazy_status(cls, Expression,Expression1, V, F,stack=[]): input_stream = InputStream(Expression) input_stream1 = InputStream(Expression1) lexer = MathExprLexer(input_stream) @@ -79,12 +83,12 @@ class GuiderMathNode(io.ComfyNode): return needed1 @classmethod - def execute(cls, V, F, Expression,Expression1): - return (MathGuider(V, F, Expression,Expression1),) + def execute(cls, V, F, Expression,Expression1,stack=[]): + return (MathGuider(V, F, Expression,Expression1),stack) class MathGuider: - def __init__(self, V, F, expression,expression1): + def __init__(self, V, F, expression,expression1,stack=[]): self.V = V self.F = F self.expression = expression @@ -94,7 +98,7 @@ class MathGuider: self.sigmas = None self.current_step = 0 self.steps = 0 - self.stck = {} + self.stck = stack @property def model_patcher(self): @@ -167,7 +171,7 @@ class MathGuider: "c": g_results.get("V2", make_zero_like(eval_samples)), "d": g_results.get("V3", make_zero_like(eval_samples)), }) - + v_stacked, v_cnt = get_v_variable(g_results) if v_stacked is not None: variables["V"] = v_stacked diff --git a/more_math/ImageMathNode.py b/more_math/ImageMathNode.py index 5a7b8c8..aa82922 100644 --- a/more_math/ImageMathNode.py +++ b/more_math/ImageMathNode.py @@ -5,6 +5,7 @@ from antlr4 import InputStream, CommonTokenStream from .Parser.MathExprLexer import MathExprLexer from .Parser.MathExprParser import MathExprParser import re +from .Stack import MrmthStack class ImageMathNode(io.ComfyNode): """ @@ -31,15 +32,17 @@ class ImageMathNode(io.ComfyNode): options=["tile", "error", "pad"], default="error", tooltip="How to handle mismatched image batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero." - ) + ), + MrmthStack.Input(id="stack", tooltip="Access stack between nodes",optional=True) ], outputs=[ io.Image.Output(), + MrmthStack.Output(), ], ) @classmethod - def check_lazy_status(cls, Expression, V, F, length_mismatch="tile"): + def check_lazy_status(cls, Expression, V, F, length_mismatch="tile",stack=[]): input_stream = InputStream(Expression) lexer = MathExprLexer(input_stream) @@ -71,7 +74,7 @@ class ImageMathNode(io.ComfyNode): return needed1 @classmethod - def execute(cls, V, F, Expression, length_mismatch="error"): + def execute(cls, V, F, Expression, length_mismatch="error",stack=[]): # I and F are Autogrow.Type which is dict[str, Any] # Identify all present tensors and their keys @@ -144,7 +147,7 @@ class ImageMathNode(io.ComfyNode): variables[k] = val if val is not None else 0.0 tree = parse_expr(Expression); - visitor = UnifiedMathVisitor(variables, ae.shape,ae.device) + visitor = UnifiedMathVisitor(variables, ae.shape,ae.device,state_storage=stack) result = visitor.visit(tree) result = as_tensor(result, ae.shape) - return (result,) + return (result,stack) diff --git a/more_math/LatentMathNode.py b/more_math/LatentMathNode.py index 7a6b419..58cffaf 100644 --- a/more_math/LatentMathNode.py +++ b/more_math/LatentMathNode.py @@ -17,6 +17,7 @@ from .Parser.MathExprLexer import MathExprLexer from .Parser.MathExprParser import MathExprParser import re from comfy.nested_tensor import NestedTensor +from .Stack import MrmthStack class LatentMathNode(io.ComfyNode): """ @@ -43,17 +44,20 @@ class LatentMathNode(io.ComfyNode): default="error", 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.Int.Input(id="batching"), + MrmthStack.Input(id="stack", tooltip="Access stack between nodes",optional=True) ], outputs=[ io.Latent.Output(is_output_list=True), + MrmthStack.Output(), + ], ) tooltip = cleandoc(__doc__) @classmethod - def check_lazy_status(cls, Expression, V, F,batching, length_mismatch="tile"): + def check_lazy_status(cls, Expression, V, F,batching, length_mismatch="tile",stack=[]): input_stream = InputStream(Expression) lexer = MathExprLexer(input_stream) @@ -85,7 +89,7 @@ class LatentMathNode(io.ComfyNode): return needed1 @classmethod - def execute(cls, V, F, Expression,batching, length_mismatch="tile") -> io.NodeOutput: + def execute(cls, V, F, Expression,batching, length_mismatch="tile",stack=[]) -> io.NodeOutput: # Determine reference latent ref_latent = None for lat in V.values(): @@ -198,7 +202,7 @@ class LatentMathNode(io.ComfyNode): for k, v in F.items(): variables[k] = v if v is not None else 0.0 - visitor = UnifiedMathVisitor(variables, ae.shape,ae.device) + visitor = UnifiedMathVisitor(variables, ae.shape,ae.device,state_storage=stack) result_t = as_tensor(visitor.visit(tree), ae.shape) result_latent = ref_latent.copy() @@ -221,7 +225,7 @@ class LatentMathNode(io.ComfyNode): else: rl["samples"] = result_t results1.append(rl) - return (results1,) + return (results1,stack) rl = result_latent.copy() rl["samples"] = result_t - return ([rl],) + return ([rl],stack) diff --git a/more_math/MaskMathNode.py b/more_math/MaskMathNode.py index 97f6000..4b416cb 100644 --- a/more_math/MaskMathNode.py +++ b/more_math/MaskMathNode.py @@ -5,6 +5,7 @@ from antlr4 import InputStream, CommonTokenStream from .Parser.MathExprLexer import MathExprLexer from .Parser.MathExprParser import MathExprParser import re +from .Stack import MrmthStack class MaskMathNode(io.ComfyNode): @@ -32,15 +33,17 @@ class MaskMathNode(io.ComfyNode): options=["tile", "error", "pad"], default="error", tooltip="How to handle mismatched mask batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero." - ) + ), + MrmthStack.Input(id="stack", tooltip="Access stack between nodes",optional=True) ], outputs=[ io.Mask.Output(), + MrmthStack.Output(), ], ) @classmethod - def check_lazy_status(cls, Expression, V, F, length_mismatch="tile"): + def check_lazy_status(cls, Expression, V, F, length_mismatch="tile",stack=[]): input_stream = InputStream(Expression) lexer = MathExprLexer(input_stream) @@ -72,7 +75,7 @@ class MaskMathNode(io.ComfyNode): return needed1 @classmethod - def execute(cls, V, F, Expression, length_mismatch="tile"): + def execute(cls, V, F, Expression, length_mismatch="tile",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: @@ -139,7 +142,7 @@ class MaskMathNode(io.ComfyNode): variables[k] = val if val is not None else 0.0 tree = parse_expr(Expression); - visitor = UnifiedMathVisitor(variables, ae.shape,ae.device) + visitor = UnifiedMathVisitor(variables, ae.shape,ae.device,state_storage=stack) result = visitor.visit(tree) result = as_tensor(result, ae.shape) - return (result,) + return (result,stack) diff --git a/more_math/ModelMathNode.py b/more_math/ModelMathNode.py index 15f2578..51836e1 100644 --- a/more_math/ModelMathNode.py +++ b/more_math/ModelMathNode.py @@ -4,7 +4,7 @@ from antlr4 import InputStream, CommonTokenStream from .Parser.MathExprLexer import MathExprLexer from .Parser.MathExprParser import MathExprParser import re - +from .Stack import MrmthStack class ModelMathNode(io.ComfyNode): """ @@ -27,17 +27,20 @@ class ModelMathNode(io.ComfyNode): options=["tile", "error", "pad"], default="error", tooltip="How to handle mismatched layer counts. For models, this usually defaults to broadcast (zero for missing layers)." - ) + ), + MrmthStack.Input(id="stack", tooltip="Access stack between nodes",optional=True) ], outputs=[ io.Model.Output(), + MrmthStack.Output(), + ], ) tooltip = cleandoc(__doc__) @classmethod - def check_lazy_status(cls, Expression, V, F, length_mismatch="tile"): + def check_lazy_status(cls, Expression, V, F, length_mismatch="tile",stack=[]): input_stream = InputStream(Expression) lexer = MathExprLexer(input_stream) @@ -69,7 +72,7 @@ class ModelMathNode(io.ComfyNode): return needed1 @classmethod - def execute(cls, V, F, Expression, length_mismatch="tile") -> io.NodeOutput: + def execute(cls, V, F, Expression, length_mismatch="tile",stack=[]) -> io.NodeOutput: # Determine reference model for cloning a = V.get("V0") if a is None: @@ -86,9 +89,9 @@ class ModelMathNode(io.ComfyNode): aliases = {"a": "V0", "b": "V1", "c": "V2", "d": "V3", "w": "F0", "x": "F1", "y": "F2", "z": "F3"} - patches = calculate_patches_autogrow(Expression, V=V, F=F, mapping=aliases) + patches = calculate_patches_autogrow(Expression, V=V, F=F, mapping=aliases,stack=stack) out_model = a.clone() if patches: out_model.add_patches(patches, 1.0, 1.0) - return (out_model,) + return (out_model,stack) diff --git a/more_math/NoiseMathNode.py b/more_math/NoiseMathNode.py index e9d2d5f..46c8a27 100644 --- a/more_math/NoiseMathNode.py +++ b/more_math/NoiseMathNode.py @@ -5,6 +5,7 @@ from .Parser.MathExprParser import MathExprParser,InputStream,CommonTokenStream from .Parser.MathExprLexer import MathExprLexer import re from .Parser.UnifiedMathVisitor import UnifiedMathVisitor +from .Stack import MrmthStack class NoiseMathNode(io.ComfyNode): """ @@ -31,16 +32,17 @@ class NoiseMathNode(io.ComfyNode): inputs=[ io.Autogrow.Input(id="V",template=io.Autogrow.TemplatePrefix(io.Noise.Input("values"), prefix="V", min=1, max=50)), io.Autogrow.Input(id="F", template=io.Autogrow.TemplatePrefix(io.Float.Input("float", default=0.0, optional=True, lazy=True, force_input=True), prefix="F", min=1, max=50)), - io.String.Input(id="Noise", default="a*(1-w)+b*w"), + MrmthStack.Input(id="stack", tooltip="Access stack between nodes",optional=True) ], outputs=[ io.Noise.Output(), + MrmthStack.Output(), ], ) @classmethod - def check_lazy_status(cls, Noise, V, F): + def check_lazy_status(cls, Noise, V, F,stack=[]): input_stream = InputStream(Noise) lexer = MathExprLexer(input_stream) stream = CommonTokenStream(lexer) @@ -71,15 +73,16 @@ class NoiseMathNode(io.ComfyNode): return needed1 @classmethod - def execute(cls, Noise, V,F): - return (NoiseExecutor(V,F, Noise),) + def execute(cls, Noise, V,F,stack=[]): + return (NoiseExecutor(V,F, Noise,stack),) class NoiseExecutor: - def __init__(self, V,F, expr): + def __init__(self, V,F, expr,stack): self.V = V self.F = F self.tree = parse_expr(expr) + self.stack = stack seed = -1 @@ -139,7 +142,7 @@ class NoiseExecutor: F = getIndexTensorAlongDim(samples, time_dim) variables.update({"frame": F, "frame_count": frame_count}) - visitor = UnifiedMathVisitor(variables, samples.shape,samples.device) + visitor = UnifiedMathVisitor(variables, samples.shape,samples.device,state_storage=self.stack) result = visitor.visit(self.tree) result = as_tensor(result, samples.shape) return result diff --git a/more_math/SigmasMathNode.py b/more_math/SigmasMathNode.py index cf707c5..c929b5c 100644 --- a/more_math/SigmasMathNode.py +++ b/more_math/SigmasMathNode.py @@ -5,6 +5,7 @@ from antlr4 import InputStream, CommonTokenStream from .Parser.MathExprLexer import MathExprLexer from .Parser.MathExprParser import MathExprParser import re +from .Stack import MrmthStack class SigmasMathNode(io.ComfyNode): """ @@ -31,15 +32,17 @@ class SigmasMathNode(io.ComfyNode): options=["tile", "error", "pad"], default="error", tooltip="How to handle mismatched image batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero." - ) + ), + MrmthStack.Input(id="stack", tooltip="Access stack between nodes",optional=True) ], outputs=[ io.Sigmas.Output(), + MrmthStack.Output(), ], ) @classmethod - def check_lazy_status(cls, Expression, V, F, length_mismatch="tile"): + def check_lazy_status(cls, Expression, V, F, length_mismatch="tile",stack=[]): input_stream = InputStream(Expression) lexer = MathExprLexer(input_stream) @@ -71,7 +74,7 @@ class SigmasMathNode(io.ComfyNode): return needed1 @classmethod - def execute(cls, V, F, Expression, length_mismatch="tile"): + def execute(cls, V, F, Expression, length_mismatch="tile",stack=[]): # I and F are Autogrow.Type which is dict[str, Any] # Determine reference image for zero-initialization (fallback for a,b,c,d) @@ -125,7 +128,7 @@ class SigmasMathNode(io.ComfyNode): variables[k] = v if v is not None else 0.0 tree = parse_expr(Expression); - visitor = UnifiedMathVisitor(variables, ae.shape,ae.device) + visitor = UnifiedMathVisitor(variables, ae.shape,ae.device,state_storage=stack) result = visitor.visit(tree) result = as_tensor(result, ae.shape) - return (result,) + return (result,stack) diff --git a/more_math/Stack.py b/more_math/Stack.py new file mode 100644 index 0000000..d6b3812 --- /dev/null +++ b/more_math/Stack.py @@ -0,0 +1,13 @@ +from comfy_api.latest import io + +@io.comfytype(io_type="STACK") +class MrmthStack(io.ComfyTypeIO): + Type = list # Python type hint + + class Input(io.Input): + def __init__(self, id: str, **kwargs): + super().__init__(id, **kwargs) + + class Output(io.Output): + def __init__(self, **kwargs): + super().__init__(**kwargs) diff --git a/more_math/VaeMathNode.py b/more_math/VaeMathNode.py index ffcc5cd..d646ead 100644 --- a/more_math/VaeMathNode.py +++ b/more_math/VaeMathNode.py @@ -2,10 +2,12 @@ from inspect import cleandoc from comfy_api.latest import io import copy from antlr4 import InputStream, CommonTokenStream + +from custom_nodes.more_math.more_math.Stack import MrmthStack from .Parser.MathExprLexer import MathExprLexer from .Parser.MathExprParser import MathExprParser import re - +from .Stack import MrmthStack class VAEMathNode(io.ComfyNode): """ @@ -27,17 +29,19 @@ class VAEMathNode(io.ComfyNode): options=["tile", "error", "pad"], default="error", tooltip="How to handle mismatched layer counts. For models, this usually defaults to broadcast (zero for missing layers)." - ) + ), + MrmthStack.Input(id="stack", tooltip="Access stack between nodes",optional=True) ], outputs=[ io.Vae.Output(), + MrmthStack.Output(), ], ) tooltip = cleandoc(__doc__) @classmethod - def check_lazy_status(cls, Expression, V, F, length_mismatch="tile"): + def check_lazy_status(cls, Expression, V, F, length_mismatch="tile",stack=[]): input_stream = InputStream(Expression) lexer = MathExprLexer(input_stream) @@ -69,7 +73,7 @@ class VAEMathNode(io.ComfyNode): return needed1 @classmethod - def execute(cls, V, F, Expression, length_mismatch="tile") -> io.NodeOutput: + def execute(cls, V, F, Expression, length_mismatch="tile",stack=[]) -> io.NodeOutput: # Determine reference VAE a = V.get("V0") if a is None: @@ -92,7 +96,7 @@ class VAEMathNode(io.ComfyNode): # Calculate patches using the patchers (weights are in patcher.model.state_dict) from .modelLikeCommon import calculate_patches_autogrow aliases = {"a": "V0", "b": "V1", "c": "V2", "d": "V3", "w": "F0", "x": "F1", "y": "F2", "z": "F3"} - patches = calculate_patches_autogrow(Expression, V=patchers_V, F=F, mapping=aliases) + patches = calculate_patches_autogrow(Expression, V=patchers_V, F=F, mapping=aliases,stack=stack) # VAE does not have a clone method, so we shallow copy and clone the patcher out_vae = copy.copy(a) diff --git a/more_math/VideoMathNode.py b/more_math/VideoMathNode.py index 6ead17c..3adfeec 100644 --- a/more_math/VideoMathNode.py +++ b/more_math/VideoMathNode.py @@ -6,6 +6,7 @@ from antlr4 import InputStream, CommonTokenStream from .Parser.MathExprLexer import MathExprLexer from .Parser.MathExprParser import MathExprParser import re +from .Stack import MrmthStack class VideoMathNode(io.ComfyNode): """ @@ -33,15 +34,17 @@ class VideoMathNode(io.ComfyNode): options=["tile", "error", "pad"], default="error", tooltip="How to handle mismatched image batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero." - ) + ), + MrmthStack.Input(id="stack", tooltip="Access stack between nodes",optional=True) ], outputs=[ io.Conditioning.Output(), + MrmthStack.Output(), ], ) @classmethod - def check_lazy_status(cls, Expression,Expression_pi, V, F, length_mismatch="tile"): + def check_lazy_status(cls, Expression,Expression_pi, V, F, length_mismatch="tile",stack=[]): input_stream = InputStream(Expression) lexer = MathExprLexer(input_stream) @@ -78,8 +81,7 @@ class VideoMathNode(io.ComfyNode): return needed1 @classmethod - def execute(cls, V, F, Expression, Expression_pi, length_mismatch="tile"): - ss = {} + def execute(cls, V, F, Expression, Expression_pi, length_mismatch="tile",stack=[]): 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.") @@ -149,7 +151,7 @@ class VideoMathNode(io.ComfyNode): variables[k] = val if val is not None else 0.0 tree = parse_expr(Expression); - visitor = UnifiedMathVisitor(variables, ae.shape,ae.device,state_storage=ss) + visitor = UnifiedMathVisitor(variables, ae.shape,ae.device,state_storage=stack) result = visitor.visit(tree) result = as_tensor(result, ae.shape) @@ -216,8 +218,8 @@ class VideoMathNode(io.ComfyNode): variables[k] = val if val is not None else 0.0 tree = parse_expr(Expression); - visitor = UnifiedMathVisitor(variables, a_w.shape,state_storage=ss) + visitor = UnifiedMathVisitor(variables, a_w.shape,state_storage=stack) result1 = visitor.visit(tree) result1 = as_tensor(result, a_w.shape) - return ([result,{"waveform":result1,"sample_rate":sample_rate}],) + return ([result,{"waveform":result1,"sample_rate":sample_rate}],stack) diff --git a/more_math/modelLikeCommon.py b/more_math/modelLikeCommon.py index 326cb05..7379f50 100644 --- a/more_math/modelLikeCommon.py +++ b/more_math/modelLikeCommon.py @@ -8,7 +8,7 @@ def calculate_patches(Model, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0 """Legacy calculate_patches for backward compatibility.""" return calculate_patches_autogrow(Model, V={"V0": a, "V1": b, "V2": c, "V3": d}, F={"F0": w, "F1": x, "F2": y, "F3": z}, mapping={"a": "V0", "b": "V1", "c": "V2", "d": "V3", "w": "F0", "x": "F1", "y": "F2", "z": "F3"}) -def calculate_patches_autogrow(Expr, V, F, mapping=None): +def calculate_patches_autogrow(Expr, V, F, mapping=None,stack = []): """ Calculate patches for model-like objects (Model, VAE, CLIP) using Autogrow inputs. Iterates over the UNION of keys from all input models to support merging disjoint architectures/patches. @@ -26,7 +26,6 @@ def calculate_patches_autogrow(Expr, V, F, mapping=None): # Collect all unique keys from all models all_keys = set() models = [v for v in V.values() if v is not None] - stck = {} if not models: return {} @@ -118,7 +117,7 @@ def calculate_patches_autogrow(Expr, V, F, mapping=None): # Execute math - visitor = UnifiedMathVisitor(variables, ref_tensor.shape,state_storage=stck) + visitor = UnifiedMathVisitor(variables, ref_tensor.shape,state_storage=stack) res = visitor.visit(tree) res = as_tensor(res, ref_tensor.shape)