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

{- | Prediction, care-point identification, node validity, and tree loss. The
batched, cache-aware variants resolve each branch condition's truth vector
once per call instead of once per row.
-}
module DataFrame.DecisionTree.Predict (
    predictWithTree,
    predictManyWithTree,
    predictManyWithTreeCached,
    identifyCarePoints,
    identifyCarePointsCached,
    countCarePointErrors,
    partitionIndices,
    partitionIndicesCached,
    majorityValueFromIndices,
    computeTreeLoss,
    computeTreeLossCached,
    isValidAtNode,
) where

import DataFrame.DecisionTree.CondVec (
    CondCache,
    countErrorsByVec,
    lookupCondVec,
 )
import DataFrame.DecisionTree.Types (
    CarePoint (..),
    Direction (..),
    Tree (..),
    TreeConfig (..),
 )
import DataFrame.Errors (DataFrameException (..))
import DataFrame.Internal.Column (Columnable, TypedColumn (..), toVector)
import DataFrame.Internal.DataFrame (DataFrame)
import DataFrame.Internal.Expression (Expr (..))
import DataFrame.Internal.Interpreter (interpret)

import Control.Exception (throw)
import Control.Monad.ST (ST)
import Data.Function (on)
import Data.List (maximumBy)
import qualified Data.Map.Strict as M
import qualified Data.Text as T
import qualified Data.Vector as V
import qualified Data.Vector.Mutable as VM
import qualified Data.Vector.Unboxed as VU

{- | A condition's truth vector over the DataFrame, or 'Nothing' on a
type/interpret failure (callers default such rows to the left child).
-}
branchBool :: DataFrame -> Expr Bool -> Maybe (VU.Vector Bool)
branchBool :: DataFrame -> Expr Bool -> Maybe (Vector Bool)
branchBool DataFrame
df Expr Bool
cond = case forall a.
Columnable a =>
DataFrame -> Expr a -> Either DataFrameException (TypedColumn a)
interpret @Bool DataFrame
df Expr Bool
cond of
    Right (TColumn Column
column) -> (DataFrameException -> Maybe (Vector Bool))
-> (Vector Bool -> Maybe (Vector Bool))
-> Either DataFrameException (Vector Bool)
-> Maybe (Vector Bool)
forall a c b. (a -> c) -> (b -> c) -> Either a b -> c
either (Maybe (Vector Bool) -> DataFrameException -> Maybe (Vector Bool)
forall a b. a -> b -> a
const Maybe (Vector Bool)
forall a. Maybe a
Nothing) Vector Bool -> Maybe (Vector Bool)
forall a. a -> Maybe a
Just (forall a (v :: * -> *).
(Vector v a, Columnable a) =>
Column -> Either DataFrameException (v a)
toVector @Bool @VU.Vector Column
column)
    Either DataFrameException (TypedColumn Bool)
_ -> Maybe (Vector Bool)
forall a. Maybe a
Nothing

-- | The target column as a label vector, or 'Nothing' on failure.
interpretLabelCol ::
    forall a. (Columnable a) => DataFrame -> T.Text -> Maybe (V.Vector a)
interpretLabelCol :: forall a. Columnable a => DataFrame -> Text -> Maybe (Vector a)
interpretLabelCol 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) -> (DataFrameException -> Maybe (Vector a))
-> (Vector a -> Maybe (Vector a))
-> Either DataFrameException (Vector a)
-> Maybe (Vector a)
forall a c b. (a -> c) -> (b -> c) -> Either a b -> c
either (Maybe (Vector a) -> DataFrameException -> Maybe (Vector a)
forall a b. a -> b -> a
const Maybe (Vector a)
forall a. Maybe a
Nothing) Vector a -> Maybe (Vector a)
forall a. a -> Maybe a
Just (forall a (v :: * -> *).
(Vector v a, Columnable a) =>
Column -> Either DataFrameException (v a)
toVector @a Column
column)
    Either DataFrameException (TypedColumn a)
_ -> Maybe (Vector a)
forall a. Maybe a
Nothing

-- | Predict the label for a single row by walking a fixed tree (@True@ → left).
predictWithTree ::
    forall a. (Columnable a) => T.Text -> DataFrame -> Int -> Tree a -> a
predictWithTree :: forall a. Columnable a => Text -> DataFrame -> Int -> Tree a -> a
predictWithTree Text
_ DataFrame
_ Int
_ (Leaf a
v) = a
v
predictWithTree Text
target DataFrame
df Int
idx (Branch Expr Bool
cond Tree a
left Tree a
right) =
    forall a. Columnable a => Text -> DataFrame -> Int -> Tree a -> a
