(AI) set value based on position
This commit is contained in:
@@ -6,7 +6,8 @@ start: (funcDef | varDef | stmt)* expr EOF;
|
||||
funcDef:
|
||||
VARIABLE LPAREN paramList? RPAREN ARROW (block | expr) SEMICOLON # FunctionDef;
|
||||
|
||||
varDef: VARIABLE EQUEALS expr SEMICOLON;
|
||||
varDef:
|
||||
VARIABLE (LBRACKET expr (COMMA expr)* RBRACKET)* EQUEALS expr SEMICOLON;
|
||||
|
||||
paramList: VARIABLE (COMMA VARIABLE)*;
|
||||
|
||||
|
||||
File diff suppressed because one or more lines are too long
@@ -99,40 +99,44 @@ FLIP=98
|
||||
COV=99
|
||||
SORT=100
|
||||
APPEND=101
|
||||
TIMESTAMP=102
|
||||
BREAK=103
|
||||
CONTINUE=104
|
||||
PLUS=105
|
||||
MINUS=106
|
||||
MULT=107
|
||||
DIV=108
|
||||
MOD=109
|
||||
POW=110
|
||||
GE=111
|
||||
GT=112
|
||||
LE=113
|
||||
LT=114
|
||||
EQ=115
|
||||
EQUEALS=116
|
||||
NE=117
|
||||
PIPE=118
|
||||
LPAREN=119
|
||||
RPAREN=120
|
||||
COMMA=121
|
||||
SEMICOLON=122
|
||||
ARROW=123
|
||||
LBRACKET=124
|
||||
RBRACKET=125
|
||||
QUESTION=126
|
||||
COLON=127
|
||||
LBRACE=128
|
||||
RBRACE=129
|
||||
NUMBER=130
|
||||
CONSTANT=131
|
||||
VARIABLE=132
|
||||
SL_COMMENT=133
|
||||
ML_COMMENT=134
|
||||
WS=135
|
||||
GET_VALUE=102
|
||||
CROP=103
|
||||
FOR=104
|
||||
IN=105
|
||||
TIMESTAMP=106
|
||||
BREAK=107
|
||||
CONTINUE=108
|
||||
PLUS=109
|
||||
MINUS=110
|
||||
MULT=111
|
||||
DIV=112
|
||||
MOD=113
|
||||
POW=114
|
||||
GE=115
|
||||
GT=116
|
||||
LE=117
|
||||
LT=118
|
||||
EQ=119
|
||||
EQUEALS=120
|
||||
NE=121
|
||||
PIPE=122
|
||||
LPAREN=123
|
||||
RPAREN=124
|
||||
COMMA=125
|
||||
SEMICOLON=126
|
||||
ARROW=127
|
||||
LBRACKET=128
|
||||
RBRACKET=129
|
||||
QUESTION=130
|
||||
COLON=131
|
||||
LBRACE=132
|
||||
RBRACE=133
|
||||
NUMBER=134
|
||||
CONSTANT=135
|
||||
VARIABLE=136
|
||||
SL_COMMENT=137
|
||||
ML_COMMENT=138
|
||||
WS=139
|
||||
'sin'=1
|
||||
'cos'=2
|
||||
'tan'=3
|
||||
@@ -212,30 +216,34 @@ WS=135
|
||||
'cov'=99
|
||||
'sort'=100
|
||||
'append'=101
|
||||
'break'=103
|
||||
'continue'=104
|
||||
'+'=105
|
||||
'-'=106
|
||||
'*'=107
|
||||
'/'=108
|
||||
'%'=109
|
||||
'^'=110
|
||||
'>='=111
|
||||
'>'=112
|
||||
'<='=113
|
||||
'<'=114
|
||||
'=='=115
|
||||
'='=116
|
||||
'!='=117
|
||||
'|'=118
|
||||
'('=119
|
||||
')'=120
|
||||
','=121
|
||||
';'=122
|
||||
'->'=123
|
||||
'['=124
|
||||
']'=125
|
||||
'?'=126
|
||||
':'=127
|
||||
'{'=128
|
||||
'}'=129
|
||||
'get_value'=102
|
||||
'crop'=103
|
||||
'for'=104
|
||||
'in'=105
|
||||
'break'=107
|
||||
'continue'=108
|
||||
'+'=109
|
||||
'-'=110
|
||||
'*'=111
|
||||
'/'=112
|
||||
'%'=113
|
||||
'^'=114
|
||||
'>='=115
|
||||
'>'=116
|
||||
'<='=117
|
||||
'<'=118
|
||||
'=='=119
|
||||
'='=120
|
||||
'!='=121
|
||||
'|'=122
|
||||
'('=123
|
||||
')'=124
|
||||
','=125
|
||||
';'=126
|
||||
'->'=127
|
||||
'['=128
|
||||
']'=129
|
||||
'?'=130
|
||||
':'=131
|
||||
'{'=132
|
||||
'}'=133
|
||||
|
||||
File diff suppressed because one or more lines are too long
+513
-497
File diff suppressed because it is too large
Load Diff
@@ -99,40 +99,44 @@ FLIP=98
|
||||
COV=99
|
||||
SORT=100
|
||||
APPEND=101
|
||||
TIMESTAMP=102
|
||||
BREAK=103
|
||||
CONTINUE=104
|
||||
PLUS=105
|
||||
MINUS=106
|
||||
MULT=107
|
||||
DIV=108
|
||||
MOD=109
|
||||
POW=110
|
||||
GE=111
|
||||
GT=112
|
||||
LE=113
|
||||
LT=114
|
||||
EQ=115
|
||||
EQUEALS=116
|
||||
NE=117
|
||||
PIPE=118
|
||||
LPAREN=119
|
||||
RPAREN=120
|
||||
COMMA=121
|
||||
SEMICOLON=122
|
||||
ARROW=123
|
||||
LBRACKET=124
|
||||
RBRACKET=125
|
||||
QUESTION=126
|
||||
COLON=127
|
||||
LBRACE=128
|
||||
RBRACE=129
|
||||
NUMBER=130
|
||||
CONSTANT=131
|
||||
VARIABLE=132
|
||||
SL_COMMENT=133
|
||||
ML_COMMENT=134
|
||||
WS=135
|
||||
GET_VALUE=102
|
||||
CROP=103
|
||||
FOR=104
|
||||
IN=105
|
||||
TIMESTAMP=106
|
||||
BREAK=107
|
||||
CONTINUE=108
|
||||
PLUS=109
|
||||
MINUS=110
|
||||
MULT=111
|
||||
DIV=112
|
||||
MOD=113
|
||||
POW=114
|
||||
GE=115
|
||||
GT=116
|
||||
LE=117
|
||||
LT=118
|
||||
EQ=119
|
||||
EQUEALS=120
|
||||
NE=121
|
||||
PIPE=122
|
||||
LPAREN=123
|
||||
RPAREN=124
|
||||
COMMA=125
|
||||
SEMICOLON=126
|
||||
ARROW=127
|
||||
LBRACKET=128
|
||||
RBRACKET=129
|
||||
QUESTION=130
|
||||
COLON=131
|
||||
LBRACE=132
|
||||
RBRACE=133
|
||||
NUMBER=134
|
||||
CONSTANT=135
|
||||
VARIABLE=136
|
||||
SL_COMMENT=137
|
||||
ML_COMMENT=138
|
||||
WS=139
|
||||
'sin'=1
|
||||
'cos'=2
|
||||
'tan'=3
|
||||
@@ -212,30 +216,34 @@ WS=135
|
||||
'cov'=99
|
||||
'sort'=100
|
||||
'append'=101
|
||||
'break'=103
|
||||
'continue'=104
|
||||
'+'=105
|
||||
'-'=106
|
||||
'*'=107
|
||||
'/'=108
|
||||
'%'=109
|
||||
'^'=110
|
||||
'>='=111
|
||||
'>'=112
|
||||
'<='=113
|
||||
'<'=114
|
||||
'=='=115
|
||||
'='=116
|
||||
'!='=117
|
||||
'|'=118
|
||||
'('=119
|
||||
')'=120
|
||||
','=121
|
||||
';'=122
|
||||
'->'=123
|
||||
'['=124
|
||||
']'=125
|
||||
'?'=126
|
||||
':'=127
|
||||
'{'=128
|
||||
'}'=129
|
||||
'get_value'=102
|
||||
'crop'=103
|
||||
'for'=104
|
||||
'in'=105
|
||||
'break'=107
|
||||
'continue'=108
|
||||
'+'=109
|
||||
'-'=110
|
||||
'*'=111
|
||||
'/'=112
|
||||
'%'=113
|
||||
'^'=114
|
||||
'>='=115
|
||||
'>'=116
|
||||
'<='=117
|
||||
'<'=118
|
||||
'=='=119
|
||||
'='=120
|
||||
'!='=121
|
||||
'|'=122
|
||||
'('=123
|
||||
')'=124
|
||||
','=125
|
||||
';'=126
|
||||
'->'=127
|
||||
'['=128
|
||||
']'=129
|
||||
'?'=130
|
||||
':'=131
|
||||
'{'=132
|
||||
'}'=133
|
||||
|
||||
@@ -107,6 +107,15 @@ class MathExprListener(ParseTreeListener):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#ForStatement.
|
||||
def enterForStatement(self, ctx:MathExprParser.ForStatementContext):
|
||||
pass
|
||||
|
||||
# Exit a parse tree produced by MathExprParser#ForStatement.
|
||||
def exitForStatement(self, ctx:MathExprParser.ForStatementContext):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#ExprStatement.
|
||||
def enterExprStatement(self, ctx:MathExprParser.ExprStatementContext):
|
||||
pass
|
||||
@@ -134,6 +143,15 @@ class MathExprListener(ParseTreeListener):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#forStmt.
|
||||
def enterForStmt(self, ctx:MathExprParser.ForStmtContext):
|
||||
pass
|
||||
|
||||
# Exit a parse tree produced by MathExprParser#forStmt.
|
||||
def exitForStmt(self, ctx:MathExprParser.ForStmtContext):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#block.
|
||||
def enterBlock(self, ctx:MathExprParser.BlockContext):
|
||||
pass
|
||||
@@ -1196,6 +1214,15 @@ class MathExprListener(ParseTreeListener):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#GetValueFunc.
|
||||
def enterGetValueFunc(self, ctx:MathExprParser.GetValueFuncContext):
|
||||
pass
|
||||
|
||||
# Exit a parse tree produced by MathExprParser#GetValueFunc.
|
||||
def exitGetValueFunc(self, ctx:MathExprParser.GetValueFuncContext):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#ClampFunc.
|
||||
def enterClampFunc(self, ctx:MathExprParser.ClampFuncContext):
|
||||
pass
|
||||
@@ -1295,6 +1322,15 @@ class MathExprListener(ParseTreeListener):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#CropFunc.
|
||||
def enterCropFunc(self, ctx:MathExprParser.CropFuncContext):
|
||||
pass
|
||||
|
||||
# Exit a parse tree produced by MathExprParser#CropFunc.
|
||||
def exitCropFunc(self, ctx:MathExprParser.CropFuncContext):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#SwapFunc.
|
||||
def enterSwapFunc(self, ctx:MathExprParser.SwapFuncContext):
|
||||
pass
|
||||
|
||||
+1785
-1417
File diff suppressed because it is too large
Load Diff
@@ -64,6 +64,11 @@ class MathExprVisitor(ParseTreeVisitor):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#ForStatement.
|
||||
def visitForStatement(self, ctx:MathExprParser.ForStatementContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#ExprStatement.
|
||||
def visitExprStatement(self, ctx:MathExprParser.ExprStatementContext):
|
||||
return self.visitChildren(ctx)
|
||||
@@ -79,6 +84,11 @@ class MathExprVisitor(ParseTreeVisitor):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#forStmt.
|
||||
def visitForStmt(self, ctx:MathExprParser.ForStmtContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#block.
|
||||
def visitBlock(self, ctx:MathExprParser.BlockContext):
|
||||
return self.visitChildren(ctx)
|
||||
@@ -669,6 +679,11 @@ class MathExprVisitor(ParseTreeVisitor):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#GetValueFunc.
|
||||
def visitGetValueFunc(self, ctx:MathExprParser.GetValueFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#ClampFunc.
|
||||
def visitClampFunc(self, ctx:MathExprParser.ClampFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
@@ -724,6 +739,11 @@ class MathExprVisitor(ParseTreeVisitor):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#CropFunc.
|
||||
def visitCropFunc(self, ctx:MathExprParser.CropFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#SwapFunc.
|
||||
def visitSwapFunc(self, ctx:MathExprParser.SwapFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
@@ -1538,9 +1538,67 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
|
||||
def visitVarDef(self, ctx):
|
||||
var_name = ctx.VARIABLE().getText()
|
||||
val = yield ctx.expr()
|
||||
self.variables[var_name] = val
|
||||
return val
|
||||
expr_list = ctx.expr()
|
||||
|
||||
if not ctx.LBRACKET():
|
||||
# Standard assignment: x = value
|
||||
val = yield expr_list[0]
|
||||
self.variables[var_name] = val
|
||||
return val
|
||||
|
||||
# Indexed assignment: x[i, j...] = value
|
||||
# The last expression is the value to assign
|
||||
val_expr = expr_list[-1]
|
||||
assigned_val = yield val_expr
|
||||
|
||||
# Evaluate indices
|
||||
indices = []
|
||||
for i in range(len(expr_list) - 1):
|
||||
indices.append((yield expr_list[i]))
|
||||
|
||||
if var_name not in self.variables:
|
||||
raise ValueError(f"Variable '{var_name}' not found for indexed assignment.")
|
||||
|
||||
target = self.variables[var_name]
|
||||
|
||||
if self._is_tensor(target):
|
||||
# Process indices for PyTorch
|
||||
torch_indices = []
|
||||
for idx in indices:
|
||||
if self._is_list(idx):
|
||||
torch_indices.append(torch.tensor(idx, device=self.device, dtype=torch.long))
|
||||
elif self._is_tensor(idx):
|
||||
torch_indices.append(idx.long())
|
||||
else:
|
||||
torch_indices.append(int(idx))
|
||||
|
||||
idx_tuple = tuple(torch_indices)
|
||||
val_t = self._promote_to_tensor(assigned_val)
|
||||
|
||||
try:
|
||||
# Target slice - used to compute expected shape
|
||||
target_slice = target[idx_tuple]
|
||||
|
||||
# Squeeze leading ones to match target slice rank if it's smaller
|
||||
# but target_slice.ndim might be 0 if it's a scalar location.
|
||||
while val_t.ndim > target_slice.ndim and val_t.shape[0] == 1:
|
||||
val_t = val_t.squeeze(0)
|
||||
|
||||
target[idx_tuple] = val_t
|
||||
return assigned_val
|
||||
except Exception as e:
|
||||
raise ValueError(f"Indexed assignment to '{var_name}' failed: {str(e)}")
|
||||
|
||||
elif self._is_list(target):
|
||||
# Recurse through nested lists if multiple indices provided
|
||||
curr = target
|
||||
for idx in indices[:-1]:
|
||||
curr = curr[int(idx + len(curr) if idx < 0 else idx)]
|
||||
last_idx = int(indices[-1])
|
||||
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)}")
|
||||
|
||||
|
||||
def visitFunctionDef(self, ctx):
|
||||
|
||||
Reference in New Issue
Block a user