gt22 icon

Untitled

gt22 | PRO | 09/21/20 06:40:55 PM UTC | 0 ⭐ | 1923 👁️ | Never ⏰ | []
Python |

6.75 KB

|

None

|

0 👍

/

0 👎

from __future__ import annotations
from abc import ABC, abstractmethod
from typing import Dict
import numpy as np
from math import log
 
 
def literal_convert(x):
    if isinstance(x, Expression):
        return x
    if isinstance(x, int):
        return Constant(float(x))
    if isinstance(x, float):
        return Constant(x)
    if isinstance(x, str):
        try:
            return Constant(float(x))
        except ValueError:
            return Variable(x)
    raise ValueError(f"Can't convert ({x})@{type(x)} to expression")
 
 
class Expression(ABC):
 
    subexprs: Dict[str, Expression]
 
    @abstractmethod
    def value(self, var: Dict[str, float]) -> float:
        raise NotImplemented
 
    @abstractmethod
    def _proper_deriv(self, sub: str) -> Expression:
        raise NotImplemented
 
    def deriv(self, by: str):
        d = Constant(0)
        for name, expr in self.subexprs.items():
            d += expr.deriv(by) * self._proper_deriv(name)
        return d
 
    def _const_reduce(self):
        try:
            return Constant(self.value({}))
        except ValueError:
            return self
 
    def reduce(self):
        return self._const_reduce()
 
    def __eq__(self, other):
        return type(self) is type(other) and all(other.subexprs[name] == e for name, e in self.subexprs.items())
 
    def __add__(self, other):
        return Plus(self, literal_convert(other))
 
    def __mul__(self, other):
        return Mul(self, literal_convert(other))
 
    def __neg__(self):
        return self * -1
 
    def __sub__(self, other):
        return self + (-other)
 
    def __pow__(self, power, modulo=None):
        if modulo is not None:
            raise NotImplemented
        return Pow(self, literal_convert(power))
 
    def __truediv__(self, other):
        return self * (other ** -1)
 
 
class Constant(Expression):
 
    val: float
 
    def __init__(self, val: float):
        self.val = val
        self.subexprs = {}
 
    def value(self, var: Dict[str, float]) -> float:
        return self.val
 
    def _proper_deriv(self, sub: str) -> Expression:
        raise NotImplemented
 
    def deriv(self, var: str) -> Expression:
        return Constant(0)
 
    def __str__(self):
        return str(self.val)
 
    def __eq__(self, other):
        return super(Constant, self).__eq__(other) and self.val == other.val
 
 
class Variable(Expression):
 
    name: str
 
    def __init__(self, name: str):
        self.name = name
        self.subexprs = {}
 
    def value(self, var: Dict[str, float]) -> float:
        if self.name not in var:
            raise ValueError(f"Variable {self.name} not found")
        return var[self.name]
 
    def _proper_deriv(self, sub: str) -> Expression:
        raise NotImplemented
 
    def deriv(self, var: str) -> Expression:
        return Constant(1 if var == self.name else 0)
 
    def __str__(self):
        return self.name
 
    def __eq__(self, other):
        return super(Variable, self).__eq__(other) and self.name == other.name
 
 
class Plus(Expression):
 
    left: Expression
    right: Expression
 
    def __init__(self, left: Expression, right: Expression):
        self.subexprs = {
            'left': left,
            'right': right
        }
        self.left = left
        self.right = right
 
    def value(self, var: Dict[str, float]) -> float:
        return self.left.value(var) + self.right.value(var)
 
    def _proper_deriv(self, sub: str) -> Expression:
        return Constant(1)
 
    def reduce(self):
        le = self.left.reduce()
        ri = self.right.reduce()
        if le == Constant(0):
            return ri
        if ri == Constant(0):
            return le
        return (le + ri)._const_reduce()
 
    def __str__(self):
        return f"({self.left} + {self.right})"
 
 
class Mul(Expression):
 
    left: Expression
    right: Expression
 
    def __init__(self, left: Expression, right: Expression):
        self.subexprs = {
            'left': left,
            'right': right
        }
        self.left = left
        self.right = right
 
    def value(self, var: Dict[str, float]) -> float:
        return self.left.value(var) * self.right.value(var)
 
    def _proper_deriv(self, sub: str) -> Expression:
        return self.right if sub == 'left' else self.left
 
    def reduce(self):
        le = self.left.reduce()
        ri = self.right.reduce()
        if le == Constant(0) or ri == Constant(0):
            return Constant(0)
        if le == Constant(1):
            return ri
        if ri == Constant(1):
            return le
        return (le * ri)._const_reduce()
 
    def __str__(self):
        return f"({self.left} * {self.right})"
 
 
class Log(Expression):
 
    x: Expression
    base: Expression
 
    def __init__(self, x: Expression, base: Expression = Variable('e')):
        self.subexprs = {
            'x': x,
            'base': base
        }
        self.x = x
        self.base = base
 
    def value(self, var: Dict[str, float]) -> float:
        return log(self.x.value(var), self.base.value(var))
 
    def _proper_deriv(self, sub: str) -> Expression:
        if sub == 'x':
            return Constant(1) / (self.x * Log(self.base))
        if sub == 'base':
            return -(self / (self.base * Log(self.base)))
 
    def reduce(self):
        b = self.base.reduce()
        xi = self.x.reduce()
        if isinstance(xi, Pow) and xi.base == b:
            return xi.power
        if b == xi:
            return Constant(1)
        return Log(xi, b)._const_reduce()
 
    def __str__(self):
        return f"log_{self.base}({self.x})"
 
 
class Pow(Expression):
 
    base: Expression
    power: Expression
 
    def __init__(self, base: Expression, power: Expression):
        self.subexprs = {
            'base': base,
            'power': power
        }
        self.base = base
        self.power = power
 
    def value(self, var: Dict[str, float]) -> float:
        return np.power(self.base.value(var), self.power.value(var))
 
    def _proper_deriv(self, sub: str) -> Expression:
        if sub == 'base':
            return self.power * (self.base ** (self.power - 1))
        if sub == 'power':
            return Log(self.base) * self
 
    def reduce(self):
        b = self.base.reduce()
        p = self.power.reduce()
        if isinstance(p, Log) and p.base == b:
            return p.x
        return (b ** p)._const_reduce()
 
    def __str__(self):
        return f"({self.base})^({self.power})"
 
 
def main():
    e = Log(Variable('x') * 'y' + 'x')
    print(e)
    print(e.value({'x': 5, 'y': 3, 'e': np.e}))
    print(e.deriv('x').reduce())
    pass
 
 
if __name__ == '__main__':
    main()
 

Comments