predictWithTree @a Text
target DataFrame
df Int
idx (Expr Bool -> Tree a -> Tree a -> Int -> DataFrame -> Tree a
forall a.
Expr Bool -> Tree a -> Tree a -> Int -> DataFrame -> Tree a
childFor Expr Bool
cond Tree a
left Tree a
right Int
idx DataFrame
df)

childFor :: Expr Bool -> Tree a -> Tree a -> Int -> DataFrame -> Tree a
childFor :: forall a.
Expr Bool -> Tree a -> Tree a -> Int -> DataFrame -> Tree a
childFor Expr Bool
cond Tree a
left Tree a
right Int
idx DataFrame
df = case DataFrame -> Expr Bool -> Maybe (Vector Bool)
branchBool DataFrame
df Expr Bool
cond of
    Maybe (Vector Bool)
Nothing -> Tree a
left
    Just Vector Bool
boolVals -> if Vector Bool
boolVals Vector Bool -> Int -> Bool
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
idx then Tree a
left else Tree a
right

predictManyWithTree ::
    forall a. (Columnable a) => Tree a -> DataFrame -> V.Vector Int -> V.Vector a
predictManyWithTree :: forall a.
Columnable a =>
Tree a -> DataFrame -> Vector Int -> Vector a
predictManyWithTree = forall a.
Columnable a =>
CondCache -> Tree a -> DataFrame -> Vector Int -> Vector a
predictManyWithTreeCached @a CondCache
forall k a. Map k a
M.empty

{- | 'predictManyWithTree' resolving each branch condition through a 'CondCache'.
Each condition is read at most once per call rather than once per row.
-}
predictManyWithTreeCached ::
    forall a.
    (Columnable a) => CondCache -> Tree a -> DataFrame -> V.Vector Int -> V.Vector a
predictManyWithTreeCached :: forall a.
Columnable a =>
CondCache -> Tree a -> DataFrame -> Vector Int -> Vector a
predictManyWithTreeCached CondCache
cache Tree a
tree DataFrame
df Vector Int
indices = (forall s. ST s (MVector s a)) -> Vector a
forall a. (forall s. ST s (MVector s a)) -> Vector a
V.create ((forall s. ST s (MVector s a)) -> Vector a)
-> (forall s. ST s (MVector s a)) -> Vector a
forall a b. (a -> b) -> a -> b
$ do
    MVector s a
mv <- Int -> ST s (MVector (PrimState (ST s)) a)
forall (m :: * -> *) a.
PrimMonad m =>
Int -> m (MVector (PrimState m) a)
VM.new (Vector Int -> Int
forall a. Vector a -> Int
V.length Vector Int
indices)
    MVector s a -> Vector (Int, Int) -> Tree a -> ST s ()
forall s. MVector s a -> Vector (Int, Int) -> Tree a -> ST s ()
fill MVector s a
mv (Vector Int -> Vector Int -> Vector (Int, Int)
forall a b. Vector a -> Vector b -> Vector (a, b)
V.zip (Int -> Int -> Vector Int
forall a. Num a => a -> Int -> Vector a
V.enumFromN Int
0 (Vector Int -> Int
forall a. Vector a -> Int
V.length Vector Int
indices)) Vector Int
indices) Tree a
tree
    MVector s a -> ST s (MVector s a)
forall a. a -> ST s a
forall (f :: * -> *) a. Applicative f => a -> f a
pure MVector s a
mv
  where
    fill :: VM.MVector s a -> V.Vector (Int, Int) -> Tree a -> ST s ()
    fill :: forall s. MVector s a -> Vector (Int, Int) -> Tree a -> ST s ()
fill MVector s a
mv Vector (Int, Int)
prs (Leaf a
v) = ((Int, Int) -> ST s ()) -> Vector (Int, Int) -> ST s ()
forall (m :: * -> *) a b. Monad m => (a -> m b) -> Vector a -> m ()
V.mapM_ (\(Int
p, Int
_) -> MVector (PrimState (ST s)) a -> Int -> a -> ST s ()
forall (m :: * -> *) a.
PrimMonad m =>
MVector (PrimState m) a -> Int -> a -> m ()
VM.write MVector s a
MVector (PrimState (ST s)) a
mv Int
p a
v) Vector (Int, Int)
prs
    fill MVector s a
