{-# LANGUAGE FlexibleContexts #-}
module DataFrame.SymbolicRegression.Expr (
SRExpr (..),
BinOp (..),
UnOp (..),
evalSR,
toDataFrameExpr,
srSize,
constants,
setConstants,
allBinOps,
allUnOps,
) where
import qualified Data.Text as T
import qualified Data.Vector as V
import qualified Data.Vector.Unboxed as VU
import qualified DataFrame.Functions as F
import DataFrame.Internal.Expression (Expr (..))
import DataFrame.Operators ((.*.), (.+.), (.-.), (./.))
data BinOp = SAdd | SSub | SMul | SDiv
deriving (BinOp -> BinOp -> Bool
(BinOp -> BinOp -> Bool) -> (BinOp -> BinOp -> Bool) -> Eq BinOp
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: BinOp -> BinOp -> Bool
== :: BinOp -> BinOp -> Bool
$c/= :: BinOp -> BinOp -> Bool
/= :: BinOp -> BinOp -> Bool
Eq, Eq BinOp
Eq BinOp =>
(BinOp -> BinOp -> Ordering)
-> (BinOp -> BinOp -> Bool)
-> (BinOp -> BinOp -> Bool)
-> (BinOp -> BinOp -> Bool)
-> (BinOp -> BinOp -> Bool)
-> (BinOp -> BinOp -> BinOp)
-> (BinOp -> BinOp -> BinOp)
-> Ord BinOp
BinOp -> BinOp -> Bool
BinOp -> BinOp -> Ordering
BinOp -> BinOp -> BinOp
forall a.
Eq a =>
(a -> a -> Ordering)
-> (a -> a -> Bool)
-> (a -> a -> Bool)
-> (a -> a -> Bool)
-> (a -> a -> Bool)
-> (a -> a -> a)
-> (a -> a -> a)
-> Ord a
$ccompare :: BinOp -> BinOp -> Ordering
compare :: BinOp -> BinOp -> Ordering
$c< :: BinOp -> BinOp -> Bool
< :: BinOp -> BinOp -> Bool
$c<= :: BinOp -> BinOp -> Bool
<= :: BinOp -> BinOp -> Bool
$c> :: BinOp -> BinOp -> Bool
> :: BinOp -> BinOp -> Bool
$c>= :: BinOp -> BinOp -> Bool
>= :: BinOp -> BinOp -> Bool
$cmax :: BinOp -> BinOp -> BinOp
max :: BinOp -> BinOp -> BinOp
$cmin :: BinOp -> BinOp -> BinOp
min :: BinOp -> BinOp -> BinOp
Ord, Int -> BinOp -> ShowS
[BinOp] -> ShowS
BinOp -> String
(Int -> BinOp -> ShowS)
-> (BinOp -> String) -> ([BinOp] -> ShowS) -> Show BinOp
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> BinOp -> ShowS
showsPrec :: Int -> BinOp -> ShowS
$cshow :: BinOp -> String
show :: BinOp -> String
$cshowList :: [BinOp] -> ShowS
showList :: [BinOp] -> ShowS
Show, Int -> BinOp
BinOp -> Int
BinOp -> [BinOp]
BinOp -> BinOp
BinOp -> BinOp -> [BinOp]
BinOp -> BinOp -> BinOp -> [BinOp]
(BinOp -> BinOp)
-> (BinOp -> BinOp)
-> (Int -> BinOp)
-> (BinOp -> Int)
-> (BinOp -> [BinOp])
-> (BinOp -> BinOp -> [BinOp])
-> (BinOp -> BinOp -> [BinOp])
-> (BinOp -> BinOp -> BinOp -> [BinOp])
-> Enum BinOp
forall a.
(a -> a)
-> (a -> a)
-> (Int -> a)
-> (a -> Int)
-> (a -> [a])
-> (a -> a -> [a])
-> (a -> a -> [a])
-> (a -> a -> a -> [a])
-> Enum a
$csucc :: BinOp -> BinOp
succ :: BinOp -> BinOp
$cpred :: BinOp -> BinOp
pred :: BinOp -> BinOp
$ctoEnum :: Int -> BinOp
toEnum :: Int -> BinOp
$cfromEnum :: BinOp -> Int
fromEnum :: BinOp -> Int
$cenumFrom :: BinOp -> [BinOp]
enumFrom :: BinOp -> [BinOp]
$cenumFromThen :: BinOp -> BinOp -> [BinOp]
enumFromThen :: BinOp -> BinOp -> [BinOp]
$cenumFromTo :: BinOp -> BinOp -> [BinOp]
enumFromTo :: BinOp -> BinOp -> [BinOp]
$cenumFromThenTo :: BinOp -> BinOp -> BinOp -> [BinOp]
enumFromThenTo :: BinOp -> BinOp -> BinOp -> [BinOp]
Enum, BinOp
BinOp -> BinOp -> Bounded BinOp
forall a. a -> a -> Bounded a
$cminBound :: BinOp
minBound :: BinOp
$cmaxBound :: BinOp
maxBound :: BinOp
Bounded)
data UnOp = SNeg | SSin | SCos | SExp | SLog | SSqrt
deriving (UnOp -> UnOp -> Bool
(UnOp -> UnOp -> Bool) -> (UnOp -> UnOp -> Bool) -> Eq UnOp
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: UnOp -> UnOp -> Bool
== :: UnOp -> UnOp -> Bool
$c/= :: UnOp -> UnOp -> Bool
/= :: UnOp -> UnOp -> Bool
Eq, Eq UnOp
Eq UnOp =>
(UnOp -> UnOp -> Ordering)
-> (UnOp -> UnOp -> Bool)
-> (UnOp -> UnOp -> Bool)
-> (UnOp -> UnOp -> Bool)
-> (UnOp -> UnOp -> Bool)
-> (UnOp -> UnOp -> UnOp)
-> (UnOp -> UnOp -> UnOp)
-> Ord UnOp
UnOp -> UnOp -> Bool
UnOp -> UnOp -> Ordering
UnOp -> UnOp -> UnOp
forall a.
Eq a =>
(a -> a -> Ordering)
-> (a -> a -> Bool)
-> (a -> a -> Bool)
-> (a -> a -> Bool)
-> (a -> a -> Bool)
-> (a -> a -> a)
-> (a -> a -> a)
-> Ord a
$ccompare :: UnOp -> UnOp -> Ordering
compare :: UnOp -> UnOp -> Ordering
$c< :: UnOp -> UnOp -> Bool
< :: UnOp -> UnOp -> Bool
$c<= :: UnOp -> UnOp -> Bool
<= :: UnOp -> UnOp -> Bool
$c> :: UnOp -> UnOp -> Bool
> :: UnOp -> UnOp -> Bool
$c>= :: UnOp -> UnOp -> Bool
>= :: UnOp -> UnOp -> Bool
$cmax :: UnOp -> UnOp -> UnOp
max :: UnOp -> UnOp -> UnOp
$cmin :: UnOp -> UnOp -> UnOp
min :: UnOp -> UnOp -> UnOp
Ord, Int -> UnOp -> ShowS
[UnOp] -> ShowS
UnOp -> String
(Int -> UnOp -> ShowS)
-> (UnOp -> String) -> ([UnOp] -> ShowS) -> Show UnOp
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> UnOp -> ShowS
showsPrec :: Int -> UnOp -> ShowS
$cshow :: UnOp -> String
show :: UnOp -> String
$cshowList :: [UnOp] -> ShowS
showList :: [UnOp] -> ShowS
Show, Int -> UnOp
UnOp -> Int
UnOp -> [UnOp]
UnOp -> UnOp
UnOp -> UnOp -> [UnOp]
UnOp -> UnOp -> UnOp -> [UnOp]
(UnOp -> UnOp)
-> (UnOp -> UnOp)
-> (Int -> UnOp)
-> (UnOp -> Int)
-> (UnOp -> [UnOp])
-> (UnOp -> UnOp -> [UnOp])
-> (UnOp -> UnOp -> [UnOp])
-> (UnOp -> UnOp -> UnOp -> [UnOp])
-> Enum UnOp
forall a.
(a -> a)
-> (a -> a)
-> (Int -> a)
-> (a -> Int)
-> (a -> [a])
-> (a -> a -> [a])
-> (a -> a -> [a])
-> (a -> a -> a -> [a])
-> Enum a
$csucc :: UnOp -> UnOp
succ :: UnOp -> UnOp
$cpred :: UnOp -> UnOp
pred :: UnOp -> UnOp
$ctoEnum :: Int -> UnOp
toEnum :: Int -> UnOp
$cfromEnum :: UnOp -> Int
fromEnum :: UnOp -> Int
$cenumFrom :: UnOp -> [UnOp]
enumFrom :: UnOp -> [UnOp]
$cenumFromThen :: UnOp -> UnOp -> [UnOp]
enumFromThen :: UnOp -> UnOp -> [UnOp]
$cenumFromTo :: UnOp -> UnOp -> [UnOp]
enumFromTo :: UnOp -> UnOp -> [UnOp]
$cenumFromThenTo :: UnOp -> UnOp -> UnOp -> [UnOp]
enumFromThenTo :: UnOp -> UnOp -> UnOp -> [UnOp]
Enum, UnOp
UnOp -> UnOp -> Bounded UnOp
forall a. a -> a -> Bounded a
$cminBound :: UnOp
minBound :: UnOp
$cmaxBound :: UnOp
maxBound :: UnOp
Bounded)
data SRExpr
= SVar !Int
| SConst !Double
| SUn !UnOp SRExpr
| SBin !BinOp SRExpr SRExpr
deriving (SRExpr -> SRExpr -> Bool
(SRExpr -> SRExpr -> Bool)
-> (SRExpr -> SRExpr -> Bool) -> Eq SRExpr
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: SRExpr -> SRExpr -> Bool
== :: SRExpr -> SRExpr -> Bool
$c/= :: SRExpr -> SRExpr -> Bool
/= :: SRExpr -> SRExpr -> Bool
Eq, Eq SRExpr
Eq SRExpr =>
(SRExpr -> SRExpr -> Ordering)
-> (SRExpr -> SRExpr -> Bool)
-> (SRExpr -> SRExpr -> Bool)
-> (SRExpr -> SRExpr -> Bool)
-> (SRExpr -> SRExpr -> Bool)
-> (SRExpr -> SRExpr -> SRExpr)
-> (SRExpr -> SRExpr -> SRExpr)
-> Ord SRExpr
SRExpr -> SRExpr -> Bool
SRExpr -> SRExpr -> Ordering
SRExpr -> SRExpr -> SRExpr
forall a.
Eq a =>
(a -> a -> Ordering)
-> (a -> a -> Bool)
-> (a -> a -> Bool)
-> (a -> a -> Bool)
-> (a -> a -> Bool)
-> (a -> a -> a)
-> (a -> a -> a)
-> Ord a
$ccompare :: SRExpr -> SRExpr -> Ordering
compare :: SRExpr -> SRExpr -> Ordering
$c< :: SRExpr -> SRExpr -> Bool
< :: SRExpr -> SRExpr -> Bool
$c<= :: SRExpr -> SRExpr -> Bool
<= :: SRExpr -> SRExpr -> Bool
$c> :: SRExpr -> SRExpr -> Bool
> :: SRExpr -> SRExpr -> Bool
$c>= :: SRExpr -> SRExpr -> Bool
>= :: SRExpr -> SRExpr -> Bool
$cmax :: SRExpr -> SRExpr -> SRExpr
max :: SRExpr -> SRExpr -> SRExpr
$cmin :: SRExpr -> SRExpr -> SRExpr
min :: SRExpr -> SRExpr -> SRExpr
Ord, Int -> SRExpr -> ShowS
[SRExpr] -> ShowS
SRExpr -> String
(Int -> SRExpr -> ShowS)
-> (SRExpr -> String) -> ([SRExpr] -> ShowS) -> Show SRExpr
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> SRExpr -> ShowS
showsPrec :: Int -> SRExpr -> ShowS
$cshow :: SRExpr -> String
show :: SRExpr -> String
$cshowList :: [SRExpr] -> ShowS
showList :: [SRExpr] -> ShowS
Show)
allBinOps :: [BinOp]
allBinOps :: [BinOp]
allBinOps = [BinOp
forall a. Bounded a => a
minBound .. BinOp
forall a. Bounded a => a
maxBound]
allUnOps :: [UnOp]
allUnOps :: [UnOp]
allUnOps = [UnOp
forall a. Bounded a => a
minBound .. UnOp
forall a. Bounded a => a
maxBound]
evalSR :: V.Vector (VU.Vector Double) -> Int -> SRExpr -> VU.Vector Double
evalSR :: Vector (Vector Double) -> Int -> SRExpr -> Vector Double
evalSR Vector (Vector Double)
feats Int
n = SRExpr -> Vector Double
go
where
go :: SRExpr -> Vector Double
go (SVar Int
j)
| Int
j Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Vector (Vector Double) -> Int
forall a. Vector a -> Int
V.length Vector (Vector Double)
feats = Vector (Vector Double)
feats Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
j
| Bool
otherwise = Int -> Double -> Vector Double
forall a. Unbox a => Int -> a -> Vector a
VU.replicate Int
n Double
0
go (SConst Double
c) = Int -> Double -> Vector Double
forall a. Unbox a => Int -> a -> Vector a
VU.replicate Int
n Double
c
go (SUn UnOp
op SRExpr
e) = (Double -> Double) -> Vector Double -> Vector Double
forall a b. (Unbox a, Unbox b) => (a -> b) -> Vector a -> Vector b
VU.map (UnOp -> Double -> Double
unFn UnOp
op) (SRExpr -> Vector Double
go SRExpr
e)
go (SBin BinOp
op SRExpr
a SRExpr
b) = (Double -> Double -> Double)
-> Vector Double -> Vector Double -> Vector Double
forall a b c.
(Unbox a, Unbox b, Unbox c) =>
(a -> b -> c) -> Vector a -> Vector b -> Vector c
VU.zipWith (BinOp -> Double -> Double -> Double
binFn BinOp
op) (SRExpr -> Vector Double
go SRExpr
a) (SRExpr -> Vector Double
go SRExpr
b)
binFn :: BinOp -> Double -> Double -> Double
binFn :: BinOp -> Double -> Double -> Double
binFn BinOp
SAdd Double
a Double
b = Double
a Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
b
binFn BinOp
SSub Double
a Double
b = Double
a Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
b
binFn BinOp
SMul Double
a Double
b = Double
a Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
b
binFn BinOp
SDiv Double
a Double
b = if Double -> Double
forall a. Num a => a -> a
abs Double
b Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
< Double
1e-9 then Double
1 else Double
a Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
b
unFn :: UnOp -> Double -> Double
unFn :: UnOp -> Double -> Double
unFn UnOp
SNeg = Double -> Double
forall a. Num a => a -> a
negate
unFn UnOp
SSin = Double -> Double
forall a. Floating a => a -> a
sin
unFn UnOp
SCos = Double -> Double
forall a. Floating a => a -> a
cos
unFn UnOp
SExp = Double -> Double
forall a. Floating a => a -> a
exp (Double -> Double) -> (Double -> Double) -> Double -> Double
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Double -> Double -> Double
forall a. Ord a => a -> a -> a
min Double
50
unFn UnOp
SLog = \Double
x -> Double -> Double
forall a. Floating a => a -> a
log (Double -> Double
forall a. Num a => a -> a
abs Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
1e-9)
unFn UnOp
SSqrt = Double -> Double
forall a. Floating a => a -> a
sqrt (Double -> Double) -> (Double -> Double) -> Double -> Double
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Double -> Double
forall a. Num a => a -> a
abs
toDataFrameExpr :: V.Vector T.Text -> SRExpr -> Expr Double
toDataFrameExpr :: Vector Text -> SRExpr -> Expr Double
toDataFrameExpr Vector Text
names = SRExpr -> Expr Double
go
where
go :: SRExpr -> Expr Double
go (SVar Int
j)
| Int
j Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Vector Text -> Int
forall a. Vector a -> Int
V.length Vector Text
names = Text -> Expr Double
forall a. Columnable a => Text -> Expr a
Col (Vector Text
names Vector Text -> Int -> Text
forall a. Vector a -> Int -> a
V.! Int
j)
| Bool
otherwise = Double -> Expr Double
forall a. Columnable a => a -> Expr a
F.lit Double
0
go (SConst Double
c) = Double -> Expr Double
forall a. Columnable a => a -> Expr a
F.lit Double
c
go (SUn UnOp
op SRExpr
e) = UnOp -> Expr Double -> Expr Double
forall {a}. Floating a => UnOp -> a -> a
unExpr UnOp
op (SRExpr -> Expr Double
go SRExpr
e)
go (SBin BinOp
op SRExpr
a SRExpr
b) = BinOp -> Expr Double -> Expr Double -> Expr Double
forall {a}.
(When (Unboxable a) (Unbox a), When (IntegralTypes a) (Integral a),
When (FloatingTypes a) (Real a, Fractional a), Typeable a, Show a,
Eq a, ColumnifyRep (KindOf a) a, SBoolI (Unboxable a),
SBoolI (Numeric a), SBoolI (IntegralTypes a),
SBoolI (FloatingTypes a), Fractional a) =>
BinOp -> Expr a -> Expr a -> Expr a
binExpr BinOp
op (SRExpr -> Expr Double
go SRExpr
a) (SRExpr -> Expr Double
go SRExpr
b)
unExpr :: UnOp -> a -> a
unExpr UnOp
SNeg = a -> a
forall a. Num a => a -> a
negate
unExpr UnOp
SSin = a -> a
forall a. Floating a => a -> a
sin
unExpr UnOp
SCos = a -> a
forall a. Floating a => a -> a
cos
unExpr UnOp
SExp = a -> a
forall a. Floating a => a -> a
exp
unExpr UnOp
SLog = a -> a
forall a. Floating a => a -> a
log
unExpr UnOp
SSqrt = a -> a
forall a. Floating a => a -> a
sqrt
binExpr :: BinOp -> Expr a -> Expr a -> Expr a
binExpr BinOp
SAdd = Expr a -> Expr a -> Expr a
forall a. (Columnable a, Num a) => Expr a -> Expr a -> Expr a
(.+.)
binExpr BinOp
SSub = Expr a -> Expr a -> Expr a
forall a. (Columnable a, Num a) => Expr a -> Expr a -> Expr a
(.-.)
binExpr BinOp
SMul = Expr a -> Expr a -> Expr a
forall a. (Columnable a, Num a) => Expr a -> Expr a -> Expr a
(.*.)
binExpr BinOp
SDiv = Expr a -> Expr a -> Expr a
forall a.
(Columnable a, Fractional a) =>
Expr a -> Expr a -> Expr a
(./.)
srSize :: SRExpr -> Int
srSize :: SRExpr -> Int
srSize (SVar Int
_) = Int
1
srSize (SConst Double
_) = Int
1
srSize (SUn UnOp
_ SRExpr
e) = Int
1 Int -> Int -> Int
forall a. Num a => a -> a -> a
+ SRExpr -> Int
srSize SRExpr
e
srSize (SBin BinOp
_ SRExpr
a SRExpr
b) = Int
1 Int -> Int -> Int
forall a. Num a => a -> a -> a
+ SRExpr -> Int
srSize SRExpr
a Int -> Int -> Int
forall a. Num a => a -> a -> a
+ SRExpr -> Int
srSize SRExpr
b
constants :: SRExpr -> [Double]
constants :: SRExpr -> [Double]
constants (SConst Double
c) = [Double
c]
constants (SVar Int
_) = []
constants (SUn UnOp
_ SRExpr
e) = SRExpr -> [Double]
constants SRExpr
e
constants (SBin BinOp
_ SRExpr
a SRExpr
b) = SRExpr -> [Double]
constants SRExpr
a [Double] -> [Double] -> [Double]
forall a. [a] -> [a] -> [a]
++ SRExpr -> [Double]
constants SRExpr
b
setConstants :: [Double] -> SRExpr -> SRExpr
setConstants :: [Double] -> SRExpr -> SRExpr
setConstants [Double]
vals SRExpr
e = (SRExpr, [Double]) -> SRExpr
forall a b. (a, b) -> a
fst ([Double] -> SRExpr -> (SRExpr, [Double])
go [Double]
vals SRExpr
e)
where
go :: [Double] -> SRExpr -> (SRExpr, [Double])
go [Double]
vs (SConst Double
_) = case [Double]
vs of
(Double
v : [Double]
rest) -> (Double -> SRExpr
SConst Double
v, [Double]
rest)
[] -> (Double -> SRExpr
SConst Double
0, [])
go [Double]
vs (SVar Int
j) = (Int -> SRExpr
SVar Int
j, [Double]
vs)
go [Double]
vs (SUn UnOp
op SRExpr
a) = let (SRExpr
a', [Double]
vs') = [Double] -> SRExpr -> (SRExpr, [Double])
go [Double]
vs SRExpr
a in (UnOp -> SRExpr -> SRExpr
SUn UnOp
op SRExpr
a', [Double]
vs')
go [Double]
vs (SBin BinOp
op SRExpr
a SRExpr
b) =
let (SRExpr
a', [Double]
vs') = [Double] -> SRExpr -> (SRExpr, [Double])
go [Double]
vs SRExpr
a
(SRExpr
b', [Double]
vs'') = [Double] -> SRExpr -> (SRExpr, [Double])
go [Double]
vs' SRExpr
b
in (BinOp -> SRExpr -> SRExpr -> SRExpr
SBin BinOp
op SRExpr
a' SRExpr
b', [Double]
vs'')