{-# LANGUAGE FlexibleInstances #-}
{-# LANGUAGE MultiParamTypeClasses #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeFamilies #-}

{- | Symbolic regression by genetic programming (modelled on the
@symbolic-regression@ library, ported dependency-light: no e-graphs, no NLOPT).
'predict' is the best discovered @Expr Double@; the search also returns the
accuracy-vs-complexity Pareto front. Deterministic given the seed.
-}
module DataFrame.SymbolicRegression (
    module DataFrame.Model,
    UnOp (..),
    SRConfig (..),
    defaultSRConfig,
    SRModel (..),
) where

import Control.Exception (throw)
import qualified Data.Vector as V
import qualified Data.Vector.Unboxed as VU

import DataFrame.Featurize.Internal (featureNames, targetDoubles)
import DataFrame.Internal.DataFrame (DataFrame)
import DataFrame.Internal.Expression (Expr (..))
import DataFrame.Model
import DataFrame.Operations.Core (columnAsDoubleVector)
import DataFrame.Random (mkGen)
import DataFrame.SymbolicRegression.Expr (
    UnOp (..),
    allUnOps,
    toDataFrameExpr,
 )
import DataFrame.SymbolicRegression.GP (GPParams (..), runGP)
import DataFrame.SymbolicRegression.Simplify (simplify)

data SRConfig = SRConfig
    { SRConfig -> Int
srSeed :: !Int
    , SRConfig -> Int
srPopSize :: !Int
    , SRConfig -> Int
srGenerations :: !Int
    , SRConfig -> Int
srMaxSize :: !Int
    , SRConfig -> Int
srTournament :: !Int
    , SRConfig -> Double
srCrossoverP :: !Double
    , SRConfig -> Double
srMutationP :: !Double
    , SRConfig -> Double
srOptimizeP :: !Double
    , SRConfig -> Double
srParsimony :: !Double
    , SRConfig -> [UnOp]
srUnaryOps :: ![UnOp]
    }
    deriving (SRConfig -> SRConfig -> Bool
(SRConfig -> SRConfig -> Bool)
-> (SRConfig -> SRConfig -> Bool) -> Eq SRConfig
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: SRConfig -> SRConfig -> Bool
== :: SRConfig -> SRConfig -> Bool
$c/= :: SRConfig -> SRConfig -> Bool
/= :: SRConfig -> SRConfig -> Bool
Eq, Int -> SRConfig -> ShowS
[SRConfig] -> ShowS
SRConfig -> String
(Int -> SRConfig -> ShowS)
-> (SRConfig -> String) -> ([SRConfig] -> ShowS) -> Show SRConfig
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> SRConfig -> ShowS
showsPrec :: Int -> SRConfig -> ShowS
$cshow :: SRConfig -> String
show :: SRConfig -> String
$cshowList :: [SRConfig] -> ShowS
showList :: [SRConfig] -> ShowS
Show)

defaultSRConfig :: SRConfig
defaultSRConfig :: SRConfig
defaultSRConfig =
    SRConfig
        { srSeed :: Int
srSeed = Int
42
        , srPopSize :: Int
srPopSize = Int
200
        , srGenerations :: Int
srGenerations = Int
40
        , srMaxSize :: Int
srMaxSize = Int
25
        , srTournament :: Int
srTournament = Int
5
        , srCrossoverP :: Double
srCrossoverP = Double
0.9
        , srMutationP :: Double
srMutationP = Double
0.3
        , srOptimizeP :: Double
srOptimizeP = Double
0.15
        , srParsimony :: Double
srParsimony = Double
1.0e-3
        , srUnaryOps :: [UnOp]
srUnaryOps = [UnOp]
allUnOps
        }

{- | A fitted symbolic regressor. 'srBest' is the lowest-error expression;
'srPareto' is the @(complexity, mse, expr)@ frontier.
-}
data SRModel = SRModel
    { SRModel -> Expr Double
srBest :: !(Expr Double)
    , SRModel -> Double
srBestMSE :: !Double
    , SRModel -> [(Int, Double, Expr Double)]
srPareto :: ![(Int, Double, Expr Double)]
    , SRModel -> Int
srGenerationsRun :: !Int
    }

instance Fit SRConfig (Expr Double) where
    type ModelOf SRConfig (Expr Double) = SRModel
    fit :: CheckFrame
  (FrameReq SRConfig (Expr Double)) (FrameFor (Expr Double)) =>
SRConfig
-> Expr Double
-> FrameFor (Expr Double)
-> FitResult
     (FrameFor (Expr Double)) (ModelOf SRConfig (Expr Double))
fit = SRConfig -> Expr Double -> DataFrame -> SRModel
SRConfig
-> Expr Double
-> FrameFor (Expr Double)
-> FitResult
     (FrameFor (Expr Double)) (ModelOf SRConfig (Expr Double))
fitSymbolicRegression

instance Predict SRModel where
    type Prediction SRModel = Expr Double
    predict :: SRModel -> Prediction SRModel
predict = SRModel -> Expr Double
SRModel -> Prediction SRModel
srBest

-- | Search for an expression predicting @target@ from the other columns.
fitSymbolicRegression :: SRConfig -> Expr Double -> DataFrame -> SRModel
fitSymbolicRegression :: SRConfig -> Expr Double -> DataFrame -> SRModel
fitSymbolicRegression SRConfig
cfg Expr Double
target DataFrame
df =
    SRModel
        { srBest :: Expr Double
srBest = SRExpr -> Expr Double
translate SRExpr
best
        , srBestMSE :: Double
srBestMSE = Double
bestMse
        , srPareto :: [(Int, Double, Expr Double)]
srPareto = [(Int
sz, Double
mse, SRExpr -> Expr Double
translate SRExpr
e) | (Int
sz, Double
mse, SRExpr
e) <- [(Int, Double, SRExpr)]
front]
        , srGenerationsRun :: Int
srGenerationsRun = Int
gens
        }
  where
    names :: [Text]
names = Expr Double -> DataFrame -> [Text]
forall a. Expr a -> DataFrame -> [Text]
featureNames Expr Double
target DataFrame
df
    nameVec :: Vector Text
nameVec = [Text] -> Vector Text
forall a. [a] -> Vector a
V.fromList [Text]
names
    cols :: Vector (Vector Double)
cols = [Vector Double] -> Vector (Vector Double)
forall a. [a] -> Vector a
V.fromList ((Text -> Vector Double) -> [Text] -> [Vector Double]
forall a b. (a -> b) -> [a] -> [b]
map (DataFrame -> Expr Double -> Vector Double
materialize DataFrame
df (Expr Double -> Vector Double)
-> (Text -> Expr Double) -> Text -> Vector Double
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Text -> Expr Double
forall a. Columnable a => Text -> Expr a
Col) [Text]
names)
    target' :: Vector Double
target' = Expr Double -> DataFrame -> Vector Double
targetDoubles Expr Double
target DataFrame
df
    n :: Int
n = Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Double
target'
    params :: GPParams
params =
        GPParams
            { gpFeats :: Vector (Vector Double)
gpFeats = Vector (Vector Double)
cols
            , gpN :: Int
gpN = Int
n
            , gpTarget :: Vector Double
gpTarget = Vector Double
target'
            , gpNVars :: Int
gpNVars = [Text] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [Text]
names
            , gpUnOps :: [UnOp]
gpUnOps = SRConfig -> [UnOp]
srUnaryOps SRConfig
cfg
            , gpPopSize :: Int
gpPopSize = SRConfig -> Int
srPopSize SRConfig
cfg
            , gpGenerations :: Int
gpGenerations = SRConfig -> Int
srGenerations SRConfig
cfg
            , gpMaxSize :: Int
gpMaxSize = SRConfig -> Int
srMaxSize SRConfig
cfg
            , gpTournament :: Int
gpTournament = SRConfig -> Int
srTournament SRConfig
cfg
            , gpCrossoverP :: Double
gpCrossoverP = SRConfig -> Double
srCrossoverP SRConfig
cfg
            , gpMutationP :: Double
gpMutationP = SRConfig -> Double
srMutationP SRConfig
cfg
            , gpOptimizeP :: Double
gpOptimizeP = SRConfig -> Double
srOptimizeP SRConfig
cfg
            , gpParsimony :: Double
gpParsimony = SRConfig -> Double
srParsimony SRConfig
cfg
            }
    (SRExpr
best, [(Int, Double, SRExpr)]
front, Int
gens) = GPParams -> Gen -> (SRExpr, [(Int, Double, SRExpr)], Int)
runGP GPParams
params (Int -> Gen
mkGen (SRConfig -> Int
srSeed SRConfig
cfg))
    bestMse :: Double
bestMse = case [Double
m | (Int
_, Double
m, SRExpr
e) <- [(Int, Double, SRExpr)]
front, SRExpr
e SRExpr -> SRExpr -> Bool
forall a. Eq a => a -> a -> Bool
== SRExpr
best] of
        (Double
m : [Double]
_) -> Double
m
        [] -> Double
1 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
0
    translate :: SRExpr -> Expr Double
translate = Vector Text -> SRExpr -> Expr Double
toDataFrameExpr Vector Text
nameVec (SRExpr -> Expr Double)
-> (SRExpr -> SRExpr) -> SRExpr -> Expr Double
forall b c a. (b -> c) -> (a -> b) -> a -> c
. SRExpr -> SRExpr
simplify

materialize :: DataFrame -> Expr Double -> VU.Vector Double
materialize :: DataFrame -> Expr Double -> Vector Double
materialize DataFrame
df Expr Double
e = case Expr Double
-> DataFrame -> Either DataFrameException (Vector Double)
forall a.
(Columnable a, Num a) =>
Expr a -> DataFrame -> Either DataFrameException (Vector Double)
columnAsDoubleVector Expr Double
e DataFrame
df of
    Right Vector Double
v -> Vector Double
v
    Left DataFrameException
err -> DataFrameException -> Vector Double
forall a e. Exception e => e -> a
throw DataFrameException
err