mv Vector (Int, Int)
prs (Branch Expr Bool
cond Tree a
left Tree a
right) = case CondCache -> DataFrame -> Expr Bool -> Maybe (Vector Bool)
lookupCondVec CondCache
cache DataFrame
df Expr Bool
cond of
        Maybe (Vector Bool)
Nothing -> MVector s a -> Vector (Int, Int) -> Tree a -> ST s ()
forall s. MVector s a -> Vector (Int, Int) -> Tree a -> ST s ()
fill MVector s a
mv Vector (Int, Int)
prs Tree a
left
        Just Vector Bool
boolVals -> MVector s a
-> (Vector (Int, Int), Vector (Int, Int))
-> Tree a
-> Tree a
-> ST s ()
forall s.
MVector s a
-> (Vector (Int, Int), Vector (Int, Int))
-> Tree a
-> Tree a
-> ST s ()
fillSplit MVector s a
mv (((Int, Int) -> Bool)
-> Vector (Int, Int) -> (Vector (Int, Int), Vector (Int, Int))
forall a. (a -> Bool) -> Vector a -> (Vector a, Vector a)
V.partition (\(Int
_, Int
i) -> Vector Bool
boolVals Vector Bool -> Int -> Bool
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i) Vector (Int, Int)
prs) Tree a
left Tree a
right

    fillSplit ::
        VM.MVector s a ->
        (V.Vector (Int, Int), V.Vector (Int, Int)) ->
        Tree a ->
        Tree a ->
        ST s ()
    fillSplit :: forall s.
MVector s a
-> (Vector (Int, Int), Vector (Int, Int))
-> Tree a
-> Tree a
-> ST s ()
fillSplit MVector s a
mv (Vector (Int, Int)
leftPrs, Vector (Int, Int)
rightPrs) Tree a
left Tree a
right = MVector s a -> Vector (Int, Int) -> Tree a -> ST s ()
forall s. MVector s a -> Vector (Int, Int) -> Tree a -> ST s ()
fill MVector s a
mv Vector (Int, Int)
leftPrs Tree a
left ST s () -> ST s () -> ST s ()
forall a b. ST s a -> ST s b -> ST s b
forall (m :: * -> *) a b. Monad m => m a -> m b -> m b
>> MVector s a -> Vector (Int, Int) -> Tree a -> ST s ()
forall s. MVector s a -> Vector (Int, Int) -> Tree a -> ST s ()
fill MVector s a
mv Vector (Int, Int)
rightPrs Tree a
right

identifyCarePoints ::
    forall a.
    (Columnable a) =>
    T.Text -> DataFrame -> V.Vector Int -> Tree a -> Tree a -> [CarePoint]
identifyCarePoints :: forall a.
Columnable a =>
Text -> DataFrame -> Vector Int -> Tree a -> Tree a -> [CarePoint]
identifyCarePoints = forall a.
Columnable a =>
CondCache
-> Text
-> DataFrame
-> Vector Int
-> Tree a
-> Tree a
-> [CarePoint]
identifyCarePointsCached @a CondCache
forall k a. Map k a
M.empty

{- | Rows the parent must route to a specific child for the (fixed) subtrees to
classify correctly; a 'CondCache' avoids re-interpreting subtree conditions.
-}
identifyCarePointsCached ::
    forall a.
    (Columnable a) =>
    CondCache ->
    T.Text ->
    DataFrame ->
    V.Vector Int ->
    Tree a ->
    Tree a ->
    [CarePoint]
identifyCarePointsCached :: forall a.
Columnable a =>
CondCache
-> Text
-> DataFrame
-> Vector Int
-> Tree a
-> Tree a
-> [CarePoint]
identifyCarePointsCached CondCache
cache Text
target DataFrame
df Vector Int
indices Tree a
leftTree Tree a
rightTree =
    [CarePoint]
-> (Vector a -> [CarePoint]) -> Maybe (Vector a) -> [CarePoint]
forall b a. b -> (a -> b) -> Maybe a -> b
maybe [] Vector a -> [CarePoint]
carePoints (forall a. Columnable a => DataFrame -> Text -> Maybe (Vector a)
interpretLabelCol @a DataFrame
df Text
target)
  where
    leftPreds :: Vector a
leftPreds = CondCache -> Tree a -> DataFrame -> Vector Int -> Vector a
forall a.
Columnable a =>
CondCache -> Tree a -> DataFrame -> Vector Int -> Vector a
predictManyWithTreeCached CondCache
cache Tree a
leftTree DataFrame
df Vector Int
indices
    rightPreds :: Vector a
