diff --git a/more_math/Parser/MathExpr.g4 b/more_math/Parser/MathExpr.g4 index 40447f5..ca871af 100644 --- a/more_math/Parser/MathExpr.g4 +++ b/more_math/Parser/MathExpr.g4 @@ -18,10 +18,12 @@ stmt: | continueStmt # ContinueStatement | returnStmt # ReturnStatement | varDef # VarDefStmt + | forStmt # ForStatement | expr SEMICOLON # ExprStatement; ifStmt: IF LPAREN expr RPAREN stmt (ELSE stmt)?; whileStmt: WHILE LPAREN expr RPAREN stmt; +forStmt: FOR LPAREN VARIABLE IN expr RPAREN stmt; block: LBRACE stmt* RBRACE; breakStmt: BREAK SEMICOLON; continueStmt: CONTINUE SEMICOLON; @@ -162,7 +164,9 @@ 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 - | PUSH LPAREN expr COMMA expr RPAREN # PushFunc; + | BOTK_IND LPAREN expr COMMA expr RPAREN # BotkIndFunc + | PUSH LPAREN expr COMMA expr RPAREN # PushFunc + | GET_VALUE LPAREN expr COMMA expr RPAREN # GetValueFunc; func3: CLAMP LPAREN expr COMMA expr COMMA expr RPAREN # ClampFunc @@ -301,6 +305,8 @@ FLIP: 'flip'; COV: 'cov'; SORT: 'sort'; APPEND: 'append'; +FOR: 'for'; +IN: 'in'; TIMESTAMP: 'timestamp' | 'now'; BREAK: 'break'; diff --git a/more_math/Parser/UnifiedMathVisitor.py b/more_math/Parser/UnifiedMathVisitor.py index b296f8c..a2ea1a3 100644 --- a/more_math/Parser/UnifiedMathVisitor.py +++ b/more_math/Parser/UnifiedMathVisitor.py @@ -1410,6 +1410,34 @@ class UnifiedMathVisitor(MathExprVisitor): if isinstance(res, ContinueSignal): continue return None + + def visitForStmt(self, ctx): + var_name = ctx.VARIABLE().getText() + iterable = yield ctx.expr() + + iterator = [] + if self._is_tensor(iterable): + if iterable.ndim == 0: + iterator = [iterable] + else: + iterator = iterable + elif self._is_list(iterable): + iterator = iterable + else: + iterator = [iterable] + + for val in iterator: + self.variables[var_name] = val + res = yield ctx.stmt() + + if isinstance(res, ReturnSignal): + return res + if isinstance(res, BreakSignal): + break + if isinstance(res, ContinueSignal): + continue + return None + def visitBreakStmt(self, ctx): return BreakSignal()