add topk, botk and convolution inverse function

This commit is contained in:
mcDandy
2026-01-09 13:28:46 +01:00
parent 1935b80212
commit f4edef96cc
14 changed files with 2657 additions and 10241 deletions
+5 -1
View File
@@ -50,7 +50,6 @@ You can also get the node from comfy manager under the name of More math.
- `fract(x)`: Returns the fractional part of x (x - floor(x)).
- `sign(x)`: Returns -1 for negative, 1 for positive, 0 for zero.
- `gamma(x)`: Gamma function.
- `range(start, end, step)`: Generates a list of values from start (inclusive) to end (exclusive) with given step.
### Trigonometric
@@ -83,6 +82,8 @@ You can also get the node from comfy manager under the name of More math.
- `tmax(x, y)`: Element-wise maximum of x and y.
- `smin(x, ...)`: **Scalar** minimum. Returns the single smallest value across all input tensors/values.
- `smax(x, ...)`: **Scalar** maximum. Returns the single largest value across all input tensors/values.
- `topk(x, k)`: Returns a **masked tensor** with only the **top K largest** values preserved at their original positions (others zeroed). For lists, returns the top K largest items sorted descending. Supports complex tensors (uses magnitude for selection).
- `botk(x, k)`: Returns a **masked tensor** with only the **bottom K smallest** values preserved at their original positions (others zeroed). For lists, returns the bottom K smallest items sorted ascending.
- `tnorm(x)`: **Tensor** Normalizes x (L2 norm along last dimension).
- `snorm(x)`: **Scalar** L2 norm of the entire tensor.
- `swap(tensor, dim, index1, index2)`: Swaps two slices of a tensor along a specified dimension. (Tensor only)
@@ -107,6 +108,9 @@ You can also get the node from comfy manager under the name of More math.
- `print(x)`: Prints the value of x to the console and returns x.
- `print_shape(x)` or `pshp`: Prints the shape of x to the console and returns x.
- `pinv(x)`: Computes the permutation inverse of list or tensor (>1D tensor can have interesting results). If `permute(i,x) = j`, then `permute(j,pinv(x)) = i`.
- `range(start, end, step)`: Generates a list of values from start (inclusive) to end (exclusive) with given step.
## Variables
+6
View File
@@ -91,6 +91,7 @@ func1
| SIGN '(' expr ')' # SignFunc
| PRINT_SHAPE_L '(' expr ')' # PrintShapeFunc
| PRINT_SHAPE '(' expr ')' # PrintShapeFunc
| PINV '(' expr ')' # PinvFunc
;
// Two-argument functions
@@ -100,6 +101,8 @@ func2
| TMIN '(' expr ',' expr ')' # TMinFunc
| TMAX '(' expr ',' expr ')' # TMaxFunc
| STEP '(' expr ',' expr ')' # StepFunc
| TOPK '(' expr ',' expr ')' # TopkFunc
| BOTK '(' expr ',' expr ')' # BotkFunc
;
func3
@@ -176,6 +179,9 @@ SWAP : 'swap';
PERM : 'permute';
RESHAPE : 'reshape';
RANGE : 'range';
TOPK : 'topk';
BOTK : 'botk';
PINV : 'pinv';
PLUS : '+';
MINUS : '-';
File diff suppressed because one or more lines are too long
+36 -30
View File
@@ -54,23 +54,26 @@ SWAP=53
PERM=54
RESHAPE=55
RANGE=56
PLUS=57
MINUS=58
MULT=59
DIV=60
MOD=61
POW=62
GE=63
GT=64
LE=65
LT=66
EQ=67
NE=68
PIPE=69
CONSTANT=70
NUMBER=71
VARIABLE=72
WS=73
TOPK=57
BOTK=58
PINV=59
PLUS=60
MINUS=61
MULT=62
DIV=63
MOD=64
POW=65
GE=66
GT=67
LE=68
LT=69
EQ=70
NE=71
PIPE=72
CONSTANT=73
NUMBER=74
VARIABLE=75
WS=76
'('=1
')'=2
'['=3
@@ -127,16 +130,19 @@ WS=73
'permute'=54
'reshape'=55
'range'=56
'+'=57
'-'=58
'*'=59
'/'=60
'%'=61
'^'=62
'>='=63
'>'=64
'<='=65
'<'=66
'=='=67
'!='=68
'|'=69
'topk'=57
'botk'=58
'pinv'=59
'+'=60
'-'=61
'*'=62
'/'=63
'%'=64
'^'=65
'>='=66
'>'=67
'<='=68
'<'=69
'=='=70
'!='=71
'|'=72
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
+36 -30
View File
@@ -54,23 +54,26 @@ SWAP=53
PERM=54
RESHAPE=55
RANGE=56
PLUS=57
MINUS=58
MULT=59
DIV=60
MOD=61
POW=62
GE=63
GT=64
LE=65
LT=66
EQ=67
NE=68
PIPE=69
CONSTANT=70
NUMBER=71
VARIABLE=72
WS=73
TOPK=57
BOTK=58
PINV=59
PLUS=60
MINUS=61
MULT=62
DIV=63
MOD=64
POW=65
GE=66
GT=67
LE=68
LT=69
EQ=70
NE=71
PIPE=72
CONSTANT=73
NUMBER=74
VARIABLE=75
WS=76
'('=1
')'=2
'['=3
@@ -127,16 +130,19 @@ WS=73
'permute'=54
'reshape'=55
'range'=56
'+'=57
'-'=58
'*'=59
'/'=60
'%'=61
'^'=62
'>='=63
'>'=64
'<='=65
'<'=66
'=='=67
'!='=68
'|'=69
'topk'=57
'botk'=58
'pinv'=59
'+'=60
'-'=61
'*'=62
'/'=63
'%'=64
'^'=65
'>='=66
'>'=67
'<='=68
'<'=69
'=='=70
'!='=71
'|'=72
+272 -165
View File
@@ -1,661 +1,768 @@
# Generated from ./MathExpr.g4 by ANTLR 4.13.2
from antlr4 import *
if "." in __name__:
from .MathExprParser import MathExprParser
else:
from MathExprParser import MathExprParser
# This class defines a complete listener for a parse tree produced by MathExprParser.
class MathExprListener(ParseTreeListener):
# Enter a parse tree produced by MathExprParser#expr.
def enterExpr(self, ctx: MathExprParser.ExprContext):
def enterExpr(self, ctx:MathExprParser.ExprContext):
pass
# Exit a parse tree produced by MathExprParser#expr.
def exitExpr(self, ctx: MathExprParser.ExprContext):
def exitExpr(self, ctx:MathExprParser.ExprContext):
pass
# Enter a parse tree produced by MathExprParser#LtExp.
def enterLtExp(self, ctx: MathExprParser.LtExpContext):
def enterLtExp(self, ctx:MathExprParser.LtExpContext):
pass
# Exit a parse tree produced by MathExprParser#LtExp.
def exitLtExp(self, ctx: MathExprParser.LtExpContext):
def exitLtExp(self, ctx:MathExprParser.LtExpContext):
pass
# Enter a parse tree produced by MathExprParser#EqExp.
def enterEqExp(self, ctx: MathExprParser.EqExpContext):
def enterEqExp(self, ctx:MathExprParser.EqExpContext):
pass
# Exit a parse tree produced by MathExprParser#EqExp.
def exitEqExp(self, ctx: MathExprParser.EqExpContext):
def exitEqExp(self, ctx:MathExprParser.EqExpContext):
pass
# Enter a parse tree produced by MathExprParser#ToAdd.
def enterToAdd(self, ctx: MathExprParser.ToAddContext):
def enterToAdd(self, ctx:MathExprParser.ToAddContext):
pass
# Exit a parse tree produced by MathExprParser#ToAdd.
def exitToAdd(self, ctx: MathExprParser.ToAddContext):
def exitToAdd(self, ctx:MathExprParser.ToAddContext):
pass
# Enter a parse tree produced by MathExprParser#GeExp.
def enterGeExp(self, ctx: MathExprParser.GeExpContext):
def enterGeExp(self, ctx:MathExprParser.GeExpContext):
pass
# Exit a parse tree produced by MathExprParser#GeExp.
def exitGeExp(self, ctx: MathExprParser.GeExpContext):
def exitGeExp(self, ctx:MathExprParser.GeExpContext):
pass
# Enter a parse tree produced by MathExprParser#LeExp.
def enterLeExp(self, ctx: MathExprParser.LeExpContext):
def enterLeExp(self, ctx:MathExprParser.LeExpContext):
pass
# Exit a parse tree produced by MathExprParser#LeExp.
def exitLeExp(self, ctx: MathExprParser.LeExpContext):
def exitLeExp(self, ctx:MathExprParser.LeExpContext):
pass
# Enter a parse tree produced by MathExprParser#NeExp.
def enterNeExp(self, ctx: MathExprParser.NeExpContext):
def enterNeExp(self, ctx:MathExprParser.NeExpContext):
pass
# Exit a parse tree produced by MathExprParser#NeExp.
def exitNeExp(self, ctx: MathExprParser.NeExpContext):
def exitNeExp(self, ctx:MathExprParser.NeExpContext):
pass
# Enter a parse tree produced by MathExprParser#GtExp.
def enterGtExp(self, ctx: MathExprParser.GtExpContext):
def enterGtExp(self, ctx:MathExprParser.GtExpContext):
pass
# Exit a parse tree produced by MathExprParser#GtExp.
def exitGtExp(self, ctx: MathExprParser.GtExpContext):
def exitGtExp(self, ctx:MathExprParser.GtExpContext):
pass
# Enter a parse tree produced by MathExprParser#AddExp.
def enterAddExp(self, ctx: MathExprParser.AddExpContext):
def enterAddExp(self, ctx:MathExprParser.AddExpContext):
pass
# Exit a parse tree produced by MathExprParser#AddExp.
def exitAddExp(self, ctx: MathExprParser.AddExpContext):
def exitAddExp(self, ctx:MathExprParser.AddExpContext):
pass
# Enter a parse tree produced by MathExprParser#ToMul.
def enterToMul(self, ctx: MathExprParser.ToMulContext):
def enterToMul(self, ctx:MathExprParser.ToMulContext):
pass
# Exit a parse tree produced by MathExprParser#ToMul.
def exitToMul(self, ctx: MathExprParser.ToMulContext):
def exitToMul(self, ctx:MathExprParser.ToMulContext):
pass
# Enter a parse tree produced by MathExprParser#SubExp.
def enterSubExp(self, ctx: MathExprParser.SubExpContext):
def enterSubExp(self, ctx:MathExprParser.SubExpContext):
pass
# Exit a parse tree produced by MathExprParser#SubExp.
def exitSubExp(self, ctx: MathExprParser.SubExpContext):
def exitSubExp(self, ctx:MathExprParser.SubExpContext):
pass
# Enter a parse tree produced by MathExprParser#MulExp.
def enterMulExp(self, ctx: MathExprParser.MulExpContext):
def enterMulExp(self, ctx:MathExprParser.MulExpContext):
pass
# Exit a parse tree produced by MathExprParser#MulExp.
def exitMulExp(self, ctx: MathExprParser.MulExpContext):
def exitMulExp(self, ctx:MathExprParser.MulExpContext):
pass
# Enter a parse tree produced by MathExprParser#ModExp.
def enterModExp(self, ctx: MathExprParser.ModExpContext):
def enterModExp(self, ctx:MathExprParser.ModExpContext):
pass
# Exit a parse tree produced by MathExprParser#ModExp.
def exitModExp(self, ctx: MathExprParser.ModExpContext):
def exitModExp(self, ctx:MathExprParser.ModExpContext):
pass
# Enter a parse tree produced by MathExprParser#DivExp.
def enterDivExp(self, ctx: MathExprParser.DivExpContext):
def enterDivExp(self, ctx:MathExprParser.DivExpContext):
pass
# Exit a parse tree produced by MathExprParser#DivExp.
def exitDivExp(self, ctx: MathExprParser.DivExpContext):
def exitDivExp(self, ctx:MathExprParser.DivExpContext):
pass
# Enter a parse tree produced by MathExprParser#ToPow.
def enterToPow(self, ctx: MathExprParser.ToPowContext):
def enterToPow(self, ctx:MathExprParser.ToPowContext):
pass
# Exit a parse tree produced by MathExprParser#ToPow.
def exitToPow(self, ctx: MathExprParser.ToPowContext):
def exitToPow(self, ctx:MathExprParser.ToPowContext):
pass
# Enter a parse tree produced by MathExprParser#PowExp.
def enterPowExp(self, ctx: MathExprParser.PowExpContext):
def enterPowExp(self, ctx:MathExprParser.PowExpContext):
pass
# Exit a parse tree produced by MathExprParser#PowExp.
def exitPowExp(self, ctx: MathExprParser.PowExpContext):
def exitPowExp(self, ctx:MathExprParser.PowExpContext):
pass
# Enter a parse tree produced by MathExprParser#ToUnary.
def enterToUnary(self, ctx: MathExprParser.ToUnaryContext):
def enterToUnary(self, ctx:MathExprParser.ToUnaryContext):
pass
# Exit a parse tree produced by MathExprParser#ToUnary.
def exitToUnary(self, ctx: MathExprParser.ToUnaryContext):
def exitToUnary(self, ctx:MathExprParser.ToUnaryContext):
pass
# Enter a parse tree produced by MathExprParser#UnaryPlus.
def enterUnaryPlus(self, ctx: MathExprParser.UnaryPlusContext):
def enterUnaryPlus(self, ctx:MathExprParser.UnaryPlusContext):
pass
# Exit a parse tree produced by MathExprParser#UnaryPlus.
def exitUnaryPlus(self, ctx: MathExprParser.UnaryPlusContext):
def exitUnaryPlus(self, ctx:MathExprParser.UnaryPlusContext):
pass
# Enter a parse tree produced by MathExprParser#UnaryMinus.
def enterUnaryMinus(self, ctx: MathExprParser.UnaryMinusContext):
def enterUnaryMinus(self, ctx:MathExprParser.UnaryMinusContext):
pass
# Exit a parse tree produced by MathExprParser#UnaryMinus.
def exitUnaryMinus(self, ctx: MathExprParser.UnaryMinusContext):
def exitUnaryMinus(self, ctx:MathExprParser.UnaryMinusContext):
pass
# Enter a parse tree produced by MathExprParser#ToAtom.
def enterToAtom(self, ctx: MathExprParser.ToAtomContext):
def enterToAtom(self, ctx:MathExprParser.ToAtomContext):
pass
# Exit a parse tree produced by MathExprParser#ToAtom.
def exitToAtom(self, ctx: MathExprParser.ToAtomContext):
def exitToAtom(self, ctx:MathExprParser.ToAtomContext):
pass
# Enter a parse tree produced by MathExprParser#Func1Exp.
def enterFunc1Exp(self, ctx: MathExprParser.Func1ExpContext):
def enterFunc1Exp(self, ctx:MathExprParser.Func1ExpContext):
pass
# Exit a parse tree produced by MathExprParser#Func1Exp.
def exitFunc1Exp(self, ctx: MathExprParser.Func1ExpContext):
def exitFunc1Exp(self, ctx:MathExprParser.Func1ExpContext):
pass
# Enter a parse tree produced by MathExprParser#Func2Exp.
def enterFunc2Exp(self, ctx: MathExprParser.Func2ExpContext):
def enterFunc2Exp(self, ctx:MathExprParser.Func2ExpContext):
pass
# Exit a parse tree produced by MathExprParser#Func2Exp.
def exitFunc2Exp(self, ctx: MathExprParser.Func2ExpContext):
def exitFunc2Exp(self, ctx:MathExprParser.Func2ExpContext):
pass
# Enter a parse tree produced by MathExprParser#Func3Exp.
def enterFunc3Exp(self, ctx: MathExprParser.Func3ExpContext):
def enterFunc3Exp(self, ctx:MathExprParser.Func3ExpContext):
pass
# Exit a parse tree produced by MathExprParser#Func3Exp.
def exitFunc3Exp(self, ctx: MathExprParser.Func3ExpContext):
def exitFunc3Exp(self, ctx:MathExprParser.Func3ExpContext):
pass
# Enter a parse tree produced by MathExprParser#Func4Exp.
def enterFunc4Exp(self, ctx: MathExprParser.Func4ExpContext):
def enterFunc4Exp(self, ctx:MathExprParser.Func4ExpContext):
pass
# Exit a parse tree produced by MathExprParser#Func4Exp.
def exitFunc4Exp(self, ctx: MathExprParser.Func4ExpContext):
def exitFunc4Exp(self, ctx:MathExprParser.Func4ExpContext):
pass
# Enter a parse tree produced by MathExprParser#FuncNExp.
def enterFuncNExp(self, ctx: MathExprParser.FuncNExpContext):
def enterFuncNExp(self, ctx:MathExprParser.FuncNExpContext):
pass
# Exit a parse tree produced by MathExprParser#FuncNExp.
def exitFuncNExp(self, ctx: MathExprParser.FuncNExpContext):
def exitFuncNExp(self, ctx:MathExprParser.FuncNExpContext):
pass
# Enter a parse tree produced by MathExprParser#VariableExp.
def enterVariableExp(self, ctx: MathExprParser.VariableExpContext):
def enterVariableExp(self, ctx:MathExprParser.VariableExpContext):
pass
# Exit a parse tree produced by MathExprParser#VariableExp.
def exitVariableExp(self, ctx: MathExprParser.VariableExpContext):
def exitVariableExp(self, ctx:MathExprParser.VariableExpContext):
pass
# Enter a parse tree produced by MathExprParser#NumberExp.
def enterNumberExp(self, ctx: MathExprParser.NumberExpContext):
def enterNumberExp(self, ctx:MathExprParser.NumberExpContext):
pass
# Exit a parse tree produced by MathExprParser#NumberExp.
def exitNumberExp(self, ctx: MathExprParser.NumberExpContext):
def exitNumberExp(self, ctx:MathExprParser.NumberExpContext):
pass
# Enter a parse tree produced by MathExprParser#ConstantExp.
def enterConstantExp(self, ctx: MathExprParser.ConstantExpContext):
def enterConstantExp(self, ctx:MathExprParser.ConstantExpContext):
pass
# Exit a parse tree produced by MathExprParser#ConstantExp.
def exitConstantExp(self, ctx: MathExprParser.ConstantExpContext):
def exitConstantExp(self, ctx:MathExprParser.ConstantExpContext):
pass
# Enter a parse tree produced by MathExprParser#ParenExp.
def enterParenExp(self, ctx: MathExprParser.ParenExpContext):
def enterParenExp(self, ctx:MathExprParser.ParenExpContext):
pass
# Exit a parse tree produced by MathExprParser#ParenExp.
def exitParenExp(self, ctx: MathExprParser.ParenExpContext):
def exitParenExp(self, ctx:MathExprParser.ParenExpContext):
pass
# Enter a parse tree produced by MathExprParser#AbsExp.
def enterAbsExp(self, ctx: MathExprParser.AbsExpContext):
def enterAbsExp(self, ctx:MathExprParser.AbsExpContext):
pass
# Exit a parse tree produced by MathExprParser#AbsExp.
def exitAbsExp(self, ctx: MathExprParser.AbsExpContext):
def exitAbsExp(self, ctx:MathExprParser.AbsExpContext):
pass
# Enter a parse tree produced by MathExprParser#ListExp.
def enterListExp(self, ctx: MathExprParser.ListExpContext):
def enterListExp(self, ctx:MathExprParser.ListExpContext):
pass
# Exit a parse tree produced by MathExprParser#ListExp.
def exitListExp(self, ctx: MathExprParser.ListExpContext):
def exitListExp(self, ctx:MathExprParser.ListExpContext):
pass
# Enter a parse tree produced by MathExprParser#SinFunc.
def enterSinFunc(self, ctx: MathExprParser.SinFuncContext):
def enterSinFunc(self, ctx:MathExprParser.SinFuncContext):
pass
# Exit a parse tree produced by MathExprParser#SinFunc.
def exitSinFunc(self, ctx: MathExprParser.SinFuncContext):
def exitSinFunc(self, ctx:MathExprParser.SinFuncContext):
pass
# Enter a parse tree produced by MathExprParser#CosFunc.
def enterCosFunc(self, ctx: MathExprParser.CosFuncContext):
def enterCosFunc(self, ctx:MathExprParser.CosFuncContext):
pass
# Exit a parse tree produced by MathExprParser#CosFunc.
def exitCosFunc(self, ctx: MathExprParser.CosFuncContext):
def exitCosFunc(self, ctx:MathExprParser.CosFuncContext):
pass
# Enter a parse tree produced by MathExprParser#TanFunc.
def enterTanFunc(self, ctx: MathExprParser.TanFuncContext):
def enterTanFunc(self, ctx:MathExprParser.TanFuncContext):
pass
# Exit a parse tree produced by MathExprParser#TanFunc.
def exitTanFunc(self, ctx: MathExprParser.TanFuncContext):
def exitTanFunc(self, ctx:MathExprParser.TanFuncContext):
pass
# Enter a parse tree produced by MathExprParser#AsinFunc.
def enterAsinFunc(self, ctx: MathExprParser.AsinFuncContext):
def enterAsinFunc(self, ctx:MathExprParser.AsinFuncContext):
pass
# Exit a parse tree produced by MathExprParser#AsinFunc.
def exitAsinFunc(self, ctx: MathExprParser.AsinFuncContext):
def exitAsinFunc(self, ctx:MathExprParser.AsinFuncContext):
pass
# Enter a parse tree produced by MathExprParser#AcosFunc.
def enterAcosFunc(self, ctx: MathExprParser.AcosFuncContext):
def enterAcosFunc(self, ctx:MathExprParser.AcosFuncContext):
pass
# Exit a parse tree produced by MathExprParser#AcosFunc.
def exitAcosFunc(self, ctx: MathExprParser.AcosFuncContext):
def exitAcosFunc(self, ctx:MathExprParser.AcosFuncContext):
pass
# Enter a parse tree produced by MathExprParser#AtanFunc.
def enterAtanFunc(self, ctx: MathExprParser.AtanFuncContext):
def enterAtanFunc(self, ctx:MathExprParser.AtanFuncContext):
pass
# Exit a parse tree produced by MathExprParser#AtanFunc.
def exitAtanFunc(self, ctx: MathExprParser.AtanFuncContext):
def exitAtanFunc(self, ctx:MathExprParser.AtanFuncContext):
pass
# Enter a parse tree produced by MathExprParser#SinhFunc.
def enterSinhFunc(self, ctx: MathExprParser.SinhFuncContext):
def enterSinhFunc(self, ctx:MathExprParser.SinhFuncContext):
pass
# Exit a parse tree produced by MathExprParser#SinhFunc.
def exitSinhFunc(self, ctx: MathExprParser.SinhFuncContext):
def exitSinhFunc(self, ctx:MathExprParser.SinhFuncContext):
pass
# Enter a parse tree produced by MathExprParser#CoshFunc.
def enterCoshFunc(self, ctx: MathExprParser.CoshFuncContext):
def enterCoshFunc(self, ctx:MathExprParser.CoshFuncContext):
pass
# Exit a parse tree produced by MathExprParser#CoshFunc.
def exitCoshFunc(self, ctx: MathExprParser.CoshFuncContext):
def exitCoshFunc(self, ctx:MathExprParser.CoshFuncContext):
pass
# Enter a parse tree produced by MathExprParser#TanhFunc.
def enterTanhFunc(self, ctx: MathExprParser.TanhFuncContext):
def enterTanhFunc(self, ctx:MathExprParser.TanhFuncContext):
pass
# Exit a parse tree produced by MathExprParser#TanhFunc.
def exitTanhFunc(self, ctx: MathExprParser.TanhFuncContext):
def exitTanhFunc(self, ctx:MathExprParser.TanhFuncContext):
pass
# Enter a parse tree produced by MathExprParser#AsinhFunc.
def enterAsinhFunc(self, ctx: MathExprParser.AsinhFuncContext):
def enterAsinhFunc(self, ctx:MathExprParser.AsinhFuncContext):
pass
# Exit a parse tree produced by MathExprParser#AsinhFunc.
def exitAsinhFunc(self, ctx: MathExprParser.AsinhFuncContext):
def exitAsinhFunc(self, ctx:MathExprParser.AsinhFuncContext):
pass
# Enter a parse tree produced by MathExprParser#AcoshFunc.
def enterAcoshFunc(self, ctx: MathExprParser.AcoshFuncContext):
def enterAcoshFunc(self, ctx:MathExprParser.AcoshFuncContext):
pass
# Exit a parse tree produced by MathExprParser#AcoshFunc.
def exitAcoshFunc(self, ctx: MathExprParser.AcoshFuncContext):
def exitAcoshFunc(self, ctx:MathExprParser.AcoshFuncContext):
pass
# Enter a parse tree produced by MathExprParser#AtanhFunc.
def enterAtanhFunc(self, ctx: MathExprParser.AtanhFuncContext):
def enterAtanhFunc(self, ctx:MathExprParser.AtanhFuncContext):
pass
# Exit a parse tree produced by MathExprParser#AtanhFunc.
def exitAtanhFunc(self, ctx: MathExprParser.AtanhFuncContext):
def exitAtanhFunc(self, ctx:MathExprParser.AtanhFuncContext):
pass
# Enter a parse tree produced by MathExprParser#AbsFunc.
def enterAbsFunc(self, ctx: MathExprParser.AbsFuncContext):
def enterAbsFunc(self, ctx:MathExprParser.AbsFuncContext):
pass
# Exit a parse tree produced by MathExprParser#AbsFunc.
def exitAbsFunc(self, ctx: MathExprParser.AbsFuncContext):
def exitAbsFunc(self, ctx:MathExprParser.AbsFuncContext):
pass
# Enter a parse tree produced by MathExprParser#SqrtFunc.
def enterSqrtFunc(self, ctx: MathExprParser.SqrtFuncContext):
def enterSqrtFunc(self, ctx:MathExprParser.SqrtFuncContext):
pass
# Exit a parse tree produced by MathExprParser#SqrtFunc.
def exitSqrtFunc(self, ctx: MathExprParser.SqrtFuncContext):
def exitSqrtFunc(self, ctx:MathExprParser.SqrtFuncContext):
pass
# Enter a parse tree produced by MathExprParser#LnFunc.
def enterLnFunc(self, ctx: MathExprParser.LnFuncContext):
def enterLnFunc(self, ctx:MathExprParser.LnFuncContext):
pass
# Exit a parse tree produced by MathExprParser#LnFunc.
def exitLnFunc(self, ctx: MathExprParser.LnFuncContext):
def exitLnFunc(self, ctx:MathExprParser.LnFuncContext):
pass
# Enter a parse tree produced by MathExprParser#LogFunc.
def enterLogFunc(self, ctx: MathExprParser.LogFuncContext):
def enterLogFunc(self, ctx:MathExprParser.LogFuncContext):
pass
# Exit a parse tree produced by MathExprParser#LogFunc.
def exitLogFunc(self, ctx: MathExprParser.LogFuncContext):
def exitLogFunc(self, ctx:MathExprParser.LogFuncContext):
pass
# Enter a parse tree produced by MathExprParser#ExpFunc.
def enterExpFunc(self, ctx: MathExprParser.ExpFuncContext):
def enterExpFunc(self, ctx:MathExprParser.ExpFuncContext):
pass
# Exit a parse tree produced by MathExprParser#ExpFunc.
def exitExpFunc(self, ctx: MathExprParser.ExpFuncContext):
def exitExpFunc(self, ctx:MathExprParser.ExpFuncContext):
pass
# Enter a parse tree produced by MathExprParser#TNormFunc.
def enterTNormFunc(self, ctx: MathExprParser.TNormFuncContext):
def enterTNormFunc(self, ctx:MathExprParser.TNormFuncContext):
pass
# Exit a parse tree produced by MathExprParser#TNormFunc.
def exitTNormFunc(self, ctx: MathExprParser.TNormFuncContext):
def exitTNormFunc(self, ctx:MathExprParser.TNormFuncContext):
pass
# Enter a parse tree produced by MathExprParser#SNormFunc.
def enterSNormFunc(self, ctx: MathExprParser.SNormFuncContext):
def enterSNormFunc(self, ctx:MathExprParser.SNormFuncContext):
pass
# Exit a parse tree produced by MathExprParser#SNormFunc.
def exitSNormFunc(self, ctx: MathExprParser.SNormFuncContext):
def exitSNormFunc(self, ctx:MathExprParser.SNormFuncContext):
pass
# Enter a parse tree produced by MathExprParser#FloorFunc.
def enterFloorFunc(self, ctx: MathExprParser.FloorFuncContext):
def enterFloorFunc(self, ctx:MathExprParser.FloorFuncContext):
pass
# Exit a parse tree produced by MathExprParser#FloorFunc.
def exitFloorFunc(self, ctx: MathExprParser.FloorFuncContext):
def exitFloorFunc(self, ctx:MathExprParser.FloorFuncContext):
pass
# Enter a parse tree produced by MathExprParser#CeilFunc.
def enterCeilFunc(self, ctx: MathExprParser.CeilFuncContext):
def enterCeilFunc(self, ctx:MathExprParser.CeilFuncContext):
pass
# Exit a parse tree produced by MathExprParser#CeilFunc.
def exitCeilFunc(self, ctx: MathExprParser.CeilFuncContext):
def exitCeilFunc(self, ctx:MathExprParser.CeilFuncContext):
pass
# Enter a parse tree produced by MathExprParser#RoundFunc.
def enterRoundFunc(self, ctx: MathExprParser.RoundFuncContext):
def enterRoundFunc(self, ctx:MathExprParser.RoundFuncContext):
pass
# Exit a parse tree produced by MathExprParser#RoundFunc.
def exitRoundFunc(self, ctx: MathExprParser.RoundFuncContext):
def exitRoundFunc(self, ctx:MathExprParser.RoundFuncContext):
pass
# Enter a parse tree produced by MathExprParser#GammaFunc.
def enterGammaFunc(self, ctx: MathExprParser.GammaFuncContext):
def enterGammaFunc(self, ctx:MathExprParser.GammaFuncContext):
pass
# Exit a parse tree produced by MathExprParser#GammaFunc.
def exitGammaFunc(self, ctx: MathExprParser.GammaFuncContext):
def exitGammaFunc(self, ctx:MathExprParser.GammaFuncContext):
pass
# Enter a parse tree produced by MathExprParser#sigmoidFunc.
def enterSigmoidFunc(self, ctx: MathExprParser.SigmoidFuncContext):
def enterSigmoidFunc(self, ctx:MathExprParser.SigmoidFuncContext):
pass
# Exit a parse tree produced by MathExprParser#sigmoidFunc.
def exitSigmoidFunc(self, ctx: MathExprParser.SigmoidFuncContext):
def exitSigmoidFunc(self, ctx:MathExprParser.SigmoidFuncContext):
pass
# Enter a parse tree produced by MathExprParser#sfftFunc.
def enterSfftFunc(self, ctx: MathExprParser.SfftFuncContext):
def enterSfftFunc(self, ctx:MathExprParser.SfftFuncContext):
pass
# Exit a parse tree produced by MathExprParser#sfftFunc.
def exitSfftFunc(self, ctx: MathExprParser.SfftFuncContext):
def exitSfftFunc(self, ctx:MathExprParser.SfftFuncContext):
pass
# Enter a parse tree produced by MathExprParser#sifftFunc.
def enterSifftFunc(self, ctx: MathExprParser.SifftFuncContext):
def enterSifftFunc(self, ctx:MathExprParser.SifftFuncContext):
pass
# Exit a parse tree produced by MathExprParser#sifftFunc.
def exitSifftFunc(self, ctx: MathExprParser.SifftFuncContext):
def exitSifftFunc(self, ctx:MathExprParser.SifftFuncContext):
pass
# Enter a parse tree produced by MathExprParser#anglFunc.
def enterAnglFunc(self, ctx: MathExprParser.AnglFuncContext):
def enterAnglFunc(self, ctx:MathExprParser.AnglFuncContext):
pass
# Exit a parse tree produced by MathExprParser#anglFunc.
def exitAnglFunc(self, ctx: MathExprParser.AnglFuncContext):
def exitAnglFunc(self, ctx:MathExprParser.AnglFuncContext):
pass
# Enter a parse tree produced by MathExprParser#printFunc.
def enterPrintFunc(self, ctx: MathExprParser.PrintFuncContext):
def enterPrintFunc(self, ctx:MathExprParser.PrintFuncContext):
pass
# Exit a parse tree produced by MathExprParser#printFunc.
def exitPrintFunc(self, ctx: MathExprParser.PrintFuncContext):
def exitPrintFunc(self, ctx:MathExprParser.PrintFuncContext):
pass
# Enter a parse tree produced by MathExprParser#FractFunc.
def enterFractFunc(self, ctx: MathExprParser.FractFuncContext):
def enterFractFunc(self, ctx:MathExprParser.FractFuncContext):
pass
# Exit a parse tree produced by MathExprParser#FractFunc.
def exitFractFunc(self, ctx: MathExprParser.FractFuncContext):
def exitFractFunc(self, ctx:MathExprParser.FractFuncContext):
pass
# Enter a parse tree produced by MathExprParser#ReluFunc.
def enterReluFunc(self, ctx: MathExprParser.ReluFuncContext):
def enterReluFunc(self, ctx:MathExprParser.ReluFuncContext):
pass
# Exit a parse tree produced by MathExprParser#ReluFunc.
def exitReluFunc(self, ctx: MathExprParser.ReluFuncContext):
def exitReluFunc(self, ctx:MathExprParser.ReluFuncContext):
pass
# Enter a parse tree produced by MathExprParser#SoftplusFunc.
def enterSoftplusFunc(self, ctx: MathExprParser.SoftplusFuncContext):
def enterSoftplusFunc(self, ctx:MathExprParser.SoftplusFuncContext):
pass
# Exit a parse tree produced by MathExprParser#SoftplusFunc.
def exitSoftplusFunc(self, ctx: MathExprParser.SoftplusFuncContext):
def exitSoftplusFunc(self, ctx:MathExprParser.SoftplusFuncContext):
pass
# Enter a parse tree produced by MathExprParser#GeluFunc.
def enterGeluFunc(self, ctx: MathExprParser.GeluFuncContext):
def enterGeluFunc(self, ctx:MathExprParser.GeluFuncContext):
pass
# Exit a parse tree produced by MathExprParser#GeluFunc.
def exitGeluFunc(self, ctx: MathExprParser.GeluFuncContext):
def exitGeluFunc(self, ctx:MathExprParser.GeluFuncContext):
pass
# Enter a parse tree produced by MathExprParser#SignFunc.
def enterSignFunc(self, ctx: MathExprParser.SignFuncContext):
def enterSignFunc(self, ctx:MathExprParser.SignFuncContext):
pass
# Exit a parse tree produced by MathExprParser#SignFunc.
def exitSignFunc(self, ctx: MathExprParser.SignFuncContext):
def exitSignFunc(self, ctx:MathExprParser.SignFuncContext):
pass
# Enter a parse tree produced by MathExprParser#PrintShapeFunc.
def enterPrintShapeFunc(self, ctx: MathExprParser.PrintShapeFuncContext):
def enterPrintShapeFunc(self, ctx:MathExprParser.PrintShapeFuncContext):
pass
# Exit a parse tree produced by MathExprParser#PrintShapeFunc.
def exitPrintShapeFunc(self, ctx: MathExprParser.PrintShapeFuncContext):
def exitPrintShapeFunc(self, ctx:MathExprParser.PrintShapeFuncContext):
pass
# Enter a parse tree produced by MathExprParser#PinvFunc.
def enterPinvFunc(self, ctx:MathExprParser.PinvFuncContext):
pass
# Exit a parse tree produced by MathExprParser#PinvFunc.
def exitPinvFunc(self, ctx:MathExprParser.PinvFuncContext):
pass
# Enter a parse tree produced by MathExprParser#PowFunc.
def enterPowFunc(self, ctx: MathExprParser.PowFuncContext):
def enterPowFunc(self, ctx:MathExprParser.PowFuncContext):
pass
# Exit a parse tree produced by MathExprParser#PowFunc.
def exitPowFunc(self, ctx: MathExprParser.PowFuncContext):
def exitPowFunc(self, ctx:MathExprParser.PowFuncContext):
pass
# Enter a parse tree produced by MathExprParser#Atan2Func.
def enterAtan2Func(self, ctx: MathExprParser.Atan2FuncContext):
def enterAtan2Func(self, ctx:MathExprParser.Atan2FuncContext):
pass
# Exit a parse tree produced by MathExprParser#Atan2Func.
def exitAtan2Func(self, ctx: MathExprParser.Atan2FuncContext):
def exitAtan2Func(self, ctx:MathExprParser.Atan2FuncContext):
pass
# Enter a parse tree produced by MathExprParser#TMinFunc.
def enterTMinFunc(self, ctx: MathExprParser.TMinFuncContext):
def enterTMinFunc(self, ctx:MathExprParser.TMinFuncContext):
pass
# Exit a parse tree produced by MathExprParser#TMinFunc.
def exitTMinFunc(self, ctx: MathExprParser.TMinFuncContext):
def exitTMinFunc(self, ctx:MathExprParser.TMinFuncContext):
pass
# Enter a parse tree produced by MathExprParser#TMaxFunc.
def enterTMaxFunc(self, ctx: MathExprParser.TMaxFuncContext):
def enterTMaxFunc(self, ctx:MathExprParser.TMaxFuncContext):
pass
# Exit a parse tree produced by MathExprParser#TMaxFunc.
def exitTMaxFunc(self, ctx: MathExprParser.TMaxFuncContext):
def exitTMaxFunc(self, ctx:MathExprParser.TMaxFuncContext):
pass
# Enter a parse tree produced by MathExprParser#StepFunc.
def enterStepFunc(self, ctx: MathExprParser.StepFuncContext):
def enterStepFunc(self, ctx:MathExprParser.StepFuncContext):
pass
# Exit a parse tree produced by MathExprParser#StepFunc.
def exitStepFunc(self, ctx: MathExprParser.StepFuncContext):
def exitStepFunc(self, ctx:MathExprParser.StepFuncContext):
pass
# Enter a parse tree produced by MathExprParser#TopkFunc.
def enterTopkFunc(self, ctx:MathExprParser.TopkFuncContext):
pass
# Exit a parse tree produced by MathExprParser#TopkFunc.
def exitTopkFunc(self, ctx:MathExprParser.TopkFuncContext):
pass
# Enter a parse tree produced by MathExprParser#BotkFunc.
def enterBotkFunc(self, ctx:MathExprParser.BotkFuncContext):
pass
# Exit a parse tree produced by MathExprParser#BotkFunc.
def exitBotkFunc(self, ctx:MathExprParser.BotkFuncContext):
pass
# Enter a parse tree produced by MathExprParser#ClampFunc.
def enterClampFunc(self, ctx: MathExprParser.ClampFuncContext):
def enterClampFunc(self, ctx:MathExprParser.ClampFuncContext):
pass
# Exit a parse tree produced by MathExprParser#ClampFunc.
def exitClampFunc(self, ctx: MathExprParser.ClampFuncContext):
def exitClampFunc(self, ctx:MathExprParser.ClampFuncContext):
pass
# Enter a parse tree produced by MathExprParser#LerpFunc.
def enterLerpFunc(self, ctx: MathExprParser.LerpFuncContext):
def enterLerpFunc(self, ctx:MathExprParser.LerpFuncContext):
pass
# Exit a parse tree produced by MathExprParser#LerpFunc.
def exitLerpFunc(self, ctx: MathExprParser.LerpFuncContext):
def exitLerpFunc(self, ctx:MathExprParser.LerpFuncContext):
pass
# Enter a parse tree produced by MathExprParser#SmoothstepFunc.
def enterSmoothstepFunc(self, ctx: MathExprParser.SmoothstepFuncContext):
def enterSmoothstepFunc(self, ctx:MathExprParser.SmoothstepFuncContext):
pass
# Exit a parse tree produced by MathExprParser#SmoothstepFunc.
def exitSmoothstepFunc(self, ctx: MathExprParser.SmoothstepFuncContext):
def exitSmoothstepFunc(self, ctx:MathExprParser.SmoothstepFuncContext):
pass
# Enter a parse tree produced by MathExprParser#RangeFunc.
def enterRangeFunc(self, ctx: MathExprParser.RangeFuncContext):
def enterRangeFunc(self, ctx:MathExprParser.RangeFuncContext):
pass
# Exit a parse tree produced by MathExprParser#RangeFunc.
def exitRangeFunc(self, ctx: MathExprParser.RangeFuncContext):
def exitRangeFunc(self, ctx:MathExprParser.RangeFuncContext):
pass
# Enter a parse tree produced by MathExprParser#SwapFunc.
def enterSwapFunc(self, ctx: MathExprParser.SwapFuncContext):
def enterSwapFunc(self, ctx:MathExprParser.SwapFuncContext):
pass
# Exit a parse tree produced by MathExprParser#SwapFunc.
def exitSwapFunc(self, ctx: MathExprParser.SwapFuncContext):
def exitSwapFunc(self, ctx:MathExprParser.SwapFuncContext):
pass
# Enter a parse tree produced by MathExprParser#SMinFunc.
def enterSMinFunc(self, ctx: MathExprParser.SMinFuncContext):
def enterSMinFunc(self, ctx:MathExprParser.SMinFuncContext):
pass
# Exit a parse tree produced by MathExprParser#SMinFunc.
def exitSMinFunc(self, ctx: MathExprParser.SMinFuncContext):
def exitSMinFunc(self, ctx:MathExprParser.SMinFuncContext):
pass
# Enter a parse tree produced by MathExprParser#SMaxFunc.
def enterSMaxFunc(self, ctx: MathExprParser.SMaxFuncContext):
def enterSMaxFunc(self, ctx:MathExprParser.SMaxFuncContext):
pass
# Exit a parse tree produced by MathExprParser#SMaxFunc.
def exitSMaxFunc(self, ctx: MathExprParser.SMaxFuncContext):
def exitSMaxFunc(self, ctx:MathExprParser.SMaxFuncContext):
pass
# Enter a parse tree produced by MathExprParser#MapFunc.
def enterMapFunc(self, ctx: MathExprParser.MapFuncContext):
def enterMapFunc(self, ctx:MathExprParser.MapFuncContext):
pass
# Exit a parse tree produced by MathExprParser#MapFunc.
def exitMapFunc(self, ctx: MathExprParser.MapFuncContext):
def exitMapFunc(self, ctx:MathExprParser.MapFuncContext):
pass
# Enter a parse tree produced by MathExprParser#ConvFunc.
def enterConvFunc(self, ctx: MathExprParser.ConvFuncContext):
def enterConvFunc(self, ctx:MathExprParser.ConvFuncContext):
pass
# Exit a parse tree produced by MathExprParser#ConvFunc.
def exitConvFunc(self, ctx: MathExprParser.ConvFuncContext):
def exitConvFunc(self, ctx:MathExprParser.ConvFuncContext):
pass
# Enter a parse tree produced by MathExprParser#PermuteFunc.
def enterPermuteFunc(self, ctx: MathExprParser.PermuteFuncContext):
def enterPermuteFunc(self, ctx:MathExprParser.PermuteFuncContext):
pass
# Exit a parse tree produced by MathExprParser#PermuteFunc.
def exitPermuteFunc(self, ctx: MathExprParser.PermuteFuncContext):
def exitPermuteFunc(self, ctx:MathExprParser.PermuteFuncContext):
pass
# Enter a parse tree produced by MathExprParser#ReshapeFunc.
def enterReshapeFunc(self, ctx: MathExprParser.ReshapeFuncContext):
def enterReshapeFunc(self, ctx:MathExprParser.ReshapeFuncContext):
pass
# Exit a parse tree produced by MathExprParser#ReshapeFunc.
def exitReshapeFunc(self, ctx: MathExprParser.ReshapeFuncContext):
def exitReshapeFunc(self, ctx:MathExprParser.ReshapeFuncContext):
pass
del MathExprParser
del MathExprParser
File diff suppressed because it is too large Load Diff
+179 -84
View File
@@ -1,6 +1,5 @@
# Generated from ./MathExpr.g4 by ANTLR 4.13.2
from antlr4 import *
if "." in __name__:
from .MathExprParser import MathExprParser
else:
@@ -8,331 +7,427 @@ else:
# This class defines a complete generic visitor for a parse tree produced by MathExprParser.
class MathExprVisitor(ParseTreeVisitor):
# Visit a parse tree produced by MathExprParser#expr.
def visitExpr(self, ctx: MathExprParser.ExprContext):
def visitExpr(self, ctx:MathExprParser.ExprContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#LtExp.
def visitLtExp(self, ctx: MathExprParser.LtExpContext):
def visitLtExp(self, ctx:MathExprParser.LtExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#EqExp.
def visitEqExp(self, ctx: MathExprParser.EqExpContext):
def visitEqExp(self, ctx:MathExprParser.EqExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ToAdd.
def visitToAdd(self, ctx: MathExprParser.ToAddContext):
def visitToAdd(self, ctx:MathExprParser.ToAddContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#GeExp.
def visitGeExp(self, ctx: MathExprParser.GeExpContext):
def visitGeExp(self, ctx:MathExprParser.GeExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#LeExp.
def visitLeExp(self, ctx: MathExprParser.LeExpContext):
def visitLeExp(self, ctx:MathExprParser.LeExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#NeExp.
def visitNeExp(self, ctx: MathExprParser.NeExpContext):
def visitNeExp(self, ctx:MathExprParser.NeExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#GtExp.
def visitGtExp(self, ctx: MathExprParser.GtExpContext):
def visitGtExp(self, ctx:MathExprParser.GtExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#AddExp.
def visitAddExp(self, ctx: MathExprParser.AddExpContext):
def visitAddExp(self, ctx:MathExprParser.AddExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ToMul.
def visitToMul(self, ctx: MathExprParser.ToMulContext):
def visitToMul(self, ctx:MathExprParser.ToMulContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#SubExp.
def visitSubExp(self, ctx: MathExprParser.SubExpContext):
def visitSubExp(self, ctx:MathExprParser.SubExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#MulExp.
def visitMulExp(self, ctx: MathExprParser.MulExpContext):
def visitMulExp(self, ctx:MathExprParser.MulExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ModExp.
def visitModExp(self, ctx: MathExprParser.ModExpContext):
def visitModExp(self, ctx:MathExprParser.ModExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#DivExp.
def visitDivExp(self, ctx: MathExprParser.DivExpContext):
def visitDivExp(self, ctx:MathExprParser.DivExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ToPow.
def visitToPow(self, ctx: MathExprParser.ToPowContext):
def visitToPow(self, ctx:MathExprParser.ToPowContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#PowExp.
def visitPowExp(self, ctx: MathExprParser.PowExpContext):
def visitPowExp(self, ctx:MathExprParser.PowExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ToUnary.
def visitToUnary(self, ctx: MathExprParser.ToUnaryContext):
def visitToUnary(self, ctx:MathExprParser.ToUnaryContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#UnaryPlus.
def visitUnaryPlus(self, ctx: MathExprParser.UnaryPlusContext):
def visitUnaryPlus(self, ctx:MathExprParser.UnaryPlusContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#UnaryMinus.
def visitUnaryMinus(self, ctx: MathExprParser.UnaryMinusContext):
def visitUnaryMinus(self, ctx:MathExprParser.UnaryMinusContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ToAtom.
def visitToAtom(self, ctx: MathExprParser.ToAtomContext):
def visitToAtom(self, ctx:MathExprParser.ToAtomContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#Func1Exp.
def visitFunc1Exp(self, ctx: MathExprParser.Func1ExpContext):
def visitFunc1Exp(self, ctx:MathExprParser.Func1ExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#Func2Exp.
def visitFunc2Exp(self, ctx: MathExprParser.Func2ExpContext):
def visitFunc2Exp(self, ctx:MathExprParser.Func2ExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#Func3Exp.
def visitFunc3Exp(self, ctx: MathExprParser.Func3ExpContext):
def visitFunc3Exp(self, ctx:MathExprParser.Func3ExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#Func4Exp.
def visitFunc4Exp(self, ctx: MathExprParser.Func4ExpContext):
def visitFunc4Exp(self, ctx:MathExprParser.Func4ExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#FuncNExp.
def visitFuncNExp(self, ctx: MathExprParser.FuncNExpContext):
def visitFuncNExp(self, ctx:MathExprParser.FuncNExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#VariableExp.
def visitVariableExp(self, ctx: MathExprParser.VariableExpContext):
def visitVariableExp(self, ctx:MathExprParser.VariableExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#NumberExp.
def visitNumberExp(self, ctx: MathExprParser.NumberExpContext):
def visitNumberExp(self, ctx:MathExprParser.NumberExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ConstantExp.
def visitConstantExp(self, ctx: MathExprParser.ConstantExpContext):
def visitConstantExp(self, ctx:MathExprParser.ConstantExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ParenExp.
def visitParenExp(self, ctx: MathExprParser.ParenExpContext):
def visitParenExp(self, ctx:MathExprParser.ParenExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#AbsExp.
def visitAbsExp(self, ctx: MathExprParser.AbsExpContext):
def visitAbsExp(self, ctx:MathExprParser.AbsExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ListExp.
def visitListExp(self, ctx: MathExprParser.ListExpContext):
def visitListExp(self, ctx:MathExprParser.ListExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#SinFunc.
def visitSinFunc(self, ctx: MathExprParser.SinFuncContext):
def visitSinFunc(self, ctx:MathExprParser.SinFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#CosFunc.
def visitCosFunc(self, ctx: MathExprParser.CosFuncContext):
def visitCosFunc(self, ctx:MathExprParser.CosFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#TanFunc.
def visitTanFunc(self, ctx: MathExprParser.TanFuncContext):
def visitTanFunc(self, ctx:MathExprParser.TanFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#AsinFunc.
def visitAsinFunc(self, ctx: MathExprParser.AsinFuncContext):
def visitAsinFunc(self, ctx:MathExprParser.AsinFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#AcosFunc.
def visitAcosFunc(self, ctx: MathExprParser.AcosFuncContext):
def visitAcosFunc(self, ctx:MathExprParser.AcosFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#AtanFunc.
def visitAtanFunc(self, ctx: MathExprParser.AtanFuncContext):
def visitAtanFunc(self, ctx:MathExprParser.AtanFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#SinhFunc.
def visitSinhFunc(self, ctx: MathExprParser.SinhFuncContext):
def visitSinhFunc(self, ctx:MathExprParser.SinhFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#CoshFunc.
def visitCoshFunc(self, ctx: MathExprParser.CoshFuncContext):
def visitCoshFunc(self, ctx:MathExprParser.CoshFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#TanhFunc.
def visitTanhFunc(self, ctx: MathExprParser.TanhFuncContext):
def visitTanhFunc(self, ctx:MathExprParser.TanhFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#AsinhFunc.
def visitAsinhFunc(self, ctx: MathExprParser.AsinhFuncContext):
def visitAsinhFunc(self, ctx:MathExprParser.AsinhFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#AcoshFunc.
def visitAcoshFunc(self, ctx: MathExprParser.AcoshFuncContext):
def visitAcoshFunc(self, ctx:MathExprParser.AcoshFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#AtanhFunc.
def visitAtanhFunc(self, ctx: MathExprParser.AtanhFuncContext):
def visitAtanhFunc(self, ctx:MathExprParser.AtanhFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#AbsFunc.
def visitAbsFunc(self, ctx: MathExprParser.AbsFuncContext):
def visitAbsFunc(self, ctx:MathExprParser.AbsFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#SqrtFunc.
def visitSqrtFunc(self, ctx: MathExprParser.SqrtFuncContext):
def visitSqrtFunc(self, ctx:MathExprParser.SqrtFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#LnFunc.
def visitLnFunc(self, ctx: MathExprParser.LnFuncContext):
def visitLnFunc(self, ctx:MathExprParser.LnFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#LogFunc.
def visitLogFunc(self, ctx: MathExprParser.LogFuncContext):
def visitLogFunc(self, ctx:MathExprParser.LogFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ExpFunc.
def visitExpFunc(self, ctx: MathExprParser.ExpFuncContext):
def visitExpFunc(self, ctx:MathExprParser.ExpFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#TNormFunc.
def visitTNormFunc(self, ctx: MathExprParser.TNormFuncContext):
def visitTNormFunc(self, ctx:MathExprParser.TNormFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#SNormFunc.
def visitSNormFunc(self, ctx: MathExprParser.SNormFuncContext):
def visitSNormFunc(self, ctx:MathExprParser.SNormFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#FloorFunc.
def visitFloorFunc(self, ctx: MathExprParser.FloorFuncContext):
def visitFloorFunc(self, ctx:MathExprParser.FloorFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#CeilFunc.
def visitCeilFunc(self, ctx: MathExprParser.CeilFuncContext):
def visitCeilFunc(self, ctx:MathExprParser.CeilFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#RoundFunc.
def visitRoundFunc(self, ctx: MathExprParser.RoundFuncContext):
def visitRoundFunc(self, ctx:MathExprParser.RoundFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#GammaFunc.
def visitGammaFunc(self, ctx: MathExprParser.GammaFuncContext):
def visitGammaFunc(self, ctx:MathExprParser.GammaFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#sigmoidFunc.
def visitSigmoidFunc(self, ctx: MathExprParser.SigmoidFuncContext):
def visitSigmoidFunc(self, ctx:MathExprParser.SigmoidFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#sfftFunc.
def visitSfftFunc(self, ctx: MathExprParser.SfftFuncContext):
def visitSfftFunc(self, ctx:MathExprParser.SfftFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#sifftFunc.
def visitSifftFunc(self, ctx: MathExprParser.SifftFuncContext):
def visitSifftFunc(self, ctx:MathExprParser.SifftFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#anglFunc.
def visitAnglFunc(self, ctx: MathExprParser.AnglFuncContext):
def visitAnglFunc(self, ctx:MathExprParser.AnglFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#printFunc.
def visitPrintFunc(self, ctx: MathExprParser.PrintFuncContext):
def visitPrintFunc(self, ctx:MathExprParser.PrintFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#FractFunc.
def visitFractFunc(self, ctx: MathExprParser.FractFuncContext):
def visitFractFunc(self, ctx:MathExprParser.FractFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ReluFunc.
def visitReluFunc(self, ctx: MathExprParser.ReluFuncContext):
def visitReluFunc(self, ctx:MathExprParser.ReluFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#SoftplusFunc.
def visitSoftplusFunc(self, ctx: MathExprParser.SoftplusFuncContext):
def visitSoftplusFunc(self, ctx:MathExprParser.SoftplusFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#GeluFunc.
def visitGeluFunc(self, ctx: MathExprParser.GeluFuncContext):
def visitGeluFunc(self, ctx:MathExprParser.GeluFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#SignFunc.
def visitSignFunc(self, ctx: MathExprParser.SignFuncContext):
def visitSignFunc(self, ctx:MathExprParser.SignFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#PrintShapeFunc.
def visitPrintShapeFunc(self, ctx: MathExprParser.PrintShapeFuncContext):
def visitPrintShapeFunc(self, ctx:MathExprParser.PrintShapeFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#PinvFunc.
def visitPinvFunc(self, ctx:MathExprParser.PinvFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#PowFunc.
def visitPowFunc(self, ctx: MathExprParser.PowFuncContext):
def visitPowFunc(self, ctx:MathExprParser.PowFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#Atan2Func.
def visitAtan2Func(self, ctx: MathExprParser.Atan2FuncContext):
def visitAtan2Func(self, ctx:MathExprParser.Atan2FuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#TMinFunc.
def visitTMinFunc(self, ctx: MathExprParser.TMinFuncContext):
def visitTMinFunc(self, ctx:MathExprParser.TMinFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#TMaxFunc.
def visitTMaxFunc(self, ctx: MathExprParser.TMaxFuncContext):
def visitTMaxFunc(self, ctx:MathExprParser.TMaxFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#StepFunc.
def visitStepFunc(self, ctx: MathExprParser.StepFuncContext):
def visitStepFunc(self, ctx:MathExprParser.StepFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#TopkFunc.
def visitTopkFunc(self, ctx:MathExprParser.TopkFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#BotkFunc.
def visitBotkFunc(self, ctx:MathExprParser.BotkFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ClampFunc.
def visitClampFunc(self, ctx: MathExprParser.ClampFuncContext):
def visitClampFunc(self, ctx:MathExprParser.ClampFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#LerpFunc.
def visitLerpFunc(self, ctx: MathExprParser.LerpFuncContext):
def visitLerpFunc(self, ctx:MathExprParser.LerpFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#SmoothstepFunc.
def visitSmoothstepFunc(self, ctx: MathExprParser.SmoothstepFuncContext):
def visitSmoothstepFunc(self, ctx:MathExprParser.SmoothstepFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#RangeFunc.
def visitRangeFunc(self, ctx: MathExprParser.RangeFuncContext):
def visitRangeFunc(self, ctx:MathExprParser.RangeFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#SwapFunc.
def visitSwapFunc(self, ctx: MathExprParser.SwapFuncContext):
def visitSwapFunc(self, ctx:MathExprParser.SwapFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#SMinFunc.
def visitSMinFunc(self, ctx: MathExprParser.SMinFuncContext):
def visitSMinFunc(self, ctx:MathExprParser.SMinFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#SMaxFunc.
def visitSMaxFunc(self, ctx: MathExprParser.SMaxFuncContext):
def visitSMaxFunc(self, ctx:MathExprParser.SMaxFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#MapFunc.
def visitMapFunc(self, ctx: MathExprParser.MapFuncContext):
def visitMapFunc(self, ctx:MathExprParser.MapFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ConvFunc.
def visitConvFunc(self, ctx: MathExprParser.ConvFuncContext):
def visitConvFunc(self, ctx:MathExprParser.ConvFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#PermuteFunc.
def visitPermuteFunc(self, ctx: MathExprParser.PermuteFuncContext):
def visitPermuteFunc(self, ctx:MathExprParser.PermuteFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ReshapeFunc.
def visitReshapeFunc(self, ctx: MathExprParser.ReshapeFuncContext):
def visitReshapeFunc(self, ctx:MathExprParser.ReshapeFuncContext):
return self.visitChildren(ctx)
del MathExprParser
del MathExprParser
+91 -5
View File
@@ -23,7 +23,7 @@ class UnifiedMathVisitor(MathExprVisitor):
def _promote_to_tensor(self, val):
if self._is_tensor(val):
return val
return val.contiguous()
if self._is_list(val):
return torch.tensor(val, device=self.device)
return torch.tensor(val, device=self.device)
@@ -50,7 +50,7 @@ class UnifiedMathVisitor(MathExprVisitor):
if self._is_tensor(a) or self._is_tensor(b):
if torch_op:
return torch_op(a, b)
return torch_op(a, b).contiguous()
return scalar_op(a, b)
return scalar_op(a, b)
@@ -59,7 +59,7 @@ class UnifiedMathVisitor(MathExprVisitor):
if self._is_list(a):
return [self._unary_op(x, torch_op, scalar_op) for x in a]
if self._is_tensor(a):
return torch_op(a) if torch_op else scalar_op(a)
return torch_op(a).contiguous() if torch_op else scalar_op(a)
return scalar_op(a)
# ========================
@@ -143,7 +143,7 @@ class UnifiedMathVisitor(MathExprVisitor):
if self._is_list(arg):
return [self._func_dispatch(x, torch_fn, scalar_fn) for x in arg]
if self._is_tensor(arg):
return torch_fn(arg)
return torch_fn(arg).contiguous()
return scalar_fn(arg)
def visitSinFunc(self, ctx):
@@ -287,12 +287,98 @@ class UnifiedMathVisitor(MathExprVisitor):
lambda x, edge: 1.0 if x >= edge else 0.0,
)
def visitTopkFunc(self, ctx):
val = self.visit(ctx.expr(0))
k = self.visit(ctx.expr(1))
if self._is_tensor(k):
k_val = int(k.flatten()[0].item())
else:
k_val = int(k)
if self._is_tensor(val):
size = val.shape[-1]
k_val = max(1, min(k_val, size))
score_val = val.abs() if torch.is_complex(val) else val
_, indices = torch.topk(score_val, k=k_val, dim=-1)
indices = indices.contiguous()
mask = torch.zeros_like(score_val, dtype=torch.bool)
mask.scatter_(dim=-1, index=indices, value=True)
result = torch.where(mask, val, torch.zeros_like(val))
return result.contiguous()
if self._is_list(val):
k_val = max(0, min(k_val, len(val)))
try:
return sorted(val, reverse=True)[:k_val]
except:
return val[:k_val]
return val
def visitBotkFunc(self, ctx):
val = self.visit(ctx.expr(0))
k = self.visit(ctx.expr(1))
if self._is_tensor(k):
k_val = int(k.flatten()[0].item())
else:
k_val = int(k)
if self._is_tensor(val):
size = val.shape[-1]
k_val = max(1, min(k_val, size))
score_val = val.abs() if torch.is_complex(val) else val
_, indices = torch.topk(score_val, k=k_val, dim=-1, largest=False)
indices = indices.contiguous()
mask = torch.zeros_like(score_val, dtype=torch.bool)
mask.scatter_(dim=-1, index=indices, value=True)
result = torch.where(mask, val, torch.zeros_like(val))
return result.contiguous()
if self._is_list(val):
k_val = max(0, min(k_val, len(val)))
try:
return sorted(val)[:k_val]
except:
return val[:k_val]
return val
def visitPinvFunc(self, ctx):
"""Permutation inverse: if input[i] = j, output[j] = i."""
val = self.visit(ctx.expr())
if self._is_list(val):
n = len(val)
result = [0] * n
for i, v in enumerate(val):
idx = int(v)
if 0 <= idx < n:
result[idx] = i
return result
if self._is_tensor(val):
val_list = val.flatten().tolist()
n = len(val_list)
result = [0] * n
for i, v in enumerate(val_list):
idx = int(v)
if 0 <= idx < n:
result[idx] = i
return torch.tensor(result, device=val.device, dtype=val.dtype).reshape(val.shape)
return val
# Three-argument functions
def visitClampFunc(self, ctx):
val = self.visit(ctx.expr(0))
min_v = self.visit(ctx.expr(1))
max_v = self.visit(ctx.expr(2))
# Handle mixed types manually or promote?
if any(self._is_tensor(x) for x in [val, min_v, max_v]):
return torch.clamp(self._promote_to_tensor(val), self._promote_to_tensor(min_v), self._promote_to_tensor(max_v))
if self._is_list(val):
-14
View File
@@ -139,17 +139,3 @@ def make_zero_like(ref):
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
+18
View File
@@ -26,6 +26,7 @@ from more_math.ConditioningMathNode import ConditioningMathNode
from more_math.LatentMathNode import LatentMathNode
from more_math.ImageMathNode import ImageMathNode
from more_math.FloatMathNode import FloatMathNode
from more_math.AudioMathNode import AudioMathNode
# ==========================================
@@ -264,6 +265,23 @@ def test_image_swap():
assert torch.allclose(res_swap, img_blue)
# ==========================================
# Audio Math Operations
# ==========================================
def test_audio_math_basic():
node = AudioMathNode()
waveform = torch.randn(1, 1, 1024)
audio = {"waveform": waveform, "sample_rate": 44100}
# result = a * 2.0
res = node.execute("a * 2.0", a=audio)[0]
assert isinstance(res, dict)
assert "waveform" in res
assert res["sample_rate"] == 44100
assert torch.allclose(res["waveform"], waveform * 2.0)
# ==========================================
# Nested Expressions
# ==========================================
+80
View File
@@ -40,6 +40,12 @@ def test_scalar_ops():
assert parse_and_visit("a * b", vars) == 6.0
assert parse_and_visit("sin(0)", vars) == 0.0
assert parse_and_visit("smax(a, b)", vars) == 3.0
assert parse_and_visit("step(2, 1)", vars) == 1.0
assert parse_and_visit("step(0, 1)", vars) == 0.0
assert parse_and_visit("clamp(5, 0, 10)", vars) == 5.0
assert parse_and_visit("clamp(-5, 0, 10)", vars) == 0.0
assert parse_and_visit("lerp(0, 10, 0.5)", vars) == 5.0
assert parse_and_visit("fract(1.25)", vars) == 0.25
# Type check - ensure they are python float/int, not tensor
res = parse_and_visit("a + b", vars)
@@ -176,6 +182,77 @@ def test_bool_ops():
assert torch.all(res == torch.tensor([1.0, 0.0]))
def test_topk():
# Tensor topk
t = torch.tensor([1.0, 5.0, 2.0, 8.0, 3.0])
vars = {"t": t}
# Tensor topk masking
t = torch.tensor([1.0, 5.0, 2.0, 8.0, 3.0])
vars = {"t": t}
res = parse_and_visit("topk(t, 3)", vars)
assert isinstance(res, torch.Tensor)
assert res.shape == t.shape
# Top 3 are 8, 5, 3. Masked result should be [0, 5, 0, 8, 3]
expected = torch.tensor([0.0, 5.0, 0.0, 8.0, 3.0])
assert torch.allclose(res, expected)
assert res.is_contiguous()
# Complex topk masking
tc = torch.tensor([1.0+1j, 5.0+5j, 2.0+2j])
vars["tc"] = tc
res_c = parse_and_visit("topk(tc, 1)", vars)
assert res_c.shape == tc.shape
# Top 1 is 5+5j. Masked should be [0, 5+5j, 0]
expected_c = torch.tensor([0.0+0j, 5.0+5j, 0.0+0j])
assert torch.allclose(res_c, expected_c)
def test_botk():
# Tensor botk masking (bottom k smallest values)
t = torch.tensor([1.0, 5.0, 2.0, 8.0, 3.0])
vars = {"t": t}
res = parse_and_visit("botk(t, 3)", vars)
assert isinstance(res, torch.Tensor)
assert res.shape == t.shape
# Bottom 3 are 1, 2, 3. Masked result should be [1, 0, 2, 0, 3]
expected = torch.tensor([1.0, 0.0, 2.0, 0.0, 3.0])
assert torch.allclose(res, expected)
assert res.is_contiguous()
# List botk
l = [1.0, 5.0, 2.0, 8.0, 3.0]
vars = {"l": l}
res_l = parse_and_visit("botk(l, 2)", vars)
assert isinstance(res_l, list)
assert res_l == [1.0, 2.0]
def test_pinv():
# List permutation inverse
perm = [2, 0, 1] # 0->2, 1->0, 2->1
vars = {"perm": perm}
res = parse_and_visit("pinv(perm)", vars)
assert isinstance(res, list)
# Inverse: if perm[i]=j, then inv[j]=i
# perm[0]=2 -> inv[2]=0
# perm[1]=0 -> inv[0]=1
# perm[2]=1 -> inv[1]=2
assert res == [1, 2, 0]
# Tensor permutation inverse
perm_t = torch.tensor([2, 0, 1])
vars["perm_t"] = perm_t
res_t = parse_and_visit("pinv(perm_t)", vars)
assert isinstance(res_t, torch.Tensor)
assert torch.equal(res_t, torch.tensor([1, 2, 0]))
def test_pinv_identity():
perm = [2,0,1,6,4,3,5]
tensor = torch.rand([11,14,32,21,4,3,1])
varbl = {'c':tensor,'a':perm}
res = parse_and_visit("permute(permute(c,a),pinv(a))",varbl)
assert torch.equal(tensor,res)
if __name__ == "__main__":
try:
test_scalar_ops()
@@ -187,6 +264,9 @@ if __name__ == "__main__":
test_hyperbolic_trig()
test_kernel_coords()
test_bool_ops()
test_topk()
test_botk()
test_pinv()
print("All UnifiedMathVisitor tests passed!")
except Exception:
import traceback