(AI) set value based on position

This commit is contained in:
mcDandy
2026-02-01 15:21:32 +01:00
parent 4c5ef98571
commit 3467cf0663
10 changed files with 2578 additions and 2042 deletions
+2 -1
View File
@@ -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
+69 -61
View File
@@ -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
File diff suppressed because it is too large Load Diff
+69 -61
View File
@@ -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
+36
View File
@@ -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
File diff suppressed because it is too large Load Diff
+20
View File
@@ -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)
+61 -3
View File
@@ -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):