Add stack passing
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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 +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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user