rightPreds = CondCache -> Tree a -> DataFrame -> Vector Int -> Vector a
forall a.
Columnable a =>
CondCache -> Tree a -> DataFrame -> Vector Int -> Vector a
predictManyWithTreeCached CondCache
cache Tree a
rightTree DataFrame
df Vector Int
indices
    carePoints :: Vector a -> [CarePoint]
carePoints Vector a
targetVals = Vector CarePoint -> [CarePoint]
forall a. Vector a -> [a]
V.toList ((Int -> Int -> Maybe CarePoint) -> Vector Int -> Vector CarePoint
forall a b. (Int -> a -> Maybe b) -> Vector a -> Vector b
V.imapMaybe (Vector a -> Vector a -> Vector a -> Int -> Int -> Maybe CarePoint
forall a.
Eq a =>
Vector a -> Vector a -> Vector a -> Int -> Int -> Maybe CarePoint
checkPoint Vector a
targetVals Vector a
leftPreds Vector a
rightPreds) Vector Int
indices)

checkPoint ::
    (Eq a) =>
    V.Vector a -> V.Vector a -> V.Vector a -> Int -> Int -> Maybe CarePoint
checkPoint :: forall a.
Eq a =>
Vector a -> Vector a -> Vector a -> Int -> Int -> Maybe CarePoint
checkPoint Vector a
targetVals Vector a
leftPreds Vector a
rightPreds Int
k Int
idx =
    case (Vector a
leftPreds Vector a -> Int -> a
forall a. Vector a -> Int -> a
V.! Int
k a -> a -> Bool
forall a. Eq a => a -> a -> Bool
== a
trueLabel, Vector a
rightPreds Vector a -> Int -> a
forall a. Vector a -> Int -> a
V.! Int
k a -> a -> Bool
forall a. Eq a => a -> a -> Bool
== a
trueLabel) of
        (Bool
True, Bool
False) -> CarePoint -> Maybe CarePoint
forall a. a -> Maybe a
Just (Int -> Direction -> CarePoint
CarePoint Int
idx Direction
GoLeft)
        (Bool
False, Bool
True) -> CarePoint -> Maybe CarePoint
forall a. a -> Maybe a
Just (Int -> Direction -> CarePoint
CarePoint Int
idx Direction
GoRight)
        (Bool, Bool)
_ -> Maybe CarePoint
forall a. Maybe a
Nothing
  where
    trueLabel :: a
trueLabel = Vector a
targetVals Vector a -> Int -> a
forall a. Vector a -> Int -> a
V.! Int
idx

-- | Care points a free condition misroutes (uncached; for the linear path).
countCarePointErrors :: Expr Bool -> DataFrame -> [CarePoint] -> Int
countCarePointErrors :: Expr Bool -> DataFrame -> [CarePoint] -> Int
countCarePointErrors Expr Bool
cond DataFrame
df [CarePoint]
carePoints =
    Int -> (Vector Bool -> Int) -> Maybe (Vector Bool) -> Int
forall b a. b -> (a -> b) -> Maybe a -> b
maybe ([CarePoint] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [CarePoint]
carePoints) (Vector Bool -> [CarePoint] -> Int
`countErrorsByVec` [CarePoint]
carePoints) (DataFrame -> Expr Bool -> Maybe (Vector Bool)
branchBool DataFrame
df Expr Bool
cond)

partitionIndices ::
    Expr Bool -> DataFrame -> V.Vector Int -> (V.Vector Int, V.Vector Int)
partitionIndices :: Expr Bool -> DataFrame -> Vector Int -> (Vector Int, Vector Int)
partitionIndices = CondCache
-> Expr Bool -> DataFrame -> Vector Int -> (Vector Int, Vector Int)
partitionIndicesCached CondCache
forall k a. Map k a
M.empty

{- | 'partitionIndices' resolving the condition through a 'CondCache'; a miss
routes every index left (matching the uncached fallback).
-}
partitionIndicesCached ::
    CondCache ->
    Expr Bool ->
    DataFrame ->
    V.Vector Int ->
    (V.Vector Int, V.Vector Int)
partitionIndicesCached :: CondCache
-> Expr Bool -> DataFrame -> Vector Int -> (Vector Int, Vector Int)
partitionIndicesCached CondCache
cache Expr Bool
cond DataFrame
df Vector Int
indices = case CondCache -> DataFrame -> Expr Bool -> Maybe (Vector Bool)
lookupCondVec CondCache
cache DataFrame
df Expr Bool
cond of
    Maybe (Vector Bool)
