Better errors + batch shuffle (subject to rename)

This commit is contained in:
mcDandy
2026-02-03 13:59:19 +01:00
parent b06e26b82e
commit 38e795eb35
10 changed files with 1612 additions and 1532 deletions
+2 -1
View File
@@ -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
+71 -70
View File
@@ -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
File diff suppressed because it is too large Load Diff
+71 -70
View File
@@ -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
+9
View File
@@ -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
File diff suppressed because it is too large Load Diff
+5
View File
@@ -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)
+50 -21
View File
@@ -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)