{-# LANGUAGE AllowAmbiguousTypes #-}
{-# LANGUAGE FlexibleContexts #-}
{-# LANGUAGE GADTs #-}
{-# LANGUAGE RankNTypes #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications #-}
module DataFrame.DecisionTree.Types (
Tree (..),
treeDepth,
TreeConfig (..),
SynthConfig (..),
defaultTreeConfig,
defaultSynthConfig,
ColumnOrdering (..),
orderable,
defaultColumnOrdering,
withOrdFrom,
CarePoint (..),
Direction (..),
ttrace,
) where
import DataFrame.Internal.Column (Columnable)
import DataFrame.Internal.Expression (Expr (..))
import qualified DataFrame.LinearSolver as LS
import Data.Int (Int16, Int32, Int64, Int8)
import qualified Data.Map.Strict as M
import Data.Proxy (Proxy (..))
import qualified Data.Text as T
import Data.Type.Equality (testEquality, (:~:) (..))
import Data.Word (Word16, Word32, Word64, Word8)
import qualified Debug.Trace as Trace
import System.Environment (lookupEnv)
import System.IO.Unsafe (unsafePerformIO)
import Type.Reflection (SomeTypeRep (..), typeRep)
data Tree a
= Leaf !a
| Branch !(Expr Bool) !(Tree a) !(Tree a)
deriving (Int -> Tree a -> ShowS
[Tree a] -> ShowS
Tree a -> String
(Int -> Tree a -> ShowS)
-> (Tree a -> String) -> ([Tree a] -> ShowS) -> Show (Tree a)
forall a. Show a => Int -> Tree a -> ShowS
forall a. Show a => [Tree a] -> ShowS
forall a. Show a => Tree a -> String
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: forall a. Show a => Int -> Tree a -> ShowS
showsPrec :: Int -> Tree a -> ShowS
$cshow :: forall a. Show a => Tree a -> String
show :: Tree a -> String
$cshowList :: forall a. Show a => [Tree a] -> ShowS
showList :: [Tree a] -> ShowS
Show)
treeDepth :: Tree a -> Int
treeDepth :: forall a. Tree a -> Int
treeDepth (Leaf a
_) = Int
0
treeDepth (Branch Expr Bool
_ Tree a
l Tree a
r) = Int
1 Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int -> Int -> Int
forall a. Ord a => a -> a -> a
max (Tree a -> Int
forall a. Tree a -> Int
treeDepth Tree a
l) (Tree a -> Int
forall a. Tree a -> Int
treeDepth Tree a
r)
data CarePoint = CarePoint
{ CarePoint -> Int
cpIndex :: !Int
, CarePoint -> Direction
cpCorrectDir :: !Direction
}
deriving (CarePoint -> CarePoint -> Bool
(CarePoint -> CarePoint -> Bool)
-> (CarePoint -> CarePoint -> Bool) -> Eq CarePoint
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: CarePoint -> CarePoint -> Bool
== :: CarePoint -> CarePoint -> Bool
$c/= :: CarePoint -> CarePoint -> Bool
/= :: CarePoint -> CarePoint -> Bool
Eq, Int -> CarePoint -> ShowS
[CarePoint] -> ShowS
CarePoint -> String
(Int -> CarePoint -> ShowS)
-> (CarePoint -> String)
-> ([CarePoint] -> ShowS)
-> Show CarePoint
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> CarePoint -> ShowS
showsPrec :: Int -> CarePoint -> ShowS
$cshow :: CarePoint -> String
show :: CarePoint -> String
$cshowList :: [CarePoint] -> ShowS
showList :: [CarePoint] -> ShowS
Show)
data Direction = GoLeft | GoRight
deriving (Direction -> Direction -> Bool
(Direction -> Direction -> Bool)
-> (Direction -> Direction -> Bool) -> Eq Direction
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: Direction -> Direction -> Bool
== :: Direction -> Direction -> Bool
$c/= :: Direction -> Direction -> Bool
/= :: Direction -> Direction -> Bool
Eq, Int -> Direction -> ShowS
[Direction] -> ShowS
Direction -> String
(Int -> Direction -> ShowS)
-> (Direction -> String)
-> ([Direction] -> ShowS)
-> Show Direction
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> Direction -> ShowS
showsPrec :: Int -> Direction -> ShowS
$cshow :: Direction -> String
show :: Direction -> String
$cshowList :: [Direction] -> ShowS
showList :: [Direction] -> ShowS
Show)
data TreeConfig = TreeConfig
{ TreeConfig -> Int
maxTreeDepth :: Int
, TreeConfig -> Int
minSamplesSplit :: Int
, TreeConfig -> Int
minLeafSize :: Int
, TreeConfig -> [Int]
percentiles :: [Int]
, TreeConfig -> Int
expressionPairs :: Int
, TreeConfig -> SynthConfig
synthConfig :: SynthConfig
, TreeConfig -> Int
taoIterations :: Int
, TreeConfig -> Double
taoConvergenceTol :: Double
, TreeConfig -> ColumnOrdering
columnOrdering :: ColumnOrdering
, TreeConfig -> Bool
useLinearSolver :: Bool
, TreeConfig -> SolverConfig
linearSolverConfig :: LS.SolverConfig
, TreeConfig -> Int
minCarePointsForLinear :: Int
, TreeConfig -> Bool
pureReplacementLinear :: Bool
}
data SynthConfig = SynthConfig
{ SynthConfig -> Int
maxExprDepth :: Int
, SynthConfig -> Int
boolExpansion :: Int
, SynthConfig -> [(Text, Text)]
disallowedCombinations :: [(T.Text, T.Text)]
, SynthConfig -> Double
complexityPenalty :: Double
, SynthConfig -> Bool
enableStringOps :: Bool
, SynthConfig -> Bool
enableCrossCols :: Bool
, SynthConfig -> Bool
enableArithOps :: Bool
, SynthConfig -> Int
maxCategoricalSubsetCardinality :: Int
, SynthConfig -> Maybe Int
perColumnQuota :: Maybe Int
}
deriving (SynthConfig -> SynthConfig -> Bool
(SynthConfig -> SynthConfig -> Bool)
-> (SynthConfig -> SynthConfig -> Bool) -> Eq SynthConfig
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: SynthConfig -> SynthConfig -> Bool
== :: SynthConfig -> SynthConfig -> Bool
$c/= :: SynthConfig -> SynthConfig -> Bool
/= :: SynthConfig -> SynthConfig -> Bool
Eq, Int -> SynthConfig -> ShowS
[SynthConfig] -> ShowS
SynthConfig -> String
(Int -> SynthConfig -> ShowS)
-> (SynthConfig -> String)
-> ([SynthConfig] -> ShowS)
-> Show SynthConfig
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> SynthConfig -> ShowS
showsPrec :: Int -> SynthConfig -> ShowS
$cshow :: SynthConfig -> String
show :: SynthConfig -> String
$cshowList :: [SynthConfig] -> ShowS
showList :: [SynthConfig] -> ShowS
Show)
defaultSynthConfig :: SynthConfig
defaultSynthConfig :: SynthConfig
defaultSynthConfig =
SynthConfig
{ maxExprDepth :: Int
maxExprDepth = Int
2
, boolExpansion :: Int
boolExpansion = Int
2
, disallowedCombinations :: [(Text, Text)]
disallowedCombinations = []
, complexityPenalty :: Double
complexityPenalty = Double
0.05
, enableStringOps :: Bool
enableStringOps = Bool
True
, enableCrossCols :: Bool
enableCrossCols = Bool
True
, enableArithOps :: Bool
enableArithOps = Bool
True
, maxCategoricalSubsetCardinality :: Int
maxCategoricalSubsetCardinality = Int
4
, perColumnQuota :: Maybe Int
perColumnQuota = Int -> Maybe Int
forall a. a -> Maybe a
Just Int
3
}
defaultTreeConfig :: TreeConfig
defaultTreeConfig :: TreeConfig
defaultTreeConfig =
TreeConfig
{ maxTreeDepth :: Int
maxTreeDepth = Int
4
, minSamplesSplit :: Int
minSamplesSplit = Int
5
, minLeafSize :: Int
minLeafSize = Int
1
, percentiles :: [Int]
percentiles = [Int
0, Int
10 .. Int
100]
, expressionPairs :: Int
expressionPairs = Int
10
, synthConfig :: SynthConfig
synthConfig = SynthConfig
defaultSynthConfig
, taoIterations :: Int
taoIterations = Int
10
, taoConvergenceTol :: Double
taoConvergenceTol = Double
1e-6
, columnOrdering :: ColumnOrdering
columnOrdering = ColumnOrdering
defaultColumnOrdering
, useLinearSolver :: Bool
useLinearSolver = Bool
True
, linearSolverConfig :: SolverConfig
linearSolverConfig = SolverConfig
LS.defaultSolverConfig
, minCarePointsForLinear :: Int
minCarePointsForLinear = Int
10
, pureReplacementLinear :: Bool
pureReplacementLinear = Bool
False
}
newtype ColumnOrdering = ColumnOrdering (M.Map SomeTypeRep OrdDict)
instance Semigroup ColumnOrdering where
ColumnOrdering Map SomeTypeRep OrdDict
a <> :: ColumnOrdering -> ColumnOrdering -> ColumnOrdering
<> ColumnOrdering Map SomeTypeRep OrdDict
b = Map SomeTypeRep OrdDict -> ColumnOrdering
ColumnOrdering (Map SomeTypeRep OrdDict
a Map SomeTypeRep OrdDict
-> Map SomeTypeRep OrdDict -> Map SomeTypeRep OrdDict
forall a. Semigroup a => a -> a -> a
<> Map SomeTypeRep OrdDict
b)
instance Monoid ColumnOrdering where
mempty :: ColumnOrdering
mempty = Map SomeTypeRep OrdDict -> ColumnOrdering
ColumnOrdering Map SomeTypeRep OrdDict
forall k a. Map k a
M.empty
orderable :: forall a. (Columnable a, Ord a) => ColumnOrdering
orderable :: forall a. (Columnable a, Ord a) => ColumnOrdering
orderable = Map SomeTypeRep OrdDict -> ColumnOrdering
ColumnOrdering (SomeTypeRep -> OrdDict -> Map SomeTypeRep OrdDict
forall k a. k -> a -> Map k a
M.singleton (TypeRep a -> SomeTypeRep
forall k (a :: k). TypeRep a -> SomeTypeRep
SomeTypeRep (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @a)) (Proxy a -> OrdDict
forall a. (Columnable a, Ord a) => Proxy a -> OrdDict
OrdDict (forall t. Proxy t
forall {k} (t :: k). Proxy t
Proxy @a)))
defaultColumnOrdering :: ColumnOrdering
defaultColumnOrdering :: ColumnOrdering
defaultColumnOrdering = [ColumnOrdering] -> ColumnOrdering
forall a. Monoid a => [a] -> a
mconcat ([ColumnOrdering]
numericOrderings [ColumnOrdering] -> [ColumnOrdering] -> [ColumnOrdering]
forall a. [a] -> [a] -> [a]
++ [ColumnOrdering]
otherOrderings)
numericOrderings :: [ColumnOrdering]
numericOrderings :: [ColumnOrdering]
numericOrderings =
[ forall a. (Columnable a, Ord a) => ColumnOrdering
orderable @Int
, forall a. (Columnable a, Ord a) => ColumnOrdering
orderable @Int8
, forall a. (Columnable a, Ord a) => ColumnOrdering
orderable @Int16
, forall a. (Columnable a, Ord a) => ColumnOrdering
orderable @Int32
, forall a. (Columnable a, Ord a) => ColumnOrdering
orderable @Int64
, forall a. (Columnable a, Ord a) => ColumnOrdering
orderable @Word
, forall a. (Columnable a, Ord a) => ColumnOrdering
orderable @Word8
, forall a. (Columnable a, Ord a) => ColumnOrdering
orderable @Word16
, forall a. (Columnable a, Ord a) => ColumnOrdering
orderable @Word32
, forall a. (Columnable a, Ord a) => ColumnOrdering
orderable @Word64
, forall a. (Columnable a, Ord a) => ColumnOrdering
orderable @Integer
, forall a. (Columnable a, Ord a) => ColumnOrdering
orderable @Double
, forall a. (Columnable a, Ord a) => ColumnOrdering
orderable @Float
]
otherOrderings :: [ColumnOrdering]
otherOrderings :: [ColumnOrdering]
otherOrderings =
[forall a. (Columnable a, Ord a) => ColumnOrdering
orderable @Bool, forall a. (Columnable a, Ord a) => ColumnOrdering
orderable @Char, forall a. (Columnable a, Ord a) => ColumnOrdering
orderable @T.Text, forall a. (Columnable a, Ord a) => ColumnOrdering
orderable @String]
data OrdDict where
OrdDict :: (Columnable a, Ord a) => Proxy a -> OrdDict
withOrdFrom ::
forall a r. (Columnable a) => ColumnOrdering -> ((Ord a) => r) -> Maybe r
withOrdFrom :: forall a r.
Columnable a =>
ColumnOrdering -> (Ord a => r) -> Maybe r
withOrdFrom (ColumnOrdering Map SomeTypeRep OrdDict
m) Ord a => r
k = case SomeTypeRep -> Map SomeTypeRep OrdDict -> Maybe OrdDict
forall k a. Ord k => k -> Map k a -> Maybe a
M.lookup (TypeRep a -> SomeTypeRep
forall k (a :: k). TypeRep a -> SomeTypeRep
SomeTypeRep (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @a)) Map SomeTypeRep OrdDict
m of
Just (OrdDict (Proxy a
_ :: Proxy b)) -> case TypeRep a -> TypeRep a -> Maybe (a :~: a)
forall a b. TypeRep a -> TypeRep b -> Maybe (a :~: b)
forall {k} (f :: k -> *) (a :: k) (b :: k).
TestEquality f =>
f a -> f b -> Maybe (a :~: b)
testEquality (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @a) (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @b) of
Just a :~: a
Refl -> r -> Maybe r
forall a. a -> Maybe a
Just r
Ord a => r
k
Maybe (a :~: a)
Nothing -> Maybe r
forall a. Maybe a
Nothing
Maybe OrdDict
Nothing -> Maybe r
forall a. Maybe a
Nothing
{-# NOINLINE taoTraceEnabled #-}
taoTraceEnabled :: Bool
taoTraceEnabled :: Bool
taoTraceEnabled = IO Bool -> Bool
forall a. IO a -> a
unsafePerformIO ((Maybe String -> Bool) -> IO (Maybe String) -> IO Bool
forall a b. (a -> b) -> IO a -> IO b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap (Maybe String -> Maybe String -> Bool
forall a. Eq a => a -> a -> Bool
== String -> Maybe String
forall a. a -> Maybe a
Just String
"1") (String -> IO (Maybe String)
lookupEnv String
"TAO_TRACE"))
ttrace :: String -> a -> a
ttrace :: forall a. String -> a -> a
ttrace String
msg a
x
| Bool
taoTraceEnabled = String -> a -> a
forall a. String -> a -> a
Trace.trace (String
"[TAO] " String -> ShowS
forall a. [a] -> [a] -> [a]
++ String
msg) a
x
| Bool
otherwise = a
x