Better errors + batch shuffle (subject to rename)
This commit is contained in:
@@ -168,7 +168,7 @@ func2:
|
||||
| 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
|
||||
| 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;
|
||||
@@ -314,6 +314,7 @@ SORT: 'sort';
|
||||
COUNT: 'count' | 'length' | 'cnt';
|
||||
APPEND: 'append';
|
||||
GET_VALUE: 'get_value';
|
||||
BATCH_SHUFFLE: 'batch_shuffle' | 'shuffle' | 'select';
|
||||
CROP: 'crop';
|
||||
FOR: 'for';
|
||||
IN: 'in';
|
||||
|
||||
File diff suppressed because one or more lines are too long
@@ -101,45 +101,46 @@ SORT=100
|
||||
COUNT=101
|
||||
APPEND=102
|
||||
GET_VALUE=103
|
||||
CROP=104
|
||||
FOR=105
|
||||
IN=106
|
||||
TIMESTAMP=107
|
||||
NONE=108
|
||||
BREAK=109
|
||||
CONTINUE=110
|
||||
TENSOR=111
|
||||
PLUS=112
|
||||
MINUS=113
|
||||
MULT=114
|
||||
DIV=115
|
||||
MOD=116
|
||||
POW=117
|
||||
GE=118
|
||||
GT=119
|
||||
LE=120
|
||||
LT=121
|
||||
EQ=122
|
||||
EQUEALS=123
|
||||
NE=124
|
||||
PIPE=125
|
||||
LPAREN=126
|
||||
RPAREN=127
|
||||
COMMA=128
|
||||
SEMICOLON=129
|
||||
ARROW=130
|
||||
LBRACKET=131
|
||||
RBRACKET=132
|
||||
QUESTION=133
|
||||
COLON=134
|
||||
LBRACE=135
|
||||
RBRACE=136
|
||||
NUMBER=137
|
||||
CONSTANT=138
|
||||
VARIABLE=139
|
||||
SL_COMMENT=140
|
||||
ML_COMMENT=141
|
||||
WS=142
|
||||
BATCH_SHUFFLE=104
|
||||
CROP=105
|
||||
FOR=106
|
||||
IN=107
|
||||
TIMESTAMP=108
|
||||
NONE=109
|
||||
BREAK=110
|
||||
CONTINUE=111
|
||||
TENSOR=112
|
||||
PLUS=113
|
||||
MINUS=114
|
||||
MULT=115
|
||||
DIV=116
|
||||
MOD=117
|
||||
POW=118
|
||||
GE=119
|
||||
GT=120
|
||||
LE=121
|
||||
LT=122
|
||||
EQ=123
|
||||
EQUEALS=124
|
||||
NE=125
|
||||
PIPE=126
|
||||
LPAREN=127
|
||||
RPAREN=128
|
||||
COMMA=129
|
||||
SEMICOLON=130
|
||||
ARROW=131
|
||||
LBRACKET=132
|
||||
RBRACKET=133
|
||||
QUESTION=134
|
||||
COLON=135
|
||||
LBRACE=136
|
||||
RBRACE=137
|
||||
NUMBER=138
|
||||
CONSTANT=139
|
||||
VARIABLE=140
|
||||
SL_COMMENT=141
|
||||
ML_COMMENT=142
|
||||
WS=143
|
||||
'sin'=1
|
||||
'cos'=2
|
||||
'tan'=3
|
||||
@@ -220,34 +221,34 @@ WS=142
|
||||
'sort'=100
|
||||
'append'=102
|
||||
'get_value'=103
|
||||
'crop'=104
|
||||
'for'=105
|
||||
'in'=106
|
||||
'break'=109
|
||||
'continue'=110
|
||||
'tensor'=111
|
||||
'+'=112
|
||||
'-'=113
|
||||
'*'=114
|
||||
'/'=115
|
||||
'%'=116
|
||||
'^'=117
|
||||
'>='=118
|
||||
'>'=119
|
||||
'<='=120
|
||||
'<'=121
|
||||
'=='=122
|
||||
'='=123
|
||||
'!='=124
|
||||
'|'=125
|
||||
'('=126
|
||||
')'=127
|
||||
','=128
|
||||
';'=129
|
||||
'->'=130
|
||||
'['=131
|
||||
']'=132
|
||||
'?'=133
|
||||
':'=134
|
||||
'{'=135
|
||||
'}'=136
|
||||
'crop'=105
|
||||
'for'=106
|
||||
'in'=107
|
||||
'break'=110
|
||||
'continue'=111
|
||||
'tensor'=112
|
||||
'+'=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
|
||||
|
||||
File diff suppressed because one or more lines are too long
+549
-535
File diff suppressed because it is too large
Load Diff
@@ -101,45 +101,46 @@ SORT=100
|
||||
COUNT=101
|
||||
APPEND=102
|
||||
GET_VALUE=103
|
||||
CROP=104
|
||||
FOR=105
|
||||
IN=106
|
||||
TIMESTAMP=107
|
||||
NONE=108
|
||||
BREAK=109
|
||||
CONTINUE=110
|
||||
TENSOR=111
|
||||
PLUS=112
|
||||
MINUS=113
|
||||
MULT=114
|
||||
DIV=115
|
||||
MOD=116
|
||||
POW=117
|
||||
GE=118
|
||||
GT=119
|
||||
LE=120
|
||||
LT=121
|
||||
EQ=122
|
||||
EQUEALS=123
|
||||
NE=124
|
||||
PIPE=125
|
||||
LPAREN=126
|
||||
RPAREN=127
|
||||
COMMA=128
|
||||
SEMICOLON=129
|
||||
ARROW=130
|
||||
LBRACKET=131
|
||||
RBRACKET=132
|
||||
QUESTION=133
|
||||
COLON=134
|
||||
LBRACE=135
|
||||
RBRACE=136
|
||||
NUMBER=137
|
||||
CONSTANT=138
|
||||
VARIABLE=139
|
||||
SL_COMMENT=140
|
||||
ML_COMMENT=141
|
||||
WS=142
|
||||
BATCH_SHUFFLE=104
|
||||
CROP=105
|
||||
FOR=106
|
||||
IN=107
|
||||
TIMESTAMP=108
|
||||
NONE=109
|
||||
BREAK=110
|
||||
CONTINUE=111
|
||||
TENSOR=112
|
||||
PLUS=113
|
||||
MINUS=114
|
||||
MULT=115
|
||||
DIV=116
|
||||
MOD=117
|
||||
POW=118
|
||||
GE=119
|
||||
GT=120
|
||||
LE=121
|
||||
LT=122
|
||||
EQ=123
|
||||
EQUEALS=124
|
||||
NE=125
|
||||
PIPE=126
|
||||
LPAREN=127
|
||||
RPAREN=128
|
||||
COMMA=129
|
||||
SEMICOLON=130
|
||||
ARROW=131
|
||||
LBRACKET=132
|
||||
RBRACKET=133
|
||||
QUESTION=134
|
||||
COLON=135
|
||||
LBRACE=136
|
||||
RBRACE=137
|
||||
NUMBER=138
|
||||
CONSTANT=139
|
||||
VARIABLE=140
|
||||
SL_COMMENT=141
|
||||
ML_COMMENT=142
|
||||
WS=143
|
||||
'sin'=1
|
||||
'cos'=2
|
||||
'tan'=3
|
||||
@@ -220,34 +221,34 @@ WS=142
|
||||
'sort'=100
|
||||
'append'=102
|
||||
'get_value'=103
|
||||
'crop'=104
|
||||
'for'=105
|
||||
'in'=106
|
||||
'break'=109
|
||||
'continue'=110
|
||||
'tensor'=111
|
||||
'+'=112
|
||||
'-'=113
|
||||
'*'=114
|
||||
'/'=115
|
||||
'%'=116
|
||||
'^'=117
|
||||
'>='=118
|
||||
'>'=119
|
||||
'<='=120
|
||||
'<'=121
|
||||
'=='=122
|
||||
'='=123
|
||||
'!='=124
|
||||
'|'=125
|
||||
'('=126
|
||||
')'=127
|
||||
','=128
|
||||
';'=129
|
||||
'->'=130
|
||||
'['=131
|
||||
']'=132
|
||||
'?'=133
|
||||
':'=134
|
||||
'{'=135
|
||||
'}'=136
|
||||
'crop'=105
|
||||
'for'=106
|
||||
'in'=107
|
||||
'break'=110
|
||||
'continue'=111
|
||||
'tensor'=112
|
||||
'+'=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
|
||||
|
||||
@@ -1241,6 +1241,15 @@ class MathExprListener(ParseTreeListener):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#BatchShuffleFunc.
|
||||
def enterBatchShuffleFunc(self, ctx:MathExprParser.BatchShuffleFuncContext):
|
||||
pass
|
||||
|
||||
# Exit a parse tree produced by MathExprParser#BatchShuffleFunc.
|
||||
def exitBatchShuffleFunc(self, ctx:MathExprParser.BatchShuffleFuncContext):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#PushFunc.
|
||||
def enterPushFunc(self, ctx:MathExprParser.PushFuncContext):
|
||||
pass
|
||||
|
||||
+848
-833
File diff suppressed because it is too large
Load Diff
@@ -694,6 +694,11 @@ class MathExprVisitor(ParseTreeVisitor):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#BatchShuffleFunc.
|
||||
def visitBatchShuffleFunc(self, ctx:MathExprParser.BatchShuffleFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#PushFunc.
|
||||
def visitPushFunc(self, ctx:MathExprParser.PushFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
@@ -179,7 +179,7 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
if var_name in self.variables:
|
||||
res = self.variables[var_name]
|
||||
return res
|
||||
raise ValueError(f"line {ctx.VARIABLE().getPayload().line}:{ctx.VARIABLE().getPayload().column}: Variable '{var_name}' not found")
|
||||
raise ValueError(f"{ctx.VARIABLE().getPayload().line}:{ctx.VARIABLE().getPayload().column}: Variable '{var_name}' not found")
|
||||
|
||||
def visitListExp(self, ctx):
|
||||
res = []
|
||||
@@ -249,7 +249,7 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
else:
|
||||
return val[int(idx + len(val) if idx < 0 else idx)]
|
||||
else:
|
||||
raise ValueError("Indexing only supported on tensors and lists.")
|
||||
raise ValueError(f"{ctx.VARIABLE().getPayload().line}:{ctx.VARIABLE().getPayload().column}: Indexing only supported on tensors and lists.")
|
||||
|
||||
def visitToAtom(self, ctx):
|
||||
return (yield ctx.atom())
|
||||
@@ -598,7 +598,7 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
pos_list = yield ctx.expr(1)
|
||||
|
||||
if not self._is_tensor(var):
|
||||
raise ValueError("get_value expects a tensor as first argument")
|
||||
raise ValueError(f"{ctx.VARIABLE().getPayload().line}:{ctx.VARIABLE().getPayload().column}: get_value expects a tensor as first argument")
|
||||
|
||||
if not self._is_list(pos_list) and not self._is_tensor(pos_list):
|
||||
pos_list = [pos_list]
|
||||
@@ -607,7 +607,7 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
pos_list = pos_list.tolist()
|
||||
|
||||
if len(pos_list) != var.ndim:
|
||||
raise ValueError(f"Position list length {len(pos_list)} does not match tensor dimensions {var.ndim}")
|
||||
raise ValueError(f"{ctx.VARIABLE().getPayload().line}:{ctx.VARIABLE().getPayload().column}: Position list length {len(pos_list)} does not match tensor dimensions {var.ndim}")
|
||||
|
||||
shape = var.shape
|
||||
c_strides = [1] * var.ndim
|
||||
@@ -619,11 +619,31 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
for i, p in enumerate(pos_list):
|
||||
idx = int(p)
|
||||
if idx < 0 or idx >= shape[i]:
|
||||
raise ValueError(f"Index {idx} out of bounds for dimension {i} with size {shape[i]}")
|
||||
raise ValueError(f"{ctx.VARIABLE().getPayload().line}:{ctx.VARIABLE().getPayload().column}: Index {idx} out of bounds for dimension {i} with size {shape[i]}")
|
||||
offset += idx * c_strides[i]
|
||||
|
||||
return var.contiguous().flatten()[offset]
|
||||
|
||||
def visitBatchShuffleFunc(self, ctx):
|
||||
tsr_val = yield ctx.expr(0)
|
||||
idx_val = yield ctx.expr(1)
|
||||
|
||||
tsr = self._promote_to_tensor(tsr_val)
|
||||
|
||||
if self._is_tensor(idx_val):
|
||||
indices = idx_val.long()
|
||||
elif self._is_list(idx_val):
|
||||
indices = torch.tensor([int(float(x)) for x in idx_val], dtype=torch.long, device=tsr.device)
|
||||
else:
|
||||
indices = torch.tensor([int(float(idx_val))], dtype=torch.long, device=tsr.device)
|
||||
|
||||
# Check bounds
|
||||
max_idx = tsr.size(0)
|
||||
if torch.any(indices < 0) or torch.any(indices >= max_idx):
|
||||
raise ValueError(f"{ctx.VARIABLE().getPayload().line}:{ctx.VARIABLE().getPayload().column}: Batch index out of bounds (0-{max_idx-1})")
|
||||
|
||||
return tsr[indices]
|
||||
|
||||
# Three-argument functions
|
||||
def visitClampFunc(self, ctx):
|
||||
val = (yield ctx.expr(0))
|
||||
@@ -711,8 +731,8 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
# Let's enforce or just take first N?
|
||||
# For robustness, we'll assume user provides correct dims or we raise error?
|
||||
# Given "lists described in get_value" implies stricter checking.
|
||||
if len(p_l) != inp.ndim: raise ValueError(f"crop: position dim {len(p_l)} != input dim {inp.ndim}")
|
||||
if len(s_l) != inp.ndim: raise ValueError(f"crop: size dim {len(s_l)} != input dim {inp.ndim}")
|
||||
if len(p_l) != inp.ndim: raise ValueError(f"{ctx.VARIABLE().getPayload().line}:{ctx.VARIABLE().getPayload().column}: crop: position dim {len(p_l)} != input dim {inp.ndim}")
|
||||
if len(s_l) != inp.ndim: raise ValueError(f"{ctx.VARIABLE().getPayload().line}:{ctx.VARIABLE().getPayload().column}: crop: size dim {len(s_l)} != input dim {inp.ndim}")
|
||||
|
||||
out_tensor = torch.zeros(tuple(s_l), dtype=inp.dtype, device=inp.device)
|
||||
|
||||
@@ -1149,7 +1169,7 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
y_flat = y.flatten()
|
||||
|
||||
if x_flat.numel() != y_flat.numel():
|
||||
raise ValueError("x and y must have the same number of elements")
|
||||
raise ValueError(f"{ctx.VARIABLE().getPayload().line}:{ctx.VARIABLE().getPayload().column}: x and y must have the same number of elements")
|
||||
|
||||
n = x_flat.numel()
|
||||
if n < 2:
|
||||
@@ -1171,7 +1191,7 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
if num_coords == 0:
|
||||
return tensor
|
||||
if num_coords > 3:
|
||||
raise ValueError("map() supports max 3 mapping functions.")
|
||||
raise ValueError(f"{ctx.VARIABLE().getPayload().line}:{ctx.VARIABLE().getPayload().column}: map() supports max 3 mapping functions.")
|
||||
|
||||
spatial_in_shape = tensor.shape[-num_coords:]
|
||||
leading_shape = tensor.shape[:-num_coords]
|
||||
@@ -1219,13 +1239,13 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
tensor = self._promote_to_tensor(input_raw)
|
||||
num_args = len(ctx.expr())
|
||||
if num_args < 3:
|
||||
raise ValueError("conv() requires at least 3 arguments")
|
||||
raise ValueError(f"{ctx.VARIABLE().getPayload().line}:{ctx.VARIABLE().getPayload().column}: conv() requires at least 3 arguments")
|
||||
|
||||
kernel_arg_idx = num_args - 1
|
||||
|
||||
spatial_dims_count = num_args - 2
|
||||
if spatial_dims_count not in [1, 2, 3]:
|
||||
raise ValueError(f"conv() supports 1D, 2D, or 3D. Found {spatial_dims_count}")
|
||||
raise ValueError(f"{ctx.VARIABLE().getPayload().line}:{ctx.VARIABLE().getPayload().column}: conv() supports 1D, 2D, or 3D. Found {spatial_dims_count}")
|
||||
|
||||
kernel_sizes = []
|
||||
for i in range(1, 1 + spatial_dims_count):
|
||||
@@ -1371,7 +1391,7 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
# Must have at least Channel + Spatial dims
|
||||
min_dims = spatial_dims_count + 1
|
||||
if tensor.ndim < min_dims:
|
||||
raise ValueError(f"convolution() input requires at least Channels + Spatial dimensions. Got shape {tensor.shape} for {spatial_dims_count}D conv.")
|
||||
raise ValueError(f"{ctx.VARIABLE().getPayload().line}:{ctx.VARIABLE().getPayload().column}: convolution() input requires at least Channels + Spatial dimensions. Got shape {tensor.shape} for {spatial_dims_count}D conv.")
|
||||
|
||||
spatial_shape = tensor.shape[-spatial_dims_count:]
|
||||
in_channels = tensor.shape[-(spatial_dims_count + 1)]
|
||||
@@ -1574,7 +1594,7 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
indices.append((yield expr_list[i]))
|
||||
|
||||
if var_name not in self.variables:
|
||||
raise ValueError(f"Variable '{var_name}' not found for indexed assignment.")
|
||||
raise ValueError(f"{ctx.VARIABLE().getPayload().line}:{ctx.VARIABLE().getPayload().column}: Variable '{var_name}' not found for indexed assignment.")
|
||||
|
||||
target = self.variables[var_name]
|
||||
|
||||
@@ -1604,7 +1624,7 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
target[idx_tuple] = val_t
|
||||
return assigned_val
|
||||
except Exception as e:
|
||||
raise ValueError(f"Indexed assignment to '{var_name}' failed: {str(e)}")
|
||||
raise ValueError(f"{ctx.VARIABLE().getPayload().line}:{ctx.VARIABLE().getPayload().column}: Indexed assignment to '{var_name}' failed: {str(e)}")
|
||||
|
||||
elif self._is_list(target):
|
||||
# Recurse through nested lists if multiple indices provided
|
||||
@@ -1615,7 +1635,7 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
curr[last_idx + len(curr) if last_idx < 0 else last_idx] = assigned_val
|
||||
return assigned_val
|
||||
else:
|
||||
raise ValueError(f"Indexed assignment not supported for {type(target)}")
|
||||
raise ValueError(f"{ctx.VARIABLE().getPayload().line}:{ctx.VARIABLE().getPayload().column}: Indexed assignment not supported for {type(target)}")
|
||||
|
||||
|
||||
def visitFunctionDef(self, ctx):
|
||||
@@ -1645,8 +1665,7 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
args.append((yield e))
|
||||
|
||||
if len(args) != len(params):
|
||||
raise ValueError(
|
||||
f"Function '{func_name}' expects {len(params)} arguments, got {len(args)}"
|
||||
raise ValueError(f"{ctx.VARIABLE().getPayload().line}:{ctx.VARIABLE().getPayload().column}: Function '{func_name}' expects {len(params)} arguments, got {len(args)}"
|
||||
)
|
||||
|
||||
# Create a new scope for function execution
|
||||
@@ -1675,7 +1694,7 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
self.variables = self._scope_stack.pop()
|
||||
self.depth -= 1
|
||||
|
||||
raise ValueError(f"Unknown function: {func_name}")
|
||||
raise ValueError(f"{ctx.VARIABLE().getPayload().line}:{ctx.VARIABLE().getPayload().column}: Unknown function: {func_name}")
|
||||
|
||||
def visitNoiseFunc(self,ctx):
|
||||
seed_val = yield ctx.expr()
|
||||
@@ -1931,18 +1950,20 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
self._state_storage[slot] = []
|
||||
value = yield ctx.expr(1)
|
||||
self._state_storage[slot].append(value)
|
||||
print(self._state_storage.keys())
|
||||
return value
|
||||
|
||||
def visitPopFunc(self, ctx):
|
||||
self._ensure_dict_storage()
|
||||
slot = int((yield ctx.expr()))
|
||||
if slot not in self._state_storage or not self._state_storage[slot]:
|
||||
raise ValueError(f"Pop from empty slot: {slot}")
|
||||
raise ValueError(f"{ctx.VARIABLE().getPayload().line}:{ctx.VARIABLE().getPayload().column}: Pop from empty slot: {slot}")
|
||||
return self._state_storage[slot].pop()
|
||||
|
||||
def visitClearFunc(self, ctx):
|
||||
self._ensure_dict_storage()
|
||||
slot = int((yield ctx.expr()))
|
||||
print(self._state_storage.keys())
|
||||
if slot in self._state_storage:
|
||||
self._state_storage[slot] = []
|
||||
return None
|
||||
@@ -1956,7 +1977,7 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
self._ensure_dict_storage()
|
||||
slot = int((yield ctx.expr()))
|
||||
if slot not in self._state_storage:
|
||||
raise ValueError(f"Get from empty slot: {slot}")
|
||||
raise ValueError(f"{ctx.VARIABLE().getPayload().line}:{ctx.VARIABLE().getPayload().column}: Get from empty slot: {slot}")
|
||||
storage_list = self._state_storage[slot]
|
||||
return storage_list[-1] if storage_list else None
|
||||
|
||||
@@ -1968,5 +1989,13 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
|
||||
def visitEmptyTensorFunc(self, ctx):
|
||||
value = (yield ctx.expr()) if ctx.expr() else 0.0
|
||||
shape = yield ctx.indexExpr()
|
||||
shape_val = yield ctx.indexExpr()
|
||||
|
||||
if self._is_tensor(shape_val):
|
||||
shape = shape_val.int().tolist()
|
||||
elif self._is_list(shape_val):
|
||||
shape = [int(float(x)) for x in shape_val]
|
||||
else:
|
||||
shape = [int(float(shape_val))]
|
||||
|
||||
return torch.full(shape, value, device=self.device)
|
||||
|
||||
Reference in New Issue
Block a user