Nothing -> (Vector Int
indices, Vector Int
forall a. Vector a
V.empty)
    Just Vector Bool
boolVals -> (Int -> Bool) -> Vector Int -> (Vector Int, Vector Int)
forall a. (a -> Bool) -> Vector a -> (Vector a, Vector a)
V.partition (Vector Bool
boolVals Vector Bool -> Int -> Bool
forall a. Unbox a => Vector a -> Int -> a
VU.!) Vector Int
indices

-- | A split is valid at a node when both children keep at least 'minLeafSize'.
isValidAtNode :: TreeConfig -> DataFrame -> V.Vector Int -> Expr Bool -> Bool
isValidAtNode :: TreeConfig -> DataFrame -> Vector Int -> Expr Bool -> Bool
isValidAtNode TreeConfig
cfg DataFrame
df Vector Int
indices Expr Bool
c =
    Vector Int -> Int
forall a. Vector a -> Int
V.length Vector Int
t Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= TreeConfig -> Int
minLeafSize TreeConfig
cfg Bool -> Bool -> Bool
&& Vector Int -> Int
forall a. Vector a -> Int
V.length Vector Int
f Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= TreeConfig -> Int
minLeafSize TreeConfig
cfg
  where
    (Vector Int
t, Vector Int
f) = Expr Bool -> DataFrame -> Vector Int -> (Vector Int, Vector Int)
partitionIndices Expr Bool
c DataFrame
df Vector Int
indices

majorityValueFromIndices ::
    forall a. (Columnable a, Ord a) => T.Text -> DataFrame -> V.Vector Int -> a
majorityValueFromIndices :: forall a.
(Columnable a, Ord a) =>
Text -> DataFrame -> Vector Int -> a
majorityValueFromIndices Text
target DataFrame
df Vector Int
indices = Map a Int -> a
forall a. Map a Int -> a
majorityOf (Vector a -> Vector Int -> Map a Int
forall a. Ord a => Vector a -> Vector Int -> Map a Int
countLabels (forall a. Columnable a => DataFrame -> Text -> Vector a
labelColOrThrow @a DataFrame
df Text
target) Vector Int
indices)

labelColOrThrow :: forall a. (Columnable a) => DataFrame -> T.Text -> V.Vector a
labelColOrThrow :: forall a. Columnable a => DataFrame -> Text -> Vector a
labelColOrThrow 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
    Left DataFrameException
e -> DataFrameException -> Vector a
forall a e. Exception e => e -> a
throw DataFrameException
e
    Right (TColumn Column
column) -> (DataFrameException -> Vector a)
-> (Vector a -> Vector a)
-> Either DataFrameException (Vector a)
-> Vector a
forall a c b. (a -> c) -> (b -> c) -> Either a b -> c
either DataFrameException -> Vector a
forall a e. Exception e => e -> a
throw Vector a -> Vector a
forall a. a -> a
id (forall a (v :: * -> *).
(Vector v a, Columnable a) =>
Column -> Either DataFrameException (v a)
toVector @a Column
column)

countLabels :: (Ord a) => V.Vector a -> V.Vector Int -> M.Map a Int
countLabels :: forall a. Ord a => Vector a -> Vector Int -> Map a Int
countLabels Vector a
vals = (Map a Int -> Int -> Map a Int)
-> Map a Int -> Vector Int -> Map a Int
forall a b. (a -> b -> a) -> a -> Vector b -> a
V.foldl' (\Map a Int
acc Int
i -> (Int -> Int -> Int) -> a -> Int -> Map a Int -> Map a 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
(+) (Vector a
vals Vector a -> Int -> a
forall a. Vector a -> Int -> a
V.! Int
i) (Int
1 :: Int) Map a Int
acc) Map a Int
forall k a. Map k a
M.empty

majorityOf :: M.Map a Int -> a
majorityOf :: forall a. Map a Int -> a
majorityOf Map a Int
counts
    | Map a Int -> Bool
forall k a. Map k a -> Bool
M.null Map a Int
counts = DataFrameException -> a
forall a e. Exception e => e -> a
throw (Text -> DataFrameException
EmptyDataSetException Text
"majorityValueFromIndices")
    | Bool
