from dataclasses import dataclass, replace import math class EscapeToZero(Exception): pass def escape_to_zero(): raise EscapeToZero() @dataclass class ASTNode: pass @dataclass class OpBinary(ASTNode): lhs: any rhs: any @dataclass class OpAdd(OpBinary): rhs_flip_sign: bool = False @dataclass class OpMultiply(OpBinary): pass @dataclass class OpDivide(OpBinary): pass @dataclass class OpMod(OpBinary): pass @dataclass class OpAnd(OpBinary): pass @dataclass class OpOr(OpBinary): pass @dataclass class OpCompare(OpBinary): comparison: str @dataclass class OpVariableAccess(ASTNode): name: any @dataclass class OpArrayAccess(ASTNode): name: any index: any @dataclass class OpIn(OpBinary): pass @dataclass class VariableStub: pass def eval(ast, env): def getenv(name): val = env.get(name, None) if callable(val): val = val() if val is None: raise Exception(f"Variable {ast.name} not found") return val if isinstance(ast, list): # OpIn rhs return [eval(a, env) for a in ast] if isinstance(ast, OpBinary): lhs_val = eval(ast.lhs, env) if isinstance(lhs_val, ASTNode): # lhs contains variables try: rhs_val = eval(ast.rhs, env) return replace(ast, lhs=lhs_val, rhs=rhs_val) except (EscapeToZero, IndexError) as e: def raiser(): raise e return replace(ast, lhs=lhs_val, rhs=raiser) # short-circuit if isinstance(ast, OpAnd): if not lhs_val: return False return eval(ast.rhs, env) elif isinstance(ast, OpOr): if lhs_val: return True return eval(ast.rhs, env) rhs_val = eval(ast.rhs, env) if isinstance(rhs_val, ASTNode): # rhs contains variables return replace(ast, lhs=lhs_val, rhs=rhs_val) if isinstance(ast, OpAdd): return lhs_val + (-rhs_val if ast.rhs_flip_sign else rhs_val) elif isinstance(ast, OpMultiply): return lhs_val * rhs_val elif isinstance(ast, OpDivide): return lhs_val / rhs_val elif isinstance(ast, OpMod): return lhs_val % rhs_val elif isinstance(ast, OpCompare): if ast.comparison == "==": if isinstance(lhs_val, float) or isinstance(rhs_val, float): return math.isclose(float(lhs_val), float(rhs_val), rel_tol=1e-5) return lhs_val == rhs_val elif ast.comparison == "!=": return lhs_val != rhs_val elif ast.comparison == "<": return lhs_val < rhs_val elif ast.comparison == "<=": return lhs_val <= rhs_val elif ast.comparison == ">": return lhs_val > rhs_val elif ast.comparison == ">=": return lhs_val >= rhs_val elif isinstance(ast, OpIn): if isinstance(lhs_val, float): return any( math.isclose(lhs_val, float(x), rel_tol=1e-5) for x in rhs_val ) return lhs_val in rhs_val elif isinstance(ast, OpVariableAccess): val = getenv(ast.name) if isinstance(val, VariableStub): return ast if isinstance(val, list): raise Exception(f"Variable {ast.name} is a list, but was used as a scalar") return val elif isinstance(ast, OpArrayAccess): val = getenv(ast.name) if not isinstance(val, list): raise Exception(f"Variable {ast.name} is a scalar, but was used as a list") return val[ast.index] elif callable(ast): return ast() return ast # pitch_bias_func = logitbias_simplify(self.bias, env, "pitch") # take everything in env to be constant # return a function that evaluates ast in terms of variable's value, # or a constant if no expressions depend on variable # # pairs is a list of (ast, value) def logitbias_simplify(pairs, env, variable): simplified_pairs = [] bias = 0 for ast, value in pairs: try: res = eval(ast, dict(env, **{variable: VariableStub()})) except EscapeToZero: continue if isinstance(res, ASTNode): simplified_pairs.append((res, value)) elif res: bias += value simplified_env = {} if len(simplified_pairs) == 0: return bias def simplified_eval(val): simplified_env[variable] = val return bias + sum( bias_val if eval(ast, simplified_env) else 0 for ast, bias_val in simplified_pairs ) return simplified_eval