add optional shape argument
This commit is contained in:
@@ -74,6 +74,7 @@ atom:
|
||||
| func4 # Func4Exp
|
||||
| func5 # Func5Exp
|
||||
| funcN # FuncNExp
|
||||
| funcNoise # FuncNoiseExp
|
||||
| VARIABLE # VariableExp
|
||||
| NUMBER # NumberExp
|
||||
| CONSTANT # ConstantExp
|
||||
@@ -90,61 +91,57 @@ exprList: expr (COMMA expr)*;
|
||||
func0: TIMESTAMP LPAREN RPAREN # TimestampFunc;
|
||||
// Single-argument functions
|
||||
func1:
|
||||
SIN LPAREN expr RPAREN # SinFunc
|
||||
| COS LPAREN expr RPAREN # CosFunc
|
||||
| TAN LPAREN expr RPAREN # TanFunc
|
||||
| ASIN LPAREN expr RPAREN # AsinFunc
|
||||
| ACOS LPAREN expr RPAREN # AcosFunc
|
||||
| ATAN LPAREN expr RPAREN # AtanFunc
|
||||
| SINH LPAREN expr RPAREN # SinhFunc
|
||||
| COSH LPAREN expr RPAREN # CoshFunc
|
||||
| TANH LPAREN expr RPAREN # TanhFunc
|
||||
| ASINH LPAREN expr RPAREN # AsinhFunc
|
||||
| ACOSH LPAREN expr RPAREN # AcoshFunc
|
||||
| ATANH LPAREN expr RPAREN # AtanhFunc
|
||||
| ABS LPAREN expr RPAREN # AbsFunc
|
||||
| SQRT LPAREN expr RPAREN # SqrtFunc
|
||||
| LN LPAREN expr RPAREN # LnFunc
|
||||
| LOG LPAREN expr RPAREN # LogFunc
|
||||
| EXP LPAREN expr RPAREN # ExpFunc
|
||||
| TNORM LPAREN expr RPAREN # TNormFunc
|
||||
| SNORM LPAREN expr RPAREN # SNormFunc
|
||||
| FLOOR LPAREN expr RPAREN # FloorFunc
|
||||
| CEIL LPAREN expr RPAREN # CeilFunc
|
||||
| ROUND LPAREN expr RPAREN # RoundFunc
|
||||
| GAMMA LPAREN expr RPAREN # GammaFunc
|
||||
| SIGM LPAREN expr RPAREN # sigmoidFunc
|
||||
| SFFT LPAREN expr RPAREN # sfftFunc
|
||||
| SIFFT LPAREN expr RPAREN # sifftFunc
|
||||
| ANGL LPAREN expr RPAREN # anglFunc
|
||||
| PRNT LPAREN expr RPAREN # printFunc
|
||||
| FRACT LPAREN expr RPAREN # FractFunc
|
||||
| RELU LPAREN expr RPAREN # ReluFunc
|
||||
| SOFTPLUS LPAREN expr RPAREN # SoftplusFunc
|
||||
| GELU LPAREN expr RPAREN # GeluFunc
|
||||
| SIGN LPAREN expr RPAREN # SignFunc
|
||||
| PRINT_SHAPE LPAREN expr RPAREN # PrintShapeFunc
|
||||
| PINV LPAREN expr RPAREN # PinvFunc
|
||||
| SUM LPAREN expr RPAREN # SumFunc
|
||||
| MEAN LPAREN expr RPAREN # MeanFunc
|
||||
| STD LPAREN expr RPAREN # StdFunc
|
||||
| VAR LPAREN expr RPAREN # VarFunc
|
||||
| SORT LPAREN expr RPAREN # SortFunc
|
||||
| NOISE LPAREN expr RPAREN # NoiseFunc
|
||||
| RAND LPAREN expr RPAREN # RandFunc
|
||||
| ANY LPAREN expr RPAREN # AnyFunc
|
||||
| ALL LPAREN expr RPAREN # AllFunc
|
||||
| EDGE LPAREN expr (COMMA expr)? RPAREN # EdgeFunc
|
||||
| MEDIAN LPAREN expr RPAREN # MedianFunc
|
||||
| MODE LPAREN expr RPAREN # ModeFunc
|
||||
| CUMSUM LPAREN expr RPAREN # CumsumFunc
|
||||
| COUNT LPAREN expr RPAREN # CountFunc
|
||||
| CUMPROD LPAREN expr RPAREN # CumprodFunc
|
||||
| POP LPAREN expr RPAREN # PopFunc
|
||||
| CLEAR LPAREN expr RPAREN # ClearFunc
|
||||
| HAS LPAREN expr RPAREN # HasFunc
|
||||
| GET LPAREN expr RPAREN # GetFunc
|
||||
| ARGSORT LPAREN expr (COMMA expr)? RPAREN # ArgsortFunc;
|
||||
SIN LPAREN expr RPAREN # SinFunc
|
||||
| COS LPAREN expr RPAREN # CosFunc
|
||||
| TAN LPAREN expr RPAREN # TanFunc
|
||||
| ASIN LPAREN expr RPAREN # AsinFunc
|
||||
| ACOS LPAREN expr RPAREN # AcosFunc
|
||||
| ATAN LPAREN expr RPAREN # AtanFunc
|
||||
| SINH LPAREN expr RPAREN # SinhFunc
|
||||
| COSH LPAREN expr RPAREN # CoshFunc
|
||||
| TANH LPAREN expr RPAREN # TanhFunc
|
||||
| ASINH LPAREN expr RPAREN # AsinhFunc
|
||||
| ACOSH LPAREN expr RPAREN # AcoshFunc
|
||||
| ATANH LPAREN expr RPAREN # AtanhFunc
|
||||
| ABS LPAREN expr RPAREN # AbsFunc
|
||||
| SQRT LPAREN expr RPAREN # SqrtFunc
|
||||
| LN LPAREN expr RPAREN # LnFunc
|
||||
| LOG LPAREN expr RPAREN # LogFunc
|
||||
| EXP LPAREN expr RPAREN # ExpFunc
|
||||
| TNORM LPAREN expr RPAREN # TNormFunc
|
||||
| SNORM LPAREN expr RPAREN # SNormFunc
|
||||
| FLOOR LPAREN expr RPAREN # FloorFunc
|
||||
| CEIL LPAREN expr RPAREN # CeilFunc
|
||||
| ROUND LPAREN expr RPAREN # RoundFunc
|
||||
| GAMMA LPAREN expr RPAREN # GammaFunc
|
||||
| SIGM LPAREN expr RPAREN # sigmoidFunc
|
||||
| ANGL LPAREN expr RPAREN # anglFunc
|
||||
| PRNT LPAREN expr RPAREN # printFunc
|
||||
| FRACT LPAREN expr RPAREN # FractFunc
|
||||
| RELU LPAREN expr RPAREN # ReluFunc
|
||||
| SOFTPLUS LPAREN expr RPAREN # SoftplusFunc
|
||||
| GELU LPAREN expr RPAREN # GeluFunc
|
||||
| SIGN LPAREN expr RPAREN # SignFunc
|
||||
| PRINT_SHAPE LPAREN expr RPAREN # PrintShapeFunc
|
||||
| PINV LPAREN expr RPAREN # PinvFunc
|
||||
| SUM LPAREN expr RPAREN # SumFunc
|
||||
| MEAN LPAREN expr RPAREN # MeanFunc
|
||||
| STD LPAREN expr RPAREN # StdFunc
|
||||
| VAR LPAREN expr RPAREN # VarFunc
|
||||
| SORT LPAREN expr RPAREN # SortFunc
|
||||
| ANY LPAREN expr RPAREN # AnyFunc
|
||||
| ALL LPAREN expr RPAREN # AllFunc
|
||||
| EDGE LPAREN expr (COMMA expr)? RPAREN # EdgeFunc
|
||||
| MEDIAN LPAREN expr RPAREN # MedianFunc
|
||||
| MODE LPAREN expr RPAREN # ModeFunc
|
||||
| CUMSUM LPAREN expr RPAREN # CumsumFunc
|
||||
| COUNT LPAREN expr RPAREN # CountFunc
|
||||
| CUMPROD LPAREN expr RPAREN # CumprodFunc
|
||||
| POP LPAREN expr RPAREN # PopFunc
|
||||
| CLEAR LPAREN expr RPAREN # ClearFunc
|
||||
| HAS LPAREN expr RPAREN # HasFunc
|
||||
| GET LPAREN expr RPAREN # GetFunc
|
||||
| ARGSORT LPAREN expr (COMMA expr)? RPAREN # ArgsortFunc;
|
||||
|
||||
// Two-argument functions Two-argument functions
|
||||
func2:
|
||||
@@ -163,16 +160,14 @@ func2:
|
||||
| FLIP LPAREN expr COMMA expr RPAREN # FlipFunc
|
||||
| COV LPAREN expr COMMA expr RPAREN # CovFunc
|
||||
| APPEND LPAREN expr COMMA expr RPAREN # AppendFunc
|
||||
| EXPONENTIAL LPAREN expr COMMA expr RPAREN # ExponentialFunc
|
||||
| BERNOULLI LPAREN expr COMMA expr RPAREN # BernoulliFunc
|
||||
| POISSON LPAREN expr COMMA expr RPAREN # PoissonFunc
|
||||
| GAUSSIAN LPAREN expr COMMA expr (COMMA expr)? RPAREN # GaussianFunc
|
||||
| TOPK_IND LPAREN expr COMMA expr RPAREN # TopkIndFunc
|
||||
| BOTK_IND LPAREN expr COMMA expr RPAREN # BotkIndFunc
|
||||
| BATCH_SHUFFLE LPAREN expr COMMA expr RPAREN # BatchShuffleFunc
|
||||
| PUSH LPAREN expr COMMA expr RPAREN # PushFunc
|
||||
| GET_VALUE LPAREN expr COMMA expr RPAREN # GetValueFunc
|
||||
| TENSOR LPAREN indexExpr (COMMA expr)? RPAREN # EmptyTensorFunc;
|
||||
| TENSOR LPAREN indexExpr (COMMA expr)? RPAREN # EmptyTensorFunc
|
||||
| PAD LPAREN expr COMMA expr RPAREN # PadFunc;
|
||||
|
||||
func3:
|
||||
CLAMP LPAREN expr COMMA expr COMMA expr RPAREN # ClampFunc
|
||||
@@ -180,15 +175,12 @@ func3:
|
||||
| SMOOTHSTEP LPAREN expr COMMA expr COMMA expr RPAREN # SmoothstepFunc
|
||||
| RANGE LPAREN expr COMMA expr COMMA expr RPAREN # RangeFunc
|
||||
| MOMENT LPAREN expr COMMA expr COMMA expr RPAREN # MomentFunc
|
||||
| CAUCHY LPAREN expr COMMA expr COMMA expr RPAREN # CauchyFunc
|
||||
| LOGNORMAL LPAREN expr COMMA expr COMMA expr RPAREN # LogNormalFunc
|
||||
| CUBIC_EASE LPAREN expr COMMA expr COMMA expr RPAREN # CubicEaseFunc
|
||||
| ELASTIC_EASE LPAREN expr COMMA expr COMMA expr RPAREN # ElasticEaseFunc
|
||||
| SINE_EASE LPAREN expr COMMA expr COMMA expr RPAREN # SineEaseFunc
|
||||
| SINE_EASE LPAREN expr COMMA expr COMMA expr RPAREN # SineEaseFunc
|
||||
| SMOOTHERSTEP LPAREN expr COMMA expr COMMA expr RPAREN # SmootherstepFunc
|
||||
| CROP LPAREN expr COMMA expr COMMA expr RPAREN # CropFunc;
|
||||
|
||||
| CROP LPAREN expr COMMA expr COMMA expr RPAREN # CropFunc
|
||||
| SIFFT LPAREN expr (COMMA expr)? RPAREN # sifftFunc;
|
||||
func4:
|
||||
SWAP LPAREN expr COMMA expr COMMA expr COMMA expr RPAREN # SwapFunc
|
||||
| NVL LPAREN expr COMMA expr COMMA expr COMMA expr RPAREN # NvlFunc
|
||||
@@ -207,6 +199,15 @@ funcN:
|
||||
| PERM LPAREN expr COMMA expr RPAREN # PermuteFunc
|
||||
| RESHAPE LPAREN expr COMMA expr RPAREN # ReshapeFunc;
|
||||
|
||||
funcNoise:
|
||||
NOISE LPAREN expr (COMMA expr)? RPAREN # NoiseFunc
|
||||
| RAND LPAREN expr (COMMA expr)? RPAREN # RandFunc
|
||||
| EXPONENTIAL LPAREN expr COMMA expr (COMMA expr)? RPAREN # ExponentialFunc
|
||||
| BERNOULLI LPAREN expr COMMA expr (COMMA expr)? RPAREN # BernoulliFunc
|
||||
| POISSON LPAREN expr COMMA expr (COMMA expr)? RPAREN # PoissonFunc
|
||||
| CAUCHY LPAREN expr COMMA expr COMMA expr (COMMA expr)? RPAREN # CauchyFunc
|
||||
| LOGNORMAL LPAREN expr COMMA expr COMMA expr (COMMA expr)? RPAREN # LogNormalFunc;
|
||||
|
||||
// LEXER RULES
|
||||
|
||||
SIN: 'sin';
|
||||
@@ -317,6 +318,7 @@ APPEND: 'append';
|
||||
GET_VALUE: 'get_value';
|
||||
BATCH_SHUFFLE: 'batch_shuffle' | 'shuffle' | 'select';
|
||||
CROP: 'crop';
|
||||
PAD: 'pad';
|
||||
ARGSORT: 'argsort';
|
||||
FOR: 'for';
|
||||
IN: 'in';
|
||||
|
||||
File diff suppressed because one or more lines are too long
@@ -103,45 +103,46 @@ APPEND=102
|
||||
GET_VALUE=103
|
||||
BATCH_SHUFFLE=104
|
||||
CROP=105
|
||||
ARGSORT=106
|
||||
FOR=107
|
||||
IN=108
|
||||
TIMESTAMP=109
|
||||
NONE=110
|
||||
BREAK=111
|
||||
CONTINUE=112
|
||||
TENSOR=113
|
||||
PLUS=114
|
||||
MINUS=115
|
||||
MULT=116
|
||||
DIV=117
|
||||
MOD=118
|
||||
POW=119
|
||||
GE=120
|
||||
GT=121
|
||||
LE=122
|
||||
LT=123
|
||||
EQ=124
|
||||
EQUEALS=125
|
||||
NE=126
|
||||
PIPE=127
|
||||
LPAREN=128
|
||||
RPAREN=129
|
||||
COMMA=130
|
||||
SEMICOLON=131
|
||||
ARROW=132
|
||||
LBRACKET=133
|
||||
RBRACKET=134
|
||||
QUESTION=135
|
||||
COLON=136
|
||||
LBRACE=137
|
||||
RBRACE=138
|
||||
NUMBER=139
|
||||
CONSTANT=140
|
||||
VARIABLE=141
|
||||
SL_COMMENT=142
|
||||
ML_COMMENT=143
|
||||
WS=144
|
||||
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
|
||||
'sin'=1
|
||||
'cos'=2
|
||||
'tan'=3
|
||||
@@ -173,8 +174,6 @@ WS=144
|
||||
'pow'=29
|
||||
'sigm'=30
|
||||
'clamp'=31
|
||||
'fft'=32
|
||||
'ifft'=33
|
||||
'angle'=34
|
||||
'print'=35
|
||||
'lerp'=38
|
||||
@@ -223,34 +222,35 @@ WS=144
|
||||
'append'=102
|
||||
'get_value'=103
|
||||
'crop'=105
|
||||
'argsort'=106
|
||||
'for'=107
|
||||
'in'=108
|
||||
'break'=111
|
||||
'continue'=112
|
||||
'tensor'=113
|
||||
'+'=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
|
||||
'pad'=106
|
||||
'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
|
||||
|
||||
File diff suppressed because one or more lines are too long
+588
-563
File diff suppressed because it is too large
Load Diff
@@ -103,45 +103,46 @@ APPEND=102
|
||||
GET_VALUE=103
|
||||
BATCH_SHUFFLE=104
|
||||
CROP=105
|
||||
ARGSORT=106
|
||||
FOR=107
|
||||
IN=108
|
||||
TIMESTAMP=109
|
||||
NONE=110
|
||||
BREAK=111
|
||||
CONTINUE=112
|
||||
TENSOR=113
|
||||
PLUS=114
|
||||
MINUS=115
|
||||
MULT=116
|
||||
DIV=117
|
||||
MOD=118
|
||||
POW=119
|
||||
GE=120
|
||||
GT=121
|
||||
LE=122
|
||||
LT=123
|
||||
EQ=124
|
||||
EQUEALS=125
|
||||
NE=126
|
||||
PIPE=127
|
||||
LPAREN=128
|
||||
RPAREN=129
|
||||
COMMA=130
|
||||
SEMICOLON=131
|
||||
ARROW=132
|
||||
LBRACKET=133
|
||||
RBRACKET=134
|
||||
QUESTION=135
|
||||
COLON=136
|
||||
LBRACE=137
|
||||
RBRACE=138
|
||||
NUMBER=139
|
||||
CONSTANT=140
|
||||
VARIABLE=141
|
||||
SL_COMMENT=142
|
||||
ML_COMMENT=143
|
||||
WS=144
|
||||
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
|
||||
'sin'=1
|
||||
'cos'=2
|
||||
'tan'=3
|
||||
@@ -173,8 +174,6 @@ WS=144
|
||||
'pow'=29
|
||||
'sigm'=30
|
||||
'clamp'=31
|
||||
'fft'=32
|
||||
'ifft'=33
|
||||
'angle'=34
|
||||
'print'=35
|
||||
'lerp'=38
|
||||
@@ -223,34 +222,35 @@ WS=144
|
||||
'append'=102
|
||||
'get_value'=103
|
||||
'crop'=105
|
||||
'argsort'=106
|
||||
'for'=107
|
||||
'in'=108
|
||||
'break'=111
|
||||
'continue'=112
|
||||
'tensor'=113
|
||||
'+'=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
|
||||
'pad'=106
|
||||
'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
|
||||
|
||||
+2141
-3259
File diff suppressed because it is too large
Load Diff
@@ -259,6 +259,11 @@ class MathExprVisitor(ParseTreeVisitor):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#FuncOptExp.
|
||||
def visitFuncOptExp(self, ctx:MathExprParser.FuncOptExpContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#VariableExp.
|
||||
def visitVariableExp(self, ctx:MathExprParser.VariableExpContext):
|
||||
return self.visitChildren(ctx)
|
||||
@@ -439,16 +444,6 @@ class MathExprVisitor(ParseTreeVisitor):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#sfftFunc.
|
||||
def visitSfftFunc(self, ctx:MathExprParser.SfftFuncContext):
|
||||
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#anglFunc.
|
||||
def visitAnglFunc(self, ctx:MathExprParser.AnglFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
@@ -519,16 +514,6 @@ class MathExprVisitor(ParseTreeVisitor):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#NoiseFunc.
|
||||
def visitNoiseFunc(self, ctx:MathExprParser.NoiseFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#RandFunc.
|
||||
def visitRandFunc(self, ctx:MathExprParser.RandFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#AnyFunc.
|
||||
def visitAnyFunc(self, ctx:MathExprParser.AnyFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
@@ -669,21 +654,6 @@ class MathExprVisitor(ParseTreeVisitor):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#ExponentialFunc.
|
||||
def visitExponentialFunc(self, ctx:MathExprParser.ExponentialFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#BernoulliFunc.
|
||||
def visitBernoulliFunc(self, ctx:MathExprParser.BernoulliFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#PoissonFunc.
|
||||
def visitPoissonFunc(self, ctx:MathExprParser.PoissonFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#GaussianFunc.
|
||||
def visitGaussianFunc(self, ctx:MathExprParser.GaussianFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
@@ -719,6 +689,11 @@ class MathExprVisitor(ParseTreeVisitor):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#PadFunc.
|
||||
def visitPadFunc(self, ctx:MathExprParser.PadFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#ClampFunc.
|
||||
def visitClampFunc(self, ctx:MathExprParser.ClampFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
@@ -744,16 +719,6 @@ class MathExprVisitor(ParseTreeVisitor):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#CauchyFunc.
|
||||
def visitCauchyFunc(self, ctx:MathExprParser.CauchyFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#LogNormalFunc.
|
||||
def visitLogNormalFunc(self, ctx:MathExprParser.LogNormalFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#CubicEaseFunc.
|
||||
def visitCubicEaseFunc(self, ctx:MathExprParser.CubicEaseFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
@@ -834,5 +799,45 @@ 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)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#RandFunc.
|
||||
def visitRandFunc(self, ctx:MathExprParser.RandFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#ExponentialFunc.
|
||||
def visitExponentialFunc(self, ctx:MathExprParser.ExponentialFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#BernoulliFunc.
|
||||
def visitBernoulliFunc(self, ctx:MathExprParser.BernoulliFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#PoissonFunc.
|
||||
def visitPoissonFunc(self, ctx:MathExprParser.PoissonFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#CauchyFunc.
|
||||
def visitCauchyFunc(self, ctx:MathExprParser.CauchyFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#LogNormalFunc.
|
||||
def visitLogNormalFunc(self, ctx:MathExprParser.LogNormalFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
|
||||
del MathExprParser
|
||||
@@ -3,6 +3,7 @@ import torch
|
||||
import math
|
||||
import inspect
|
||||
import torch.nn.functional as F
|
||||
from antlr4 import TerminalNode
|
||||
from .MathExprVisitor import MathExprVisitor
|
||||
from ..helper_functions import generate_dim_variables
|
||||
|
||||
@@ -976,7 +977,8 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
device = self.device
|
||||
|
||||
shape_to_use = self.shape if self.shape else (1, 1, 1, 1)
|
||||
|
||||
if len(ctx.expr()) > 1:
|
||||
shape_to_use = (yield ctx.expr(1))
|
||||
ndim = len(shape_to_use)
|
||||
dim_names = ["x", "y", "z", "w", "v", "u"]
|
||||
|
||||
@@ -1509,7 +1511,6 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
|
||||
return a + b
|
||||
def visitStart(self, ctx):
|
||||
from antlr4.tree.Tree import TerminalNode
|
||||
count = ctx.getChildCount()
|
||||
last_res = None
|
||||
|
||||
@@ -1752,33 +1753,45 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
|
||||
def visitNoiseFunc(self,ctx):
|
||||
seed_val = yield ctx.expr()
|
||||
shape_arg = self.shape;
|
||||
if len(ctx.expr(0)) > 1:
|
||||
shape_arg = (yield ctx.expr(1))
|
||||
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
||||
generator = torch.Generator(device=self.device).manual_seed(seed)
|
||||
return torch.randn(self.shape, generator=generator, device=self.device)
|
||||
return torch.randn(shape_arg, generator=generator, device=self.device)
|
||||
|
||||
def visitRandFunc(self, ctx):
|
||||
seed_val = yield ctx.expr()
|
||||
seed_val = yield ctx.expr(0)
|
||||
shape_arg = self.shape;
|
||||
if len(ctx.expr()) > 1:
|
||||
shape_arg = (yield ctx.expr(1))
|
||||
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
||||
generator = torch.Generator(device=self.device).manual_seed(seed)
|
||||
return torch.rand(self.shape, generator=generator, device=self.device)
|
||||
return torch.rand(shape_arg, generator=generator, device=self.device)
|
||||
|
||||
def visitExponentialFunc(self, ctx):
|
||||
seed_val = yield ctx.expr(0)
|
||||
shape_arg = self.shape;
|
||||
if len(ctx.expr()) > 2:
|
||||
shape_arg = (yield ctx.expr(2))
|
||||
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
||||
lambd_val = yield ctx.expr(1)
|
||||
lambd = float(lambd_val.item()) if self._is_tensor(lambd_val) else float(lambd_val)
|
||||
generator = torch.Generator(device=self.device).manual_seed(seed)
|
||||
return torch.empty(self.shape, device=self.device).exponential_(lambd, generator=generator)
|
||||
return torch.empty(shape_arg, device=self.device).exponential_(lambd, generator=generator)
|
||||
|
||||
def visitCauchyFunc(self, ctx):
|
||||
seed_val = yield ctx.expr(0)
|
||||
shape_arg = self.shape;
|
||||
if len(ctx.expr()) > 3:
|
||||
shape_arg = (yield ctx.expr(3))
|
||||
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
||||
median_val = yield ctx.expr(1)
|
||||
median = float(median_val.item()) if self._is_tensor(median_val) else float(median_val)
|
||||
sigma_val = yield ctx.expr(2)
|
||||
sigma = float(sigma_val.item()) if self._is_tensor(sigma_val) else float(sigma_val)
|
||||
generator = torch.Generator(device=self.device).manual_seed(seed)
|
||||
return torch.empty(self.shape, device=self.device).cauchy_(median, sigma, generator=generator)
|
||||
return torch.empty(shape_arg, device=self.device).cauchy_(median, sigma, generator=generator)
|
||||
|
||||
def visitLogNormalFunc(self, ctx):
|
||||
seed_val = yield ctx.expr(0)
|
||||
@@ -1787,26 +1800,35 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
mean = float(mean_val.item()) if self._is_tensor(mean_val) else float(mean_val)
|
||||
std_val = yield ctx.expr(2)
|
||||
std = float(std_val.item()) if self._is_tensor(std_val) else float(std_val)
|
||||
shape_arg = self.shape;
|
||||
if len(ctx.expr()) > 3:
|
||||
shape_arg = (yield ctx.expr(3))
|
||||
generator = torch.Generator(device=self.device).manual_seed(seed)
|
||||
return torch.empty(self.shape, device=self.device).log_normal_(mean, std, generator=generator)
|
||||
return torch.empty(shape_arg, device=self.device).log_normal_(mean, std, generator=generator)
|
||||
|
||||
def visitBernoulliFunc(self, ctx):
|
||||
seed_val = yield ctx.expr(0)
|
||||
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
||||
p = yield ctx.expr(1)
|
||||
generator = torch.Generator(device=self.device).manual_seed(seed)
|
||||
shape_arg = self.shape;
|
||||
if len(ctx.expr()) > 2:
|
||||
shape_arg = (yield ctx.expr(2))
|
||||
if self._is_tensor(p):
|
||||
return torch.bernoulli(p, generator=generator).to(device=self.device)
|
||||
return torch.bernoulli(torch.full(self.shape, p, device=self.device), generator=generator)
|
||||
return torch.bernoulli(torch.full(shape_arg, p, device=self.device), generator=generator)
|
||||
|
||||
def visitPoissonFunc(self, ctx):
|
||||
seed_val = yield ctx.expr(0)
|
||||
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
||||
lam = yield ctx.expr(1)
|
||||
generator = torch.Generator(device=self.device).manual_seed(seed)
|
||||
shape_arg = self.shape;
|
||||
if len(ctx.expr()) > 2:
|
||||
shape_arg = (yield ctx.expr(2))
|
||||
if self._is_tensor(lam):
|
||||
return torch.poisson(lam, generator=generator).to(device=self.device)
|
||||
return torch.poisson(torch.full(self.shape, lam, device=self.device), generator=generator)
|
||||
return torch.poisson(torch.full(shape_arg, lam, device=self.device), generator=generator)
|
||||
|
||||
def visitNvlFunc(self, ctx):
|
||||
v = yield ctx.expr(0)
|
||||
|
||||
Reference in New Issue
Block a user