Add tensor ops and parser support

Expose new tensor operations (softmax, softmin, argmin, argmax, unique, flatten, matmul, cross, and related tensor helpers) in the README and add corresponding grammar rules/tokens. Updated MathExpr.g4 to include new func productions and token definitions, and regenerated parser artifacts (MathExprLexer/Parser/Listener/Visitor/UnifiedMathVisitor and .interp/.tokens files). Also includes minor README formatting/typo fixes.
This commit is contained in:
mcDandy
2026-02-10 15:38:47 +01:00
parent f73a472324
commit 5072b191b7
11 changed files with 4510 additions and 2520 deletions
+15 -6
View File
@@ -100,6 +100,8 @@ You can also get the node from comfy manager under the name of More math.
- `gelu(x)`: Gaussian Error Linear Unit.
- `softplus(x)`: Softplus function (log(1 + e^x)).
- `sigm(x)`: Sigmoid function (1 / (1 + e^-x)).
- `softmax(x, dim)`: Softmax normalization along specified dimension (converts to probabilities).
- `softmin(x, dim)`: Softmin normalization along specified dimension (inverse softmax).
### Interpolation
@@ -133,6 +135,9 @@ You can also get the node from comfy manager under the name of More math.
- `botk_ind(x, k)` or `botk_indices`: Returns the **indices** of the bottom K smallest values in the flattened tensor.
- `sort(x)`: Sorts elements in ascending order along the last dimension.
- `argsort(x)` or `argsort(x, descending)`: Returns the **indices** that would sort the tensor/list. Optional second parameter for descending order.
- `argmin(x)`: Returns the **index** of the minimum value in the flattened tensor/list.
- `argmax(x)`: Returns the **index** of the maximum value in the flattened tensor/list.
- `unique(x)`: Returns **unique elements** from tensor/list in sorted order.
- `tnorm(x)`: Tensor normalisation. Normalises x (L2 norm along last dimension).
- `snorm(x)`: The same as |x| for tensors.
- `swap(tensor, dim, index1, index2)`: Swaps two slices of a tensor along a specified dimension.
@@ -144,6 +149,8 @@ You can also get the node from comfy manager under the name of More math.
- `all(x)`: Returns 1.0 if all elements in `x` are non-zero (True), else 0.0.
- `cumsum(x)`: Returns the cumulative sum of elements along the batch dimension (dim 0).
- `cumprod(x)`: Returns the cumulative product of elements along the batch dimension (dim 0).
- `tensor(shape,value)`: Createss a tensor of given shape filled with value. Value can be omittend and defaults to zero.
- `flatten(value)`: Flattens a tensor to 1D. If input is list, it flattens nested lists into a single list.
### Advanced Tensor Operations
@@ -155,14 +162,14 @@ You can also get the node from comfy manager under the name of More math.
- `k_expr` can be a math expression (using `kX`, `kY`, `kZ`) or a list literal.
- `get_value(tensor, position)`: Retrieves a value from a tensor at the specified N-dimensional position (provided as a list or tensor). Uses the formula `pos0*strides[0] + pos1*strides[1] + ...` to find the linear index.
- `crop(tensor, position, size)`: Extracts a sub-tensor of specified `size` starting at `position` (both provided as lists/tensors). Areas outside the input tensor are filled with zeros.
- `permute(tensor, dims)` or `perm`: Rearranges the dimensions of the tensor. (e.g., `perm(a, [2, 3, 0, 1])`)
- `reshape(tensor, shape)` or `rshp`: Reshapes the tensor to a new shape. (e.g., `rshp(a, [S0*S1, S2, S3])`)
- `blur(x, sigma)` or `gaussian`: Applies a Gaussian blur with given `sigma` along last two or spatial dimensions (toggleable by optional parameter) - default use last 2 dimensions.
- `edge(x)`: Applies a Sobel edge detection filter along the last two dimension or spatial dimensions (Height and Width) - can be selected by optional value (0 or missing = use last 2 dimensions).
- `batch_shuffle(tensor, indices)` or `shuffle` or `select`: Reorders or gathers slices along the 0th dimension of a tensor based on a list of indices. (e.g., `shuffle(V0, [0, 0, 1])` repeats the first frame twice and then the second).
- `tensor([shape],value)` Createss a tensor of given shape filled with value. Value can be omittend and defaults to zero.
- `
- `matmul(a, b)`: Matrix multiplication. For 1D vectors, performs dot product. For 2D+ tensors, performs standard matrix multiplication following NumPy rules.
- `cross(a, b)`: Computes the cross product (vector product) of two 3D vectors. Both inputs must have last dimension = 3. Returns a vector perpendicular to both inputs.
### FFT (Tensor Only)
@@ -218,7 +225,8 @@ Generates random noise with default shape of aither first input or maximum of in
- `w`, `x`, `y`, `z`
- **INSIDE IFFT**
- `F` or `frequency_count` – frequency count (freq domain, iFFT only)
- `K` or `frequency` – isotropic frequency (Euclidean norm of indices, iFFT only)
- `F` or `frequency_count` � frequency count (freq domain, iFFT only)
- `K` or `frequency` � isotropic frequency (Euclidean norm of indices, iFFT only)
- `Kx`, `Ky`, `K_dimN` - frequency index for specific dimension
- `Fx`, `Fy`, `F_dimN` - frequency count for specific dimension
- **IMAGE and LATENT**:
@@ -241,15 +249,16 @@ Generates random noise with default shape of aither first input or maximum of in
- `N` or `channel_count` - count of channels
- `C` or `channel` - channel of audio
- `S` or `sample` – current audio sample
- `S` or `sample` � current audio sample
- `T` or `sample_count` - audio lenght in samples
- `R` or `sample_rate` – sample rate
- `R` or `sample_rate` � sample rate
- **VIDEO**
- refer to `IMAGE and LATENT` for visual part (but `batch` is `frame` and `batch_count` is `frame_count`)
- refer to `AUDIO` for sound part
- **NOISE**
- refer to `IMAGE and LATENT` for most variables
- `I` or `input_latent` – latent used as input to generate noise before noise is generated into it
- `I` or `input_latent` � latent used as input to generate noise before noise is generated into it
- **GUIDER**
- refer to `IMAGE and LATENT`
- `sigma` - current sigma value
+19 -3
View File
@@ -141,9 +141,14 @@ func1:
| CLEAR LPAREN expr RPAREN # ClearFunc
| HAS LPAREN expr RPAREN # HasFunc
| GET LPAREN expr RPAREN # GetFunc
| ARGSORT LPAREN expr (COMMA expr)? RPAREN # ArgsortFunc;
| ARGSORT LPAREN expr (COMMA expr)? RPAREN # ArgsortFunc
| ARGMIN LPAREN expr RPAREN # ArgminFunc
| ARGMAX LPAREN expr RPAREN # ArgmaxFunc
| SOFTMAX LPAREN expr RPAREN # SoftmaxFunc
| SOFTMIN LPAREN expr RPAREN # SoftminFunc
| UNIQUE LPAREN expr RPAREN # UniqueFunc
| FLATTEN LPAREN expr RPAREN # FlattenFunc;
// Two-argument functions Two-argument functions
func2:
POWE LPAREN expr COMMA expr RPAREN # PowFunc
| ATAN2 LPAREN expr COMMA expr RPAREN # Atan2Func
@@ -167,7 +172,9 @@ func2:
| PUSH LPAREN expr COMMA expr RPAREN # PushFunc
| GET_VALUE LPAREN expr COMMA expr RPAREN # GetValueFunc
| TENSOR LPAREN indexExpr (COMMA expr)? RPAREN # EmptyTensorFunc
| PAD LPAREN expr COMMA expr RPAREN # PadFunc;
| PAD LPAREN expr COMMA expr RPAREN # PadFunc
| CROSS LPAREN expr COMMA expr RPAREN # CrossFunc
| MATMUL LPAREN expr COMMA expr RPAREN # MatmulFunc;
func3:
CLAMP LPAREN expr COMMA expr COMMA expr RPAREN # ClampFunc
@@ -323,6 +330,15 @@ ARGSORT: 'argsort';
FOR: 'for';
IN: 'in';
ARGMIN: 'argmin' | 'arg_min';
ARGMAX: 'argmax' | 'arg_max';
UNIQUE: 'unique';
SOFTMAX: 'softmax';
SOFTMIN: 'softmin';
FLATTEN: 'flatten';
CROSS: 'cross' | 'cross_product';
MATMUL: 'matmul' | 'matrix_multiply' | 'mat_mul';
TIMESTAMP: 'timestamp' | 'now';
NONE: 'None' | 'none' | 'NULL' | 'null';
BREAK: 'break';
File diff suppressed because one or more lines are too long
+78 -64
View File
@@ -107,42 +107,50 @@ PAD=106
ARGSORT=107
FOR=108
IN=109
TIMESTAMP=110
NONE=111
BREAK=112
CONTINUE=113
TENSOR=114
PLUS=115
MINUS=116
MULT=117
DIV=118
MOD=119
POW=120
GE=121
GT=122
LE=123
LT=124
EQ=125
EQUEALS=126
NE=127
PIPE=128
LPAREN=129
RPAREN=130
COMMA=131
SEMICOLON=132
ARROW=133
LBRACKET=134
RBRACKET=135
QUESTION=136
COLON=137
LBRACE=138
RBRACE=139
NUMBER=140
CONSTANT=141
VARIABLE=142
SL_COMMENT=143
ML_COMMENT=144
WS=145
ARGMIN=110
ARGMAX=111
UNIQUE=112
SOFTMAX=113
SOFTMIN=114
FLATTEN=115
CROSS=116
MATMUL=117
TIMESTAMP=118
NONE=119
BREAK=120
CONTINUE=121
TENSOR=122
PLUS=123
MINUS=124
MULT=125
DIV=126
MOD=127
POW=128
GE=129
GT=130
LE=131
LT=132
EQ=133
EQUEALS=134
NE=135
PIPE=136
LPAREN=137
RPAREN=138
COMMA=139
SEMICOLON=140
ARROW=141
LBRACKET=142
RBRACKET=143
QUESTION=144
COLON=145
LBRACE=146
RBRACE=147
NUMBER=148
CONSTANT=149
VARIABLE=150
SL_COMMENT=151
ML_COMMENT=152
WS=153
'sin'=1
'cos'=2
'tan'=3
@@ -174,6 +182,8 @@ WS=145
'pow'=29
'sigm'=30
'clamp'=31
'fft'=32
'ifft'=33
'angle'=34
'print'=35
'lerp'=38
@@ -226,31 +236,35 @@ WS=145
'argsort'=107
'for'=108
'in'=109
'break'=112
'continue'=113
'tensor'=114
'+'=115
'-'=116
'*'=117
'/'=118
'%'=119
'^'=120
'>='=121
'>'=122
'<='=123
'<'=124
'=='=125
'='=126
'!='=127
'|'=128
'('=129
')'=130
','=131
';'=132
'->'=133
'['=134
']'=135
'?'=136
':'=137
'{'=138
'}'=139
'unique'=112
'softmax'=113
'softmin'=114
'flatten'=115
'break'=120
'continue'=121
'tensor'=122
'+'=123
'-'=124
'*'=125
'/'=126
'%'=127
'^'=128
'>='=129
'>'=130
'<='=131
'<'=132
'=='=133
'='=134
'!='=135
'|'=136
'('=137
')'=138
','=139
';'=140
'->'=141
'['=142
']'=143
'?'=144
':'=145
'{'=146
'}'=147
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
+78 -64
View File
@@ -107,42 +107,50 @@ PAD=106
ARGSORT=107
FOR=108
IN=109
TIMESTAMP=110
NONE=111
BREAK=112
CONTINUE=113
TENSOR=114
PLUS=115
MINUS=116
MULT=117
DIV=118
MOD=119
POW=120
GE=121
GT=122
LE=123
LT=124
EQ=125
EQUEALS=126
NE=127
PIPE=128
LPAREN=129
RPAREN=130
COMMA=131
SEMICOLON=132
ARROW=133
LBRACKET=134
RBRACKET=135
QUESTION=136
COLON=137
LBRACE=138
RBRACE=139
NUMBER=140
CONSTANT=141
VARIABLE=142
SL_COMMENT=143
ML_COMMENT=144
WS=145
ARGMIN=110
ARGMAX=111
UNIQUE=112
SOFTMAX=113
SOFTMIN=114
FLATTEN=115
CROSS=116
MATMUL=117
TIMESTAMP=118
NONE=119
BREAK=120
CONTINUE=121
TENSOR=122
PLUS=123
MINUS=124
MULT=125
DIV=126
MOD=127
POW=128
GE=129
GT=130
LE=131
LT=132
EQ=133
EQUEALS=134
NE=135
PIPE=136
LPAREN=137
RPAREN=138
COMMA=139
SEMICOLON=140
ARROW=141
LBRACKET=142
RBRACKET=143
QUESTION=144
COLON=145
LBRACE=146
RBRACE=147
NUMBER=148
CONSTANT=149
VARIABLE=150
SL_COMMENT=151
ML_COMMENT=152
WS=153
'sin'=1
'cos'=2
'tan'=3
@@ -174,6 +182,8 @@ WS=145
'pow'=29
'sigm'=30
'clamp'=31
'fft'=32
'ifft'=33
'angle'=34
'print'=35
'lerp'=38
@@ -226,31 +236,35 @@ WS=145
'argsort'=107
'for'=108
'in'=109
'break'=112
'continue'=113
'tensor'=114
'+'=115
'-'=116
'*'=117
'/'=118
'%'=119
'^'=120
'>='=121
'>'=122
'<='=123
'<'=124
'=='=125
'='=126
'!='=127
'|'=128
'('=129
')'=130
','=131
';'=132
'->'=133
'['=134
']'=135
'?'=136
':'=137
'{'=138
'}'=139
'unique'=112
'softmax'=113
'softmin'=114
'flatten'=115
'break'=120
'continue'=121
'tensor'=122
'+'=123
'-'=124
'*'=125
'/'=126
'%'=127
'^'=128
'>='=129
'>'=130
'<='=131
'<'=132
'=='=133
'='=134
'!='=135
'|'=136
'('=137
')'=138
','=139
';'=140
'->'=141
'['=142
']'=143
'?'=144
':'=145
'{'=146
'}'=147
+162 -81
View File
@@ -458,6 +458,15 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#FuncNoiseExp.
def enterFuncNoiseExp(self, ctx:MathExprParser.FuncNoiseExpContext):
pass
# Exit a parse tree produced by MathExprParser#FuncNoiseExp.
def exitFuncNoiseExp(self, ctx:MathExprParser.FuncNoiseExpContext):
pass
# Enter a parse tree produced by MathExprParser#VariableExp.
def enterVariableExp(self, ctx:MathExprParser.VariableExpContext):
pass
@@ -782,24 +791,6 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#sfftFunc.
def enterSfftFunc(self, ctx:MathExprParser.SfftFuncContext):
pass
# Exit a parse tree produced by MathExprParser#sfftFunc.
def exitSfftFunc(self, ctx:MathExprParser.SfftFuncContext):
pass
# Enter a parse tree produced by MathExprParser#sifftFunc.
def enterSifftFunc(self, ctx:MathExprParser.SifftFuncContext):
pass
# Exit a parse tree produced by MathExprParser#sifftFunc.
def exitSifftFunc(self, ctx:MathExprParser.SifftFuncContext):
pass
# Enter a parse tree produced by MathExprParser#anglFunc.
def enterAnglFunc(self, ctx:MathExprParser.AnglFuncContext):
pass
@@ -926,24 +917,6 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#NoiseFunc.
def enterNoiseFunc(self, ctx:MathExprParser.NoiseFuncContext):
pass
# Exit a parse tree produced by MathExprParser#NoiseFunc.
def exitNoiseFunc(self, ctx:MathExprParser.NoiseFuncContext):
pass
# Enter a parse tree produced by MathExprParser#RandFunc.
def enterRandFunc(self, ctx:MathExprParser.RandFuncContext):
pass
# Exit a parse tree produced by MathExprParser#RandFunc.
def exitRandFunc(self, ctx:MathExprParser.RandFuncContext):
pass
# Enter a parse tree produced by MathExprParser#AnyFunc.
def enterAnyFunc(self, ctx:MathExprParser.AnyFuncContext):
pass
@@ -1061,6 +1034,60 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#ArgminFunc.
def enterArgminFunc(self, ctx:MathExprParser.ArgminFuncContext):
pass
# Exit a parse tree produced by MathExprParser#ArgminFunc.
def exitArgminFunc(self, ctx:MathExprParser.ArgminFuncContext):
pass
# Enter a parse tree produced by MathExprParser#ArgmaxFunc.
def enterArgmaxFunc(self, ctx:MathExprParser.ArgmaxFuncContext):
pass
# Exit a parse tree produced by MathExprParser#ArgmaxFunc.
def exitArgmaxFunc(self, ctx:MathExprParser.ArgmaxFuncContext):
pass
# Enter a parse tree produced by MathExprParser#SoftmaxFunc.
def enterSoftmaxFunc(self, ctx:MathExprParser.SoftmaxFuncContext):
pass
# Exit a parse tree produced by MathExprParser#SoftmaxFunc.
def exitSoftmaxFunc(self, ctx:MathExprParser.SoftmaxFuncContext):
pass
# Enter a parse tree produced by MathExprParser#SoftminFunc.
def enterSoftminFunc(self, ctx:MathExprParser.SoftminFuncContext):
pass
# Exit a parse tree produced by MathExprParser#SoftminFunc.
def exitSoftminFunc(self, ctx:MathExprParser.SoftminFuncContext):
pass
# Enter a parse tree produced by MathExprParser#UniqueFunc.
def enterUniqueFunc(self, ctx:MathExprParser.UniqueFuncContext):
pass
# Exit a parse tree produced by MathExprParser#UniqueFunc.
def exitUniqueFunc(self, ctx:MathExprParser.UniqueFuncContext):
pass
# Enter a parse tree produced by MathExprParser#FlattenFunc.
def enterFlattenFunc(self, ctx:MathExprParser.FlattenFuncContext):
pass
# Exit a parse tree produced by MathExprParser#FlattenFunc.
def exitFlattenFunc(self, ctx:MathExprParser.FlattenFuncContext):
pass
# Enter a parse tree produced by MathExprParser#PowFunc.
def enterPowFunc(self, ctx:MathExprParser.PowFuncContext):
pass
@@ -1196,33 +1223,6 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#ExponentialFunc.
def enterExponentialFunc(self, ctx:MathExprParser.ExponentialFuncContext):
pass
# Exit a parse tree produced by MathExprParser#ExponentialFunc.
def exitExponentialFunc(self, ctx:MathExprParser.ExponentialFuncContext):
pass
# Enter a parse tree produced by MathExprParser#BernoulliFunc.
def enterBernoulliFunc(self, ctx:MathExprParser.BernoulliFuncContext):
pass
# Exit a parse tree produced by MathExprParser#BernoulliFunc.
def exitBernoulliFunc(self, ctx:MathExprParser.BernoulliFuncContext):
pass
# Enter a parse tree produced by MathExprParser#PoissonFunc.
def enterPoissonFunc(self, ctx:MathExprParser.PoissonFuncContext):
pass
# Exit a parse tree produced by MathExprParser#PoissonFunc.
def exitPoissonFunc(self, ctx:MathExprParser.PoissonFuncContext):
pass
# Enter a parse tree produced by MathExprParser#GaussianFunc.
def enterGaussianFunc(self, ctx:MathExprParser.GaussianFuncContext):
pass
@@ -1286,6 +1286,33 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#PadFunc.
def enterPadFunc(self, ctx:MathExprParser.PadFuncContext):
pass
# Exit a parse tree produced by MathExprParser#PadFunc.
def exitPadFunc(self, ctx:MathExprParser.PadFuncContext):
pass
# Enter a parse tree produced by MathExprParser#CrossFunc.
def enterCrossFunc(self, ctx:MathExprParser.CrossFuncContext):
pass
# Exit a parse tree produced by MathExprParser#CrossFunc.
def exitCrossFunc(self, ctx:MathExprParser.CrossFuncContext):
pass
# Enter a parse tree produced by MathExprParser#MatmulFunc.
def enterMatmulFunc(self, ctx:MathExprParser.MatmulFuncContext):
pass
# Exit a parse tree produced by MathExprParser#MatmulFunc.
def exitMatmulFunc(self, ctx:MathExprParser.MatmulFuncContext):
pass
# Enter a parse tree produced by MathExprParser#ClampFunc.
def enterClampFunc(self, ctx:MathExprParser.ClampFuncContext):
pass
@@ -1331,24 +1358,6 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#CauchyFunc.
def enterCauchyFunc(self, ctx:MathExprParser.CauchyFuncContext):
pass
# Exit a parse tree produced by MathExprParser#CauchyFunc.
def exitCauchyFunc(self, ctx:MathExprParser.CauchyFuncContext):
pass
# Enter a parse tree produced by MathExprParser#LogNormalFunc.
def enterLogNormalFunc(self, ctx:MathExprParser.LogNormalFuncContext):
pass
# Exit a parse tree produced by MathExprParser#LogNormalFunc.
def exitLogNormalFunc(self, ctx:MathExprParser.LogNormalFuncContext):
pass
# Enter a parse tree produced by MathExprParser#CubicEaseFunc.
def enterCubicEaseFunc(self, ctx:MathExprParser.CubicEaseFuncContext):
pass
@@ -1394,6 +1403,15 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#sifftFunc.
def enterSifftFunc(self, ctx:MathExprParser.SifftFuncContext):
pass
# Exit a parse tree produced by MathExprParser#sifftFunc.
def exitSifftFunc(self, ctx:MathExprParser.SifftFuncContext):
pass
# Enter a parse tree produced by MathExprParser#SwapFunc.
def enterSwapFunc(self, ctx:MathExprParser.SwapFuncContext):
pass
@@ -1493,5 +1511,68 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#NoiseFunc.
def enterNoiseFunc(self, ctx:MathExprParser.NoiseFuncContext):
pass
# Exit a parse tree produced by MathExprParser#NoiseFunc.
def exitNoiseFunc(self, ctx:MathExprParser.NoiseFuncContext):
pass
# Enter a parse tree produced by MathExprParser#RandFunc.
def enterRandFunc(self, ctx:MathExprParser.RandFuncContext):
pass
# Exit a parse tree produced by MathExprParser#RandFunc.
def exitRandFunc(self, ctx:MathExprParser.RandFuncContext):
pass
# Enter a parse tree produced by MathExprParser#ExponentialFunc.
def enterExponentialFunc(self, ctx:MathExprParser.ExponentialFuncContext):
pass
# Exit a parse tree produced by MathExprParser#ExponentialFunc.
def exitExponentialFunc(self, ctx:MathExprParser.ExponentialFuncContext):
pass
# Enter a parse tree produced by MathExprParser#BernoulliFunc.
def enterBernoulliFunc(self, ctx:MathExprParser.BernoulliFuncContext):
pass
# Exit a parse tree produced by MathExprParser#BernoulliFunc.
def exitBernoulliFunc(self, ctx:MathExprParser.BernoulliFuncContext):
pass
# Enter a parse tree produced by MathExprParser#PoissonFunc.
def enterPoissonFunc(self, ctx:MathExprParser.PoissonFuncContext):
pass
# Exit a parse tree produced by MathExprParser#PoissonFunc.
def exitPoissonFunc(self, ctx:MathExprParser.PoissonFuncContext):
pass
# Enter a parse tree produced by MathExprParser#CauchyFunc.
def enterCauchyFunc(self, ctx:MathExprParser.CauchyFuncContext):
pass
# Exit a parse tree produced by MathExprParser#CauchyFunc.
def exitCauchyFunc(self, ctx:MathExprParser.CauchyFuncContext):
pass
# Enter a parse tree produced by MathExprParser#LogNormalFunc.
def enterLogNormalFunc(self, ctx:MathExprParser.LogNormalFuncContext):
pass
# Exit a parse tree produced by MathExprParser#LogNormalFunc.
def exitLogNormalFunc(self, ctx:MathExprParser.LogNormalFuncContext):
pass
del MathExprParser
File diff suppressed because it is too large Load Diff
+47 -7
View File
@@ -259,8 +259,8 @@ class MathExprVisitor(ParseTreeVisitor):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#FuncOptExp.
def visitFuncOptExp(self, ctx:MathExprParser.FuncOptExpContext):
# Visit a parse tree produced by MathExprParser#FuncNoiseExp.
def visitFuncNoiseExp(self, ctx:MathExprParser.FuncNoiseExpContext):
return self.visitChildren(ctx)
@@ -579,6 +579,36 @@ class MathExprVisitor(ParseTreeVisitor):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ArgminFunc.
def visitArgminFunc(self, ctx:MathExprParser.ArgminFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ArgmaxFunc.
def visitArgmaxFunc(self, ctx:MathExprParser.ArgmaxFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#SoftmaxFunc.
def visitSoftmaxFunc(self, ctx:MathExprParser.SoftmaxFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#SoftminFunc.
def visitSoftminFunc(self, ctx:MathExprParser.SoftminFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#UniqueFunc.
def visitUniqueFunc(self, ctx:MathExprParser.UniqueFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#FlattenFunc.
def visitFlattenFunc(self, ctx:MathExprParser.FlattenFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#PowFunc.
def visitPowFunc(self, ctx:MathExprParser.PowFuncContext):
return self.visitChildren(ctx)
@@ -694,6 +724,16 @@ class MathExprVisitor(ParseTreeVisitor):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#CrossFunc.
def visitCrossFunc(self, ctx:MathExprParser.CrossFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#MatmulFunc.
def visitMatmulFunc(self, ctx:MathExprParser.MatmulFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ClampFunc.
def visitClampFunc(self, ctx:MathExprParser.ClampFuncContext):
return self.visitChildren(ctx)
@@ -744,6 +784,11 @@ class MathExprVisitor(ParseTreeVisitor):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#sifftFunc.
def visitSifftFunc(self, ctx:MathExprParser.SifftFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#SwapFunc.
def visitSwapFunc(self, ctx:MathExprParser.SwapFuncContext):
return self.visitChildren(ctx)
@@ -799,11 +844,6 @@ class MathExprVisitor(ParseTreeVisitor):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#sifftFunc.
def visitSifftFunc(self, ctx:MathExprParser.SifftFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#NoiseFunc.
def visitNoiseFunc(self, ctx:MathExprParser.NoiseFuncContext):
return self.visitChildren(ctx)
+53
View File
@@ -2037,3 +2037,56 @@ class UnifiedMathVisitor(MathExprVisitor):
shape = [int(float(shape_val))]
return torch.full(shape, value, device=self.device)
def visitSoftmaxFunc(self, ctx):
val = self._promote_to_tensor((yield ctx.expr()))
return F.softmax(val.float())
def visitSoftminFunc(self, ctx):
val = self._promote_to_tensor((yield ctx.expr()))
return F.softmax(-val.float())
def visitArgminFunc(self, ctx):
val = self._promote_to_tensor((yield ctx.expr()))
if self._is_tensor(val):
return torch.argmin(val.flatten())
if self._is_list(val):
return float(val.index(min(val)))
return 0.0
def visitArgmaxFunc(self, ctx):
val = self._promote_to_tensor((yield ctx.expr()))
if self._is_tensor(val):
return torch.argmax(val.flatten())
if self._is_list(val):
return float(val.index(max(val)))
return 0.0
def visitUniqueFunc(self, ctx):
val = self._promote_to_tensor((yield ctx.expr()))
if self._is_tensor(val):
unique_vals, _ = torch.unique(val.flatten(), return_counts=False, sorted=True)
return unique_vals
if self._is_list(val):
return sorted(list(set(val)))
return val
def visitFlattenFunc(self, ctx):
val = (yield ctx.expr())
if self._is_tensor(val):
return val.flatten()
if self._is_list(val):
return self._flatten_list(val)
return val
def _flatten_list(self, lst):
"""Recursivly flatten list"""
result = []
for item in lst:
if self._is_list(item):
result.extend(self._flatten_list(item))
else:
result.append(item)
return result