add lambda
This commit is contained in:
@@ -0,0 +1,6 @@
|
||||
class LambdaFunction:
|
||||
__slots__ = ("params", "body", "env")
|
||||
def __init__(self, params, body, env):
|
||||
self.params = params
|
||||
self.body = body
|
||||
self.env = env.copy() if env is not None else {}
|
||||
@@ -82,15 +82,17 @@ atom:
|
||||
| funcNoise # FuncNoiseExp
|
||||
| VARIABLE # VariableExp
|
||||
| NUMBER # NumberExp
|
||||
| CONSTANT # ConstantExp
|
||||
| STRING # StringExp
|
||||
| LPAREN expr RPAREN # ParenExp
|
||||
| PIPE expr PIPE # AbsExp
|
||||
| CONSTANT # ConstantExp
|
||||
| STRING # StringExp
|
||||
| LPAREN paramList RPAREN ARROW (block | expr) # LambdaExp
|
||||
| LPAREN expr RPAREN # ParenExp
|
||||
| LPAREN expr RPAREN # ParenExp
|
||||
| PIPE expr PIPE # AbsExp
|
||||
| LBRACKET expr (COMMA expr)* RBRACKET # ListExp
|
||||
| VARIABLE LPAREN exprList? RPAREN # CallExp
|
||||
| NONE # NoneExp
|
||||
| BREAK # BreakExp
|
||||
| CONTINUE # ContinueExp;
|
||||
| NONE # NoneExp
|
||||
| BREAK # BreakExp
|
||||
| CONTINUE # ContinueExp;
|
||||
|
||||
exprList: expr (COMMA expr)*;
|
||||
|
||||
@@ -173,8 +175,7 @@ func1:
|
||||
| INT LPAREN expr RPAREN # IntFunc
|
||||
| FLOAT LPAREN expr RPAREN # FloatFunc
|
||||
| FLOW_MAG LPAREN expr RPAREN # FlowMagFunc
|
||||
| FLOW_ANG LPAREN expr RPAREN # FlowAngFunc
|
||||
;
|
||||
| FLOW_ANG LPAREN expr RPAREN # FlowAngFunc;
|
||||
|
||||
func2:
|
||||
POWE LPAREN expr COMMA expr RPAREN # PowFunc
|
||||
|
||||
@@ -5,6 +5,7 @@ import math
|
||||
import inspect
|
||||
import torch.nn.functional as F
|
||||
from . import optical_flow_utils as ofu
|
||||
from .Func import LambdaFunction
|
||||
from antlr4 import TerminalNode
|
||||
|
||||
from .antlr_router import get_antlr_modules
|
||||
@@ -1924,7 +1925,35 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
def visitCallExp(self, ctx):
|
||||
func_name = ctx.VARIABLE().getText()
|
||||
|
||||
# Check if it is a user-defined function
|
||||
# 1) Pokud proměnná existuje a je to lambda uložená v variables, aplikuj ji
|
||||
if func_name in self.variables and isinstance(self.variables[func_name], LambdaFunction):
|
||||
lam = self.variables[func_name]
|
||||
# vyhodnotit argumenty
|
||||
args = []
|
||||
if ctx.exprList():
|
||||
for e in ctx.exprList().expr():
|
||||
args.append((yield e))
|
||||
|
||||
# připrav nový scope na základě uzávěrky
|
||||
new_vars = lam.env.copy()
|
||||
for i, p in enumerate(lam.params):
|
||||
new_vars[p] = args[i] if i < len(args) else None
|
||||
|
||||
# push/pop scope stejným stylem jako pro pojmenované funkce
|
||||
self._scope_stack.append(self.variables)
|
||||
self.variables = new_vars
|
||||
self.depth += 1
|
||||
self.variables["depth"] = float(self.depth)
|
||||
try:
|
||||
res = yield lam.body
|
||||
if isinstance(res, ReturnSignal):
|
||||
return res.value
|
||||
return res
|
||||
finally:
|
||||
self.variables = self._scope_stack.pop()
|
||||
self.depth -= 1
|
||||
|
||||
# 2) existující uživatelské funkce (bez změn)
|
||||
if func_name in self.functions:
|
||||
func_def = self.functions[func_name]
|
||||
params = func_def["params"]
|
||||
@@ -1965,7 +1994,7 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
self.variables = self._scope_stack.pop()
|
||||
self.depth -= 1
|
||||
|
||||
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Unknown function: {func_name}")
|
||||
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Unknown function or variable: {func_name}")
|
||||
|
||||
def visitNoiseFunc(self,ctx):
|
||||
seed_val = yield ctx.expr(0)
|
||||
@@ -3347,9 +3376,9 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
b = rgb[..., 2]
|
||||
else:
|
||||
# Separate r, g, b mode
|
||||
r = self._promote_to_tensor((yield ctx.expr(0)))
|
||||
g = self._promote_to_tensor((yield ctx.expr(1)))
|
||||
b = self._promote_to_tensor((yield ctx.expr(2)))
|
||||
r = self._promote_to_tensor((yield ctx.expr(0))).float()
|
||||
g = self._promote_to_tensor((yield ctx.expr(1))).float()
|
||||
b = self._promote_to_tensor((yield ctx.expr(2))).float()
|
||||
|
||||
use_degrees = False
|
||||
if num_args == 4:
|
||||
@@ -3895,4 +3924,13 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
a = self._promote_to_tensor((yield ctx.expr(1))).float()
|
||||
b = self._promote_to_tensor((yield ctx.expr(2))).float()
|
||||
|
||||
return self._cielab_to_rgb(L*100, a*100, b*100)
|
||||
return self._cielab_to_rgb(L*100, a*100, b*100)
|
||||
|
||||
def visitLambdaExp(self, ctx):
|
||||
# ctx.paramList() nebo None; body je ctx.block() nebo ctx.expr()
|
||||
params = []
|
||||
if ctx.paramList():
|
||||
params = [node.getText() for node in ctx.paramList().VARIABLE()]
|
||||
body = ctx.block() if ctx.block() else ctx.expr()
|
||||
closure = self.variables.copy()
|
||||
return LambdaFunction(params, body, closure)
|
||||
File diff suppressed because one or more lines are too long
@@ -530,6 +530,15 @@ class MathExprListener(ParseTreeListener):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#LambdaExp.
|
||||
def enterLambdaExp(self, ctx:MathExprParser.LambdaExpContext):
|
||||
pass
|
||||
|
||||
# Exit a parse tree produced by MathExprParser#LambdaExp.
|
||||
def exitLambdaExp(self, ctx:MathExprParser.LambdaExpContext):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#ParenExp.
|
||||
def enterParenExp(self, ctx:MathExprParser.ParenExpContext):
|
||||
pass
|
||||
|
||||
+1840
-1759
File diff suppressed because it is too large
Load Diff
@@ -299,6 +299,11 @@ class MathExprVisitor(ParseTreeVisitor):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#LambdaExp.
|
||||
def visitLambdaExp(self, ctx:MathExprParser.LambdaExpContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#ParenExp.
|
||||
def visitParenExp(self, ctx:MathExprParser.ParenExpContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
File diff suppressed because one or more lines are too long
@@ -530,6 +530,15 @@ class MathExprListener(ParseTreeListener):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#LambdaExp.
|
||||
def enterLambdaExp(self, ctx:MathExprParser.LambdaExpContext):
|
||||
pass
|
||||
|
||||
# Exit a parse tree produced by MathExprParser#LambdaExp.
|
||||
def exitLambdaExp(self, ctx:MathExprParser.LambdaExpContext):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#ParenExp.
|
||||
def enterParenExp(self, ctx:MathExprParser.ParenExpContext):
|
||||
pass
|
||||
|
||||
+2087
-2004
File diff suppressed because one or more lines are too long
@@ -299,6 +299,11 @@ class MathExprVisitor(ParseTreeVisitor):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#LambdaExp.
|
||||
def visitLambdaExp(self, ctx:MathExprParser.LambdaExpContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#ParenExp.
|
||||
def visitParenExp(self, ctx:MathExprParser.ParenExpContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
Reference in New Issue
Block a user