{-# LANGUAGE FlexibleContexts #-}
{-# LANGUAGE GADTs #-}
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications #-}
module DataFrame.DecisionTree.Cart (
CartFeature (..),
CartNode (..),
sortIndicesByValue,
buildCartTree,
cartFeatures,
cartTargetLabels,
) where
import DataFrame.DecisionTree.Types (Tree (..), TreeConfig (..))
import DataFrame.Errors (DataFrameException (..), TypeErrorContext (..))
import qualified DataFrame.Functions as F
import DataFrame.Internal.Column
import DataFrame.Internal.DataFrame (
DataFrame,
columnNames,
getColumn,
unsafeGetColumn,
)
import DataFrame.Internal.Expression (Expr (..))
import DataFrame.Internal.Interpreter (interpret)
import DataFrame.Internal.Types
import DataFrame.Operations.Core (nRows)
import DataFrame.Operators
import Control.Exception (throw)
import Data.Either (fromRight)
import Data.Function (on)
import Data.List (foldl')
import qualified Data.Map.Strict as M
import qualified Data.Set as Set
import qualified Data.Text as T
import Data.Type.Equality (testEquality, (:~:) (..))
import qualified Data.Vector as V
import qualified Data.Vector.Algorithms.Merge as VA
import qualified Data.Vector.Unboxed as VU
import Type.Reflection (TypeRep, typeRep)
data CartFeature = CartFeature
{ CartFeature -> Vector Double
cfValues :: !(VU.Vector Double)
, CartFeature -> Double -> Expr Bool
cfPred :: !(Double -> Expr Bool)
}
data CartNode = CLeaf !Int | CSplit !Int !Double !CartNode !CartNode
data CartCtx = CartCtx
{ CartCtx -> Vector CartFeature
ctxFeats :: !(V.Vector CartFeature)
, CartCtx -> Int
ctxNFeats :: !Int
, CartCtx -> Vector Int
ctxCodes :: !(VU.Vector Int)
, CartCtx -> Int
ctxNClasses :: !Int
, CartCtx -> Int
ctxMaxDepth :: !Int
, CartCtx -> Int
ctxMinLeaf :: !Int
}
sortIndicesByValue :: VU.Vector Double -> VU.Vector Int
sortIndicesByValue :: Vector Double -> Vector Int
sortIndicesByValue Vector Double
vs =
(forall s. ST s (MVector s Int)) -> Vector Int
forall a. Unbox a => (forall s. ST s (MVector s a)) -> Vector a
VU.create ((forall s. ST s (MVector s Int)) -> Vector Int)
-> (forall s. ST s (MVector s Int)) -> Vector Int
forall a b. (a -> b) -> a -> b
$ do
MVector s Int
mv <- Vector Int -> ST s (MVector (PrimState (ST s)) Int)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
Vector a -> m (MVector (PrimState m) a)
VU.thaw (Int -> Int -> Vector Int
forall a. (Unbox a, Num a) => a -> Int -> Vector a
VU.enumFromN Int
0 (Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Double
vs))
Comparison Int -> MVector (PrimState (ST s)) Int -> ST s ()
forall (m :: * -> *) (v :: * -> * -> *) e.
(PrimMonad m, MVector v e) =>
Comparison e -> v (PrimState m) e -> m ()
VA.sortBy (Double -> Double -> Ordering
forall a. Ord a => a -> a -> Ordering
compare (Double -> Double -> Ordering) -> (Int -> Double) -> Comparison Int
forall b c a. (b -> b -> c) -> (a -> b) -> a -> a -> c
`on` (Vector Double
vs Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.!)) MVector s Int
MVector (PrimState (ST s)) Int
mv
MVector s Int -> ST s (MVector s Int)
forall a. a -> ST s a
forall (f :: * -> *) a. Applicative f => a -> f a
pure MVector s Int
mv
buildCartTree ::
forall a. (Columnable a, Ord a) => TreeConfig -> T.Text -> DataFrame -> Tree a
buildCartTree :: forall a.
(Columnable a, Ord a) =>
TreeConfig -> Text -> DataFrame -> Tree a
buildCartTree TreeConfig
cfg Text
target DataFrame
df =
Vector CartFeature -> Vector a -> CartNode -> Tree a
forall a. Vector CartFeature -> Vector a -> CartNode -> Tree a
cartToTree Vector CartFeature
feats Vector a
classes (CartCtx -> Int -> Vector Int -> Vector (Vector Int) -> CartNode
buildCartNode CartCtx
ctx Int
0 (Int -> Int -> Vector Int
forall a. (Unbox a, Num a) => a -> Int -> Vector a
VU.enumFromN Int
0 Int
nAll) Vector (Vector Int)
featSorted)
where
nAll :: Int
nAll = DataFrame -> Int
nRows DataFrame
df
feats :: Vector CartFeature
feats = [CartFeature] -> Vector CartFeature
forall a. [a] -> Vector a
V.fromList (Text -> DataFrame -> [CartFeature]
cartFeatures Text
target DataFrame
df)
featSorted :: Vector (Vector Int)
featSorted = (CartFeature -> Vector Int)
-> Vector CartFeature -> Vector (Vector Int)
forall a b. (a -> b) -> Vector a -> Vector b
V.map (Vector Double -> Vector Int
sortIndicesByValue (Vector Double -> Vector Int)
-> (CartFeature -> Vector Double) -> CartFeature -> Vector Int
forall b c a. (b -> c) -> (a -> b) -> a -> c
. CartFeature -> Vector Double
cfValues) Vector CartFeature
feats
labels :: Vector a
labels = forall a. Columnable a => DataFrame -> Text -> Vector a
cartLabels @a DataFrame
df Text
target
classes :: Vector a
classes = Vector a -> Vector a
forall a. Ord a => Vector a -> Vector a
cartClasses Vector a
labels
ctx :: CartCtx
ctx =
Vector CartFeature
-> Int -> Vector Int -> Int -> Int -> Int -> CartCtx
CartCtx
Vector CartFeature
feats
(Vector CartFeature -> Int
forall a. Vector a -> Int
V.length Vector CartFeature
feats)
(Vector a -> Vector a -> Vector Int
forall a. Ord a => Vector a -> Vector a -> Vector Int
classCodes Vector a
classes Vector a
labels)
(Vector a -> Int
forall a. Vector a -> Int
V.length Vector a
classes)
(TreeConfig -> Int
maxTreeDepth TreeConfig
cfg)
(Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
1 (TreeConfig -> Int
minLeafSize TreeConfig
cfg))
cartLabels :: forall a. (Columnable a) => DataFrame -> T.Text -> V.Vector a
cartLabels :: forall a. Columnable a => DataFrame -> Text -> Vector a
cartLabels DataFrame
df Text
target = case forall a.
Columnable a =>
DataFrame -> Expr a -> Either DataFrameException (TypedColumn a)
interpret @a DataFrame
df (Text -> Expr a
forall a. Columnable a => Text -> Expr a
Col Text
target) of
Right (TColumn Column
column) -> Vector a -> Either DataFrameException (Vector a) -> Vector a
forall b a. b -> Either a b -> b
fromRight (DataFrameException -> Vector a
forall a e. Exception e => e -> a
throw DataFrameException
err) (forall a (v :: * -> *).
(Vector v a, Columnable a) =>
Column -> Either DataFrameException (v a)
toVector @a Column
column)
Left DataFrameException
e -> DataFrameException -> Vector a
forall a e. Exception e => e -> a
throw DataFrameException
e
where
err :: DataFrameException
err =
TypeErrorContext a a -> DataFrameException
forall a b.
(Typeable a, Typeable b) =>
TypeErrorContext a b -> DataFrameException
TypeMismatchException
( Either String (TypeRep a)
-> Either String (TypeRep a)
-> Maybe String
-> Maybe String
-> TypeErrorContext a a
forall a b.
Either String (TypeRep a)
-> Either String (TypeRep b)
-> Maybe String
-> Maybe String
-> TypeErrorContext a b
MkTypeErrorContext
(TypeRep a -> Either String (TypeRep a)
forall a b. b -> Either a b
Right (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @a))
( String -> Either String (TypeRep a)
forall a b. a -> Either a b
Left (String -> (Column -> String) -> Maybe Column -> String
forall b a. b -> (a -> b) -> Maybe a -> b
maybe String
"missing" Column -> String
columnTypeString (Text -> DataFrame -> Maybe Column
getColumn Text
target DataFrame
df)) ::
Either String (TypeRep a)
)
(String -> Maybe String
forall a. a -> Maybe a
Just (Text -> String
T.unpack Text
target))
(String -> Maybe String
forall a. a -> Maybe a
Just String
"buildCartTree")
)
cartClasses :: (Ord a) => V.Vector a -> V.Vector a
cartClasses :: forall a. Ord a => Vector a -> Vector a
cartClasses = [a] -> Vector a
forall a. [a] -> Vector a
V.fromList ([a] -> Vector a) -> (Vector a -> [a]) -> Vector a -> Vector a
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Set a -> [a]
forall a. Set a -> [a]
Set.toList (Set a -> [a]) -> (Vector a -> Set a) -> Vector a -> [a]
forall b c a. (b -> c) -> (a -> b) -> a -> c
. [a] -> Set a
forall a. Ord a => [a] -> Set a
Set.fromList ([a] -> Set a) -> (Vector a -> [a]) -> Vector a -> Set a
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Vector a -> [a]
forall a. Vector a -> [a]
V.toList
classCodes :: (Ord a) => V.Vector a -> V.Vector a -> VU.Vector Int
classCodes :: forall a. Ord a => Vector a -> Vector a -> Vector Int
classCodes Vector a
classes Vector a
labels = Int -> (Int -> Int) -> Vector Int
forall a. Unbox a => Int -> (Int -> a) -> Vector a
VU.generate (Vector a -> Int
forall a. Vector a -> Int
V.length Vector a
labels) (\Int
i -> Int -> a -> Map a Int -> Int
forall k a. Ord k => a -> k -> Map k a -> a
M.findWithDefault Int
0 (Vector a
labels Vector a -> Int -> a
forall a. Vector a -> Int -> a
V.! Int
i) Map a Int
ix)
where
ix :: Map a Int
ix = [(a, Int)] -> Map a Int
forall k a. Ord k => [(k, a)] -> Map k a
M.fromList ([a] -> [Int] -> [(a, Int)]
forall a b. [a] -> [b] -> [(a, b)]
zip (Vector a -> [a]
forall a. Vector a -> [a]
V.toList Vector a
classes) [Int
0 ..])
cartToTree :: V.Vector CartFeature -> V.Vector a -> CartNode -> Tree a
cartToTree :: forall a. Vector CartFeature -> Vector a -> CartNode -> Tree a
cartToTree Vector CartFeature
feats Vector a
classes = CartNode -> Tree a
go
where
go :: CartNode -> Tree a
go (CLeaf Int
cid) = a -> Tree a
forall a. a -> Tree a
Leaf (Vector a
classes Vector a -> Int -> a
forall a. Vector a -> Int -> a
V.! Int
cid)
go (CSplit Int
fj Double
thr CartNode
l CartNode
r) = Expr Bool -> Tree a -> Tree a -> Tree a
forall a. Expr Bool -> Tree a -> Tree a -> Tree a
Branch (CartFeature -> Double -> Expr Bool
cfPred (Vector CartFeature
feats Vector CartFeature -> Int -> CartFeature
forall a. Vector a -> Int -> a
V.! Int
fj) Double
thr) (CartNode -> Tree a
go CartNode
l) (CartNode -> Tree a
go CartNode
r)
classCounts :: CartCtx -> VU.Vector Int -> VU.Vector Int
classCounts :: CartCtx -> Vector Int -> Vector Int
classCounts CartCtx
ctx Vector Int
idxs =
(Int -> Int -> Int)
-> Vector Int -> Vector (Int, Int) -> Vector Int
forall a b.
(Unbox a, Unbox b) =>
(a -> b -> a) -> Vector a -> Vector (Int, b) -> Vector a
VU.accumulate
Int -> Int -> Int
forall a. Num a => a -> a -> a
(+)
(Int -> Int -> Vector Int
forall a. Unbox a => Int -> a -> Vector a
VU.replicate (CartCtx -> Int
ctxNClasses CartCtx
ctx) Int
0)
((Int -> (Int, Int)) -> Vector Int -> Vector (Int, Int)
forall a b. (Unbox a, Unbox b) => (a -> b) -> Vector a -> Vector b
VU.map (\Int
i -> (CartCtx -> Vector Int
ctxCodes CartCtx
ctx Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i, Int
1)) Vector Int
idxs)
isPure :: VU.Vector Int -> Bool
isPure :: Vector Int -> Bool
isPure Vector Int
counts = Vector Int -> Int
forall a. Unbox a => Vector a -> Int
VU.length ((Int -> Bool) -> Vector Int -> Vector Int
forall a. Unbox a => (a -> Bool) -> Vector a -> Vector a
VU.filter (Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Int
0) Vector Int
counts) Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
<= Int
1
buildCartNode ::
CartCtx -> Int -> VU.Vector Int -> V.Vector (VU.Vector Int) -> CartNode
buildCartNode :: CartCtx -> Int -> Vector Int -> Vector (Vector Int) -> CartNode
buildCartNode CartCtx
ctx Int
depth Vector Int
idxs Vector (Vector Int)
sortedByFeat
| Vector Int -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Int
idxs Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
2 Bool -> Bool -> Bool
|| Int
depth Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= CartCtx -> Int
ctxMaxDepth CartCtx
ctx Bool -> Bool -> Bool
|| Vector Int -> Bool
isPure Vector Int
counts = CartNode
leaf
| Bool
otherwise =
CartNode
-> ((Int, Double) -> CartNode) -> Maybe (Int, Double) -> CartNode
forall b a. b -> (a -> b) -> Maybe a -> b
maybe
CartNode
leaf
(CartCtx
-> Int
-> Vector Int
-> Vector (Vector Int)
-> (Int, Double)
-> CartNode
splitNode CartCtx
ctx Int
depth Vector Int
idxs Vector (Vector Int)
sortedByFeat)
(CartCtx
-> Vector (Vector Int) -> Vector Int -> Int -> Maybe (Int, Double)
bestSplit CartCtx
ctx Vector (Vector Int)
sortedByFeat Vector Int
counts Int
n)
where
n :: Int
n = Vector Int -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Int
idxs
counts :: Vector Int
counts = CartCtx -> Vector Int -> Vector Int
classCounts CartCtx
ctx Vector Int
idxs
leaf :: CartNode
leaf = Int -> CartNode
CLeaf (Vector Int -> Int
forall a. (Unbox a, Ord a) => Vector a -> Int
VU.maxIndex Vector Int
counts)
splitNode ::
CartCtx ->
Int ->
VU.Vector Int ->
V.Vector (VU.Vector Int) ->
(Int, Double) ->
CartNode
splitNode :: CartCtx
-> Int
-> Vector Int
-> Vector (Vector Int)
-> (Int, Double)
-> CartNode
splitNode CartCtx
ctx Int
depth Vector Int
idxs Vector (Vector Int)
sortedByFeat (Int
fj, Double
thr) =
Int -> Double -> CartNode -> CartNode -> CartNode
CSplit Int
fj Double
thr (Vector Int -> Vector (Vector Int) -> CartNode
rec Vector Int
leftIdx Vector (Vector Int)
leftSorted) (Vector Int -> Vector (Vector Int) -> CartNode
rec Vector Int
rightIdx Vector (Vector Int)
rightSorted)
where
vals :: Vector Double
vals = CartFeature -> Vector Double
cfValues (CartCtx -> Vector CartFeature
ctxFeats CartCtx
ctx Vector CartFeature -> Int -> CartFeature
forall a. Vector a -> Int -> a
V.! Int
fj)
leftIdx :: Vector Int
leftIdx = (Int -> Bool) -> Vector Int -> Vector Int
forall a. Unbox a => (a -> Bool) -> Vector a -> Vector a
VU.filter (\Int
i -> Vector Double
vals Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
<= Double
thr) Vector Int
idxs
rightIdx :: Vector Int
rightIdx = (Int -> Bool) -> Vector Int -> Vector Int
forall a. Unbox a => (a -> Bool) -> Vector a -> Vector a
VU.filter (\Int
i -> Vector Double
vals Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
> Double
thr) Vector Int
idxs
leftSorted :: Vector (Vector Int)
leftSorted = (Vector Int -> Vector Int)
-> Vector (Vector Int) -> Vector (Vector Int)
forall a b. (a -> b) -> Vector a -> Vector b
V.map ((Int -> Bool) -> Vector Int -> Vector Int
forall a. Unbox a => (a -> Bool) -> Vector a -> Vector a
VU.filter (\Int
i -> Vector Double
vals Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
<= Double
thr)) Vector (Vector Int)
sortedByFeat
rightSorted :: Vector (Vector Int)
rightSorted = (Vector Int -> Vector Int)
-> Vector (Vector Int) -> Vector (Vector Int)
forall a b. (a -> b) -> Vector a -> Vector b
V.map ((Int -> Bool) -> Vector Int -> Vector Int
forall a. Unbox a => (a -> Bool) -> Vector a -> Vector a
VU.filter (\Int
i -> Vector Double
vals Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
> Double
thr)) Vector (Vector Int)
sortedByFeat
rec :: Vector Int -> Vector (Vector Int) -> CartNode
rec = CartCtx -> Int -> Vector Int -> Vector (Vector Int) -> CartNode
buildCartNode CartCtx
ctx (Int
depth Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
bestSplit ::
CartCtx ->
V.Vector (VU.Vector Int) ->
VU.Vector Int ->
Int ->
Maybe (Int, Double)
bestSplit :: CartCtx
-> Vector (Vector Int) -> Vector Int -> Int -> Maybe (Int, Double)
bestSplit CartCtx
ctx Vector (Vector Int)
sortedByFeat Vector Int
counts Int
n =
((Double, Int, Double) -> (Int, Double))
-> Maybe (Double, Int, Double) -> Maybe (Int, Double)
forall a b. (a -> b) -> Maybe a -> Maybe b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap (\(Double
_, Int
j, Double
t) -> (Int
j, Double
t)) ((Maybe (Double, Int, Double) -> Int -> Maybe (Double, Int, Double))
-> Maybe (Double, Int, Double)
-> [Int]
-> Maybe (Double, Int, Double)
forall b a. (b -> a -> b) -> b -> [a] -> b
forall (t :: * -> *) b a.
Foldable t =>
(b -> a -> b) -> b -> t a -> b
foldl' Maybe (Double, Int, Double) -> Int -> Maybe (Double, Int, Double)
consider Maybe (Double, Int, Double)
forall a. Maybe a
Nothing [Int
0 .. CartCtx -> Int
ctxNFeats CartCtx
ctx Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1])
where
total :: [Int]
total = Vector Int -> [Int]
forall a. Unbox a => Vector a -> [a]
VU.toList Vector Int
counts
consider :: Maybe (Double, Int, Double) -> Int -> Maybe (Double, Int, Double)
consider Maybe (Double, Int, Double)
acc Int
fj = case CartCtx
-> [Int]
-> Vector Int
-> CartFeature
-> Int
-> Maybe (Double, Double)
sweepFeature CartCtx
ctx [Int]
total (Vector (Vector Int)
sortedByFeat Vector (Vector Int) -> Int -> Vector Int
forall a. Vector a -> Int -> a
V.! Int
fj) (CartCtx -> Vector CartFeature
ctxFeats CartCtx
ctx Vector CartFeature -> Int -> CartFeature
forall a. Vector a -> Int -> a
V.! Int
fj) Int
n of
Just (Double
g, Double
thr) | Bool
-> ((Double, Int, Double) -> Bool)
-> Maybe (Double, Int, Double)
-> Bool
forall b a. b -> (a -> b) -> Maybe a -> b
maybe Bool
True (\(Double
gB, Int
_, Double
_) -> Double
g Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
< Double
gB) Maybe (Double, Int, Double)
acc -> (Double, Int, Double) -> Maybe (Double, Int, Double)
forall a. a -> Maybe a
Just (Double
g, Int
fj, Double
thr)
Maybe (Double, Double)
_ -> Maybe (Double, Int, Double)
acc
data Sweep = Sweep
{ Sweep -> Maybe (Double, Double)
swBest :: !(Maybe (Double, Double))
, Sweep -> [Int]
swLeft :: ![Int]
, Sweep -> Int
swMoved :: !Int
, Sweep -> Double
swPrev :: !Double
}
sweepFeature ::
CartCtx ->
[Int] ->
VU.Vector Int ->
CartFeature ->
Int ->
Maybe (Double, Double)
sweepFeature :: CartCtx
-> [Int]
-> Vector Int
-> CartFeature
-> Int
-> Maybe (Double, Double)
sweepFeature CartCtx
ctx [Int]
total Vector Int
si CartFeature
feat Int
n =
Sweep -> Maybe (Double, Double)
swBest
( (Sweep -> Int -> Sweep) -> Sweep -> [Int] -> Sweep
forall b a. (b -> a -> b) -> b -> [a] -> b
forall (t :: * -> *) b a.
Foldable t =>
(b -> a -> b) -> b -> t a -> b
foldl'
Sweep -> Int -> Sweep
step
(Maybe (Double, Double) -> [Int] -> Int -> Double -> Sweep
Sweep Maybe (Double, Double)
forall a. Maybe a
Nothing (Int -> Int -> [Int]
forall a. Int -> a -> [a]
replicate (CartCtx -> Int
ctxNClasses CartCtx
ctx) Int
0) Int
0 (Double
0 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
0))
[Int
0 .. Vector Int -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Int
si Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]
)
where
vals :: Vector Double
vals = CartFeature -> Vector Double
cfValues CartFeature
feat
step :: Sweep -> Int -> Sweep
step Sweep
s Int
k = CartCtx -> [Int] -> Int -> Double -> Int -> Sweep -> Sweep
advance CartCtx
ctx [Int]
total Int
n (Vector Double
vals Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i) (CartCtx -> Vector Int
ctxCodes CartCtx
ctx Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i) Sweep
s
where
i :: Int
i = Vector Int
si Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
k
advance :: CartCtx -> [Int] -> Int -> Double -> Int -> Sweep -> Sweep
advance :: CartCtx -> [Int] -> Int -> Double -> Int -> Sweep -> Sweep
advance CartCtx
ctx [Int]
total Int
n Double
v Int
c Sweep
s =
Maybe (Double, Double) -> [Int] -> Int -> Double -> Sweep
Sweep
(CartCtx
-> [Int] -> Int -> Double -> Sweep -> Maybe (Double, Double)
considerThreshold CartCtx
ctx [Int]
total Int
n Double
v Sweep
s)
(Int -> [Int] -> [Int]
bumpClass Int
c (Sweep -> [Int]
swLeft Sweep
s))
(Sweep -> Int
swMoved Sweep
s Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
Double
v
considerThreshold ::
CartCtx -> [Int] -> Int -> Double -> Sweep -> Maybe (Double, Double)
considerThreshold :: CartCtx
-> [Int] -> Int -> Double -> Sweep -> Maybe (Double, Double)
considerThreshold CartCtx
ctx [Int]
total Int
n Double
v Sweep
s
| Sweep -> Int
swMoved Sweep
s Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= CartCtx -> Int
ctxMinLeaf CartCtx
ctx
, Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Sweep -> Int
swMoved Sweep
s Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= CartCtx -> Int
ctxMinLeaf CartCtx
ctx
, Double
v Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
> Sweep -> Double
swPrev Sweep
s Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
1e-7 =
Maybe (Double, Double)
-> Double -> Double -> Maybe (Double, Double)
keepBetter
(Sweep -> Maybe (Double, Double)
swBest Sweep
s)
([Int] -> [Int] -> Int -> Int -> Double
weightedGini [Int]
total (Sweep -> [Int]
swLeft Sweep
s) (Sweep -> Int
swMoved Sweep
s) Int
n)
((Sweep -> Double
swPrev Sweep
s Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
v) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
2)
| Bool
otherwise = Sweep -> Maybe (Double, Double)
swBest Sweep
s
keepBetter ::
Maybe (Double, Double) -> Double -> Double -> Maybe (Double, Double)
keepBetter :: Maybe (Double, Double)
-> Double -> Double -> Maybe (Double, Double)
keepBetter Maybe (Double, Double)
best Double
g Double
thr = case Maybe (Double, Double)
best of
Just (Double
wb, Double
_) | Double
wb Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
<= Double
g -> Maybe (Double, Double)
best
Maybe (Double, Double)
_ -> (Double, Double) -> Maybe (Double, Double)
forall a. a -> Maybe a
Just (Double
g, Double
thr)
weightedGini :: [Int] -> [Int] -> Int -> Int -> Double
weightedGini :: [Int] -> [Int] -> Int -> Int -> Double
weightedGini [Int]
total [Int]
leftAcc Int
nl Int
n =
( Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
nl Double -> Double -> Double
forall a. Num a => a -> a -> a
* [Int] -> Int -> Double
giniImpurity [Int]
leftAcc Int
nl
Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
nr Double -> Double -> Double
forall a. Num a => a -> a -> a
* [Int] -> Int -> Double
giniImpurity [Int]
rightAcc Int
nr
)
Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
n
where
nr :: Int
nr = Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
nl
rightAcc :: [Int]
rightAcc = (Int -> Int -> Int) -> [Int] -> [Int] -> [Int]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith (-) [Int]
total [Int]
leftAcc
giniImpurity :: [Int] -> Int -> Double
giniImpurity :: [Int] -> Int -> Double
giniImpurity [Int]
_ Int
0 = Double
0
giniImpurity [Int]
cs Int
m = Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- [Double] -> Double
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [let p :: Double
p = Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
c Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
m in Double
p Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
p | Int
c <- [Int]
cs]
bumpClass :: Int -> [Int] -> [Int]
bumpClass :: Int -> [Int] -> [Int]
bumpClass Int
c = (Int -> Int -> Int) -> [Int] -> [Int] -> [Int]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith (\Int
j Int
x -> if Int
j Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
c then Int
x Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1 else Int
x) [Int
0 ..]
cartFeatures :: T.Text -> DataFrame -> [CartFeature]
cartFeatures :: Text -> DataFrame -> [CartFeature]
cartFeatures Text
target DataFrame
df = (Text -> [CartFeature]) -> [Text] -> [CartFeature]
forall (t :: * -> *) a b. Foldable t => (a -> [b]) -> t a -> [b]
concatMap (DataFrame -> Text -> [CartFeature]
featuresOfColumn DataFrame
df) ((Text -> Bool) -> [Text] -> [Text]
forall a. (a -> Bool) -> [a] -> [a]
filter (Text -> Text -> Bool
forall a. Eq a => a -> a -> Bool
/= Text
target) (DataFrame -> [Text]
columnNames DataFrame
df))
featuresOfColumn :: DataFrame -> T.Text -> [CartFeature]
featuresOfColumn :: DataFrame -> Text -> [CartFeature]
featuresOfColumn DataFrame
df Text
c = case Text -> DataFrame -> Column
unsafeGetColumn Text
c DataFrame
df of
UnboxedColumn Maybe Bitmap
_ (Vector a
v :: VU.Vector b) -> forall b.
(Columnable b, Unbox b) =>
Text -> Vector b -> [CartFeature]
numericFeature @b Text
c Vector a
v
BoxedColumn Maybe Bitmap
_ (Vector a
v :: V.Vector b) -> forall b. Columnable b => Int -> Text -> Vector b -> [CartFeature]
oneHotFeatures @b (DataFrame -> Int
nRows DataFrame
df) Text
c Vector a
v
pt :: Column
pt@(PackedText Maybe Bitmap
_ PackedTextData
_) -> case Column -> Column
materializePacked Column
pt of
BoxedColumn Maybe Bitmap
_ (Vector a
v :: V.Vector b) -> forall b. Columnable b => Int -> Text -> Vector b -> [CartFeature]
oneHotFeatures @b (DataFrame -> Int
nRows DataFrame
df) Text
c Vector a
v
Column
_ -> []
mc :: Column
mc@(MergedColumn Column
_ Column
_) -> case Column -> Column
materializeMerged Column
mc of
BoxedColumn Maybe Bitmap
_ (Vector a
v :: V.Vector b) -> forall b. Columnable b => Int -> Text -> Vector b -> [CartFeature]
oneHotFeatures @b (DataFrame -> Int
nRows DataFrame
df) Text
c Vector a
v
Column
_ -> []
numericFeature ::
forall b. (Columnable b, VU.Unbox b) => T.Text -> VU.Vector b -> [CartFeature]
numericFeature :: forall b.
(Columnable b, Unbox b) =>
Text -> Vector b -> [CartFeature]
numericFeature Text
c Vector b
v = case TypeRep b -> TypeRep Double -> Maybe (b :~: Double)
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 @b) (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @Double) of
Just b :~: Double
Refl -> [Vector Double -> (Double -> Expr Bool) -> CartFeature
CartFeature Vector b
Vector Double
v (\Double
t -> forall a. Columnable a => Text -> Expr a
F.col @Double Text
c Expr Double -> Expr Double -> Expr Bool
forall a. (Columnable a, Ord a) => Expr a -> Expr a -> Expr Bool
.<=. Double -> Expr Double
forall a. Columnable a => a -> Expr a
F.lit Double
t)]
Maybe (b :~: Double)
Nothing -> case forall a. SBoolI (IntegralTypes a) => SBool (IntegralTypes a)
sIntegral @b of
SBool (IntegralTypes b)
STrue ->
[ Vector Double -> (Double -> Expr Bool) -> CartFeature
CartFeature ((b -> Double) -> Vector b -> Vector Double
forall a b. (Unbox a, Unbox b) => (a -> b) -> Vector a -> Vector b
VU.map b -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Vector b
v) (\Double
t -> Expr b -> Expr Double
forall a. (Columnable a, Real a) => Expr a -> Expr Double
F.toDouble (forall a. Columnable a => Text -> Expr a
F.col @b Text
c) Expr Double -> Expr Double -> Expr Bool
forall a. (Columnable a, Ord a) => Expr a -> Expr a -> Expr Bool
.<=. Double -> Expr Double
forall a. Columnable a => a -> Expr a
F.lit Double
t)
]
SBool (IntegralTypes b)
SFalse -> []
oneHotFeatures ::
forall b. (Columnable b) => Int -> T.Text -> V.Vector b -> [CartFeature]
oneHotFeatures :: forall b. Columnable b => Int -> Text -> Vector b -> [CartFeature]
oneHotFeatures Int
nAll Text
c Vector b
v = case TypeRep b -> TypeRep Text -> Maybe (b :~: Text)
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 @b) (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @T.Text) of
Just b :~: Text
Refl -> [Int -> Text -> Vector Text -> Text -> CartFeature
oneHot Int
nAll Text
c Vector b
Vector Text
v b
Text
cat | b
cat <- Set b -> [b]
forall a. Set a -> [a]
Set.toList ([b] -> Set b
forall a. Ord a => [a] -> Set a
Set.fromList (Vector b -> [b]
forall a. Vector a -> [a]
V.toList Vector b
v))]
Maybe (b :~: Text)
Nothing -> []
oneHot :: Int -> T.Text -> V.Vector T.Text -> T.Text -> CartFeature
oneHot :: Int -> Text -> Vector Text -> Text -> CartFeature
oneHot Int
nAll Text
c Vector Text
v Text
cat =
Vector Double -> (Double -> Expr Bool) -> CartFeature
CartFeature
(Int -> (Int -> Double) -> Vector Double
forall a. Unbox a => Int -> (Int -> a) -> Vector a
VU.generate Int
nAll (\Int
i -> if Vector Text
v Vector Text -> Int -> Text
forall a. Vector a -> Int -> a
V.! Int
i Text -> Text -> Bool
forall a. Eq a => a -> a -> Bool
== Text
cat then Double
1 else Double
0))
(Expr Bool -> Double -> Expr Bool
forall a b. a -> b -> a
const (forall a. Columnable a => Text -> Expr a
F.col @T.Text Text
c Expr Text -> Expr Text -> Expr Bool
forall a. (Columnable a, Eq a) => Expr a -> Expr a -> Expr Bool
./=. Text -> Expr Text
forall a. Columnable a => a -> Expr a
F.lit Text
cat))
cartTargetLabels :: T.Text -> DataFrame -> V.Vector T.Text
cartTargetLabels :: Text -> DataFrame -> Vector Text
cartTargetLabels Text
target DataFrame
df = case Text -> DataFrame -> Column
unsafeGetColumn Text
target DataFrame
df of
BoxedColumn Maybe Bitmap
_ (Vector a
v :: V.Vector b) -> case TypeRep a -> TypeRep Text -> Maybe (a :~: Text)
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 @b) (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @T.Text) of
Just a :~: Text
Refl -> Vector a
Vector Text
v
Maybe (a :~: Text)
Nothing -> (a -> Text) -> Vector a -> Vector Text
forall a b. (a -> b) -> Vector a -> Vector b
V.map (String -> Text
T.pack (String -> Text) -> (a -> String) -> a -> Text
forall b c a. (b -> c) -> (a -> b) -> a -> c
. a -> String
forall a. Show a => a -> String
show) Vector a
v
UnboxedColumn Maybe Bitmap
_ (Vector a
v :: VU.Vector b) -> (a -> Text) -> Vector a -> Vector Text
forall a b. (a -> b) -> Vector a -> Vector b
V.map (String -> Text
T.pack (String -> Text) -> (a -> String) -> a -> Text
forall b c a. (b -> c) -> (a -> b) -> a -> c
. a -> String
forall a. Show a => a -> String
show) (Vector a -> Vector a
forall (v :: * -> *) a (w :: * -> *).
(Vector v a, Vector w a) =>
v a -> w a
V.convert Vector a
v)
pt :: Column
pt@(PackedText Maybe Bitmap
_ PackedTextData
_) -> case Column -> Column
materializePacked Column
pt of
BoxedColumn Maybe Bitmap
_ (Vector a
v :: V.Vector b) -> case TypeRep a -> TypeRep Text -> Maybe (a :~: Text)
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 @b) (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @T.Text) of
Just a :~: Text
Refl -> Vector a
Vector Text
v
Maybe (a :~: Text)
Nothing -> (a -> Text) -> Vector a -> Vector Text
forall a b. (a -> b) -> Vector a -> Vector b
V.map (String -> Text
T.pack (String -> Text) -> (a -> String) -> a -> Text
forall b c a. (b -> c) -> (a -> b) -> a -> c
. a -> String
forall a. Show a => a -> String
show) Vector a
v
Column
_ -> Vector Text
forall a. Vector a
V.empty
mc :: Column
mc@(MergedColumn Column
_ Column
_) -> case Column -> Column
materializeMerged Column
mc of
BoxedColumn Maybe Bitmap
_ (Vector a
v :: V.Vector b) -> case TypeRep a -> TypeRep Text -> Maybe (a :~: Text)
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 @b) (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @T.Text) of
Just a :~: Text
Refl -> Vector a
Vector Text
v
Maybe (a :~: Text)
Nothing -> (a -> Text) -> Vector a -> Vector Text
forall a b. (a -> b) -> Vector a -> Vector b
V.map (String -> Text
T.pack (String -> Text) -> (a -> String) -> a -> Text
forall b c a. (b -> c) -> (a -> b) -> a -> c
. a -> String
forall a. Show a => a -> String
show) Vector a
v
Column
_ -> Vector Text
forall a. Vector a
V.empty