{-# LANGUAGE FlexibleContexts #-}
{-# LANGUAGE GADTs #-}
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications #-}

{- | sklearn-faithful CART initializer used to seed TAO. One-hot encodes
categoricals and splits on exact (unsmoothed) Gini over midpoint thresholds
(@<=@ routes left), matching @DecisionTreeClassifier(criterion='gini')@.
-}
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)

{- | A one-hot feature column: per-row Double values plus the sklearn LEFT
predicate (@x <= threshold@) over the ORIGINAL DataFrame.
-}
data CartFeature = CartFeature
    { CartFeature -> Vector Double
cfValues :: !(VU.Vector Double)
    , CartFeature -> Double -> Expr Bool
cfPred :: !(Double -> Expr Bool)
    }

-- | Pre-'Tree' CART node: a leaf class id, or a split on feature @j@.
data CartNode = CLeaf !Int | CSplit !Int !Double !CartNode !CartNode

-- | Immutable per-fit context for the CART recursion.
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
    }

{- | Indices @0..n-1@ stably sorted by their value (ascending), ties keeping
ascending index. In-place unboxed merge sort — no boxed-list allocation.
-}
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))

{- | Read the target column at the type the tree is being fitted at. Names the
column and both types on failure: a bare @fromIntegral@ defaults to 'Integer'
and lands here, and the old message said only that something went wrong.
-}
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)

{- | Minimum weighted-child-Gini @(feature, threshold)@; the first feature wins
ties; 'Nothing' when no feature has a leaf-size-respecting threshold.
-}
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

{- | Accumulator while sweeping a feature's sorted rows: best @(gini, thr)@ so
far, per-class left counts, rows moved left, and the previous value seen.
-}
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

-- | Gini impurity @1 - Σ (c/m)²@ of a class-count list of total @m@.
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 ..]

-- | One-hot features in @pd.get_dummies(drop_first=False)@ column order.
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))

-- | Target column as string labels (matches pandas @y.astype(str)@).
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