{-# LANGUAGE FlexibleContexts #-}
{-# LANGUAGE FlexibleInstances #-}
{-# LANGUAGE MultiParamTypeClasses #-}
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE TypeFamilies #-}
{-# LANGUAGE UndecidableInstances #-}
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
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)
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