otherwise = (a, Int) -> a
forall a b. (a, b) -> a
fst (((a, Int) -> (a, Int) -> Ordering) -> [(a, Int)] -> (a, Int)
forall (t :: * -> *) a.
Foldable t =>
(a -> a -> Ordering) -> t a -> a
maximumBy (Int -> Int -> Ordering
forall a. Ord a => a -> a -> Ordering
compare (Int -> Int -> Ordering)
-> ((a, Int) -> Int) -> (a, Int) -> (a, Int) -> Ordering
forall b c a. (b -> b -> c) -> (a -> b) -> a -> a -> c
`on` (a, Int) -> Int
forall a b. (a, b) -> b
snd) (Map a Int -> [(a, Int)]
forall k a. Map k a -> [(k, a)]
M.toList Map a Int
counts))

computeTreeLoss ::
    forall a.
    (Columnable a) => T.Text -> DataFrame -> V.Vector Int -> Tree a -> Double
computeTreeLoss :: forall a.
Columnable a =>
Text -> DataFrame -> Vector Int -> Tree a -> Double
computeTreeLoss = forall a.
Columnable a =>
CondCache -> Text -> DataFrame -> Vector Int -> Tree a -> Double
computeTreeLossCached @a CondCache
forall k a. Map k a
M.empty

-- | 0/1 loss of a tree over @indices@, with a 'CondCache' for the predictions.
computeTreeLossCached ::
    forall a.
    (Columnable a) =>
    CondCache -> T.Text -> DataFrame -> V.Vector Int -> Tree a -> Double
computeTreeLossCached :: forall a.
Columnable a =>
CondCache -> Text -> DataFrame -> Vector Int -> Tree a -> Double
computeTreeLossCached CondCache
cache Text
target DataFrame
df Vector Int
indices Tree a
tree
    | Vector Int -> Bool
forall a. Vector a -> Bool
V.null Vector Int
indices = Double
0
    | Bool
otherwise =
        Double -> (Vector a -> Double) -> Maybe (Vector a) -> Double
forall b a. b -> (a -> b) -> Maybe a -> b
maybe Double
1.0 (CondCache
-> Tree a -> DataFrame -> Vector Int -> Vector a -> Double
forall a.
Columnable a =>
CondCache
-> Tree a -> DataFrame -> Vector Int -> Vector a -> Double
treeLoss CondCache
cache Tree a
tree DataFrame
df Vector Int
indices) (forall a. Columnable a => DataFrame -> Text -> Maybe (Vector a)
interpretLabelCol @a DataFrame
df Text
target)

treeLoss ::
    (Columnable a) =>
    CondCache -> Tree a -> DataFrame -> V.Vector Int -> V.Vector a -> Double
treeLoss :: forall a.
Columnable a =>
CondCache
-> Tree a -> DataFrame -> Vector Int -> Vector a -> Double
treeLoss CondCache
cache Tree a
tree DataFrame
df Vector Int
indices Vector a
targetVals =
    Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Vector a -> Vector Int -> Vector a -> Int
forall a. Eq a => Vector a -> Vector Int -> Vector a -> Int
countMismatches Vector a
targetVals Vector Int
indices Vector a
preds)
        Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Vector Int -> Int
forall a. Vector a -> Int
V.length Vector Int
indices)
  where
    preds :: Vector a
preds = CondCache -> Tree a -> DataFrame -> Vector Int -> Vector a
forall a.
Columnable a =>
CondCache -> Tree a -> DataFrame -> Vector Int -> Vector a
predictManyWithTreeCached CondCache
cache Tree a
tree DataFrame
df Vector Int
indices

countMismatches :: (Eq a) => V.Vector a -> V.Vector Int -> V.Vector a -> Int
countMismatches :: forall a. Eq a => Vector a -> Vector Int -> Vector a -> Int
countMismatches Vector a
targetVals Vector Int
indices Vector a
preds =
    Vector a -> Int
forall a. Vector a -> Int
V.length
        ((Int -> a -> Bool) -> Vector a -> Vector a
forall a. (Int -> a -> Bool) -> Vector a -> Vector a
V.ifilter (\Int
k a
_ -> Vector a
targetVals Vector a -> Int -> a
forall a. Vector a -> Int -> a
V.! (Vector Int
indices Vector Int -> Int -> Int
forall a. Vector a -> Int -> a
V.! Int
k) a -> a -> Bool
forall a. Eq a => a -> a -> Bool
/= Vector a
preds Vector a -> Int -> a
forall a. Vector a -> Int -> a
V.! Int
k) Vector a
preds)