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:
@@ -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
|
||||
|
||||
@@ -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
@@ -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
+627
-585
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
|
||||
@@ -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
|
||||
+3384
-1703
File diff suppressed because it is too large
Load Diff
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user