add optional shape argument

This commit is contained in:
mcDandy
2026-02-07 16:37:39 +01:00
parent 74e5859366
commit b8b98cb758
9 changed files with 3033 additions and 4091 deletions
+66 -64
View File
@@ -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
+72 -72
View File
@@ -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
File diff suppressed because it is too large Load Diff
+72 -72
View File
@@ -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 it is too large Load Diff
+50 -45
View File
@@ -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
+32 -10
View File
@@ -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)