{-# LANGUAGE FlexibleContexts #-}

{- | The symbolic-regression expression tree: a small first-order ADT with
vectorized evaluation and a total translation to a dataframe 'Expr Double'.
Division, log, and sqrt are protected so evaluation never produces @NaN@.
-}
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)

-- | A symbolic-regression expression over feature variables and constants.
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]

{- | Evaluate over a feature matrix given column-major (@feats ! j@ is feature
@j@ across all rows). Protected operators keep results finite.
-}
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

-- | Translate to a dataframe expression over the named feature columns.
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

-- | The constant values in left-to-right traversal order.
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

-- | Replace the constants in traversal order; extra values are ignored.
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'')