{-# LANGUAGE FlexibleContexts #-}
{-# LANGUAGE FlexibleInstances #-}
{-# LANGUAGE MultiParamTypeClasses #-}
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE TypeFamilies #-}
{-# LANGUAGE UndecidableInstances #-}

{- | sklearn-style standalone tree estimators returning inspectable records
(depth, leaf count, per-feature split usage). 'fit' trains a classifier (from a
'TreeConfig') or a regressor (from a 'RegTreeConfig'); 'predict' is the compiled
tree expression, and the record exposes the raw 'Tree' too. The bare
'DataFrame.DecisionTree.Fit.fitDecisionTree' remains for callers that only want
the classifier @Expr@.
-}
module DataFrame.DecisionTree.Model (
    module DataFrame.Model,
    DecisionTreeClassifier (..),
    DecisionTreeRegressor (..),
) where

import Control.Exception (throw)
import qualified Data.Map.Strict as M
import qualified Data.Text as T
import DataFrame.Errors (DataFrameException (..))

import qualified Data.Vector as V

import DataFrame.DecisionTree.Cart (cartFeatures)
import DataFrame.DecisionTree.Fit (fitDecisionTree, treeToExpr)
import DataFrame.DecisionTree.Regression (RegTreeConfig, fitRegTreeOn)
import DataFrame.DecisionTree.Types (Tree (..), TreeConfig)
import DataFrame.Featurize.Internal (targetDoubles)
import DataFrame.Internal.Column (Columnable)
import DataFrame.Internal.Expression (Expr (..), getColumns)
import DataFrame.Model

-- | A fitted classification tree with structural diagnostics.
data DecisionTreeClassifier a = DecisionTreeClassifier
    { forall a. DecisionTreeClassifier a -> Expr a
dtcExpr :: !(Expr a)
    , forall a. DecisionTreeClassifier a -> Int
dtcDepth :: !Int
    , forall a. DecisionTreeClassifier a -> Int
dtcNLeaves :: !Int
    , forall a. DecisionTreeClassifier a -> Map Text Int
dtcFeatureUsage :: !(M.Map T.Text Int)
    }
    deriving (Int -> DecisionTreeClassifier a -> ShowS
[DecisionTreeClassifier a] -> ShowS
DecisionTreeClassifier a -> String
(Int -> DecisionTreeClassifier a -> ShowS)
-> (DecisionTreeClassifier a -> String)
-> ([DecisionTreeClassifier a] -> ShowS)
-> Show (DecisionTreeClassifier a)
forall a. Show a => Int -> DecisionTreeClassifier a -> ShowS
forall a. Show a => [DecisionTreeClassifier a] -> ShowS
forall a. Show a => DecisionTreeClassifier a -> String
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: forall a. Show a => Int -> DecisionTreeClassifier a -> ShowS
showsPrec :: Int -> DecisionTreeClassifier a -> ShowS
$cshow :: forall a. Show a => DecisionTreeClassifier a -> String
show :: DecisionTreeClassifier a -> String
$cshowList :: forall a. Show a => [DecisionTreeClassifier a] -> ShowS
showList :: [DecisionTreeClassifier a] -> ShowS
Show)

-- | A fitted regression tree with structural diagnostics.
data DecisionTreeRegressor = DecisionTreeRegressor
    { DecisionTreeRegressor -> Tree Double
dtrTree :: !(Tree Double)
    , DecisionTreeRegressor -> Expr Double
dtrExpr :: !(Expr Double)
    , DecisionTreeRegressor -> Int
dtrDepth :: !Int
    , DecisionTreeRegressor -> Int
dtrNLeaves :: !Int
    , DecisionTreeRegressor -> Map Text Int
dtrFeatureUsage :: !(M.Map T.Text Int)
    }
    deriving (Int -> DecisionTreeRegressor -> ShowS
[DecisionTreeRegressor] -> ShowS
DecisionTreeRegressor -> String
(Int -> DecisionTreeRegressor -> ShowS)
-> (DecisionTreeRegressor -> String)
-> ([DecisionTreeRegressor] -> ShowS)
-> Show DecisionTreeRegressor
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> DecisionTreeRegressor -> ShowS
showsPrec :: Int -> DecisionTreeRegressor -> ShowS
$cshow :: DecisionTreeRegressor -> String
show :: DecisionTreeRegressor -> String
$cshowList :: [DecisionTreeRegressor] -> ShowS
showList :: [DecisionTreeRegressor] -> ShowS
Show)

instance (Columnable a, Ord a) => Fit TreeConfig (Expr a) where
    type ModelOf TreeConfig (Expr a) = (DecisionTreeClassifier a)
    fit :: CheckFrame (FrameReq TreeConfig (Expr a)) (FrameFor (Expr a)) =>
TreeConfig
-> Expr a
-> FrameFor (Expr a)
-> FitResult (FrameFor (Expr a)) (ModelOf TreeConfig (Expr a))
fit TreeConfig
cfg Expr a
target FrameFor (Expr a)
df =
        Expr a -> Int -> Int -> Map Text Int -> DecisionTreeClassifier a
forall a.
Expr a -> Int -> Int -> Map Text Int -> DecisionTreeClassifier a
DecisionTreeClassifier
            Expr a
e
            (Expr a -> Int
forall a. Expr a -> Int
exprDepth Expr a
e)
            (Expr a -> Int
forall a. Expr a -> Int
exprLeaves Expr a
e)
            ([Text] -> Map Text Int
usageCounts (Expr a -> [Text]
forall a. Expr a -> [Text]
exprUsage Expr a
e))
      where
        e :: Expr a
e = TreeConfig -> Expr a -> DataFrame -> Expr a
forall a.
(Columnable a, Ord a) =>
TreeConfig -> Expr a -> DataFrame -> Expr a
fitDecisionTree TreeConfig
cfg Expr a
target DataFrame
FrameFor (Expr a)
df

instance Predict (DecisionTreeClassifier a) where
    type Prediction (DecisionTreeClassifier a) = Expr a
    predict :: DecisionTreeClassifier a -> Prediction (DecisionTreeClassifier a)
predict = DecisionTreeClassifier a -> Expr a
DecisionTreeClassifier a -> Prediction (DecisionTreeClassifier a)
forall a. DecisionTreeClassifier a -> Expr a
dtcExpr

instance Fit RegTreeConfig (Expr Double) where
    type ModelOf RegTreeConfig (Expr Double) = DecisionTreeRegressor
    fit :: CheckFrame
  (FrameReq RegTreeConfig (Expr Double)) (FrameFor (Expr Double)) =>
RegTreeConfig
-> Expr Double
-> FrameFor (Expr Double)
-> FitResult
     (FrameFor (Expr Double)) (ModelOf RegTreeConfig (Expr Double))
fit RegTreeConfig
cfg Expr Double
target FrameFor (Expr Double)
df =
        Tree Double
-> Expr Double
-> Int
-> Int
-> Map Text Int
-> DecisionTreeRegressor
DecisionTreeRegressor
            Tree Double
t
            Expr Double
e
            (Expr Double -> Int
forall a. Expr a -> Int
exprDepth Expr Double
e)
            (Expr Double -> Int
forall a. Expr a -> Int
exprLeaves Expr Double
e)
            ([Text] -> Map Text Int
usageCounts (Expr Double -> [Text]
forall a. Expr a -> [Text]
exprUsage Expr Double
e))
      where
        t :: Tree Double
t = case Expr Double
target of
            Col Text
name ->
                RegTreeConfig
-> Vector CartFeature
-> Vector Double
-> Maybe (Vector Double)
-> Tree Double
fitRegTreeOn
                    RegTreeConfig
cfg
                    ([CartFeature] -> Vector CartFeature
forall a. [a] -> Vector a
V.fromList (Text -> DataFrame -> [CartFeature]
cartFeatures Text
name DataFrame
FrameFor (Expr Double)
df))
                    (Expr Double -> DataFrame -> Vector Double
targetDoubles Expr Double
target DataFrame
FrameFor (Expr Double)
df)
                    Maybe (Vector Double)
forall a. Maybe a
Nothing
            Expr Double
_ ->
                DataFrameException -> Tree Double
forall a e. Exception e => e -> a
throw
                    ( Text -> DataFrameException
NonColumnReferenceException
                        (Text
"fit @DecisionTreeRegressor: " Text -> Text -> Text
forall a. Semigroup a => a -> a -> a
<> String -> Text
T.pack (Expr Double -> String
forall a. Show a => a -> String
show Expr Double
target))
                    )
        e :: Expr Double
e = Tree Double -> Expr Double
forall a. Columnable a => Tree a -> Expr a
treeToExpr Tree Double
t

instance Predict DecisionTreeRegressor where
    type Prediction DecisionTreeRegressor = Expr Double
    predict :: DecisionTreeRegressor -> Prediction DecisionTreeRegressor
predict = DecisionTreeRegressor -> Expr Double
DecisionTreeRegressor -> Prediction DecisionTreeRegressor
dtrExpr

usageCounts :: [T.Text] -> M.Map T.Text Int
usageCounts :: [Text] -> Map Text Int
usageCounts = (Text -> Map Text Int -> Map Text Int)
-> Map Text Int -> [Text] -> Map Text Int
forall a b. (a -> b -> b) -> b -> [a] -> b
forall (t :: * -> *) a b.
Foldable t =>
(a -> b -> b) -> b -> t a -> b
foldr (\Text
c -> (Int -> Int -> Int) -> Text -> Int -> Map Text Int -> Map Text Int
forall k a. Ord k => (a -> a -> a) -> k -> a -> Map k a -> Map k a
M.insertWith Int -> Int -> Int
forall a. Num a => a -> a -> a
(+) Text
c Int
1) Map Text Int
forall k a. Map k a
M.empty

exprUsage :: Expr a -> [T.Text]
exprUsage :: forall a. Expr a -> [Text]
exprUsage (If Expr Bool
c Expr a
t Expr a
e) = Expr Bool -> [Text]
forall a. Expr a -> [Text]
getColumns Expr Bool
c [Text] -> [Text] -> [Text]
forall a. [a] -> [a] -> [a]
++ Expr a -> [Text]
forall a. Expr a -> [Text]
exprUsage Expr a
t [Text] -> [Text] -> [Text]
forall a. [a] -> [a] -> [a]
++ Expr a -> [Text]
forall a. Expr a -> [Text]
exprUsage Expr a
e
exprUsage Expr a
_ = []

exprLeaves :: Expr a -> Int
exprLeaves :: forall a. Expr a -> Int
exprLeaves (If Expr Bool
_ Expr a
t Expr a
e) = Expr a -> Int
forall a. Expr a -> Int
exprLeaves Expr a
t Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Expr a -> Int
forall a. Expr a -> Int
exprLeaves Expr a
e
exprLeaves Expr a
_ = Int
1

exprDepth :: Expr a -> Int
exprDepth :: forall a. Expr a -> Int
exprDepth (If Expr Bool
_ Expr a
t Expr a
e) = Int
1 Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int -> Int -> Int
forall a. Ord a => a -> a -> a
max (Expr a -> Int
forall a. Expr a -> Int
exprDepth Expr a
t) (Expr a -> Int
forall a. Expr a -> Int
exprDepth Expr a
e)
exprDepth Expr a
_ = Int
0