From be29ad7f67cdaa023609bc8a7746733c29c9c699 Mon Sep 17 00:00:00 2001 From: mcDandy Date: Thu, 25 Dec 2025 19:43:09 +0100 Subject: [PATCH] massive refactor (AI) --- more_math/AudioMathNode.py | 98 ++++++++-------- more_math/ConditioningMathNode.py | 139 ++++++++--------------- more_math/FloatMathNode.py | 81 ++++---------- more_math/ImageMathNode.py | 106 ++++++------------ more_math/LatentMathNode.py | 27 ++--- more_math/NoiseMathNode.py | 34 ++---- more_math/VideoMathNode.py | 178 ++++++++++-------------------- more_math/helper_functions.py | 160 ++++++++++++++++++++++----- more_math/modelLikeCommon.py | 72 +++++------- tests/conftest.py | 20 ++-- 10 files changed, 391 insertions(+), 524 deletions(-) diff --git a/more_math/AudioMathNode.py b/more_math/AudioMathNode.py index 35756af..18f7226 100644 --- a/more_math/AudioMathNode.py +++ b/more_math/AudioMathNode.py @@ -1,29 +1,21 @@ import torch -from antlr4 import CommonTokenStream, InputStream -from .Parser.MathExprParser import MathExprParser -from .Parser.MathExprLexer import MathExprLexer -from .Parser.TensorEvalVisitor import TensorEvalVisitor -from .helper_functions import getIndexTensorAlongDim, comonLazy +from .helper_functions import getIndexTensorAlongDim, comonLazy, eval_tensor_expr, make_zero_like from comfy_api.latest import io + class AudioMathNode(io.ComfyNode): """ - This node enables the use of math expressions on AUDIO tensors. - inputs: - a, b, c, d: - AUDIO, bound to variables with the same name. Defaults to zero AUDIO if not provided. - w, x, y, z: - Floats, bound to variables of the expression. Defaults to 0.0 if not provided. - Audio expression: - String, describing expression to apply to audio tensors. - - outputs: - AUDIO: - Returns an AUDIO object that contains the result of the math expression applied to the input audio tensors. + Enables math expressions on Audio tensors. + + Inputs: + a, b, c, d: Audio inputs (b, c, d default to zero if not provided) + w, x, y, z: Float variables for expressions + AudioExpr: Expression to apply on audio tensors + + Outputs: + AUDIO: Result of applying expression to input audio """ - def __init__(self): - pass @classmethod def define_schema(cls) -> io.Schema: @@ -33,56 +25,52 @@ class AudioMathNode(io.ComfyNode): display_name="Audio math", inputs=[ io.Audio.Input(id="a", tooltip="Input audio tensor"), - io.Audio.Input(id="b", optional=True,lazy=True, tooltip="Second input audio tensor"), - io.Audio.Input(id="c", optional=True,lazy=True, tooltip="Third input audio tensor"), - io.Audio.Input(id="d", optional=True,lazy=True, tooltip="Fourth input audio tensor"), - io.Float.Input(id="w", default=0.0, optional=True,lazy=True, force_input=True), - io.Float.Input(id="x", default=0.0, optional=True,lazy=True, force_input=True), - io.Float.Input(id="y", default=0.0, optional=True,lazy=True, force_input=True), - io.Float.Input(id="z", default=0.0, optional=True,lazy=True, force_input=True), + io.Audio.Input(id="b", optional=True, lazy=True, tooltip="Second input audio tensor"), + io.Audio.Input(id="c", optional=True, lazy=True, tooltip="Third input audio tensor"), + io.Audio.Input(id="d", optional=True, lazy=True, tooltip="Fourth input audio tensor"), + io.Float.Input(id="w", default=0.0, optional=True, lazy=True, force_input=True), + io.Float.Input(id="x", default=0.0, optional=True, lazy=True, force_input=True), + io.Float.Input(id="y", default=0.0, optional=True, lazy=True, force_input=True), + io.Float.Input(id="z", default=0.0, optional=True, lazy=True, force_input=True), io.String.Input(id="AudioExpr", default="a*(1-w)+b*w", tooltip="Expression to apply on input audio tensors"), ], outputs=[ io.Audio.Output(), ], ) + @classmethod - def check_lazy_status(cls, AudioExpr, a, b=[], c=[], d=[],w=0,x=0,y=0,z=0): - return comonLazy(AudioExpr, a, b, c, d,w,x,y,z) + def check_lazy_status(cls, AudioExpr, a, b=[], c=[], d=[], w=0, x=0, y=0, z=0): + return comonLazy(AudioExpr, a, b, c, d, w, x, y, z) + @classmethod def execute(cls, a, AudioExpr, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0): + waveform = a['waveform'] + sample_rate = a['sample_rate'] + b = b if b else make_zero_like(a) + c = c if c else make_zero_like(a) + d = d if d else make_zero_like(a) - bv = b if b else {'waveform':torch.zeros_like(a['waveform']),'sample_rate':a['sample_rate']} - cv = c if c else {'waveform':torch.zeros_like(a['waveform']),'sample_rate':a['sample_rate']} - dv = d if d else {'waveform':torch.zeros_like(a['waveform']),'sample_rate':a['sample_rate']} - - B = getIndexTensorAlongDim(a['waveform'], 0) - C = getIndexTensorAlongDim(a['waveform'], 1) - S = getIndexTensorAlongDim(a['waveform'], 2) - R = torch.full_like(S, a['sample_rate'], dtype=torch.float32) - T = torch.full_like(S, a['waveform'].shape[2], dtype=torch.float32) + bv, cv, dv = b['waveform'], c['waveform'], d['waveform'] variables = { - 'a': a['waveform'], 'b': bv['waveform'], 'c': cv['waveform'], 'd': dv['waveform'], + 'a': waveform, 'b': bv, 'c': cv, 'd': dv, 'w': w, 'x': x, 'y': y, 'z': z, - 'B': B, 'C': C, 'S': S,'R': R, 'T' : T, 'N': a['waveform'].shape[1], - 'batch': B, 'channel': C, 'sample': S, 'sample_rate': R, 'sample_count': T,'channel_count': a['waveform'].shape[1] + 'B': getIndexTensorAlongDim(waveform, 0), + 'C': getIndexTensorAlongDim(waveform, 1), + 'S': getIndexTensorAlongDim(waveform, 2), + 'R': torch.full_like(waveform, sample_rate, dtype=torch.float32), + 'T': torch.full_like(waveform, waveform.shape[2], dtype=torch.float32), + 'N': waveform.shape[1], + 'batch': getIndexTensorAlongDim(waveform, 0), + 'channel': getIndexTensorAlongDim(waveform, 1), + 'sample': getIndexTensorAlongDim(waveform, 2), + 'sample_rate': torch.full_like(waveform, sample_rate, dtype=torch.float32), + 'sample_count': torch.full_like(waveform, waveform.shape[2], dtype=torch.float32), + 'channel_count': waveform.shape[1], } - input_stream = InputStream(AudioExpr) - lexer = MathExprLexer(input_stream) - stream = CommonTokenStream(lexer) - parser = MathExprParser(stream) - tree = parser.expr() + result_tensor = eval_tensor_expr(AudioExpr, variables, waveform.shape) - visitor = TensorEvalVisitor(variables, a['waveform'].shape) - result_tensor = visitor.visit(tree) - - # Create output dictionary with the same sample rate - output = { - 'waveform': result_tensor, - 'sample_rate': a['sample_rate'] - } - - return (output,) + return ({'waveform': result_tensor, 'sample_rate': sample_rate},) diff --git a/more_math/ConditioningMathNode.py b/more_math/ConditioningMathNode.py index 973b73e..e6cf844 100644 --- a/more_math/ConditioningMathNode.py +++ b/more_math/ConditioningMathNode.py @@ -1,34 +1,25 @@ from inspect import cleandoc -from antlr4 import CommonTokenStream, InputStream import torch -from .Parser.MathExprParser import MathExprParser -from .Parser.MathExprLexer import MathExprLexer -from .Parser.TensorEvalVisitor import TensorEvalVisitor -from .helper_functions import ThrowingErrorListener, comonLazy - +from .helper_functions import comonLazy, eval_tensor_expr, make_zero_like from comfy_api.latest import io + + class ConditioningMathNode(io.ComfyNode): """ - This node enables the use of math on conditionings. It is recommended to keep the Tensor expression the same as the pooled_output expression. - inputs: - a, b, c, d: - Conditioning, bound to variables with the same name. Defaults to zero conditioning if not provided. - w, x, y, z: - Floats, bound to variables of the expression. Defaults to 0.0 if not provided. - Tensor expression: - String, describing expression to mix tensor part of conditioning. Tensor probably describes the composition of the image as it has the say what is on the image. - Pooled output expression: - String, describing expression to mix pooled_output part of conditioning. pooled_output is Tensor compressed into 1 token. - - outputs: - CONDITIONING: - Returns a CONDITIONING object that contains the result of the math expression applied to the input conditionings. + Enables math operations on conditionings. + + Inputs: + a, b, c, d: Conditioning inputs (b, c, d default to zero if not provided) + w, x, y, z: Float variables for expressions + Tensor: Expression for the tensor part (describes image composition) + pooled_output: Expression for the pooled output (condensed representation) + + Outputs: + CONDITIONING: Result of applying expressions to input conditionings """ - def __init__(self): - pass - + @classmethod def define_schema(cls) -> io.Schema: return io.Schema( @@ -37,15 +28,15 @@ class ConditioningMathNode(io.ComfyNode): category="More math", inputs=[ io.Conditioning.Input(id="a"), - io.Conditioning.Input(id="b", optional=True,lazy=True), - io.Conditioning.Input(id="c", optional=True,lazy=True), - io.Conditioning.Input(id="d", optional=True,lazy=True), - io.Float.Input(id="w", default=0.0,optional=True,lazy=True, force_input=True), - io.Float.Input(id="x", default=0.0,optional=True,lazy=True, force_input=True), - io.Float.Input(id="y", default=0.0,optional=True,lazy=True, force_input=True), - io.Float.Input(id="z", default=0.0,optional=True,lazy=True, force_input=True), - io.String.Input(id="Tensor", default="a*(1-w)+b*w", tooltip="Describes composition of the image."), - io.String.Input(id="pooled_output", default="a*(1-w)+b*w", tooltip="Composition of the image condensed into one vector"), + io.Conditioning.Input(id="b", optional=True, lazy=True), + io.Conditioning.Input(id="c", optional=True, lazy=True), + io.Conditioning.Input(id="d", optional=True, lazy=True), + io.Float.Input(id="w", default=0.0, optional=True, lazy=True, force_input=True), + io.Float.Input(id="x", default=0.0, optional=True, lazy=True, force_input=True), + io.Float.Input(id="y", default=0.0, optional=True, lazy=True, force_input=True), + io.Float.Input(id="z", default=0.0, optional=True, lazy=True, force_input=True), + io.String.Input(id="Tensor", default="a*(1-w)+b*w", tooltip="Expression for tensor part (image composition)"), + io.String.Input(id="pooled_output", default="a*(1-w)+b*w", tooltip="Expression for pooled output (condensed representation)"), ], outputs=[ io.Conditioning.Output(), @@ -54,71 +45,35 @@ class ConditioningMathNode(io.ComfyNode): tooltip = cleandoc(__doc__) - #OUTPUT_NODE = False - #OUTPUT_TOOLTIPS = ("",) # Tooltips for the output node @classmethod - def check_lazy_status(cls, Tensor,pooled_output, a, b=[], c=[], d=[],w=0,x=0,y=0,z=0): - return list(set(comonLazy(Tensor, a, b, c, d)).union(comonLazy(pooled_output, a, b, c, d))) + def check_lazy_status(cls, Tensor, pooled_output, a, b=[], c=[], d=[], w=0, x=0, y=0, z=0): + tensor_needs = set(comonLazy(Tensor, a, b, c, d)) + pooled_needs = set(comonLazy(pooled_output, a, b, c, d)) + return list(tensor_needs.union(pooled_needs)) + @classmethod - def execute(cls, Tensor,pooled_output, a, b=None, c=None, d=None,w=0.0,x=0.0,y=0.0,z=0.0): - if b is None: - b = [[torch.zeros_like(a[0][0]), {"pooled_output": torch.zeros_like(a[0][1]["pooled_output"]) if a[0][1]["pooled_output"] is not None else None}]] - if c is None: - c = [[torch.zeros_like(a[0][0]), {"pooled_output": torch.zeros_like(a[0][1]["pooled_output"]) if a[0][1]["pooled_output"] is not None else None}]] - if d is None: - d = [[torch.zeros_like(a[0][0]), {"pooled_output": torch.zeros_like(a[0][1]["pooled_output"]) if a[0][1]["pooled_output"] is not None else None}]] - - - ta = a[0][0].clone() - tb = b[0][0].clone() - tc = c[0][0].clone() - td = d[0][0].clone() - pa = None - pb = None - pc = None - pd = None - if a[0][1]["pooled_output"] is not None: - pa = a[0][1]["pooled_output"].clone() - pb = b[0][1]["pooled_output"].clone() - pc = c[0][1]["pooled_output"].clone() - pd = d[0][1]["pooled_output"].clone() - + def execute(cls, Tensor, pooled_output, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0): + # Default missing conditionings to zero + b = b if b is not None else make_zero_like(a) + c = c if c is not None else make_zero_like(a) + d = d if d is not None else make_zero_like(a) + # Extract tensors + ta, tb, tc, td = a[0][0], b[0][0], c[0][0], d[0][0] + + # Evaluate tensor expression variables = {'a': ta, 'b': tb, 'c': tc, 'd': td, 'w': w, 'x': x, 'y': y, 'z': z} + result_tensor = eval_tensor_expr(Tensor, variables, ta.shape) - input_stream = InputStream(Tensor) - lexer = MathExprLexer(input_stream) - stream = CommonTokenStream(lexer) - parser = MathExprParser(stream) - parser.addErrorListener(ThrowingErrorListener()) - tree = parser.expr() - visitor = TensorEvalVisitor(variables,ta.shape) - result1 = visitor.visit(tree) - - if a[0][1]["pooled_output"] is not None: + # Evaluate pooled_output expression if available + pa = a[0][1].get("pooled_output") + if pa is not None: + pb = b[0][1].get("pooled_output") + pc = c[0][1].get("pooled_output") + pd = d[0][1].get("pooled_output") variables = {'a': pa, 'b': pb, 'c': pc, 'd': pd, 'w': w, 'x': x, 'y': y, 'z': z} - input_stream = InputStream(Tensor) - lexer = MathExprLexer(input_stream) - stream = CommonTokenStream(lexer) - parser = MathExprParser(stream) - tree = parser.expr() - visitor = TensorEvalVisitor(variables,pa.shape) - result2 = visitor.visit(tree) + result_pooled = eval_tensor_expr(pooled_output, variables, pa.shape) else: - result2 = None - print("No pooled_output found in input conditioning, skipping pooled_output calculation.") + result_pooled = None - result = [[result1, {"pooled_output": result2}]] - return (result,) - - """ - The node will always be re executed if any of the inputs change but - this method can be used to force the node to execute again even when the inputs don't change. - You can make this node return a number or a string. This value will be compared to the one returned the last time the node was - executed, if it is different the node will be executed again. - This method is used in the core repo for the LoadImage node where they return the image hash as a string, if the image hash - changes between executions the LoadImage node is executed again. - """ - #@classmethod - #def IS_CHANGED(s, image, string_field, int_field, float_field, print_to_screen): - # return "" + return ([[result_tensor, {"pooled_output": result_pooled}]],) diff --git a/more_math/FloatMathNode.py b/more_math/FloatMathNode.py index 11d4383..65396fd 100644 --- a/more_math/FloatMathNode.py +++ b/more_math/FloatMathNode.py @@ -1,51 +1,38 @@ from inspect import cleandoc -from antlr4 import CommonTokenStream, InputStream - -from .helper_functions import ThrowingErrorListener, comonLazy - -from .Parser.MathExprParser import MathExprParser -from .Parser.MathExprLexer import MathExprLexer -from .Parser.FloatEvalVisitor import FloatEvalVisitor +from .helper_functions import comonLazy, eval_float_expr from comfy_api.latest import io class FloatMathNode(io.ComfyNode): """ - This node enables the use of math expressions on Latents. - inputs: - a, b, c, d: - Floats, bound to variables with the same name. Defaults to 0.0 if not provided. - w, x, y, z: - Floats, bound to variables of the expression. Defaults to 0.0 if not provided. - Latent expression: - String, describing expression to mix latents. Valid functions are sin, cos, tan, abs, sqrt, min, max, norm. Valid operators are +, -, *, /, ^, %. Usable constants are e and pi. - - outputs: - LATENT: - Returns a LATENT object that contains the result of the math expression applied to the input conditionings. + This node enables the use of math expressions on Floats. + + Inputs: + a, b, c, d: Floats, bound to variables with the same name. + w, x, y, z: Floats, bound to variables of the expression. + FloatFunc: String, describing math expression. + Valid functions: sin, cos, tan, abs, sqrt, min, max, norm, etc. + Operators: +, -, *, /, ^, %. + Constants: e, pi. + + Outputs: + FLOAT: The result of evaluating the math expression. """ - def __init__(self): - pass - @classmethod - def check_lazy_status(cls, Model, a, b=[], c=[], d=[],w=0,x=0,y=0,z=0): - return comonLazy(Model, a, b, c, d,w,x,y,z) @classmethod def define_schema(cls) -> io.Schema: - """ - """ return io.Schema( node_id="mrmth_FloatMathNode", category="More math", display_name="Float math", inputs=[ io.Float.Input(id="a", force_input=True), - io.Float.Input(id="b", default=0.0,optional=True,lazy=True, force_input=True), - io.Float.Input(id="c", default=0.0,optional=True,lazy=True, force_input=True), - io.Float.Input(id="d", default=0.0,optional=True,lazy=True, force_input=True), - io.Float.Input(id="w", default=0.0,optional=True,lazy=True, force_input=True), - io.Float.Input(id="x", default=0.0,optional=True,lazy=True, force_input=True), - io.Float.Input(id="y", default=0.0,optional=True,lazy=True, force_input=True), - io.Float.Input(id="z", default=0.0,optional=True,lazy=True, force_input=True), + io.Float.Input(id="b", default=0.0, optional=True, lazy=True, force_input=True), + io.Float.Input(id="c", default=0.0, optional=True, lazy=True, force_input=True), + io.Float.Input(id="d", default=0.0, optional=True, lazy=True, force_input=True), + io.Float.Input(id="w", default=0.0, optional=True, lazy=True, force_input=True), + io.Float.Input(id="x", default=0.0, optional=True, lazy=True, force_input=True), + io.Float.Input(id="y", default=0.0, optional=True, lazy=True, force_input=True), + io.Float.Input(id="z", default=0.0, optional=True, lazy=True, force_input=True), io.String.Input(id="FloatFunc", default="a*(1-w)+b*w", tooltip="Expression to use on inputs"), ], outputs=[ @@ -53,36 +40,14 @@ class FloatMathNode(io.ComfyNode): ], ) - #RETURN_NAMES = ("image_output_name",) tooltip = cleandoc(__doc__) - #OUTPUT_NODE = False - #OUTPUT_TOOLTIPS = ("",) # Tooltips for the output node @classmethod - def check_lazy_status(cls, FloatFunc, a, b=[], c=[], d=[],w=0,x=0,y=0,z=0): + def check_lazy_status(cls, FloatFunc, a, b=[], c=[], d=[], w=0, x=0, y=0, z=0): return comonLazy(FloatFunc, a, b, c, d) + @classmethod def execute(cls, FloatFunc, a, b=0.0, c=0.0, d=0.0, w=0.0, x=0.0, y=0.0, z=0.0): - variables = {'a': a, 'b': b, 'c': c, 'd': d, 'w': w, 'x': x, 'y': y, 'z': z} - input_stream = InputStream(FloatFunc) - lexer = MathExprLexer(input_stream) - stream = CommonTokenStream(lexer) - parser = MathExprParser(stream) - parser.addErrorListener(ThrowingErrorListener()) - tree = parser.expr() - visitor = FloatEvalVisitor(variables) - result = visitor.visit(tree) + result = eval_float_expr(FloatFunc, variables) return (result,) - - """ - The node will always be re executed if any of the inputs change but - this method can be used to force the node to execute again even when the inputs don't change. - You can make this node return a number or a string. This value will be compared to the one returned the last time the node was - executed, if it is different the node will be executed again. - This method is used in the core repo for the LoadImage node where they return the image hash as a string, if the image hash - changes between executions the LoadImage node is executed again. - """ - #@classmethod - #def IS_CHANGED(s, image, string_field, int_field, float_field, print_to_screen): - # return "" diff --git a/more_math/ImageMathNode.py b/more_math/ImageMathNode.py index c1e0096..5b3ea1b 100644 --- a/more_math/ImageMathNode.py +++ b/more_math/ImageMathNode.py @@ -1,109 +1,75 @@ - -from antlr4 import CommonTokenStream -from antlr4.atn.LexerActionExecutor import InputStream import torch -from .helper_functions import ThrowingErrorListener, getIndexTensorAlongDim, comonLazy - -from .Parser.MathExprParser import MathExprParser -from .Parser.MathExprLexer import MathExprLexer -from .Parser.TensorEvalVisitor import TensorEvalVisitor +from .helper_functions import getIndexTensorAlongDim, comonLazy, eval_tensor_expr, make_zero_like from comfy_api.latest import io + class ImageMathNode(io.ComfyNode): """ - This node enables the use of math expressions on Latents. - inputs: - a, b, c, d: - Latent, bound to variables with the same name. Defaults to zero latent if not provided. - w, x, y, z: - Floats, bound to variables of the expression. Defaults to 0.0 if not provided. - Image expression: - String, describing expression to mix images. - - outputs: - LATENT: - Returns a LATENT object that contains the result of the math expression applied to the input conditionings. + Enables math expressions on Images. + + Inputs: + a, b, c, d: Image inputs (b, c, d default to zero if not provided) + w, x, y, z: Float variables for expressions + Image: Expression to apply on input images + + Outputs: + IMAGE: Result of applying expression to input images """ - def __init__(self): - pass @classmethod def define_schema(cls) -> io.Schema: - """ - """ return io.Schema( node_id="mrmth_ImageMathNode", category="More math", display_name="Image math", inputs=[ io.Image.Input(id="a"), - io.Image.Input(id="b", optional=True,lazy=True), - io.Image.Input(id="c", optional=True,lazy=True), - io.Image.Input(id="d", optional=True,lazy=True), - io.Float.Input(id="w", default=0.0,optional=True,lazy=True, force_input=True), - io.Float.Input(id="x", default=0.0,optional=True,lazy=True, force_input=True), - io.Float.Input(id="y", default=0.0,optional=True,lazy=True, force_input=True), - io.Float.Input(id="z", default=0.0,optional=True,lazy=True, force_input=True), + io.Image.Input(id="b", optional=True, lazy=True), + io.Image.Input(id="c", optional=True, lazy=True), + io.Image.Input(id="d", optional=True, lazy=True), + io.Float.Input(id="w", default=0.0, optional=True, lazy=True, force_input=True), + io.Float.Input(id="x", default=0.0, optional=True, lazy=True, force_input=True), + io.Float.Input(id="y", default=0.0, optional=True, lazy=True, force_input=True), + io.Float.Input(id="z", default=0.0, optional=True, lazy=True, force_input=True), io.String.Input(id="Image", default="a*(1-w)+b*w", tooltip="Expression to apply on input images"), ], outputs=[ io.Image.Output(), ], ) - @classmethod - def check_lazy_status(cls, Image, a, b=[], c=[], d=[],w=0,x=0,y=0,z=0): - return comonLazy(Image, a, b, c, d,w,x,y,z) - @classmethod - def execute(scls, Image, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0): - b = torch.zeros_like(a) if b is None else b - c = torch.zeros_like(a) if c is None else c - d = torch.zeros_like(a) if d is None else d - # permute to B, C, H, W + @classmethod + def check_lazy_status(cls, Image, a, b=[], c=[], d=[], w=0, x=0, y=0, z=0): + return comonLazy(Image, a, b, c, d, w, x, y, z) + + @classmethod + def execute(cls, Image, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0): + b = make_zero_like(a) if b is None else b + c = make_zero_like(a) if c is None else c + d = make_zero_like(a) if d is None else d + + # Permute to B, C, H, W for processing a = a.permute(0, 3, 1, 2) b = b.permute(0, 3, 1, 2) c = c.permute(0, 3, 1, 2) d = d.permute(0, 3, 1, 2) - B = getIndexTensorAlongDim(a, 0) - C = getIndexTensorAlongDim(a, 1) - H = getIndexTensorAlongDim(a, 2) - W = getIndexTensorAlongDim(a, 3) - variables = { 'a': a, 'b': b, 'c': c, 'd': d, 'w': w, 'x': x, 'y': y, 'z': z, - 'X': W, 'Y': H, - 'B': B,'batch': B, - 'C': C,'channel': C, + 'X': getIndexTensorAlongDim(a, 3), + 'Y': getIndexTensorAlongDim(a, 2), + 'B': getIndexTensorAlongDim(a, 0), 'batch': getIndexTensorAlongDim(a, 0), + 'C': getIndexTensorAlongDim(a, 1), 'channel': getIndexTensorAlongDim(a, 1), 'W': a.shape[3], 'width': a.shape[3], 'H': a.shape[2], 'height': a.shape[2], 'T': a.shape[0], 'batch_count': a.shape[0], 'N': a.shape[1], 'channel_count': a.shape[1], } - input_stream = InputStream(Image) - lexer = MathExprLexer(input_stream) - stream = CommonTokenStream(lexer) - parser = MathExprParser(stream) - parser.addErrorListener(ThrowingErrorListener()) - tree = parser.expr() - visitor = TensorEvalVisitor(variables, a.shape) - result = visitor.visit(tree) - # permute back to B, H, W, C + result = eval_tensor_expr(Image, variables, a.shape) + + # Permute back to B, H, W, C result = result.permute(0, 2, 3, 1) return (result,) - - - """ - The node will always be re executed if any of the inputs change but - this method can be used to force the node to execute again even when the inputs don't change. - You can make this node return a number or a string. This value will be compared to the one returned the last time the node was - executed, if it is different the node will be executed again. - This method is used in the core repo for the LoadImage node where they return the image hash as a string, if the image hash - changes between executions the LoadImage node is executed again. - """ - #@classmethod - #def IS_CHANGED(s, image, string_field, int_field, float_field, print_to_screen): - # return "" diff --git a/more_math/LatentMathNode.py b/more_math/LatentMathNode.py index cba526a..17ada1c 100644 --- a/more_math/LatentMathNode.py +++ b/more_math/LatentMathNode.py @@ -2,14 +2,9 @@ from inspect import cleandoc from comfy_api.latest import io -from antlr4 import CommonTokenStream, InputStream import torch -from .helper_functions import ThrowingErrorListener, getIndexTensorAlongDim,comonLazy - -from .Parser.MathExprParser import MathExprParser -from .Parser.MathExprLexer import MathExprLexer -from .Parser.TensorEvalVisitor import TensorEvalVisitor +from .helper_functions import getIndexTensorAlongDim, comonLazy, parse_expr, eval_tensor_expr_with_tree, make_zero_like # try to import NestedTensor type if available try: @@ -19,6 +14,8 @@ except Exception: _nested_tensor_module = None _NESTED_TENSOR_AVAILABLE = False + + class LatentMathNode(io.ComfyNode): """ This node enables the use of math expressions on Latents. @@ -78,12 +75,7 @@ class LatentMathNode(io.ComfyNode): d_in = None if d is None else d["samples"] # parse expression once - input_stream = InputStream(Latent) - lexer = MathExprLexer(input_stream) - stream = CommonTokenStream(lexer) - parser = MathExprParser(stream) - parser.addErrorListener(ThrowingErrorListener()) - tree = parser.expr() + tree = parse_expr(Latent) # Helper to evaluate for a single tensor def eval_single_tensor(a_t, b_t, c_t, d_t): @@ -131,8 +123,7 @@ class LatentMathNode(io.ComfyNode): F = getIndexTensorAlongDim(a_t, time_dim) variables.update({'frame_idx': F, 'frame': F, 'frame_count': frame_count}) - visitor = TensorEvalVisitor(variables, a_t.shape) - return visitor.visit(tree) + return eval_tensor_expr_with_tree(tree, variables, a_t.shape) # If input is a NestedTensor (from comfy), evaluate per-subtensor and return NestedTensor result if hasattr(a_in, 'is_nested') and getattr(a_in, 'is_nested'): @@ -145,7 +136,7 @@ class LatentMathNode(io.ComfyNode): def merge_to_tensor(val, ref): # ref is merged_a if val is None: - return torch.zeros_like(ref) + return make_zero_like(ref) if hasattr(val, 'is_nested') and getattr(val, 'is_nested'): lst = val.unbind() return torch.cat(lst, dim=0) @@ -166,7 +157,7 @@ class LatentMathNode(io.ComfyNode): if val.shape[0] == sum(sizes): return val # fallback - return torch.zeros_like(ref) + return make_zero_like(ref) merged_b = merge_to_tensor(b_in, merged_a) merged_c = merge_to_tensor(c_in, merged_a) @@ -187,7 +178,7 @@ class LatentMathNode(io.ComfyNode): # ensure b/c/d are set appropriately (zeros_like if None) def to_tensor(val, ref): if val is None: - return torch.zeros_like(ref) + return make_zero_like(ref) if hasattr(val, 'is_nested') and getattr(val, 'is_nested'): lst = val.unbind() elif isinstance(val, (list, tuple)): @@ -196,7 +187,7 @@ class LatentMathNode(io.ComfyNode): return val if len(lst) == 0: - return torch.zeros_like(ref) + return make_zero_like(ref) if len(lst) == 1: return lst[0] try: diff --git a/more_math/NoiseMathNode.py b/more_math/NoiseMathNode.py index d2856d2..44ef140 100644 --- a/more_math/NoiseMathNode.py +++ b/more_math/NoiseMathNode.py @@ -1,17 +1,11 @@ from inspect import cleandoc -from antlr4 import CommonTokenStream, InputStream import torch -from .helper_functions import ThrowingErrorListener, getIndexTensorAlongDim,comonLazy - -from .Parser.MathExprParser import MathExprParser -from .Parser.MathExprLexer import MathExprLexer -from .Parser.TensorEvalVisitor import TensorEvalVisitor +from .helper_functions import getIndexTensorAlongDim, comonLazy, parse_expr, eval_tensor_expr_with_tree, make_zero_like from comfy_api.latest import io -# try to import NestedTensor type if available import comfy.nested_tensor as _nested_tensor_module @@ -95,27 +89,20 @@ class NoiseExecutor(): self.y = y self.z = z self.expr = Noise - # parse expression once - input_stream = InputStream(Noise) - lexer = MathExprLexer(input_stream) - stream = CommonTokenStream(lexer) - parser = MathExprParser(stream) - parser.addErrorListener(ThrowingErrorListener()) - self.tree = parser.expr() + self.tree = parse_expr(Noise) seed = -1; def generate_noise(self, input_latent:torch.Tensor) -> torch.Tensor: samples = input_latent["samples"] - # evaluate generators / default zeros - a_val = self.a.generate_noise(input_latent) if self.a is not None else torch.zeros_like(samples) - b_val = self.b.generate_noise(input_latent) if self.b is not None else torch.zeros_like(samples) - c_val = self.c.generate_noise(input_latent) if self.c is not None else torch.zeros_like(samples) - d_val = self.d.generate_noise(input_latent) if self.d is not None else torch.zeros_like(samples) + a_val = self.a.generate_noise(input_latent) if self.a is not None else make_zero_like(samples) + b_val = self.b.generate_noise(input_latent) if self.b is not None else make_zero_like(samples) + c_val = self.c.generate_noise(input_latent) if self.c is not None else make_zero_like(samples) + d_val = self.d.generate_noise(input_latent) if self.d is not None else make_zero_like(samples) # helper to convert a returned value into a list matching ref_list def to_list(val, ref_list): if val is None: - return [torch.zeros_like(r) for r in ref_list] + return [make_zero_like(r) for r in ref_list] # If val is a NestedTensor-like, return underlying list if hasattr(val, 'is_nested') and getattr(val, 'is_nested'): return val.unbind() @@ -141,7 +128,7 @@ class NoiseExecutor(): def merge_to_tensor(val, ref): if val is None: - return torch.zeros_like(ref) + return make_zero_like(ref) if hasattr(val, 'is_nested') and getattr(val, 'is_nested'): lst = val.unbind() return torch.cat(lst, dim=0) @@ -158,7 +145,7 @@ class NoiseExecutor(): return torch.cat([val[i].unsqueeze(0).expand(sample_list[i].shape[0], *val.shape[1:]) for i in range(len(sample_list))], dim=0) except Exception: pass - return torch.zeros_like(ref) + return make_zero_like(ref) merged_a = merge_to_tensor(a_val, merged_samples) merged_b = merge_to_tensor(b_val, merged_samples) @@ -201,8 +188,7 @@ class NoiseExecutor(): F = getIndexTensorAlongDim(merged_samples, time_dim) variables.update({'frame': F, 'frame_count': frame_count}) - visitor = TensorEvalVisitor(variables, variables['a'].shape) - merged_result = visitor.visit(self.tree) + merged_result = eval_tensor_expr_with_tree(self.tree, variables, variables['a'].shape) if hasattr(samples, 'is_nested') and getattr(samples, 'is_nested'): split_results = list(merged_result.split(sizes, dim=0)) diff --git a/more_math/VideoMathNode.py b/more_math/VideoMathNode.py index 01e4ed2..f3006ac 100644 --- a/more_math/VideoMathNode.py +++ b/more_math/VideoMathNode.py @@ -4,38 +4,27 @@ from comfy_api.latest import io from comfy_api.input_impl import VideoFromComponents from comfy_api.util import VideoComponents - -from antlr4 import CommonTokenStream, InputStream import torch -from .helper_functions import ThrowingErrorListener, getIndexTensorAlongDim +from .helper_functions import getIndexTensorAlongDim, eval_tensor_expr, make_zero_like -from .Parser.MathExprParser import MathExprParser -from .Parser.MathExprLexer import MathExprLexer -from .Parser.TensorEvalVisitor import TensorEvalVisitor class VideoMathNode(io.ComfyNode): """ - This node enables the use of math expressions on Latents. - inputs: - a, b, c, d: - Latent, bound to variables with the same name. Defaults to zero latent if not provided. - w, x, y, z: - Floats, bound to variables of the expression. Defaults to 0.0 if not provided. - Latent expression: - String, describing expression to aply to latents. - - outputs: - LATENT: - Returns a LATENT object that contains the result of the math expression applied to the input conditionings. + Enables math expressions on Video (images + audio). + + Inputs: + a, b, c, d: Video inputs (b, c, d default to zero if not provided) + w, x, y, z: Float variables for expressions + Audio: Expression for audio component + Images: Expression for image component + + Outputs: + VIDEO: Result of applying expressions to input videos """ - def __init__(self): - pass @classmethod def define_schema(cls) -> io.Schema: - """ - """ return io.Schema( node_id="mrmth_VideoMathNode", display_name="Video math", @@ -45,10 +34,10 @@ class VideoMathNode(io.ComfyNode): io.Video.Input(id="b", optional=True), io.Video.Input(id="c", optional=True), io.Video.Input(id="d", optional=True), - io.Float.Input(id="w", default=0.0,optional=True, force_input=True), - io.Float.Input(id="x", default=0.0,optional=True, force_input=True), - io.Float.Input(id="y", default=0.0,optional=True, force_input=True), - io.Float.Input(id="z", default=0.0,optional=True, force_input=True), + io.Float.Input(id="w", default=0.0, optional=True, force_input=True), + io.Float.Input(id="x", default=0.0, optional=True, force_input=True), + io.Float.Input(id="y", default=0.0, optional=True, lazy=True, force_input=True), + io.Float.Input(id="z", default=0.0, optional=True, force_input=True), io.String.Input(id="Audio", default="a*(1-w)+b*w", tooltip="Expression to apply on audio part of video"), io.String.Input(id="Images", default="a*(1-w)+b*w", tooltip="Expression to apply on image part of video"), ], @@ -57,109 +46,62 @@ class VideoMathNode(io.ComfyNode): ], ) - #RETURN_NAMES = ("image_output_name",) tooltip = cleandoc(__doc__) - #OUTPUT_NODE = False - #OUTPUT_TOOLTIPS = ("",) # Tooltips for the output node - - @classmethod - def execute(cls, Audio,Images, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0) -> io.NodeOutput: - + def execute(cls, Audio, Images, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0) -> io.NodeOutput: ac = a.get_components() - bc = b.get_components() if b is not None else VideoComponents(images=torch.zeros_like(ac.images), audio={'waveform':torch.zeros_like(ac.audio['waveform']),'sample_rate':ac.audio['sample_rate']}, frame_rate=ac.frame_rate,metadata=None) - cc = c.get_components() if c is not None else VideoComponents(images=torch.zeros_like(ac.images), audio={'waveform':torch.zeros_like(ac.audio['waveform']),'sample_rate':ac.audio['sample_rate']}, frame_rate=ac.frame_rate,metadata=None) - dc = d.get_components() if d is not None else VideoComponents(images=torch.zeros_like(ac.images), audio={'waveform':torch.zeros_like(ac.audio['waveform']),'sample_rate':ac.audio['sample_rate']}, frame_rate=ac.frame_rate,metadata=None) + + bc = b.get_components() if b is not None else make_zero_like(ac) + cc = c.get_components() if c is not None else make_zero_like(ac) + dc = d.get_components() if d is not None else make_zero_like(ac) + # Process images (permute to B, C, H, W) + imgs_a = ac.images.permute(0, 3, 1, 2) + imgs_b = bc.images.permute(0, 3, 1, 2) + imgs_c = cc.images.permute(0, 3, 1, 2) + imgs_d = dc.images.permute(0, 3, 1, 2) - # permute images to B, C, H, W - ac.images = ac.images.permute(0, 3, 1, 2) - bc.images = bc.images.permute(0, 3, 1, 2) - cc.images = cc.images.permute(0, 3, 1, 2) - dc.images = dc.images.permute(0, 3, 1, 2) - - B = getIndexTensorAlongDim(ac.images, 0) - C = getIndexTensorAlongDim(ac.images, 1) - X = getIndexTensorAlongDim(ac.images, 3) # W - Y = getIndexTensorAlongDim(ac.images, 2) # H - W = torch.full_like(Y, ac.images.shape[3], dtype=torch.float32) - H = torch.full_like(Y, ac.images.shape[2], dtype=torch.float32) - R = torch.full_like(Y, float(ac.frame_rate), dtype=torch.float32) - T = torch.full_like(Y, ac.images.shape[0], dtype=torch.float32) - - variables = {'a': ac.images, 'b': bc.images, 'c': cc.images, 'd': dc.images, 'w': w, 'x': x, 'y': y, 'z': z, - 'X':X,'Y':Y, - 'B':B,'frame':B, - 'W':W,'width':W, - 'H':H,'height':H, - 'C':C,'channel':C, - 'R':R,'frame_rate':R, - 'frame_count':ac.images.shape[0], - 'N':ac.images.shape[1],'channel_count':ac.images.shape[1]} - - input_stream = InputStream(Images) - lexer = MathExprLexer(input_stream) - stream = CommonTokenStream(lexer) - parser = MathExprParser(stream) - parser.addErrorListener(ThrowingErrorListener()) - tree = parser.expr() - visitor = TensorEvalVisitor(variables, ac.images.shape) - imgs = visitor.visit(tree) - # permute back to B, H, W, C - imgs = imgs.permute(0, 2, 3, 1) - - - B = getIndexTensorAlongDim(ac.audio['waveform'], 0) - C = getIndexTensorAlongDim(ac.audio['waveform'], 1) - S = getIndexTensorAlongDim(ac.audio['waveform'], 2) - R = torch.full_like(S, ac.audio['sample_rate'], dtype=torch.float32) - T = torch.full_like(S, ac.audio['waveform'].shape[2], dtype=torch.float32) - N= ac.audio['waveform'].shape[1] - - variables = { - 'a': ac.audio['waveform'], 'b': bc.audio['waveform'], 'c': cc.audio['waveform'], 'd': dc.audio['waveform'], + img_vars = { + 'a': imgs_a, 'b': imgs_b, 'c': imgs_c, 'd': imgs_d, 'w': w, 'x': x, 'y': y, 'z': z, - - 'B': B, 'batch': B, - 'C': C, 'channel': C, - 'S': S, 'sample': S, - 'R': R, 'sample_rate': R, - 'T': T, 'sample_count': T, - 'N': N, 'channel_count': N + 'X': getIndexTensorAlongDim(imgs_a, 3), + 'Y': getIndexTensorAlongDim(imgs_a, 2), + 'B': getIndexTensorAlongDim(imgs_a, 0), 'frame': getIndexTensorAlongDim(imgs_a, 0), + 'C': getIndexTensorAlongDim(imgs_a, 1), 'channel': getIndexTensorAlongDim(imgs_a, 1), + 'W': imgs_a.shape[3], 'width': imgs_a.shape[3], + 'H': imgs_a.shape[2], 'height': imgs_a.shape[2], + 'R': float(ac.frame_rate), 'frame_rate': float(ac.frame_rate), + 'T': imgs_a.shape[0], 'frame_count': imgs_a.shape[0], + 'N': imgs_a.shape[1], 'channel_count': imgs_a.shape[1], } - input_stream = InputStream(Audio) - lexer = MathExprLexer(input_stream) - stream = CommonTokenStream(lexer) - parser = MathExprParser(stream) - tree = parser.expr() + result_imgs = eval_tensor_expr(Images, img_vars, imgs_a.shape) + result_imgs = result_imgs.permute(0, 2, 3, 1) # Back to B, H, W, C - visitor = TensorEvalVisitor(variables, ac.audio['waveform'].shape) - result_tensor = visitor.visit(tree) + # Process audio + audio_a = ac.audio['waveform'] + audio_b = bc.audio['waveform'] + audio_c = cc.audio['waveform'] + audio_d = dc.audio['waveform'] - # Create output dictionary with the same sample rate - audioo = { - 'waveform': result_tensor, - 'sample_rate': ac.audio['sample_rate'] + audio_vars = { + 'a': audio_a, 'b': audio_b, 'c': audio_c, 'd': audio_d, + 'w': w, 'x': x, 'y': y, 'z': z, + 'B': getIndexTensorAlongDim(audio_a, 0), 'batch': getIndexTensorAlongDim(audio_a, 0), + 'C': getIndexTensorAlongDim(audio_a, 1), 'channel': getIndexTensorAlongDim(audio_a, 1), + 'S': getIndexTensorAlongDim(audio_a, 2), 'sample': getIndexTensorAlongDim(audio_a, 2), + 'R': ac.audio['sample_rate'], 'sample_rate': ac.audio['sample_rate'], + 'T': audio_a.shape[2], 'sample_count': audio_a.shape[2], + 'N': audio_a.shape[1], 'channel_count': audio_a.shape[1], } + result_audio = eval_tensor_expr(Audio, audio_vars, audio_a.shape) - - - - out = VideoFromComponents(VideoComponents(images=imgs, audio=audioo, frame_rate=ac.frame_rate,metadata=ac.metadata)) - return (out,) - - - """ - The node will always be re executed if any of the inputs change but - this method can be used to force the node to execute again even when the inputs don't change. - You can make this node return a number or a string. This value will be compared to the one returned the last time the node was - executed, if it is different the node will be executed again. - This method is used in the core repo for the LoadImage node where they return the image hash as a string, if the image hash - changes between executions the LoadImage node is executed again. - """ - #@classmethod - #def IS_CHANGED(s, image, string_field, int_field, float_field, print_to_screen): - # return "" + output = VideoFromComponents(VideoComponents( + images=result_imgs, + audio={'waveform': result_audio, 'sample_rate': ac.audio['sample_rate']}, + frame_rate=ac.frame_rate, + metadata=ac.metadata + )) + return (output,) diff --git a/more_math/helper_functions.py b/more_math/helper_functions.py index 6542c1b..760d856 100644 --- a/more_math/helper_functions.py +++ b/more_math/helper_functions.py @@ -1,51 +1,151 @@ from antlr4.error.ErrorListener import ErrorListener -from antlr4 import InputStream -from .Parser.MathExprLexer import MathExprLexer -from .Parser.MathExprParser import MathExprParser -from antlr4 import CommonTokenStream +from antlr4 import InputStream, CommonTokenStream import torch -def getIndexTensorAlongDim(tensor, dim): - shape = tensor.shape +from .Parser.MathExprLexer import MathExprLexer +from .Parser.MathExprParser import MathExprParser - # Create values: shape (size of dim) - values = torch.arange(shape[dim], dtype=torch.float32) - # Reshape values to align with the target dimension - view_shape = [1] * len(shape) - view_shape[dim] = shape[dim] - values = values.view(*view_shape) - - # Broadcast to full shape - return values.expand(*shape) - -def time_to_freq(element: torch.Tensor) -> torch.Tensor: - if element.ndim < 2: - raise ValueError("FFT requires at least 2 dimensions (Batch, Channel)") - dims = tuple(range(2, element.ndim)) - return torch.fft.fftn(element, dim=dims) - -def freq_to_time(element: torch.Tensor) -> torch.Tensor: - if element.ndim < 2: - raise ValueError("IFFT requires at least 2 dimensions (Batch, Channel)") - dims = tuple(range(2, element.ndim)) - return torch.fft.ifftn(element, dim=dims).real class ThrowingErrorListener(ErrorListener): + """Error listener that raises ValueError on syntax errors.""" def syntaxError(self, recognizer, offendingSymbol, line, column, msg, e): raise ValueError(f"Syntax error in expression at line {line}, col {column}: {msg}") +def parse_expr(expr: str): + """Parse a math expression and return the parse tree.""" + input_stream = InputStream(expr) + lexer = MathExprLexer(input_stream) + stream = CommonTokenStream(lexer) + parser = MathExprParser(stream) + parser.addErrorListener(ThrowingErrorListener()) + return parser.expr() + + +def eval_tensor_expr(expr: str, variables: dict, shape: tuple, device=None): + """Parse and evaluate a tensor math expression. + + Args: + expr: Math expression string + variables: Dict of variable names to tensor/scalar values + shape: Shape tuple for the TensorEvalVisitor + device: Optional device override + + Returns: + Result tensor from evaluating the expression + """ + from .Parser.TensorEvalVisitor import TensorEvalVisitor + tree = parse_expr(expr) + visitor = TensorEvalVisitor(variables, shape, device=device) + return visitor.visit(tree) + + +def eval_tensor_expr_with_tree(tree, variables: dict, shape: tuple, device=None): + """Evaluate a pre-parsed expression tree with TensorEvalVisitor.""" + from .Parser.TensorEvalVisitor import TensorEvalVisitor + visitor = TensorEvalVisitor(variables, shape, device=device) + return visitor.visit(tree) + + +def eval_float_expr(expr: str, variables: dict): + """Parse and evaluate a float math expression.""" + from .Parser.FloatEvalVisitor import FloatEvalVisitor + tree = parse_expr(expr) + visitor = FloatEvalVisitor(variables) + return visitor.visit(tree) + + +def eval_float_expr_with_tree(tree, variables: dict): + """Evaluate a pre-parsed expression tree with FloatEvalVisitor.""" + from .Parser.FloatEvalVisitor import FloatEvalVisitor + visitor = FloatEvalVisitor(variables) + return visitor.visit(tree) + + +def getIndexTensorAlongDim(tensor, dim): + """Create a tensor of indices along a dimension, broadcasted to full shape.""" + shape = tensor.shape + values = torch.arange(shape[dim], dtype=torch.float32, device=tensor.device) + view_shape = [1] * len(shape) + view_shape[dim] = shape[dim] + values = values.view(*view_shape) + return values.expand(*shape) + + def comonLazy(expr, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0): - variables = {'a':a,'b':b,'c':c,'d':d,'w':w,'x':x,'y':y,'z':z} + """Determine which lazy inputs are needed based on expression variables.""" + variables = {'a': a, 'b': b, 'c': c, 'd': d, 'w': w, 'x': x, 'y': y, 'z': z} need_eval = [] input_stream = InputStream(expr) - lexer = MathExprLexer(input_stream) stream = CommonTokenStream(lexer) stream.fill() for token in filter(lambda t: t.type == MathExprParser.VARIABLE, stream.tokens): if token.text in variables and variables[token.text] is None: need_eval.append(token.text) - return need_eval \ No newline at end of file + return need_eval + + +def make_zero_like(ref): + """ + Create a zero-initialized version of the reference object, maintaining its structure. + Handles Conditioning, Audio, Latent, VideoComponents, and Tensors. + """ + if ref is None: + return None + + # raw torch tensor + if torch.is_tensor(ref): + return torch.zeros_like(ref) + + # Conditioning: list of lists [[tensor, dict]] + if isinstance(ref, list) and len(ref) > 0 and isinstance(ref[0], list) and len(ref[0]) >= 2: + ref_tensor = ref[0][0] + # Ensure it's a tensor-like structure + if torch.is_tensor(ref_tensor): + ref_pooled = ref[0][1].get("pooled_output") + return [[ + torch.zeros_like(ref_tensor), + {"pooled_output": torch.zeros_like(ref_pooled) if ref_pooled is not None else None} + ]] + + # Audio or Latent: dict + if isinstance(ref, dict): + if 'waveform' in ref: # Audio + return { + 'waveform': torch.zeros_like(ref['waveform']), + 'sample_rate': ref['sample_rate'] + } + if 'samples' in ref: # Latent + return { + 'samples': torch.zeros_like(ref['samples']) + } + + # VideoComponents or other objects with images/audio attributes + if hasattr(ref, 'images') and hasattr(ref, 'audio'): + # Dynamically create same type (e.g. VideoComponents) + return type(ref)( + images=torch.zeros_like(ref.images), + audio={'waveform': torch.zeros_like(ref.audio['waveform']), 'sample_rate': ref.audio['sample_rate']}, + frame_rate=ref.frame_rate, + metadata=None + ) + + return None + + +# Legacy FFT functions (kept for backward compatibility, but now unused) +def time_to_freq(element: torch.Tensor) -> torch.Tensor: + if element.ndim < 2: + raise ValueError("FFT requires at least 2 dimensions (Batch, Channel)") + dims = tuple(range(2, element.ndim)) + return torch.fft.fftn(element, dim=dims) + + +def freq_to_time(element: torch.Tensor) -> torch.Tensor: + if element.ndim < 2: + raise ValueError("IFFT requires at least 2 dimensions (Batch, Channel)") + dims = tuple(range(2, element.ndim)) + return torch.fft.ifftn(element, dim=dims).real \ No newline at end of file diff --git a/more_math/modelLikeCommon.py b/more_math/modelLikeCommon.py index 4d909bb..e0541b9 100644 --- a/more_math/modelLikeCommon.py +++ b/more_math/modelLikeCommon.py @@ -1,21 +1,13 @@ -from antlr4.error.ErrorListener import ErrorListener -from antlr4 import InputStream -from .Parser.MathExprLexer import MathExprLexer -from .Parser.MathExprParser import MathExprParser -from .Parser.TensorEvalVisitor import TensorEvalVisitor -from antlr4 import CommonTokenStream import torch -from .helper_functions import getIndexTensorAlongDim, ThrowingErrorListener import comfy.utils +from .helper_functions import getIndexTensorAlongDim, parse_expr, eval_tensor_expr_with_tree + + def calculate_patches(Model, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0): - # Parse expression - input_stream = InputStream(Model) - lexer = MathExprLexer(input_stream) - stream = CommonTokenStream(lexer) - parser = MathExprParser(stream) - parser.addErrorListener(ThrowingErrorListener()) - tree = parser.expr() + """Calculate model weight patches by applying math expression to state dicts.""" + # Parse expression once + tree = parse_expr(Model) sd_a = a.model.state_dict() sd_b = b.model.state_dict() if b is not None else {} @@ -24,55 +16,41 @@ def calculate_patches(Model, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0 patches = {} layer_count = len(sd_a) - pbar = comfy.utils.ProgressBar(layer_count) - # Iterate over all keys in the main model 'a' for i, (key, tens_a) in enumerate(sd_a.items()): - # Get corresponding tensors from other models, defaulting to zeros if missing or models not provided - tens_b = sd_b.get(key, None) - if tens_b is None: tens_b = torch.zeros_like(tens_a,device=tens_a.device) - else: tens_b = tens_b.to(tens_a.device) + # Get tensors from other models, default to zeros + tens_b = sd_b.get(key) + tens_b = tens_b.to(tens_a.device) if tens_b is not None else torch.zeros_like(tens_a) - tens_c = sd_c.get(key, None) - if tens_c is None: tens_c = torch.zeros_like(tens_a,device=tens_a.device) - else: tens_c = tens_c.to(tens_a.device) + tens_c = sd_c.get(key) + tens_c = tens_c.to(tens_a.device) if tens_c is not None else torch.zeros_like(tens_a) - tens_d = sd_d.get(key, None) - if tens_d is None: tens_d = torch.zeros_like(tens_a,device=tens_a.device) - else: tens_d = tens_d.to(tens_a.device) + tens_d = sd_d.get(key) + tens_d = tens_d.to(tens_a.device) if tens_d is not None else torch.zeros_like(tens_a) - # Variables for the visitor + # Build variables variables = { - 'a': tens_a, - 'b': tens_b, - 'c': tens_c, - 'd': tens_d, + 'a': tens_a, 'b': tens_b, 'c': tens_c, 'd': tens_d, 'w': w, 'x': x, 'y': y, 'z': z, 'L': i, 'layer': i, - 'LC': layer_count, 'layer_count': layer_count + 'LC': layer_count, 'layer_count': layer_count, } + # Add dimension index tensors for dim_idx in range(tens_a.ndim): - idx_tensor = getIndexTensorAlongDim(tens_a, dim_idx) - idx_tensor = idx_tensor.to(tens_a.device) - variables[f'D{dim_idx}'] = idx_tensor - variables[f'dim_{dim_idx}'] = idx_tensor + idx_tensor = getIndexTensorAlongDim(tens_a, dim_idx) + variables[f'D{dim_idx}'] = idx_tensor + variables[f'dim_{dim_idx}'] = idx_tensor - visitor = TensorEvalVisitor(variables, tens_a.shape) - result_tensor = visitor.visit(tree) + result_tensor = eval_tensor_expr_with_tree(tree, variables, tens_a.shape) - # Calculate difference for patching - # The patch should be: result - original - # Because ComfyUI applies: original + patch + # Calculate patch (diff from original) diff = result_tensor - tens_a - # Allow skipping zero patches to save memory - if torch.all(diff == 0): - continue - - # Store patch. ComfyUI expects { key: (tensor,) } usually - patches[key] = (diff,) + # Skip zero patches to save memory + if not torch.all(diff == 0): + patches[key] = (diff,) pbar.update(1) diff --git a/tests/conftest.py b/tests/conftest.py index f481911..5ea0eb3 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,21 +1,17 @@ import os import sys -# Ensure test discovery (Visual Studio / pytest) can import the package. -# conftest.py is imported during collection, so top-level path changes affect discovery. _here = os.path.abspath(os.path.dirname(__file__)) _project_root = os.path.abspath(os.path.join(_here, os.pardir)) -_src_path = os.path.join(_project_root, "src") +_comfy_root = os.path.abspath(os.path.join(_project_root, os.pardir, os.pardir)) -if os.path.isdir(_src_path): - _path_to_add = _src_path -else: - _path_to_add = _project_root - -if _path_to_add not in sys.path: - sys.path.insert(0, _path_to_add) - # also make it visible to subprocesses that inspect PYTHONPATH - os.environ["PYTHONPATH"] = _path_to_add + os.pathsep + os.environ.get("PYTHONPATH", "") +# Add project root and ComfyUI root to sys.path +# This ensures import comfy_api and import comfy work. +for p in [_project_root, _comfy_root]: + if p not in sys.path: + sys.path.insert(0, p) + # also make it visible to subprocesses that inspect PYTHONPATH + os.environ["PYTHONPATH"] = p + os.pathsep + os.environ.get("PYTHONPATH", "") import pytest