Add stack passing

This commit is contained in:
mcDandy
2026-02-02 12:48:13 +01:00
parent ac0c2f43d3
commit eae6b03f25
15 changed files with 126 additions and 74 deletions
+8 -5
View File
@@ -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)
+8 -5
View File
@@ -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)
+8 -6
View File
@@ -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)
+7 -4
View File
@@ -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)
+10 -6
View File
@@ -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
+8 -5
View File
@@ -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)
+10 -6
View File
@@ -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)
+8 -5
View File
@@ -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)
+9 -6
View File
@@ -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)
+9 -6
View File
@@ -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
+8 -5
View File
@@ -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)
+13
View File
@@ -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)
+9 -5
View File
@@ -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)
+9 -7
View File
@@ -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)
+2 -3
View File
@@ -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)