From 0269617a92f373ffa7ded130ab343695e3ca5571 Mon Sep 17 00:00:00 2001 From: mcDandy Date: Tue, 29 Jul 2025 17:08:19 +0200 Subject: [PATCH] add float math node --- src/more_math/FloatMathNode.py | 118 +++++++++++++++++++ src/more_math/Parser/FloatEvalVisitor.py | 138 +++++++++++++++++++++++ src/more_math/nodes.py | 6 +- 3 files changed, 261 insertions(+), 1 deletion(-) create mode 100644 src/more_math/FloatMathNode.py create mode 100644 src/more_math/Parser/FloatEvalVisitor.py diff --git a/src/more_math/FloatMathNode.py b/src/more_math/FloatMathNode.py new file mode 100644 index 0000000..0cdeaaf --- /dev/null +++ b/src/more_math/FloatMathNode.py @@ -0,0 +1,118 @@ +from inspect import cleandoc +from math import e + +from antlr4 import CommonTokenStream, InputStream + +from .Parser.MathExprParser import MathExprParser +from .Parser.MathExprLexer import MathExprLexer +from .Parser.FloatEvalVisitor import FloatEvalVisitor + +class FloatMathNode: + """ + 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. + """ + def __init__(self): + print("FloatMathNode initialized") + + @classmethod + def INPUT_TYPES(s): + """ + """ + return { + "required": { + "a": ("FLOAT", { + "default": 0, + "forceInput":True + }), + + "FloatFunc": ("STRING", { + "multiline": False, #True if you want the field to look like the one on the ClipTextEncode node + "default": "a*(1-w)+b*w", + "description": "Describes composition of the image. Valid functions are sin, cos, tan, asin, acos, atan, atan2, sinh, cosh, tanh, asinh, acosh, atanh, abs, sqrt, ln, log, exp, pow, min, max, norm, floor, ceil, round, gamma. Valid operators are +, -, *, /, %, ^,!˛&,|. Usable constants are e and pi." + + }), + }, + "optional": { + "b": ("FLOAT", { + "default": 0, + "forceInput":True + }), + "c": ("FLOAT", { + "default": 0, + "forceInput":True + }), + "d": ("FLOAT", { + "default": 0, + "forceInput":True + }), + "w": ("FLOAT", { + "default": 0, + "forceInput":True + }), + "x": ("FLOAT", { + "default": 0, + "forceInput":True + }), + "y": ("FLOAT", { + "default": 0, + "forceInput":True + }), + "z": ("FLOAT", { + "default": 0, + "forceInput":True + }), + + + # "int_field": ("INT", {"default": 0, "min": 0, "max": 100, "step": 1}), + # "float_field": ("FLOAT", {"default": 0.5, "min": -10.0, "max": 10.0, "step": 0.001}), + } + } + + RETURN_TYPES = ("FLOAT",) + #RETURN_NAMES = ("image_output_name",) + DESCRIPTION = cleandoc(__doc__) + FUNCTION = "fltMathNode" + + #OUTPUT_NODE = False + #OUTPUT_TOOLTIPS = ("",) # Tooltips for the output node + + CATEGORY = "More math" + + def fltMathNode(self, 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) + tree = parser.expr() + print("Tensor\n"+tree.toStringTree(recog=parser)) + visitor = FloatEvalVisitor(variables) + result = visitor.visit(tree) + print("Result:", result) + 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/src/more_math/Parser/FloatEvalVisitor.py b/src/more_math/Parser/FloatEvalVisitor.py new file mode 100644 index 0000000..8c19af4 --- /dev/null +++ b/src/more_math/Parser/FloatEvalVisitor.py @@ -0,0 +1,138 @@ +import math +from tkinter import SE +from sympy import true +import torch +from .MathExprVisitor import MathExprVisitor + +class FloatEvalVisitor(MathExprVisitor): + def __init__(self, variables): + self.variables = variables + + def visitNumberExp(self, ctx): + print("Visiting number expression:", ctx.getText()) + return float(ctx.getText()) + + def visitConstantExp(self, ctx): + name = ctx.getText().lower() + if name == "pi": + return 3.141592653589793 + if name == "e": + return 2.718281828459045 + raise ValueError(f"Unknown constant: {name}") + + def visitVariableExp(self, ctx): + name = ctx.getText() + if name not in self.variables: + raise ValueError(f"Variable '{name}' not found") + return self.variables[name] + + def visitParenExp(self, ctx): + return self.visit(ctx.expr()) + + def visitUnaryPlus(self, ctx): + return +self.visit(ctx.unaryExpr()) + + def visitUnaryMinus(self, ctx): + return -self.visit(ctx.unaryExpr()) + + def visitAddExp(self, ctx): + return self.visit(ctx.addExpr()) + self.visit(ctx.mulExpr()) + + def visitSubExp(self, ctx): + return self.visit(ctx.addExpr()) - self.visit(ctx.mulExpr()) + + def visitMulExp(self, ctx): + print("Visiting multiplication expression:", ctx.getText()) + if ctx.mulExpr() is None or ctx.powExpr() is None: + raise ValueError("Invalid multiplication expression") + return self.visit(ctx.mulExpr()) * self.visit(ctx.powExpr()) + + def visitDivExp(self, ctx): + return self.visit(ctx.mulExpr()) / self.visit(ctx.powExpr()) + + def visitModExp(self, ctx): + return self.visit(ctx.mulExpr()) % self.visit(ctx.powExpr()) + + def visitPowExp(self, ctx): + return math.pow(self.visit(ctx.unaryExpr()), self.visit(ctx.powExpr())) + + def visitToUnary(self, ctx): + return self.visit(ctx.unaryExpr()) + + def visitToPow(self, ctx): + return self.visit(ctx.powExpr()) + + def visitToMul(self, ctx): + return self.visit(ctx.mulExpr()) + + def visitToAdd(self, ctx): + return self.visit(ctx.addExpr()) + + def visitToAnd(self, ctx): + return self.visit(ctx.andExpr()) + + def visitToXor(self, ctx): + return self.visit(ctx.xorExpr()) + + def visitOrExp(self, ctx): + return self.visit(ctx.orExpr()).bool() | self.visit(ctx.xorExpr()).bool() + + def visitXorExp(self, ctx): + return self.visit(ctx.xorExpr()).bool() ^ self.visit(ctx.andExpr()).bool() + + def visitAndExp(self, ctx): + return self.visit(ctx.andExpr()).bool() & self.visit(ctx.addExpr()).bool() + + # Single-argument functions + def visitSinFunc(self, ctx): return math.sin(self.visit(ctx.expr())) + def visitCosFunc(self, ctx): return math.cos(self.visit(ctx.expr())) + def visitTanFunc(self, ctx): return math.tan(self.visit(ctx.expr())) + def visitAsinFunc(self, ctx): return math.asin(self.visit(ctx.expr())) + def visitAcosFunc(self, ctx): return math.acos(self.visit(ctx.expr())) + def visitAtanFunc(self, ctx): return math.atan(self.visit(ctx.expr())) + def visitSinhFunc(self, ctx): return math.sinh(self.visit(ctx.expr())) + def visitCoshFunc(self, ctx): return math.cosh(self.visit(ctx.expr())) + def visitTanhFunc(self, ctx): return math.tanh(self.visit(ctx.expr())) + def visitAsinhFunc(self, ctx): return math.asinh(self.visit(ctx.expr())) + def visitAcoshFunc(self, ctx): return math.acosh(self.visit(ctx.expr())) + def visitAtanhFunc(self, ctx): return math.atanh(self.visit(ctx.expr())) + def visitAbsFunc(self, ctx): return math.abs(self.visit(ctx.expr())) + def visitSqrtFunc(self, ctx): return math.sqrt(self.visit(ctx.expr())) + def visitLnFunc(self, ctx): return math.log(self.visit(ctx.expr())) + def visitLogFunc(self, ctx): return math.log10(self.visit(ctx.expr())) + def visitExpFunc(self, ctx): return math.exp(self.visit(ctx.expr())) + def visitNormFunc(self, ctx): return math.sqrt(math.avg(x**2 for x in self.visit(ctx.expr()))) + def visitFloorFunc(self, ctx): return math.floor(self.visit(ctx.expr())) + def visitCeilFunc(self, ctx): return math.ceil(self.visit(ctx.expr())) + def visitRoundFunc(self, ctx): return math.round(self.visit(ctx.expr())) + def visitGammaFunc(self, ctx): return math.gamma(self.visit(ctx.expr())).exp() + + # Two-argument functions + def visitPowFunc(self, ctx): + return math.pow(self.visit(ctx.expr(0)), self.visit(ctx.expr(1))) + def visitAtan2Func(self, ctx): + return math.atan2(self.visit(ctx.expr(0)), self.visit(ctx.expr(1))) + + # N-argument functions + def visitMinFunc(self, ctx): + args = [self.visit(e) for e in ctx.expr()] + return math.min(args) + def visitMaxFunc(self, ctx): + args = [self.visit(e) for e in ctx.expr()] + return math.max(args) + + def visitFunc1Exp(self, ctx): + return self.visitChildren(ctx) + def visitFunc2Exp(self, ctx): + return self.visitChildren(ctx) + def visitFuncNExp(self, ctx): + return self.visitChildren(ctx) + def visitAtomExp(self, ctx): + return self.visitChildren(ctx) + + def visitFunc2Expr(self, ctx): + return self.visit(ctx.getChild(0)) # forward to Atan2Func, PowFunc, etc. + + def visitExpr(self, ctx): + print("Visiting expression:", ctx.getText()) + return self.visitChildren(ctx) diff --git a/src/more_math/nodes.py b/src/more_math/nodes.py index c974dc5..aac0a56 100644 --- a/src/more_math/nodes.py +++ b/src/more_math/nodes.py @@ -1,3 +1,4 @@ +from .FloatMathNode import FloatMathNode from .ConditioningMathNode import ConditioningMathNode from .LatentMathNode import LatentMathNode from .ImageMathNode import ImageMathNode @@ -46,17 +47,20 @@ NODE_CLASS_MAPPINGS = { "mrmth_ImageMathNode": ImageMathNode, "mrmth_IntToFloat": IntToFloatNode, "mrmth_FloatToInt": FloatToIntNode, + "mrmth_FloatMathNode": FloatMathNode, } NODE_DISPLAY_NAME_MAPPINGS = { "mrmth_ConditioningMathNode": "Conditioning math node", "mrmth_LatentMathNode": "Latent math node", "mrmth_ImageMathNode": "Image math node", - "mrmth_IntToFloat": "Int → Float", + "mrmth_FloatMathNode": "Float math node", + "mrmth_IntToFloat": "Int → Float", "mrmth_FloatToInt": "Float → Int", "Tensor": "Tensor expression", "Latent": "Latent expression", "Image": "Image expression", + "FloatFunc": "Float expression", "pooled_output": "Pooled output tensor expression" }