Added more functions

This commit is contained in:
mcDandy
2025-12-05 20:20:35 +01:00
parent b1c96340da
commit fa8c415eb9
10 changed files with 1474 additions and 1151 deletions
+9
View File
@@ -28,6 +28,9 @@ You can also get the node from comfy manager under the name of More math.
- Hyperbolic: `sinh`, `cosh`, `tanh`, `asinh`, `acosh`, `atanh`
- Aggregates: `smin`, `smax` , `snorm` (scalar), `tmin`, `tmax`, `tnorm` (elementwise)
- Other: `floor`, `ceil`, `round`, `gamma`, `clamp`, `sigm` (sigmoid) `fft` (N-D FFT on non-batch/channel dims), `ifft` (Inverse N-D FFT, returns real component), `angle` (in ifft only)
- Shaders: `lerp(a, b, w)`, `step(edge, x)`, `smoothstep(edge0, edge1, x)`, `fract(x)`
- ML: `relu(x)`, `softplus(x)`, `gelu(x)`, `sign(x)`
- Tensor: `swap(tensor, dim, index1, index2)` (swaps two slices of a tensor along a dimension)
## Variables
- **common inputs** (matches node input type):
@@ -66,3 +69,9 @@ You can also get the node from comfy manager under the name of More math.
- no additional variables
- Constants: `e`, `pi`
## Examples
### Swap Red and Blue Channels
`swap(a, 1, 0, 2)`
*(for ImageMathNode, where dim 1 is Channel)*
+49
View File
@@ -114,7 +114,24 @@ class FloatEvalVisitor(MathExprVisitor):
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 visitFractFunc(self, ctx):
val = self.visit(ctx.expr())
return val - math.floor(val)
def visitSigmoidFunc(self, ctx): return 1/(1+math.exp(-self.visit(ctx.expr())))
def visitReluFunc(self, ctx): return max(0.0, self.visit(ctx.expr()))
def visitSoftplusFunc(self, ctx):
# log(1 + exp(x))
x = self.visit(ctx.expr())
# stability check? math.log1p(math.exp(x)) is better but might overflow for large x
if x > 20: return x
return math.log(1 + math.exp(x))
def visitGeluFunc(self, ctx):
# 0.5 * x * (1 + erf(x / sqrt(2)))
x = self.visit(ctx.expr())
return 0.5 * x * (1 + math.erf(x / 1.4142135623730951))
def visitSignFunc(self, ctx):
x = self.visit(ctx.expr())
return math.copysign(1.0, x) if x != 0 else 0.0
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()
@@ -128,6 +145,10 @@ class FloatEvalVisitor(MathExprVisitor):
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)))
def visitStepFunc(self, ctx):
edge = self.visit(ctx.expr(0))
x = self.visit(ctx.expr(1))
return 1.0 if x >= edge else 0.0
# N-argument functions
def visitSMinFunc(self, ctx):
@@ -137,10 +158,38 @@ class FloatEvalVisitor(MathExprVisitor):
args = [self.visit(e) for e in ctx.expr()]
return math.max(args)
def visitClampFunc(self, ctx):
x = self.visit(ctx.expr(0))
min_val = self.visit(ctx.expr(1))
max_val = self.visit(ctx.expr(2))
return max(min(x, max_val), min_val)
def visitLerpFunc(self, ctx):
# a + (b - a) * w
a = self.visit(ctx.expr(0))
b = self.visit(ctx.expr(1))
w = self.visit(ctx.expr(2))
return a + (b - a) * w
def visitSmoothstepFunc(self, ctx):
edge0 = self.visit(ctx.expr(0))
edge1 = self.visit(ctx.expr(1))
x = self.visit(ctx.expr(2))
# Scale, bias and saturate x to 0..1 range
t = (x - edge0) / (edge1 - edge0)
t = max(0.0, min(1.0, t))
# Evaluate polynomial
return t * t * (3.0 - 2.0 * t)
def visitFunc1Exp(self, ctx):
return self.visitChildren(ctx)
def visitFunc2Exp(self, ctx):
return self.visitChildren(ctx)
def visitFunc3Exp(self, ctx):
return self.visitChildren(ctx)
def visitFunc4Exp(self, ctx):
return self.visitChildren(ctx)
def visitFuncNExp(self, ctx):
return self.visitChildren(ctx)
def visitAtomExp(self, ctx):
+25 -3
View File
@@ -44,6 +44,7 @@ atom
: func1 # Func1Exp
| func2 # Func2Exp
| func3 # Func3Exp
| func4 # Func4Exp
| funcN # FuncNExp
| VARIABLE # VariableExp
| NUMBER # NumberExp
@@ -79,8 +80,13 @@ func1
| SIGM '(' expr ')' # sigmoidFunc
| SFFT '(' expr ')' # sfftFunc
| SIFFT '(' expr ')' # sifftFunc
| ANGL '(' expr ')' # anglFunc
| PRNT '(' expr ')' # printFunc
| ANGL '(' expr ')' # anglFunc
| PRNT '(' expr ')' # printFunc
| FRACT '(' expr ')' # FractFunc
| RELU '(' expr ')' # ReluFunc
| SOFTPLUS '(' expr ')' # SoftplusFunc
| GELU '(' expr ')' # GeluFunc
| SIGN '(' expr ')' # SignFunc
;
@@ -90,9 +96,16 @@ func2
| ATAN2 '(' expr ',' expr ')' # Atan2Func
| TMIN '(' expr ',' expr ')' # TMinFunc
| TMAX '(' expr ',' expr ')' # TMaxFunc
| STEP '(' expr ',' expr ')' # StepFunc
;
func3
: CLAMP '(' expr ',' expr ',' expr ')' # ClampFunc
: CLAMP '(' expr ',' expr ',' expr ')' # ClampFunc
| LERP '(' expr ',' expr ',' expr ')' # LerpFunc
| SMOOTHSTEP '(' expr ',' expr ',' expr ')' # SmoothstepFunc
;
func4
: SWAP '(' expr ',' expr ',' expr ',' expr ')' # SwapFunc
;
// N-argument functions (at least 2 arguments)
funcN
@@ -138,6 +151,15 @@ SFFT : 'fft';
SIFFT : 'ifft';
ANGL : 'angle';
PRNT : 'print';
LERP : 'lerp';
STEP : 'step';
SMOOTHSTEP : 'smoothstep';
FRACT : 'fract';
RELU : 'relu';
SOFTPLUS : 'softplus';
GELU : 'gelu';
SIGN : 'sign';
SWAP : 'swap';
PLUS : '+';
MINUS : '-';
+46 -28
View File
@@ -36,22 +36,31 @@ SFFT=35
SIFFT=36
ANGL=37
PRNT=38
PLUS=39
MINUS=40
MULT=41
DIV=42
MOD=43
POW=44
GE=45
GT=46
LE=47
LT=48
EQ=49
NE=50
CONSTANT=51
NUMBER=52
VARIABLE=53
WS=54
LERP=39
STEP=40
SMOOTHSTEP=41
FRACT=42
RELU=43
SOFTPLUS=44
GELU=45
SIGN=46
SWAP=47
PLUS=48
MINUS=49
MULT=50
DIV=51
MOD=52
POW=53
GE=54
GT=55
LE=56
LT=57
EQ=58
NE=59
CONSTANT=60
NUMBER=61
VARIABLE=62
WS=63
'('=1
')'=2
','=3
@@ -90,15 +99,24 @@ WS=54
'ifft'=36
'angle'=37
'print'=38
'+'=39
'-'=40
'*'=41
'/'=42
'%'=43
'^'=44
'>='=45
'>'=46
'<='=47
'<'=48
'=='=49
'!='=50
'lerp'=39
'step'=40
'smoothstep'=41
'fract'=42
'relu'=43
'softplus'=44
'gelu'=45
'sign'=46
'swap'=47
'+'=48
'-'=49
'*'=50
'/'=51
'%'=52
'^'=53
'>='=54
'>'=55
'<='=56
'<'=57
'=='=58
'!='=59
+181 -143
View File
@@ -1,4 +1,4 @@
# Generated from MathExpr.g4 by ANTLR 4.13.2
# Generated from src/more_math/Parser/MathExpr.g4 by ANTLR 4.13.2
from antlr4 import *
from io import StringIO
import sys
@@ -10,7 +10,7 @@ else:
def serializedATN():
return [
4,0,54,354,6,-1,2,0,7,0,2,1,7,1,2,2,7,2,2,3,7,3,2,4,7,4,2,5,7,5,
4,0,63,428,6,-1,2,0,7,0,2,1,7,1,2,2,7,2,2,3,7,3,2,4,7,4,2,5,7,5,
2,6,7,6,2,7,7,7,2,8,7,8,2,9,7,9,2,10,7,10,2,11,7,11,2,12,7,12,2,
13,7,13,2,14,7,14,2,15,7,15,2,16,7,16,2,17,7,17,2,18,7,18,2,19,7,
19,2,20,7,20,2,21,7,21,2,22,7,22,2,23,7,23,2,24,7,24,2,25,7,25,2,
@@ -18,122 +18,148 @@ def serializedATN():
32,2,33,7,33,2,34,7,34,2,35,7,35,2,36,7,36,2,37,7,37,2,38,7,38,2,
39,7,39,2,40,7,40,2,41,7,41,2,42,7,42,2,43,7,43,2,44,7,44,2,45,7,
45,2,46,7,46,2,47,7,47,2,48,7,48,2,49,7,49,2,50,7,50,2,51,7,51,2,
52,7,52,2,53,7,53,1,0,1,0,1,1,1,1,1,2,1,2,1,3,1,3,1,3,1,3,1,4,1,
4,1,4,1,4,1,5,1,5,1,5,1,5,1,6,1,6,1,6,1,6,1,6,1,7,1,7,1,7,1,7,1,
7,1,8,1,8,1,8,1,8,1,8,1,9,1,9,1,9,1,9,1,9,1,9,1,10,1,10,1,10,1,10,
1,10,1,11,1,11,1,11,1,11,1,11,1,12,1,12,1,12,1,12,1,12,1,13,1,13,
1,13,1,13,1,13,1,13,1,14,1,14,1,14,1,14,1,14,1,14,1,15,1,15,1,15,
1,15,1,15,1,15,1,16,1,16,1,16,1,16,1,17,1,17,1,17,1,17,1,17,1,18,
1,18,1,18,1,19,1,19,1,19,1,19,1,20,1,20,1,20,1,20,1,21,1,21,1,21,
1,21,1,21,1,22,1,22,1,22,1,22,1,22,1,23,1,23,1,23,1,23,1,23,1,24,
1,24,1,24,1,24,1,24,1,25,1,25,1,25,1,25,1,25,1,25,1,26,1,26,1,26,
1,26,1,26,1,26,1,27,1,27,1,27,1,27,1,27,1,27,1,28,1,28,1,28,1,28,
1,28,1,29,1,29,1,29,1,29,1,29,1,29,1,30,1,30,1,30,1,30,1,30,1,30,
1,31,1,31,1,31,1,31,1,32,1,32,1,32,1,32,1,32,1,33,1,33,1,33,1,33,
1,33,1,33,1,34,1,34,1,34,1,34,1,35,1,35,1,35,1,35,1,35,1,36,1,36,
1,36,1,36,1,36,1,36,1,37,1,37,1,37,1,37,1,37,1,37,1,38,1,38,1,39,
1,39,1,40,1,40,1,41,1,41,1,42,1,42,1,43,1,43,1,44,1,44,1,44,1,45,
1,45,1,46,1,46,1,46,1,47,1,47,1,48,1,48,1,48,1,49,1,49,1,49,1,50,
1,50,1,50,1,50,1,50,3,50,326,8,50,1,51,4,51,329,8,51,11,51,12,51,
330,1,51,1,51,4,51,335,8,51,11,51,12,51,336,3,51,339,8,51,1,52,1,
52,5,52,343,8,52,10,52,12,52,346,9,52,1,53,4,53,349,8,53,11,53,12,
53,350,1,53,1,53,0,0,54,1,1,3,2,5,3,7,4,9,5,11,6,13,7,15,8,17,9,
19,10,21,11,23,12,25,13,27,14,29,15,31,16,33,17,35,18,37,19,39,20,
41,21,43,22,45,23,47,24,49,25,51,26,53,27,55,28,57,29,59,30,61,31,
63,32,65,33,67,34,69,35,71,36,73,37,75,38,77,39,79,40,81,41,83,42,
85,43,87,44,89,45,91,46,93,47,95,48,97,49,99,50,101,51,103,52,105,
53,107,54,1,0,5,2,0,69,69,101,101,1,0,48,57,3,0,65,90,95,95,97,122,
4,0,48,57,65,90,95,95,97,122,3,0,9,10,13,13,32,32,360,0,1,1,0,0,
0,0,3,1,0,0,0,0,5,1,0,0,0,0,7,1,0,0,0,0,9,1,0,0,0,0,11,1,0,0,0,0,
13,1,0,0,0,0,15,1,0,0,0,0,17,1,0,0,0,0,19,1,0,0,0,0,21,1,0,0,0,0,
23,1,0,0,0,0,25,1,0,0,0,0,27,1,0,0,0,0,29,1,0,0,0,0,31,1,0,0,0,0,
33,1,0,0,0,0,35,1,0,0,0,0,37,1,0,0,0,0,39,1,0,0,0,0,41,1,0,0,0,0,
43,1,0,0,0,0,45,1,0,0,0,0,47,1,0,0,0,0,49,1,0,0,0,0,51,1,0,0,0,0,
53,1,0,0,0,0,55,1,0,0,0,0,57,1,0,0,0,0,59,1,0,0,0,0,61,1,0,0,0,0,
63,1,0,0,0,0,65,1,0,0,0,0,67,1,0,0,0,0,69,1,0,0,0,0,71,1,0,0,0,0,
73,1,0,0,0,0,75,1,0,0,0,0,77,1,0,0,0,0,79,1,0,0,0,0,81,1,0,0,0,0,
83,1,0,0,0,0,85,1,0,0,0,0,87,1,0,0,0,0,89,1,0,0,0,0,91,1,0,0,0,0,
93,1,0,0,0,0,95,1,0,0,0,0,97,1,0,0,0,0,99,1,0,0,0,0,101,1,0,0,0,
0,103,1,0,0,0,0,105,1,0,0,0,0,107,1,0,0,0,1,109,1,0,0,0,3,111,1,
0,0,0,5,113,1,0,0,0,7,115,1,0,0,0,9,119,1,0,0,0,11,123,1,0,0,0,13,
127,1,0,0,0,15,132,1,0,0,0,17,137,1,0,0,0,19,142,1,0,0,0,21,148,
1,0,0,0,23,153,1,0,0,0,25,158,1,0,0,0,27,163,1,0,0,0,29,169,1,0,
0,0,31,175,1,0,0,0,33,181,1,0,0,0,35,185,1,0,0,0,37,190,1,0,0,0,
39,193,1,0,0,0,41,197,1,0,0,0,43,201,1,0,0,0,45,206,1,0,0,0,47,211,
1,0,0,0,49,216,1,0,0,0,51,221,1,0,0,0,53,227,1,0,0,0,55,233,1,0,
0,0,57,239,1,0,0,0,59,244,1,0,0,0,61,250,1,0,0,0,63,256,1,0,0,0,
65,260,1,0,0,0,67,265,1,0,0,0,69,271,1,0,0,0,71,275,1,0,0,0,73,280,
1,0,0,0,75,286,1,0,0,0,77,292,1,0,0,0,79,294,1,0,0,0,81,296,1,0,
0,0,83,298,1,0,0,0,85,300,1,0,0,0,87,302,1,0,0,0,89,304,1,0,0,0,
91,307,1,0,0,0,93,309,1,0,0,0,95,312,1,0,0,0,97,314,1,0,0,0,99,317,
1,0,0,0,101,325,1,0,0,0,103,328,1,0,0,0,105,340,1,0,0,0,107,348,
1,0,0,0,109,110,5,40,0,0,110,2,1,0,0,0,111,112,5,41,0,0,112,4,1,
0,0,0,113,114,5,44,0,0,114,6,1,0,0,0,115,116,5,115,0,0,116,117,5,
105,0,0,117,118,5,110,0,0,118,8,1,0,0,0,119,120,5,99,0,0,120,121,
5,111,0,0,121,122,5,115,0,0,122,10,1,0,0,0,123,124,5,116,0,0,124,
125,5,97,0,0,125,126,5,110,0,0,126,12,1,0,0,0,127,128,5,97,0,0,128,
129,5,115,0,0,129,130,5,105,0,0,130,131,5,110,0,0,131,14,1,0,0,0,
132,133,5,97,0,0,133,134,5,99,0,0,134,135,5,111,0,0,135,136,5,115,
0,0,136,16,1,0,0,0,137,138,5,97,0,0,138,139,5,116,0,0,139,140,5,
97,0,0,140,141,5,110,0,0,141,18,1,0,0,0,142,143,5,97,0,0,143,144,
5,116,0,0,144,145,5,97,0,0,145,146,5,110,0,0,146,147,5,50,0,0,147,
20,1,0,0,0,148,149,5,115,0,0,149,150,5,105,0,0,150,151,5,110,0,0,
151,152,5,104,0,0,152,22,1,0,0,0,153,154,5,99,0,0,154,155,5,111,
0,0,155,156,5,115,0,0,156,157,5,104,0,0,157,24,1,0,0,0,158,159,5,
116,0,0,159,160,5,97,0,0,160,161,5,110,0,0,161,162,5,104,0,0,162,
26,1,0,0,0,163,164,5,97,0,0,164,165,5,115,0,0,165,166,5,105,0,0,
166,167,5,110,0,0,167,168,5,104,0,0,168,28,1,0,0,0,169,170,5,97,
0,0,170,171,5,99,0,0,171,172,5,111,0,0,172,173,5,115,0,0,173,174,
5,104,0,0,174,30,1,0,0,0,175,176,5,97,0,0,176,177,5,116,0,0,177,
178,5,97,0,0,178,179,5,110,0,0,179,180,5,104,0,0,180,32,1,0,0,0,
181,182,5,97,0,0,182,183,5,98,0,0,183,184,5,115,0,0,184,34,1,0,0,
0,185,186,5,115,0,0,186,187,5,113,0,0,187,188,5,114,0,0,188,189,
5,116,0,0,189,36,1,0,0,0,190,191,5,108,0,0,191,192,5,110,0,0,192,
38,1,0,0,0,193,194,5,108,0,0,194,195,5,111,0,0,195,196,5,103,0,0,
196,40,1,0,0,0,197,198,5,101,0,0,198,199,5,120,0,0,199,200,5,112,
0,0,200,42,1,0,0,0,201,202,5,115,0,0,202,203,5,109,0,0,203,204,5,
105,0,0,204,205,5,110,0,0,205,44,1,0,0,0,206,207,5,115,0,0,207,208,
5,109,0,0,208,209,5,97,0,0,209,210,5,120,0,0,210,46,1,0,0,0,211,
212,5,116,0,0,212,213,5,109,0,0,213,214,5,105,0,0,214,215,5,110,
0,0,215,48,1,0,0,0,216,217,5,116,0,0,217,218,5,109,0,0,218,219,5,
97,0,0,219,220,5,120,0,0,220,50,1,0,0,0,221,222,5,116,0,0,222,223,
5,110,0,0,223,224,5,111,0,0,224,225,5,114,0,0,225,226,5,109,0,0,
226,52,1,0,0,0,227,228,5,115,0,0,228,229,5,110,0,0,229,230,5,111,
0,0,230,231,5,114,0,0,231,232,5,109,0,0,232,54,1,0,0,0,233,234,5,
102,0,0,234,235,5,108,0,0,235,236,5,111,0,0,236,237,5,111,0,0,237,
238,5,114,0,0,238,56,1,0,0,0,239,240,5,99,0,0,240,241,5,101,0,0,
241,242,5,105,0,0,242,243,5,108,0,0,243,58,1,0,0,0,244,245,5,114,
0,0,245,246,5,111,0,0,246,247,5,117,0,0,247,248,5,110,0,0,248,249,
5,100,0,0,249,60,1,0,0,0,250,251,5,103,0,0,251,252,5,97,0,0,252,
253,5,109,0,0,253,254,5,109,0,0,254,255,5,97,0,0,255,62,1,0,0,0,
256,257,5,112,0,0,257,258,5,111,0,0,258,259,5,119,0,0,259,64,1,0,
0,0,260,261,5,115,0,0,261,262,5,105,0,0,262,263,5,103,0,0,263,264,
5,109,0,0,264,66,1,0,0,0,265,266,5,99,0,0,266,267,5,108,0,0,267,
268,5,97,0,0,268,269,5,109,0,0,269,270,5,112,0,0,270,68,1,0,0,0,
271,272,5,102,0,0,272,273,5,102,0,0,273,274,5,116,0,0,274,70,1,0,
0,0,275,276,5,105,0,0,276,277,5,102,0,0,277,278,5,102,0,0,278,279,
5,116,0,0,279,72,1,0,0,0,280,281,5,97,0,0,281,282,5,110,0,0,282,
283,5,103,0,0,283,284,5,108,0,0,284,285,5,101,0,0,285,74,1,0,0,0,
286,287,5,112,0,0,287,288,5,114,0,0,288,289,5,105,0,0,289,290,5,
110,0,0,290,291,5,116,0,0,291,76,1,0,0,0,292,293,5,43,0,0,293,78,
1,0,0,0,294,295,5,45,0,0,295,80,1,0,0,0,296,297,5,42,0,0,297,82,
1,0,0,0,298,299,5,47,0,0,299,84,1,0,0,0,300,301,5,37,0,0,301,86,
1,0,0,0,302,303,5,94,0,0,303,88,1,0,0,0,304,305,5,62,0,0,305,306,
5,61,0,0,306,90,1,0,0,0,307,308,5,62,0,0,308,92,1,0,0,0,309,310,
5,60,0,0,310,311,5,61,0,0,311,94,1,0,0,0,312,313,5,60,0,0,313,96,
1,0,0,0,314,315,5,61,0,0,315,316,5,61,0,0,316,98,1,0,0,0,317,318,
5,33,0,0,318,319,5,61,0,0,319,100,1,0,0,0,320,321,5,112,0,0,321,
326,5,105,0,0,322,323,5,80,0,0,323,326,5,73,0,0,324,326,7,0,0,0,
325,320,1,0,0,0,325,322,1,0,0,0,325,324,1,0,0,0,326,102,1,0,0,0,
327,329,7,1,0,0,328,327,1,0,0,0,329,330,1,0,0,0,330,328,1,0,0,0,
330,331,1,0,0,0,331,338,1,0,0,0,332,334,5,46,0,0,333,335,7,1,0,0,
334,333,1,0,0,0,335,336,1,0,0,0,336,334,1,0,0,0,336,337,1,0,0,0,
337,339,1,0,0,0,338,332,1,0,0,0,338,339,1,0,0,0,339,104,1,0,0,0,
340,344,7,2,0,0,341,343,7,3,0,0,342,341,1,0,0,0,343,346,1,0,0,0,
344,342,1,0,0,0,344,345,1,0,0,0,345,106,1,0,0,0,346,344,1,0,0,0,
347,349,7,4,0,0,348,347,1,0,0,0,349,350,1,0,0,0,350,348,1,0,0,0,
350,351,1,0,0,0,351,352,1,0,0,0,352,353,6,53,0,0,353,108,1,0,0,0,
7,0,325,330,336,338,344,350,1,6,0,0
52,7,52,2,53,7,53,2,54,7,54,2,55,7,55,2,56,7,56,2,57,7,57,2,58,7,
58,2,59,7,59,2,60,7,60,2,61,7,61,2,62,7,62,1,0,1,0,1,1,1,1,1,2,1,
2,1,3,1,3,1,3,1,3,1,4,1,4,1,4,1,4,1,5,1,5,1,5,1,5,1,6,1,6,1,6,1,
6,1,6,1,7,1,7,1,7,1,7,1,7,1,8,1,8,1,8,1,8,1,8,1,9,1,9,1,9,1,9,1,
9,1,9,1,10,1,10,1,10,1,10,1,10,1,11,1,11,1,11,1,11,1,11,1,12,1,12,
1,12,1,12,1,12,1,13,1,13,1,13,1,13,1,13,1,13,1,14,1,14,1,14,1,14,
1,14,1,14,1,15,1,15,1,15,1,15,1,15,1,15,1,16,1,16,1,16,1,16,1,17,
1,17,1,17,1,17,1,17,1,18,1,18,1,18,1,19,1,19,1,19,1,19,1,20,1,20,
1,20,1,20,1,21,1,21,1,21,1,21,1,21,1,22,1,22,1,22,1,22,1,22,1,23,
1,23,1,23,1,23,1,23,1,24,1,24,1,24,1,24,1,24,1,25,1,25,1,25,1,25,
1,25,1,25,1,26,1,26,1,26,1,26,1,26,1,26,1,27,1,27,1,27,1,27,1,27,
1,27,1,28,1,28,1,28,1,28,1,28,1,29,1,29,1,29,1,29,1,29,1,29,1,30,
1,30,1,30,1,30,1,30,1,30,1,31,1,31,1,31,1,31,1,32,1,32,1,32,1,32,
1,32,1,33,1,33,1,33,1,33,1,33,1,33,1,34,1,34,1,34,1,34,1,35,1,35,
1,35,1,35,1,35,1,36,1,36,1,36,1,36,1,36,1,36,1,37,1,37,1,37,1,37,
1,37,1,37,1,38,1,38,1,38,1,38,1,38,1,39,1,39,1,39,1,39,1,39,1,40,
1,40,1,40,1,40,1,40,1,40,1,40,1,40,1,40,1,40,1,40,1,41,1,41,1,41,
1,41,1,41,1,41,1,42,1,42,1,42,1,42,1,42,1,43,1,43,1,43,1,43,1,43,
1,43,1,43,1,43,1,43,1,44,1,44,1,44,1,44,1,44,1,45,1,45,1,45,1,45,
1,45,1,46,1,46,1,46,1,46,1,46,1,47,1,47,1,48,1,48,1,49,1,49,1,50,
1,50,1,51,1,51,1,52,1,52,1,53,1,53,1,53,1,54,1,54,1,55,1,55,1,55,
1,56,1,56,1,57,1,57,1,57,1,58,1,58,1,58,1,59,1,59,1,59,1,59,1,59,
3,59,400,8,59,1,60,4,60,403,8,60,11,60,12,60,404,1,60,1,60,4,60,
409,8,60,11,60,12,60,410,3,60,413,8,60,1,61,1,61,5,61,417,8,61,10,
61,12,61,420,9,61,1,62,4,62,423,8,62,11,62,12,62,424,1,62,1,62,0,
0,63,1,1,3,2,5,3,7,4,9,5,11,6,13,7,15,8,17,9,19,10,21,11,23,12,25,
13,27,14,29,15,31,16,33,17,35,18,37,19,39,20,41,21,43,22,45,23,47,
24,49,25,51,26,53,27,55,28,57,29,59,30,61,31,63,32,65,33,67,34,69,
35,71,36,73,37,75,38,77,39,79,40,81,41,83,42,85,43,87,44,89,45,91,
46,93,47,95,48,97,49,99,50,101,51,103,52,105,53,107,54,109,55,111,
56,113,57,115,58,117,59,119,60,121,61,123,62,125,63,1,0,5,2,0,69,
69,101,101,1,0,48,57,3,0,65,90,95,95,97,122,4,0,48,57,65,90,95,95,
97,122,3,0,9,10,13,13,32,32,434,0,1,1,0,0,0,0,3,1,0,0,0,0,5,1,0,
0,0,0,7,1,0,0,0,0,9,1,0,0,0,0,11,1,0,0,0,0,13,1,0,0,0,0,15,1,0,0,
0,0,17,1,0,0,0,0,19,1,0,0,0,0,21,1,0,0,0,0,23,1,0,0,0,0,25,1,0,0,
0,0,27,1,0,0,0,0,29,1,0,0,0,0,31,1,0,0,0,0,33,1,0,0,0,0,35,1,0,0,
0,0,37,1,0,0,0,0,39,1,0,0,0,0,41,1,0,0,0,0,43,1,0,0,0,0,45,1,0,0,
0,0,47,1,0,0,0,0,49,1,0,0,0,0,51,1,0,0,0,0,53,1,0,0,0,0,55,1,0,0,
0,0,57,1,0,0,0,0,59,1,0,0,0,0,61,1,0,0,0,0,63,1,0,0,0,0,65,1,0,0,
0,0,67,1,0,0,0,0,69,1,0,0,0,0,71,1,0,0,0,0,73,1,0,0,0,0,75,1,0,0,
0,0,77,1,0,0,0,0,79,1,0,0,0,0,81,1,0,0,0,0,83,1,0,0,0,0,85,1,0,0,
0,0,87,1,0,0,0,0,89,1,0,0,0,0,91,1,0,0,0,0,93,1,0,0,0,0,95,1,0,0,
0,0,97,1,0,0,0,0,99,1,0,0,0,0,101,1,0,0,0,0,103,1,0,0,0,0,105,1,
0,0,0,0,107,1,0,0,0,0,109,1,0,0,0,0,111,1,0,0,0,0,113,1,0,0,0,0,
115,1,0,0,0,0,117,1,0,0,0,0,119,1,0,0,0,0,121,1,0,0,0,0,123,1,0,
0,0,0,125,1,0,0,0,1,127,1,0,0,0,3,129,1,0,0,0,5,131,1,0,0,0,7,133,
1,0,0,0,9,137,1,0,0,0,11,141,1,0,0,0,13,145,1,0,0,0,15,150,1,0,0,
0,17,155,1,0,0,0,19,160,1,0,0,0,21,166,1,0,0,0,23,171,1,0,0,0,25,
176,1,0,0,0,27,181,1,0,0,0,29,187,1,0,0,0,31,193,1,0,0,0,33,199,
1,0,0,0,35,203,1,0,0,0,37,208,1,0,0,0,39,211,1,0,0,0,41,215,1,0,
0,0,43,219,1,0,0,0,45,224,1,0,0,0,47,229,1,0,0,0,49,234,1,0,0,0,
51,239,1,0,0,0,53,245,1,0,0,0,55,251,1,0,0,0,57,257,1,0,0,0,59,262,
1,0,0,0,61,268,1,0,0,0,63,274,1,0,0,0,65,278,1,0,0,0,67,283,1,0,
0,0,69,289,1,0,0,0,71,293,1,0,0,0,73,298,1,0,0,0,75,304,1,0,0,0,
77,310,1,0,0,0,79,315,1,0,0,0,81,320,1,0,0,0,83,331,1,0,0,0,85,337,
1,0,0,0,87,342,1,0,0,0,89,351,1,0,0,0,91,356,1,0,0,0,93,361,1,0,
0,0,95,366,1,0,0,0,97,368,1,0,0,0,99,370,1,0,0,0,101,372,1,0,0,0,
103,374,1,0,0,0,105,376,1,0,0,0,107,378,1,0,0,0,109,381,1,0,0,0,
111,383,1,0,0,0,113,386,1,0,0,0,115,388,1,0,0,0,117,391,1,0,0,0,
119,399,1,0,0,0,121,402,1,0,0,0,123,414,1,0,0,0,125,422,1,0,0,0,
127,128,5,40,0,0,128,2,1,0,0,0,129,130,5,41,0,0,130,4,1,0,0,0,131,
132,5,44,0,0,132,6,1,0,0,0,133,134,5,115,0,0,134,135,5,105,0,0,135,
136,5,110,0,0,136,8,1,0,0,0,137,138,5,99,0,0,138,139,5,111,0,0,139,
140,5,115,0,0,140,10,1,0,0,0,141,142,5,116,0,0,142,143,5,97,0,0,
143,144,5,110,0,0,144,12,1,0,0,0,145,146,5,97,0,0,146,147,5,115,
0,0,147,148,5,105,0,0,148,149,5,110,0,0,149,14,1,0,0,0,150,151,5,
97,0,0,151,152,5,99,0,0,152,153,5,111,0,0,153,154,5,115,0,0,154,
16,1,0,0,0,155,156,5,97,0,0,156,157,5,116,0,0,157,158,5,97,0,0,158,
159,5,110,0,0,159,18,1,0,0,0,160,161,5,97,0,0,161,162,5,116,0,0,
162,163,5,97,0,0,163,164,5,110,0,0,164,165,5,50,0,0,165,20,1,0,0,
0,166,167,5,115,0,0,167,168,5,105,0,0,168,169,5,110,0,0,169,170,
5,104,0,0,170,22,1,0,0,0,171,172,5,99,0,0,172,173,5,111,0,0,173,
174,5,115,0,0,174,175,5,104,0,0,175,24,1,0,0,0,176,177,5,116,0,0,
177,178,5,97,0,0,178,179,5,110,0,0,179,180,5,104,0,0,180,26,1,0,
0,0,181,182,5,97,0,0,182,183,5,115,0,0,183,184,5,105,0,0,184,185,
5,110,0,0,185,186,5,104,0,0,186,28,1,0,0,0,187,188,5,97,0,0,188,
189,5,99,0,0,189,190,5,111,0,0,190,191,5,115,0,0,191,192,5,104,0,
0,192,30,1,0,0,0,193,194,5,97,0,0,194,195,5,116,0,0,195,196,5,97,
0,0,196,197,5,110,0,0,197,198,5,104,0,0,198,32,1,0,0,0,199,200,5,
97,0,0,200,201,5,98,0,0,201,202,5,115,0,0,202,34,1,0,0,0,203,204,
5,115,0,0,204,205,5,113,0,0,205,206,5,114,0,0,206,207,5,116,0,0,
207,36,1,0,0,0,208,209,5,108,0,0,209,210,5,110,0,0,210,38,1,0,0,
0,211,212,5,108,0,0,212,213,5,111,0,0,213,214,5,103,0,0,214,40,1,
0,0,0,215,216,5,101,0,0,216,217,5,120,0,0,217,218,5,112,0,0,218,
42,1,0,0,0,219,220,5,115,0,0,220,221,5,109,0,0,221,222,5,105,0,0,
222,223,5,110,0,0,223,44,1,0,0,0,224,225,5,115,0,0,225,226,5,109,
0,0,226,227,5,97,0,0,227,228,5,120,0,0,228,46,1,0,0,0,229,230,5,
116,0,0,230,231,5,109,0,0,231,232,5,105,0,0,232,233,5,110,0,0,233,
48,1,0,0,0,234,235,5,116,0,0,235,236,5,109,0,0,236,237,5,97,0,0,
237,238,5,120,0,0,238,50,1,0,0,0,239,240,5,116,0,0,240,241,5,110,
0,0,241,242,5,111,0,0,242,243,5,114,0,0,243,244,5,109,0,0,244,52,
1,0,0,0,245,246,5,115,0,0,246,247,5,110,0,0,247,248,5,111,0,0,248,
249,5,114,0,0,249,250,5,109,0,0,250,54,1,0,0,0,251,252,5,102,0,0,
252,253,5,108,0,0,253,254,5,111,0,0,254,255,5,111,0,0,255,256,5,
114,0,0,256,56,1,0,0,0,257,258,5,99,0,0,258,259,5,101,0,0,259,260,
5,105,0,0,260,261,5,108,0,0,261,58,1,0,0,0,262,263,5,114,0,0,263,
264,5,111,0,0,264,265,5,117,0,0,265,266,5,110,0,0,266,267,5,100,
0,0,267,60,1,0,0,0,268,269,5,103,0,0,269,270,5,97,0,0,270,271,5,
109,0,0,271,272,5,109,0,0,272,273,5,97,0,0,273,62,1,0,0,0,274,275,
5,112,0,0,275,276,5,111,0,0,276,277,5,119,0,0,277,64,1,0,0,0,278,
279,5,115,0,0,279,280,5,105,0,0,280,281,5,103,0,0,281,282,5,109,
0,0,282,66,1,0,0,0,283,284,5,99,0,0,284,285,5,108,0,0,285,286,5,
97,0,0,286,287,5,109,0,0,287,288,5,112,0,0,288,68,1,0,0,0,289,290,
5,102,0,0,290,291,5,102,0,0,291,292,5,116,0,0,292,70,1,0,0,0,293,
294,5,105,0,0,294,295,5,102,0,0,295,296,5,102,0,0,296,297,5,116,
0,0,297,72,1,0,0,0,298,299,5,97,0,0,299,300,5,110,0,0,300,301,5,
103,0,0,301,302,5,108,0,0,302,303,5,101,0,0,303,74,1,0,0,0,304,305,
5,112,0,0,305,306,5,114,0,0,306,307,5,105,0,0,307,308,5,110,0,0,
308,309,5,116,0,0,309,76,1,0,0,0,310,311,5,108,0,0,311,312,5,101,
0,0,312,313,5,114,0,0,313,314,5,112,0,0,314,78,1,0,0,0,315,316,5,
115,0,0,316,317,5,116,0,0,317,318,5,101,0,0,318,319,5,112,0,0,319,
80,1,0,0,0,320,321,5,115,0,0,321,322,5,109,0,0,322,323,5,111,0,0,
323,324,5,111,0,0,324,325,5,116,0,0,325,326,5,104,0,0,326,327,5,
115,0,0,327,328,5,116,0,0,328,329,5,101,0,0,329,330,5,112,0,0,330,
82,1,0,0,0,331,332,5,102,0,0,332,333,5,114,0,0,333,334,5,97,0,0,
334,335,5,99,0,0,335,336,5,116,0,0,336,84,1,0,0,0,337,338,5,114,
0,0,338,339,5,101,0,0,339,340,5,108,0,0,340,341,5,117,0,0,341,86,
1,0,0,0,342,343,5,115,0,0,343,344,5,111,0,0,344,345,5,102,0,0,345,
346,5,116,0,0,346,347,5,112,0,0,347,348,5,108,0,0,348,349,5,117,
0,0,349,350,5,115,0,0,350,88,1,0,0,0,351,352,5,103,0,0,352,353,5,
101,0,0,353,354,5,108,0,0,354,355,5,117,0,0,355,90,1,0,0,0,356,357,
5,115,0,0,357,358,5,105,0,0,358,359,5,103,0,0,359,360,5,110,0,0,
360,92,1,0,0,0,361,362,5,115,0,0,362,363,5,119,0,0,363,364,5,97,
0,0,364,365,5,112,0,0,365,94,1,0,0,0,366,367,5,43,0,0,367,96,1,0,
0,0,368,369,5,45,0,0,369,98,1,0,0,0,370,371,5,42,0,0,371,100,1,0,
0,0,372,373,5,47,0,0,373,102,1,0,0,0,374,375,5,37,0,0,375,104,1,
0,0,0,376,377,5,94,0,0,377,106,1,0,0,0,378,379,5,62,0,0,379,380,
5,61,0,0,380,108,1,0,0,0,381,382,5,62,0,0,382,110,1,0,0,0,383,384,
5,60,0,0,384,385,5,61,0,0,385,112,1,0,0,0,386,387,5,60,0,0,387,114,
1,0,0,0,388,389,5,61,0,0,389,390,5,61,0,0,390,116,1,0,0,0,391,392,
5,33,0,0,392,393,5,61,0,0,393,118,1,0,0,0,394,395,5,112,0,0,395,
400,5,105,0,0,396,397,5,80,0,0,397,400,5,73,0,0,398,400,7,0,0,0,
399,394,1,0,0,0,399,396,1,0,0,0,399,398,1,0,0,0,400,120,1,0,0,0,
401,403,7,1,0,0,402,401,1,0,0,0,403,404,1,0,0,0,404,402,1,0,0,0,
404,405,1,0,0,0,405,412,1,0,0,0,406,408,5,46,0,0,407,409,7,1,0,0,
408,407,1,0,0,0,409,410,1,0,0,0,410,408,1,0,0,0,410,411,1,0,0,0,
411,413,1,0,0,0,412,406,1,0,0,0,412,413,1,0,0,0,413,122,1,0,0,0,
414,418,7,2,0,0,415,417,7,3,0,0,416,415,1,0,0,0,417,420,1,0,0,0,
418,416,1,0,0,0,418,419,1,0,0,0,419,124,1,0,0,0,420,418,1,0,0,0,
421,423,7,4,0,0,422,421,1,0,0,0,423,424,1,0,0,0,424,422,1,0,0,0,
424,425,1,0,0,0,425,426,1,0,0,0,426,427,6,62,0,0,427,126,1,0,0,0,
7,0,399,404,410,412,418,424,1,6,0,0
]
class MathExprLexer(Lexer):
@@ -180,22 +206,31 @@ class MathExprLexer(Lexer):
SIFFT = 36
ANGL = 37
PRNT = 38
PLUS = 39
MINUS = 40
MULT = 41
DIV = 42
MOD = 43
POW = 44
GE = 45
GT = 46
LE = 47
LT = 48
EQ = 49
NE = 50
CONSTANT = 51
NUMBER = 52
VARIABLE = 53
WS = 54
LERP = 39
STEP = 40
SMOOTHSTEP = 41
FRACT = 42
RELU = 43
SOFTPLUS = 44
GELU = 45
SIGN = 46
SWAP = 47
PLUS = 48
MINUS = 49
MULT = 50
DIV = 51
MOD = 52
POW = 53
GE = 54
GT = 55
LE = 56
LT = 57
EQ = 58
NE = 59
CONSTANT = 60
NUMBER = 61
VARIABLE = 62
WS = 63
channelNames = [ u"DEFAULT_TOKEN_CHANNEL", u"HIDDEN" ]
@@ -207,27 +242,30 @@ class MathExprLexer(Lexer):
"'acosh'", "'atanh'", "'abs'", "'sqrt'", "'ln'", "'log'", "'exp'",
"'smin'", "'smax'", "'tmin'", "'tmax'", "'tnorm'", "'snorm'",
"'floor'", "'ceil'", "'round'", "'gamma'", "'pow'", "'sigm'",
"'clamp'", "'fft'", "'ifft'", "'angle'", "'print'", "'+'", "'-'",
"'*'", "'/'", "'%'", "'^'", "'>='", "'>'", "'<='", "'<'", "'=='",
"'!='" ]
"'clamp'", "'fft'", "'ifft'", "'angle'", "'print'", "'lerp'",
"'step'", "'smoothstep'", "'fract'", "'relu'", "'softplus'",
"'gelu'", "'sign'", "'swap'", "'+'", "'-'", "'*'", "'/'", "'%'",
"'^'", "'>='", "'>'", "'<='", "'<'", "'=='", "'!='" ]
symbolicNames = [ "<INVALID>",
"SIN", "COS", "TAN", "ASIN", "ACOS", "ATAN", "ATAN2", "SINH",
"COSH", "TANH", "ASINH", "ACOSH", "ATANH", "ABS", "SQRT", "LN",
"LOG", "EXP", "SMIN", "SMAX", "TMIN", "TMAX", "TNORM", "SNORM",
"FLOOR", "CEIL", "ROUND", "GAMMA", "POWE", "SIGM", "CLAMP",
"SFFT", "SIFFT", "ANGL", "PRNT", "PLUS", "MINUS", "MULT", "DIV",
"MOD", "POW", "GE", "GT", "LE", "LT", "EQ", "NE", "CONSTANT",
"NUMBER", "VARIABLE", "WS" ]
"SFFT", "SIFFT", "ANGL", "PRNT", "LERP", "STEP", "SMOOTHSTEP",
"FRACT", "RELU", "SOFTPLUS", "GELU", "SIGN", "SWAP", "PLUS",
"MINUS", "MULT", "DIV", "MOD", "POW", "GE", "GT", "LE", "LT",
"EQ", "NE", "CONSTANT", "NUMBER", "VARIABLE", "WS" ]
ruleNames = [ "T__0", "T__1", "T__2", "SIN", "COS", "TAN", "ASIN", "ACOS",
"ATAN", "ATAN2", "SINH", "COSH", "TANH", "ASINH", "ACOSH",
"ATANH", "ABS", "SQRT", "LN", "LOG", "EXP", "SMIN", "SMAX",
"TMIN", "TMAX", "TNORM", "SNORM", "FLOOR", "CEIL", "ROUND",
"GAMMA", "POWE", "SIGM", "CLAMP", "SFFT", "SIFFT", "ANGL",
"PRNT", "PLUS", "MINUS", "MULT", "DIV", "MOD", "POW",
"GE", "GT", "LE", "LT", "EQ", "NE", "CONSTANT", "NUMBER",
"VARIABLE", "WS" ]
"PRNT", "LERP", "STEP", "SMOOTHSTEP", "FRACT", "RELU",
"SOFTPLUS", "GELU", "SIGN", "SWAP", "PLUS", "MINUS", "MULT",
"DIV", "MOD", "POW", "GE", "GT", "LE", "LT", "EQ", "NE",
"CONSTANT", "NUMBER", "VARIABLE", "WS" ]
grammarFileName = "MathExpr.g4"
+46 -28
View File
@@ -36,22 +36,31 @@ SFFT=35
SIFFT=36
ANGL=37
PRNT=38
PLUS=39
MINUS=40
MULT=41
DIV=42
MOD=43
POW=44
GE=45
GT=46
LE=47
LT=48
EQ=49
NE=50
CONSTANT=51
NUMBER=52
VARIABLE=53
WS=54
LERP=39
STEP=40
SMOOTHSTEP=41
FRACT=42
RELU=43
SOFTPLUS=44
GELU=45
SIGN=46
SWAP=47
PLUS=48
MINUS=49
MULT=50
DIV=51
MOD=52
POW=53
GE=54
GT=55
LE=56
LT=57
EQ=58
NE=59
CONSTANT=60
NUMBER=61
VARIABLE=62
WS=63
'('=1
')'=2
','=3
@@ -90,15 +99,24 @@ WS=54
'ifft'=36
'angle'=37
'print'=38
'+'=39
'-'=40
'*'=41
'/'=42
'%'=43
'^'=44
'>='=45
'>'=46
'<='=47
'<'=48
'=='=49
'!='=50
'lerp'=39
'step'=40
'smoothstep'=41
'fract'=42
'relu'=43
'softplus'=44
'gelu'=45
'sign'=46
'swap'=47
'+'=48
'-'=49
'*'=50
'/'=51
'%'=52
'^'=53
'>='=54
'>'=55
'<='=56
'<'=57
'=='=58
'!='=59
File diff suppressed because it is too large Load Diff
+51 -1
View File
@@ -1,4 +1,4 @@
# Generated from MathExpr.g4 by ANTLR 4.13.2
# Generated from src/more_math/Parser/MathExpr.g4 by ANTLR 4.13.2
from antlr4 import *
if "." in __name__:
from .MathExprParser import MathExprParser
@@ -124,6 +124,11 @@ class MathExprVisitor(ParseTreeVisitor):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#Func4Exp.
def visitFunc4Exp(self, ctx:MathExprParser.Func4ExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#FuncNExp.
def visitFuncNExp(self, ctx:MathExprParser.FuncNExpContext):
return self.visitChildren(ctx)
@@ -289,6 +294,31 @@ class MathExprVisitor(ParseTreeVisitor):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#FractFunc.
def visitFractFunc(self, ctx:MathExprParser.FractFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ReluFunc.
def visitReluFunc(self, ctx:MathExprParser.ReluFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#SoftplusFunc.
def visitSoftplusFunc(self, ctx:MathExprParser.SoftplusFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#GeluFunc.
def visitGeluFunc(self, ctx:MathExprParser.GeluFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#SignFunc.
def visitSignFunc(self, ctx:MathExprParser.SignFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#PowFunc.
def visitPowFunc(self, ctx:MathExprParser.PowFuncContext):
return self.visitChildren(ctx)
@@ -309,11 +339,31 @@ class MathExprVisitor(ParseTreeVisitor):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#StepFunc.
def visitStepFunc(self, ctx:MathExprParser.StepFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ClampFunc.
def visitClampFunc(self, ctx:MathExprParser.ClampFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#LerpFunc.
def visitLerpFunc(self, ctx:MathExprParser.LerpFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#SmoothstepFunc.
def visitSmoothstepFunc(self, ctx:MathExprParser.SmoothstepFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#SwapFunc.
def visitSwapFunc(self, ctx:MathExprParser.SwapFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#SMinFunc.
def visitSMinFunc(self, ctx:MathExprParser.SMinFuncContext):
return self.visitChildren(ctx)
+63
View File
@@ -124,6 +124,14 @@ class TensorEvalVisitor(MathExprVisitor):
def visitRoundFunc(self, ctx): return torch.round(self.visit(ctx.expr()))
def visitGammaFunc(self, ctx): return torch.special.gamma(self.visit(ctx.expr())).exp()
def visitSigmoidFunc(self, ctx): return torch.sigmoid(self.visit(ctx.expr()))
def visitReluFunc(self, ctx): return torch.relu(self.visit(ctx.expr()))
def visitSoftplusFunc(self, ctx): return torch.nn.functional.softplus(self.visit(ctx.expr()))
def visitGeluFunc(self, ctx): return torch.nn.functional.gelu(self.visit(ctx.expr()))
def visitSignFunc(self, ctx): return torch.sign(self.visit(ctx.expr()))
def visitFractFunc(self, ctx):
val = self.visit(ctx.expr())
return val - torch.floor(val)
def visitAnglFunc(self, ctx): return torch.angle(self.visit(ctx.expr()))
def visitPrintFunc(self, ctx):
val = self.visit(ctx.expr())
@@ -138,7 +146,37 @@ class TensorEvalVisitor(MathExprVisitor):
return time_to_freq(val)
finally:
self.variables = old_vars
def visitSwapFunc(self, ctx):
tsr = self.visit(ctx.expr(0))
# Evaluate arguments for dim, idx1, idx2. They return full tensors, so we take scalar value.
# We use .data.flatten()[0] to get the scalar safely from any shape
dim_t = self.visit(ctx.expr(1))
idx1_t = self.visit(ctx.expr(2))
idx2_t = self.visit(ctx.expr(3))
dim = int(dim_t.flatten()[0].item())
i = int(idx1_t.flatten()[0].item())
j = int(idx2_t.flatten()[0].item())
# Handle negative dim
if dim < 0: dim += tsr.ndim
# Create permuted index
indices = torch.arange(tsr.shape[dim], device=tsr.device)
# Swap
# Check bounds? Torch index_select will check bounds or crash.
# Support python style negative indexing for indices
if i < 0: i += tsr.shape[dim]
if j < 0: j += tsr.shape[dim]
val_i = indices[i].clone()
indices[i] = indices[j]
indices[j] = val_i
return torch.index_select(tsr, dim, indices)
def visitSifftFunc(self, ctx):
old_vars = self.variables
# Switch to freq variables
@@ -208,6 +246,27 @@ class TensorEvalVisitor(MathExprVisitor):
return torch.atan2(self.visit(ctx.expr(0)), self.visit(ctx.expr(1)))
def visitClampFunc(self, ctx): return torch.clamp(self.visit(ctx.expr(0)), self.visit(ctx.expr(1)), self.visit(ctx.expr(2)))
def visitLerpFunc(self, ctx):
a = self.visit(ctx.expr(0))
b = self.visit(ctx.expr(1))
w = self.visit(ctx.expr(2))
return torch.lerp(a, b, w)
def visitSmoothstepFunc(self, ctx):
edge0 = self.visit(ctx.expr(0))
edge1 = self.visit(ctx.expr(1))
x = self.visit(ctx.expr(2))
# Scale, bias and saturate x to 0..1 range
t = torch.clamp((x - edge0) / (edge1 - edge0), 0.0, 1.0)
# Evaluate polynomial
return t * t * (3.0 - 2.0 * t)
def visitStepFunc(self, ctx):
edge = self.visit(ctx.expr(0))
x = self.visit(ctx.expr(1))
# step(edge, x) = 1 if x >= edge else 0
return torch.where(x >= edge, 1.0, 0.0)
# N-argument functions
def visitSMinFunc(self, ctx):
args = [self.visit(e) for e in ctx.expr()]
@@ -227,6 +286,10 @@ class TensorEvalVisitor(MathExprVisitor):
return self.visitChildren(ctx)
def visitFuncNExp(self, ctx):
return self.visitChildren(ctx)
def visitFunc3Exp(self, ctx):
return self.visitChildren(ctx)
def visitFunc4Exp(self, ctx):
return self.visitChildren(ctx)
def visitAtomExp(self, ctx):
return self.visitChildren(ctx)
+138
View File
@@ -2,6 +2,83 @@
"""Tests for `more_math` package."""
import unittest
import torch
from more_math.ConditioningMathNode import ConditioningMathNode
from more_math.LatentMathNode import LatentMathNode
from more_math.LatentMathNode import LatentMathNode
from more_math.ImageMathNode import ImageMathNode
from more_math.FloatMathNode import FloatMathNode
import tokenize
from io import StringIO
def tokenize_expression(expr):
if not expr.endswith('\n'):
expr = expr + '\n'
f = StringIO(expr)
tokens = tokenize.generate_tokens(f.readline)
filtered_tokens = []
for toktype, tokval, _, _, _ in tokens:
token_name = tokenize.tok_name[toktype]
# Opravíme ERRORTOKEN na OP pro !, &, |, ^
if token_name == 'ERRORTOKEN' and tokval in {'!', '&', '|', '^'}:
token_name = 'OP'
if token_name in {'COMMENT', 'NL', 'NEWLINE', 'INDENT', 'DEDENT'}:
continue
filtered_tokens.append((token_name, tokval.strip()))
return filtered_tokens
class TestMoreMath(unittest.TestCase):
def test_conditioning_math_node_initialization(self):
node = ConditioningMathNode()
self.assertIsInstance(node, ConditioningMathNode)
def test_conditioning_math_node_metadata(self):
self.assertEqual(ConditioningMathNode.RETURN_TYPES, ["CONDITIONING"])
self.assertEqual(ConditioningMathNode.FUNCTION, "EXECUTE_NORMALIZED")
self.assertEqual(ConditioningMathNode.CATEGORY, "More math")
def test_latent_math_node_initialization(self):
node = LatentMathNode()
self.assertIsInstance(node, LatentMathNode)
def test_latent_math_node_metadata(self):
self.assertEqual(LatentMathNode.RETURN_TYPES, ["LATENT"])
self.assertEqual(LatentMathNode.FUNCTION, "EXECUTE_NORMALIZED")
self.assertEqual(LatentMathNode.CATEGORY, "More math")
def test_image_math_node_initialization(self):
node = ImageMathNode()
self.assertIsInstance(node, ImageMathNode)
def test_image_math_node_metadata(self):
self.assertEqual(ImageMathNode.RETURN_TYPES, ["IMAGE"])
self.assertEqual(ImageMathNode.FUNCTION, "EXECUTE_NORMALIZED")
self.assertEqual(ImageMathNode.CATEGORY, "More math")
def test_fft_invertibility(self):
# 1. Create random input latent (Batch, Channel, Height, Width)
input_tensor = torch.randn(1, 4, 32, 32, dtype=torch.float32)
input_dict = {"samples": input_tensor}
# 2. Execute ifft(fft(a))
# Note: execute is a classmethod
result = LatentMathNode.execute(
Latent="ifft(fft(a))",
a=input_dict
)
# 3. Get output tensor
output_tensor = result[0]["samples"]
# 4. Check correctness (approximate equality)
self.assertTrue(torch.allclose(input_tensor, output_tensor, atol=1e-5), \
f"Max difference: {(input_tensor - output_tensor).abs().max()}")
#!/usr/bin/env python
"""Tests for `more_math` package."""
import unittest
import torch
from more_math.ConditioningMathNode import ConditioningMathNode
@@ -89,4 +166,65 @@ class TestMoreMath(unittest.TestCase):
self.assertTrue(torch.allclose(input_tensor, output_tensor, atol=1e-5), \
f"Image FFT round trip failed. Max diff: {(input_tensor - output_tensor).abs().max()}")
def test_new_math_functions(self):
node = LatentMathNode()
# Test Lerp
# lerp(a, b, 0.5) where a=0, b=10 -> 5
l_a = {"samples": torch.zeros(1, 4, 32, 32)}
l_b = {"samples": torch.full((1, 4, 32, 32), 10.0)}
res_lerp = node.execute("lerp(a, b, 0.5)", a=l_a, b=l_b)[0]["samples"]
self.assertTrue(torch.allclose(res_lerp, torch.full_like(res_lerp, 5.0)))
# Test Step
# step(0.5, a) where a=0.8 -> 1
res_step = node.execute("step(0.5, a)", a={"samples": torch.full((1,1,1,1), 0.8)})[0]["samples"]
self.assertTrue(torch.allclose(res_step, torch.ones_like(res_step)))
# step(0.5, a) where a=0.2 -> 0
res_step2 = node.execute("step(0.5, a)", a={"samples": torch.full((1,1,1,1), 0.2)})[0]["samples"]
self.assertTrue(torch.allclose(res_step2, torch.zeros_like(res_step2)))
# Test Swap
# tensor 1x3: [0, 1, 2]. Swap(dim=1, 0, 2) -> [2, 1, 0]
t = torch.tensor([[[0.0, 1.0, 2.0]]]) # 1x1x3
l_t = {"samples": t}
# dim 2 (channel dim in 1x1x3? No, dims are B,C,H,W usually but here shape is 1x1x3)
# LatentMathNode uses "samples" directly.
# eval_single_tensor exposes 'a'.
# Let's use simple Latent 1x4x1x1 (B,C,H,W)
t_lat = torch.tensor([0.0, 10.0, 20.0, 30.0]).view(1,4,1,1)
# Swap channels 0 and 3 -> 30, 10, 20, 0
l_swap = {"samples": t_lat}
# swap(a, 1, 0, 3) . Dim 1 is channel (B=0, C=1)
res_swap = node.execute("swap(a, 1, 0, 3)", a=l_swap)[0]["samples"]
expected = torch.tensor([30.0, 10.0, 20.0, 0.0]).view(1,4,1,1)
self.assertTrue(torch.allclose(res_swap, expected))
# Test Relu, Sign, Fract
res_relu = node.execute("relu(-5.0)", a=l_a)[0]["samples"]
self.assertTrue(torch.allclose(res_relu, torch.zeros_like(res_relu)))
res_sign = node.execute("sign(-5.0)", a=l_a)[0]["samples"]
self.assertTrue(torch.allclose(res_sign, torch.full_like(res_sign, -1.0)))
res_fract = node.execute("fract(1.5)", a=l_a)[0]["samples"]
self.assertTrue(torch.allclose(res_fract, torch.full_like(res_fract, 0.5)))
def test_float_math_functions(self):
node = FloatMathNode()
# Test Lerp: lerp(0, 10, 0.5) -> 5.0
res = node.execute("lerp(a, b, 0.5)", a=0.0, b=10.0)[0]
self.assertAlmostEqual(res, 5.0)
# Test Step: step(0.5, 0.8) -> 1.0
res = node.execute("step(0.5, a)", a=0.8)[0]
self.assertAlmostEqual(res, 1.0)
# Test Relu: relu(-5) -> 0.0
res = node.execute("relu(a)", a=-5.0)[0]
self.assertAlmostEqual(res, 0.0)
# Test Smoothstep: smoothstep(0, 1, 0.5) -> 0.5
# 0.5*0.5*(3 - 2*0.5) = 0.25 * 2 = 0.5
res = node.execute("smoothstep(0, 1, a)", a=0.5)[0]
self.assertAlmostEqual(res, 0.5)