add float math node

This commit is contained in:
mcDandy
2025-07-29 17:08:19 +02:00
parent 431026c4ad
commit 0269617a92
3 changed files with 261 additions and 1 deletions
+118
View File
@@ -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 ""
+138
View File
@@ -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)
+5 -1
View File
@@ -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"
}