fixed an issue when (probably) multiplying list and tensor

This commit is contained in:
mcDandy
2026-01-30 21:22:58 +01:00
parent 9e2e4d4db6
commit e13e3ac013
9 changed files with 1658 additions and 60 deletions
+7 -7
View File
@@ -23,9 +23,9 @@ stmt:
ifStmt: IF LPAREN expr RPAREN stmt (ELSE stmt)?;
whileStmt: WHILE LPAREN expr RPAREN stmt;
block: LBRACE stmt* RBRACE;
breakStmt: BRK SEMICOLON;
continueStmt: CONT SEMICOLON;
returnStmt: RET expr? SEMICOLON;
breakStmt: BREAK SEMICOLON;
continueStmt: CONTINUE SEMICOLON;
returnStmt: RETURN expr? SEMICOLON;
expr: ternaryExpr | atom | compExpr;
@@ -281,7 +281,7 @@ REMAP: 'remap';
IF: 'if';
ELSE: 'else';
WHILE: 'while';
RET: 'return';
RETURN: 'return';
PUSH: 'push';
POP: 'pop';
CLEAR: 'clear';
@@ -303,8 +303,8 @@ SORT: 'sort';
APPEND: 'append';
TIMESTAMP: 'timestamp' | 'now';
BRK: 'break';
CONT: 'continue';
BREAK: 'break';
CONTINUE: 'continue';
PLUS: '+';
MINUS: '-';
@@ -339,4 +339,4 @@ VARIABLE: [a-zA-Z_] [a-zA-Z_0-9]*;
SL_COMMENT : '#' ~[\r\n]* -> skip;
ML_COMMENT : '/*' .*? '*/' -> skip;
WS: [ \t\r\n]+ -> skip;
WS: [ \t\r\n]+ -> skip;
+3 -3
View File
@@ -221,7 +221,7 @@ REMAP
IF
ELSE
WHILE
RET
RETURN
PUSH
POP
CLEAR
@@ -240,8 +240,8 @@ COV
SORT
APPEND
TIMESTAMP
BRK
CONT
BREAK
CONTINUE
PLUS
MINUS
MULT
+3 -3
View File
@@ -81,7 +81,7 @@ REMAP=80
IF=81
ELSE=82
WHILE=83
RET=84
RETURN=84
PUSH=85
POP=86
CLEAR=87
@@ -100,8 +100,8 @@ COV=99
SORT=100
APPEND=101
TIMESTAMP=102
BRK=103
CONT=104
BREAK=103
CONTINUE=104
PLUS=105
MINUS=106
MULT=107
+6 -6
View File
@@ -221,7 +221,7 @@ REMAP
IF
ELSE
WHILE
RET
RETURN
PUSH
POP
CLEAR
@@ -240,8 +240,8 @@ COV
SORT
APPEND
TIMESTAMP
BRK
CONT
BREAK
CONTINUE
PLUS
MINUS
MULT
@@ -358,7 +358,7 @@ REMAP
IF
ELSE
WHILE
RET
RETURN
PUSH
POP
CLEAR
@@ -377,8 +377,8 @@ COV
SORT
APPEND
TIMESTAMP
BRK
CONT
BREAK
CONTINUE
PLUS
MINUS
MULT
+14 -14
View File
@@ -562,7 +562,7 @@ class MathExprLexer(Lexer):
IF = 81
ELSE = 82
WHILE = 83
RET = 84
RETURN = 84
PUSH = 85
POP = 86
CLEAR = 87
@@ -581,8 +581,8 @@ class MathExprLexer(Lexer):
SORT = 100
APPEND = 101
TIMESTAMP = 102
BRK = 103
CONT = 104
BREAK = 103
CONTINUE = 104
PLUS = 105
MINUS = 106
MULT = 107
@@ -650,15 +650,15 @@ class MathExprLexer(Lexer):
"PERCENTILE", "QUANTILE", "DOT", "MOMENT", "ANY", "ALL", "EDGE",
"GAUSSIAN", "MEDIAN", "MODE", "CUMSUM", "CUMPROD", "TOPK_IND",
"BOTK_IND", "CUBIC_EASE", "ELASTIC_EASE", "SINE_EASE", "SMOOTHERSTEP",
"DIST", "REMAP", "IF", "ELSE", "WHILE", "RET", "PUSH", "POP",
"DIST", "REMAP", "IF", "ELSE", "WHILE", "RETURN", "PUSH", "POP",
"CLEAR", "HAS", "GET", "NOISE", "RAND", "CAUCHY", "EXPONENTIAL",
"LOGNORMAL", "BERNOULLI", "POISSON", "COSSIM", "FLIP", "COV",
"SORT", "APPEND", "TIMESTAMP", "BRK", "CONT", "PLUS", "MINUS",
"MULT", "DIV", "MOD", "POW", "GE", "GT", "LE", "LT", "EQ", "EQUEALS",
"NE", "PIPE", "LPAREN", "RPAREN", "COMMA", "SEMICOLON", "ARROW",
"LBRACKET", "RBRACKET", "QUESTION", "COLON", "LBRACE", "RBRACE",
"CONSTANT", "NUMBER", "VARIABLE", "SL_COMMENT", "ML_COMMENT",
"WS" ]
"SORT", "APPEND", "TIMESTAMP", "BREAK", "CONTINUE", "PLUS",
"MINUS", "MULT", "DIV", "MOD", "POW", "GE", "GT", "LE", "LT",
"EQ", "EQUEALS", "NE", "PIPE", "LPAREN", "RPAREN", "COMMA",
"SEMICOLON", "ARROW", "LBRACKET", "RBRACKET", "QUESTION", "COLON",
"LBRACE", "RBRACE", "CONSTANT", "NUMBER", "VARIABLE", "SL_COMMENT",
"ML_COMMENT", "WS" ]
ruleNames = [ "SIN", "COS", "TAN", "ASIN", "ACOS", "ATAN", "ATAN2",
"SINH", "COSH", "TANH", "ASINH", "ACOSH", "ATANH", "ABS",
@@ -672,12 +672,12 @@ class MathExprLexer(Lexer):
"DOT", "MOMENT", "ANY", "ALL", "EDGE", "GAUSSIAN", "MEDIAN",
"MODE", "CUMSUM", "CUMPROD", "TOPK_IND", "BOTK_IND", "CUBIC_EASE",
"ELASTIC_EASE", "SINE_EASE", "SMOOTHERSTEP", "DIST", "REMAP",
"IF", "ELSE", "WHILE", "RET", "PUSH", "POP", "CLEAR",
"IF", "ELSE", "WHILE", "RETURN", "PUSH", "POP", "CLEAR",
"HAS", "GET", "NOISE", "RAND", "CAUCHY", "EXPONENTIAL",
"LOGNORMAL", "BERNOULLI", "POISSON", "COSSIM", "FLIP",
"COV", "SORT", "APPEND", "TIMESTAMP", "BRK", "CONT", "PLUS",
"MINUS", "MULT", "DIV", "MOD", "POW", "GE", "GT", "LE",
"LT", "EQ", "EQUEALS", "NE", "PIPE", "LPAREN", "RPAREN",
"COV", "SORT", "APPEND", "TIMESTAMP", "BREAK", "CONTINUE",
"PLUS", "MINUS", "MULT", "DIV", "MOD", "POW", "GE", "GT",
"LE", "LT", "EQ", "EQUEALS", "NE", "PIPE", "LPAREN", "RPAREN",
"COMMA", "SEMICOLON", "ARROW", "LBRACKET", "RBRACKET",
"QUESTION", "COLON", "LBRACE", "RBRACE", "CONSTANT", "NUMBER",
"VARIABLE", "SL_COMMENT", "ML_COMMENT", "WS" ]
+3 -3
View File
@@ -81,7 +81,7 @@ REMAP=80
IF=81
ELSE=82
WHILE=83
RET=84
RETURN=84
PUSH=85
POP=86
CLEAR=87
@@ -100,8 +100,8 @@ COV=99
SORT=100
APPEND=101
TIMESTAMP=102
BRK=103
CONT=104
BREAK=103
CONTINUE=104
PLUS=105
MINUS=106
MULT=107
+351
View File
@@ -44,6 +44,132 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#IfStatement.
def enterIfStatement(self, ctx:MathExprParser.IfStatementContext):
pass
# Exit a parse tree produced by MathExprParser#IfStatement.
def exitIfStatement(self, ctx:MathExprParser.IfStatementContext):
pass
# Enter a parse tree produced by MathExprParser#WhileStatement.
def enterWhileStatement(self, ctx:MathExprParser.WhileStatementContext):
pass
# Exit a parse tree produced by MathExprParser#WhileStatement.
def exitWhileStatement(self, ctx:MathExprParser.WhileStatementContext):
pass
# Enter a parse tree produced by MathExprParser#BlockStatement.
def enterBlockStatement(self, ctx:MathExprParser.BlockStatementContext):
pass
# Exit a parse tree produced by MathExprParser#BlockStatement.
def exitBlockStatement(self, ctx:MathExprParser.BlockStatementContext):
pass
# Enter a parse tree produced by MathExprParser#BreakStatement.
def enterBreakStatement(self, ctx:MathExprParser.BreakStatementContext):
pass
# Exit a parse tree produced by MathExprParser#BreakStatement.
def exitBreakStatement(self, ctx:MathExprParser.BreakStatementContext):
pass
# Enter a parse tree produced by MathExprParser#ContinueStatement.
def enterContinueStatement(self, ctx:MathExprParser.ContinueStatementContext):
pass
# Exit a parse tree produced by MathExprParser#ContinueStatement.
def exitContinueStatement(self, ctx:MathExprParser.ContinueStatementContext):
pass
# Enter a parse tree produced by MathExprParser#ReturnStatement.
def enterReturnStatement(self, ctx:MathExprParser.ReturnStatementContext):
pass
# Exit a parse tree produced by MathExprParser#ReturnStatement.
def exitReturnStatement(self, ctx:MathExprParser.ReturnStatementContext):
pass
# Enter a parse tree produced by MathExprParser#VarDefStmt.
def enterVarDefStmt(self, ctx:MathExprParser.VarDefStmtContext):
pass
# Exit a parse tree produced by MathExprParser#VarDefStmt.
def exitVarDefStmt(self, ctx:MathExprParser.VarDefStmtContext):
pass
# Enter a parse tree produced by MathExprParser#ExprStatement.
def enterExprStatement(self, ctx:MathExprParser.ExprStatementContext):
pass
# Exit a parse tree produced by MathExprParser#ExprStatement.
def exitExprStatement(self, ctx:MathExprParser.ExprStatementContext):
pass
# Enter a parse tree produced by MathExprParser#ifStmt.
def enterIfStmt(self, ctx:MathExprParser.IfStmtContext):
pass
# Exit a parse tree produced by MathExprParser#ifStmt.
def exitIfStmt(self, ctx:MathExprParser.IfStmtContext):
pass
# Enter a parse tree produced by MathExprParser#whileStmt.
def enterWhileStmt(self, ctx:MathExprParser.WhileStmtContext):
pass
# Exit a parse tree produced by MathExprParser#whileStmt.
def exitWhileStmt(self, ctx:MathExprParser.WhileStmtContext):
pass
# Enter a parse tree produced by MathExprParser#block.
def enterBlock(self, ctx:MathExprParser.BlockContext):
pass
# Exit a parse tree produced by MathExprParser#block.
def exitBlock(self, ctx:MathExprParser.BlockContext):
pass
# Enter a parse tree produced by MathExprParser#breakStmt.
def enterBreakStmt(self, ctx:MathExprParser.BreakStmtContext):
pass
# Exit a parse tree produced by MathExprParser#breakStmt.
def exitBreakStmt(self, ctx:MathExprParser.BreakStmtContext):
pass
# Enter a parse tree produced by MathExprParser#continueStmt.
def enterContinueStmt(self, ctx:MathExprParser.ContinueStmtContext):
pass
# Exit a parse tree produced by MathExprParser#continueStmt.
def exitContinueStmt(self, ctx:MathExprParser.ContinueStmtContext):
pass
# Enter a parse tree produced by MathExprParser#returnStmt.
def enterReturnStmt(self, ctx:MathExprParser.ReturnStmtContext):
pass
# Exit a parse tree produced by MathExprParser#returnStmt.
def exitReturnStmt(self, ctx:MathExprParser.ReturnStmtContext):
pass
# Enter a parse tree produced by MathExprParser#expr.
def enterExpr(self, ctx:MathExprParser.ExprContext):
pass
@@ -53,6 +179,15 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#TernaryExp.
def enterTernaryExp(self, ctx:MathExprParser.TernaryExpContext):
pass
# Exit a parse tree produced by MathExprParser#TernaryExp.
def exitTernaryExp(self, ctx:MathExprParser.TernaryExpContext):
pass
# Enter a parse tree produced by MathExprParser#LtExp.
def enterLtExp(self, ctx:MathExprParser.LtExpContext):
pass
@@ -242,6 +377,15 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#Func0Exp.
def enterFunc0Exp(self, ctx:MathExprParser.Func0ExpContext):
pass
# Exit a parse tree produced by MathExprParser#Func0Exp.
def exitFunc0Exp(self, ctx:MathExprParser.Func0ExpContext):
pass
# Enter a parse tree produced by MathExprParser#Func1Exp.
def enterFunc1Exp(self, ctx:MathExprParser.Func1ExpContext):
pass
@@ -278,6 +422,15 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#Func5Exp.
def enterFunc5Exp(self, ctx:MathExprParser.Func5ExpContext):
pass
# Exit a parse tree produced by MathExprParser#Func5Exp.
def exitFunc5Exp(self, ctx:MathExprParser.Func5ExpContext):
pass
# Enter a parse tree produced by MathExprParser#FuncNExp.
def enterFuncNExp(self, ctx:MathExprParser.FuncNExpContext):
pass
@@ -359,6 +512,15 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#TimestampFunc.
def enterTimestampFunc(self, ctx:MathExprParser.TimestampFuncContext):
pass
# Exit a parse tree produced by MathExprParser#TimestampFunc.
def exitTimestampFunc(self, ctx:MathExprParser.TimestampFuncContext):
pass
# Enter a parse tree produced by MathExprParser#SinFunc.
def enterSinFunc(self, ctx:MathExprParser.SinFuncContext):
pass
@@ -737,6 +899,105 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#AnyFunc.
def enterAnyFunc(self, ctx:MathExprParser.AnyFuncContext):
pass
# Exit a parse tree produced by MathExprParser#AnyFunc.
def exitAnyFunc(self, ctx:MathExprParser.AnyFuncContext):
pass
# Enter a parse tree produced by MathExprParser#AllFunc.
def enterAllFunc(self, ctx:MathExprParser.AllFuncContext):
pass
# Exit a parse tree produced by MathExprParser#AllFunc.
def exitAllFunc(self, ctx:MathExprParser.AllFuncContext):
pass
# Enter a parse tree produced by MathExprParser#EdgeFunc.
def enterEdgeFunc(self, ctx:MathExprParser.EdgeFuncContext):
pass
# Exit a parse tree produced by MathExprParser#EdgeFunc.
def exitEdgeFunc(self, ctx:MathExprParser.EdgeFuncContext):
pass
# Enter a parse tree produced by MathExprParser#MedianFunc.
def enterMedianFunc(self, ctx:MathExprParser.MedianFuncContext):
pass
# Exit a parse tree produced by MathExprParser#MedianFunc.
def exitMedianFunc(self, ctx:MathExprParser.MedianFuncContext):
pass
# Enter a parse tree produced by MathExprParser#ModeFunc.
def enterModeFunc(self, ctx:MathExprParser.ModeFuncContext):
pass
# Exit a parse tree produced by MathExprParser#ModeFunc.
def exitModeFunc(self, ctx:MathExprParser.ModeFuncContext):
pass
# Enter a parse tree produced by MathExprParser#CumsumFunc.
def enterCumsumFunc(self, ctx:MathExprParser.CumsumFuncContext):
pass
# Exit a parse tree produced by MathExprParser#CumsumFunc.
def exitCumsumFunc(self, ctx:MathExprParser.CumsumFuncContext):
pass
# Enter a parse tree produced by MathExprParser#CumprodFunc.
def enterCumprodFunc(self, ctx:MathExprParser.CumprodFuncContext):
pass
# Exit a parse tree produced by MathExprParser#CumprodFunc.
def exitCumprodFunc(self, ctx:MathExprParser.CumprodFuncContext):
pass
# Enter a parse tree produced by MathExprParser#PopFunc.
def enterPopFunc(self, ctx:MathExprParser.PopFuncContext):
pass
# Exit a parse tree produced by MathExprParser#PopFunc.
def exitPopFunc(self, ctx:MathExprParser.PopFuncContext):
pass
# Enter a parse tree produced by MathExprParser#ClearFunc.
def enterClearFunc(self, ctx:MathExprParser.ClearFuncContext):
pass
# Exit a parse tree produced by MathExprParser#ClearFunc.
def exitClearFunc(self, ctx:MathExprParser.ClearFuncContext):
pass
# Enter a parse tree produced by MathExprParser#HasFunc.
def enterHasFunc(self, ctx:MathExprParser.HasFuncContext):
pass
# Exit a parse tree produced by MathExprParser#HasFunc.
def exitHasFunc(self, ctx:MathExprParser.HasFuncContext):
pass
# Enter a parse tree produced by MathExprParser#GetFunc.
def enterGetFunc(self, ctx:MathExprParser.GetFuncContext):
pass
# Exit a parse tree produced by MathExprParser#GetFunc.
def exitGetFunc(self, ctx:MathExprParser.GetFuncContext):
pass
# Enter a parse tree produced by MathExprParser#PowFunc.
def enterPowFunc(self, ctx:MathExprParser.PowFuncContext):
pass
@@ -899,6 +1160,42 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#GaussianFunc.
def enterGaussianFunc(self, ctx:MathExprParser.GaussianFuncContext):
pass
# Exit a parse tree produced by MathExprParser#GaussianFunc.
def exitGaussianFunc(self, ctx:MathExprParser.GaussianFuncContext):
pass
# Enter a parse tree produced by MathExprParser#TopkIndFunc.
def enterTopkIndFunc(self, ctx:MathExprParser.TopkIndFuncContext):
pass
# Exit a parse tree produced by MathExprParser#TopkIndFunc.
def exitTopkIndFunc(self, ctx:MathExprParser.TopkIndFuncContext):
pass
# Enter a parse tree produced by MathExprParser#BotkIndFunc.
def enterBotkIndFunc(self, ctx:MathExprParser.BotkIndFuncContext):
pass
# Exit a parse tree produced by MathExprParser#BotkIndFunc.
def exitBotkIndFunc(self, ctx:MathExprParser.BotkIndFuncContext):
pass
# Enter a parse tree produced by MathExprParser#PushFunc.
def enterPushFunc(self, ctx:MathExprParser.PushFuncContext):
pass
# Exit a parse tree produced by MathExprParser#PushFunc.
def exitPushFunc(self, ctx:MathExprParser.PushFuncContext):
pass
# Enter a parse tree produced by MathExprParser#ClampFunc.
def enterClampFunc(self, ctx:MathExprParser.ClampFuncContext):
pass
@@ -962,6 +1259,42 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#CubicEaseFunc.
def enterCubicEaseFunc(self, ctx:MathExprParser.CubicEaseFuncContext):
pass
# Exit a parse tree produced by MathExprParser#CubicEaseFunc.
def exitCubicEaseFunc(self, ctx:MathExprParser.CubicEaseFuncContext):
pass
# Enter a parse tree produced by MathExprParser#ElasticEaseFunc.
def enterElasticEaseFunc(self, ctx:MathExprParser.ElasticEaseFuncContext):
pass
# Exit a parse tree produced by MathExprParser#ElasticEaseFunc.
def exitElasticEaseFunc(self, ctx:MathExprParser.ElasticEaseFuncContext):
pass
# Enter a parse tree produced by MathExprParser#SineEaseFunc.
def enterSineEaseFunc(self, ctx:MathExprParser.SineEaseFuncContext):
pass
# Exit a parse tree produced by MathExprParser#SineEaseFunc.
def exitSineEaseFunc(self, ctx:MathExprParser.SineEaseFuncContext):
pass
# Enter a parse tree produced by MathExprParser#SmootherstepFunc.
def enterSmootherstepFunc(self, ctx:MathExprParser.SmootherstepFuncContext):
pass
# Exit a parse tree produced by MathExprParser#SmootherstepFunc.
def exitSmootherstepFunc(self, ctx:MathExprParser.SmootherstepFuncContext):
pass
# Enter a parse tree produced by MathExprParser#SwapFunc.
def enterSwapFunc(self, ctx:MathExprParser.SwapFuncContext):
pass
@@ -980,6 +1313,24 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#DistFunc.
def enterDistFunc(self, ctx:MathExprParser.DistFuncContext):
pass
# Exit a parse tree produced by MathExprParser#DistFunc.
def exitDistFunc(self, ctx:MathExprParser.DistFuncContext):
pass
# Enter a parse tree produced by MathExprParser#RemapFunc.
def enterRemapFunc(self, ctx:MathExprParser.RemapFuncContext):
pass
# Exit a parse tree produced by MathExprParser#RemapFunc.
def exitRemapFunc(self, ctx:MathExprParser.RemapFuncContext):
pass
# Enter a parse tree produced by MathExprParser#SMinFunc.
def enterSMinFunc(self, ctx:MathExprParser.SMinFuncContext):
pass
File diff suppressed because it is too large Load Diff
+18 -3
View File
@@ -91,6 +91,11 @@ class UnifiedMathVisitor(MathExprVisitor):
"""
Generic binary operation handler.
"""
if self._is_tensor(a) and a.numel() == 1:
a = float(a.flatten()[0].item())
if self._is_tensor(b) and b.numel() == 1:
b = float(b.flatten()[0].item())
# one of them is a list and one is tensor
if self._is_tensor(a) and self._is_list(b):
if(a.shape[0]==len(b)):
@@ -121,6 +126,9 @@ class UnifiedMathVisitor(MathExprVisitor):
return scalar_op(a, b)
def _unary_op(self, a, torch_op, scalar_op):
if self._is_tensor(a) and a.numel() == 1:
a = float(a.flatten()[0].item())
if self._is_list(a):
return [self._unary_op(x, torch_op, scalar_op) for x in a]
if self._is_tensor(a):
@@ -128,7 +136,10 @@ class UnifiedMathVisitor(MathExprVisitor):
return scalar_op(a)
def _reduction_op(self, val, torch_op, list_op):
# If it's a scalar tensor, return Python float to avoid 0-dim tensor propagation
if self._is_tensor(val):
if val.numel() == 1:
return float(val.flatten()[0].item())
return torch_op(val)
if self._is_list(val):
return list_op(val)
@@ -592,7 +603,11 @@ class UnifiedMathVisitor(MathExprVisitor):
return t * t * (3.0 - 2.0 * t)
def visitRangeFunc(self, ctx):
return list(torch.arange((yield ctx.expr(0)), (yield ctx.expr(1)), (yield ctx.expr(2))))
s = (yield ctx.expr(0))
e = (yield ctx.expr(1))
st = (yield ctx.expr(2))
arr = torch.arange(s, e, st, device=self.device, dtype=torch.float32)
return [float(x) for x in arr.tolist()]
def visitSmootherstepFunc(self, ctx):
x = (yield ctx.expr(0))
@@ -1311,11 +1326,11 @@ class UnifiedMathVisitor(MathExprVisitor):
return res
def visitVarDefStmt(self, ctx):
res = yield ctx.expr()
res = yield ctx.varDef()
return res
def visitBlockStatement(self, ctx):
res = yield ctx.expr()
res = yield ctx.block()
return res
def visitBlock(self, ctx):