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

{- | Tree Alternating Optimization: hold the tree fixed and re-optimize one node
at a time, bottom-up, minimizing care-point misroutes. Sibling subtrees at a
depth level are independent and optimized in parallel.
-}
module DataFrame.DecisionTree.Tao (
    taoOptimize,
    taoOptimizeCV,
    taoIteration,
    taoIterationCV,
    optimizeNode,
    findBestSplitTAO,
) where

import DataFrame.DecisionTree.CondVec
import DataFrame.DecisionTree.Linear (bestLinearCandidate)
import DataFrame.DecisionTree.Pool (
    bestDiscreteCandidate,
    candidateParChunk,
    evalWithPenaltyVec,
 )
import DataFrame.DecisionTree.Predict
import DataFrame.DecisionTree.Prune (pruneDead)
import DataFrame.DecisionTree.Types
import DataFrame.Internal.Column (Columnable)
import DataFrame.Internal.DataFrame (DataFrame)
import DataFrame.Internal.Expression (Expr)

import Control.Parallel (par, pseq)
import Control.Parallel.Strategies (parListChunk, rdeepseq, using)
import Data.Function (on)
import Data.List (foldl', minimumBy)
import Data.Maybe (catMaybes, mapMaybe)
import qualified Data.Text as T
import qualified Data.Vector as V
import qualified Data.Vector.Unboxed as VU

{- | The constant per-fit context threaded through the node-optimization
recursion (the cache is rebuilt each iteration).
-}
data TaoEnv = TaoEnv
    { TaoEnv -> CondCache
teCache :: !CondCache
    , TaoEnv -> TreeConfig
teCfg :: !TreeConfig
    , TaoEnv -> Text
teTarget :: !T.Text
    , TaoEnv -> [CondVec]
teConds :: ![CondVec]
    , TaoEnv -> DataFrame
teDf :: !DataFrame
    }

-- | Public TAO entry point over raw conditions; materializes each once.
taoOptimize ::
    forall a.
    (Columnable a, Ord a) =>
    TreeConfig ->
    T.Text ->
    [Expr Bool] ->
    DataFrame ->
    V.Vector Int ->
    Tree a ->
    Tree a
taoOptimize :: forall a.
(Columnable a, Ord a) =>
TreeConfig
-> Text
-> [Expr Bool]
-> DataFrame
-> Vector Int
-> Tree a
-> Tree a
taoOptimize TreeConfig
cfg Text
target [Expr Bool]
conds DataFrame
df =
    forall a.
(Columnable a, Ord a) =>
TreeConfig
-> Text -> [CondVec] -> DataFrame -> Vector Int -> Tree a -> Tree a
taoOptimizeCV @a TreeConfig
cfg Text
target ((Expr Bool -> Maybe CondVec) -> [Expr Bool] -> [CondVec]
forall a b. (a -> Maybe b) -> [a] -> [b]
mapMaybe (DataFrame -> Expr Bool -> Maybe CondVec
materializeCondVec DataFrame
df) [Expr Bool]
conds) DataFrame
df

{- | TAO outer loop over pre-evaluated candidates: iterate until the iteration
budget or convergence tolerance is reached, then prune dead branches.
-}
taoOptimizeCV ::
    forall a.
    (Columnable a, Ord a) =>
    TreeConfig ->
    T.Text ->
    [CondVec] ->
    DataFrame ->
    V.Vector Int ->
    Tree a ->
    Tree a
taoOptimizeCV :: forall a.
(Columnable a, Ord a) =>
TreeConfig
-> Text -> [CondVec] -> DataFrame -> Vector Int -> Tree a -> Tree a
taoOptimizeCV TreeConfig
cfg Text
target [CondVec]
condVecs DataFrame
df Vector Int
rootIndices Tree a
initialTree =
    Int -> Tree a -> Double -> Tree a
go Int
0 Tree a
initialTree (CondCache -> Tree a -> Double
lossWith CondCache
baseCache Tree a
initialTree)
  where
    baseCache :: CondCache
baseCache = [CondVec] -> CondCache
condCacheFromVecs [CondVec]
condVecs
    lossWith :: CondCache -> Tree a -> Double
lossWith CondCache
cache = forall a.
Columnable a =>
CondCache -> Text -> DataFrame -> Vector Int -> Tree a -> Double
computeTreeLossCached @a CondCache
cache Text
target DataFrame
df Vector Int
rootIndices
    go :: Int -> Tree a -> Double -> Tree a
go Int
iter Tree a
tree Double
prevLoss
        | Int
iter Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= TreeConfig -> Int
taoIterations TreeConfig
cfg = Tree a -> Tree a
forall a. Columnable a => Tree a -> Tree a
pruneDead Tree a
tree
        | Double
prevLoss Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
newLoss Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
< TreeConfig -> Double
taoConvergenceTol TreeConfig
cfg = Tree a -> Tree a
forall a. Columnable a => Tree a -> Tree a
pruneDead Tree a
tree'
        | Bool
otherwise = Int -> Tree a -> Double -> Tree a
go (Int
iter Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Tree a
tree' Double
newLoss
      where
        cache :: CondCache
cache = DataFrame -> Tree a -> CondCache -> CondCache
forall a. DataFrame -> Tree a -> CondCache -> CondCache
addTreeCondsToCache DataFrame
df Tree a
tree CondCache
baseCache
        tree' :: Tree a
tree' = forall a.
(Columnable a, Ord a) =>
CondCache
-> TreeConfig
-> Text
-> [CondVec]
-> DataFrame
-> Vector Int
-> Tree a
-> Tree a
taoIterationCV @a CondCache
cache TreeConfig
cfg Text
target [CondVec]
condVecs DataFrame
df Vector Int
rootIndices Tree a
tree
        newLoss :: Double
newLoss = CondCache -> Tree a -> Double
lossWith CondCache
cache Tree a
tree'

-- | Public single-iteration entry point.
taoIteration ::
    forall a.
    (Columnable a, Ord a) =>
    TreeConfig ->
    T.Text ->
    [Expr Bool] ->
    DataFrame ->
    V.Vector Int ->
    Tree a ->
    Tree a
taoIteration :: forall a.
(Columnable a, Ord a) =>
TreeConfig
-> Text
-> [Expr Bool]
-> DataFrame
-> Vector Int
-> Tree a
-> Tree a
taoIteration TreeConfig
cfg Text
target [Expr Bool]
conds DataFrame
df Vector Int
rootIndices Tree a
tree =
    let condVecs :: [CondVec]
condVecs = (Expr Bool -> Maybe CondVec) -> [Expr Bool] -> [CondVec]
forall a b. (a -> Maybe b) -> [a] -> [b]
mapMaybe (DataFrame -> Expr Bool -> Maybe CondVec
materializeCondVec DataFrame
df) [Expr Bool]
conds
        cache :: CondCache
cache = DataFrame -> Tree a -> CondCache -> CondCache
forall a. DataFrame -> Tree a -> CondCache -> CondCache
addTreeCondsToCache DataFrame
df Tree a
tree ([CondVec] -> CondCache
condCacheFromVecs [CondVec]
condVecs)
     in forall a.
(Columnable a, Ord a) =>
CondCache
-> TreeConfig
-> Text
-> [CondVec]
-> DataFrame
-> Vector Int
-> Tree a
-> Tree a
taoIterationCV @a CondCache
cache TreeConfig
cfg Text
target [CondVec]
condVecs DataFrame
df Vector Int
rootIndices Tree a
tree

-- | One bottom-to-top sweep: re-optimize every node level by level.
taoIterationCV ::
    forall a.
    (Columnable a, Ord a) =>
    CondCache ->
    TreeConfig ->
    T.Text ->
    [CondVec] ->
    DataFrame ->
    V.Vector Int ->
    Tree a ->
    Tree a
taoIterationCV :: forall a.
(Columnable a, Ord a) =>
CondCache
-> TreeConfig
-> Text
-> [CondVec]
-> DataFrame
-> Vector Int
-> Tree a
-> Tree a
taoIterationCV CondCache
cache TreeConfig
cfg Text
target [CondVec]
condVecs DataFrame
df Vector Int
rootIndices Tree a
tree =
    (Tree a -> Int -> Tree a) -> Tree a -> [Int] -> Tree a
forall b a. (b -> a -> b) -> b -> [a] -> b
forall (t :: * -> *) b a.
Foldable t =>
(b -> a -> b) -> b -> t a -> b
foldl'
        (TaoEnv -> Vector Int -> Tree a -> Int -> Tree a
forall a.
(Columnable a, Ord a) =>
TaoEnv -> Vector Int -> Tree a -> Int -> Tree a
optimizeDepthLevel TaoEnv
env Vector Int
rootIndices)
        Tree a
tree
        [Tree a -> Int
forall a. Tree a -> Int
treeDepth Tree a
tree, Tree a -> Int
forall a. Tree a -> Int
treeDepth Tree a
tree Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1 .. Int
0]
  where
    env :: TaoEnv
env = CondCache -> TreeConfig -> Text -> [CondVec] -> DataFrame -> TaoEnv
TaoEnv CondCache
cache TreeConfig
cfg Text
target [CondVec]
condVecs DataFrame
df

optimizeDepthLevel ::
    forall a.
    (Columnable a, Ord a) => TaoEnv -> V.Vector Int -> Tree a -> Int -> Tree a
optimizeDepthLevel :: forall a.
(Columnable a, Ord a) =>
TaoEnv -> Vector Int -> Tree a -> Int -> Tree a
optimizeDepthLevel TaoEnv
env Vector Int
rootIndices Tree a
tree = forall a.
(Columnable a, Ord a) =>
TaoEnv -> Vector Int -> Tree a -> Int -> Int -> Tree a
optimizeAtDepth @a TaoEnv
env Vector Int
rootIndices Tree a
tree Int
0

optimizeAtDepth ::
    forall a.
    (Columnable a, Ord a) =>
    TaoEnv -> V.Vector Int -> Tree a -> Int -> Int -> Tree a
optimizeAtDepth :: forall a.
(Columnable a, Ord a) =>
TaoEnv -> Vector Int -> Tree a -> Int -> Int -> Tree a
optimizeAtDepth TaoEnv
env Vector Int
indices Tree a
tree Int
currentDepth Int
targetDepth
    | Int
currentDepth Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
targetDepth = forall a.
(Columnable a, Ord a) =>
TaoEnv -> Vector Int -> Tree a -> Tree a
optimizeNode @a TaoEnv
env Vector Int
indices Tree a
tree
    | Bool
otherwise = case Tree a
tree of
        Leaf a
v -> a -> Tree a
forall a. a -> Tree a
Leaf a
v
        Branch Expr Bool
cond Tree a
left Tree a
right -> forall a.
(Columnable a, Ord a) =>
TaoEnv
-> Vector Int
-> Expr Bool
-> Tree a
-> Tree a
-> Int
-> Int
-> Tree a
optimizeChildren @a TaoEnv
env Vector Int
indices Expr Bool
cond Tree a
left Tree a
right Int
currentDepth Int
targetDepth

{- | Optimize the two subtrees over their disjoint index sets, scoring the left
in parallel with the right (the cache is read-only, so this is pure).
-}
optimizeChildren ::
    forall a.
    (Columnable a, Ord a) =>
    TaoEnv -> V.Vector Int -> Expr Bool -> Tree a -> Tree a -> Int -> Int -> Tree a
optimizeChildren :: forall a.
(Columnable a, Ord a) =>
TaoEnv
-> Vector Int
-> Expr Bool
-> Tree a
-> Tree a
-> Int
-> Int
-> Tree a
optimizeChildren TaoEnv
env Vector Int
indices Expr Bool
cond Tree a
left Tree a
right Int
currentDepth Int
targetDepth =
    Tree a -> ()
forall a. Tree a -> ()
forceTreeWork Tree a
left' () -> Tree a -> Tree a
forall a b. a -> b -> b
`par` (Tree a -> ()
forall a. Tree a -> ()
forceTreeWork Tree a
right' () -> Tree a -> Tree a
forall a b. a -> b -> b
`pseq` Expr Bool -> Tree a -> Tree a -> Tree a
forall a. Expr Bool -> Tree a -> Tree a -> Tree a
Branch Expr Bool
cond Tree a
left' Tree a
right')
  where
    (Vector Int
indicesL, Vector Int
indicesR) = CondCache
-> Expr Bool -> DataFrame -> Vector Int -> (Vector Int, Vector Int)
partitionIndicesCached (TaoEnv -> CondCache
teCache TaoEnv
env) Expr Bool
cond (TaoEnv -> DataFrame
teDf TaoEnv
env) Vector Int
indices
    left' :: Tree a
left' = forall a.
(Columnable a, Ord a) =>
TaoEnv -> Vector Int -> Tree a -> Int -> Int -> Tree a
optimizeAtDepth @a TaoEnv
env Vector Int
indicesL Tree a
left (Int
currentDepth Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Int
targetDepth
    right' :: Tree a
right' = forall a.
(Columnable a, Ord a) =>
TaoEnv -> Vector Int -> Tree a -> Int -> Int -> Tree a
optimizeAtDepth @a TaoEnv
env Vector Int
indicesR Tree a
right (Int
currentDepth Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Int
targetDepth

{- | Force a subtree's optimization work to WHNF so the parallel scheduler has
something substantial to evaluate; pure and value-preserving.
-}
forceTreeWork :: Tree a -> ()
forceTreeWork :: forall a. Tree a -> ()
forceTreeWork (Leaf a
v) = a
v a -> () -> ()
forall a b. a -> b -> b
`seq` ()
forceTreeWork (Branch Expr Bool
c Tree a
l Tree a
r) = Expr Bool
c Expr Bool -> () -> ()
forall a b. a -> b -> b
`seq` Tree a -> ()
forall a. Tree a -> ()
forceTreeWork Tree a
l () -> () -> ()
forall a b. a -> b -> b
`seq` Tree a -> ()
forall a. Tree a -> ()
forceTreeWork Tree a
r

{- | Re-optimize one node: pick its best split, or collapse to a leaf when the
node is empty or the chosen split underflows 'minLeafSize'.
-}
optimizeNode ::
    forall a. (Columnable a, Ord a) => TaoEnv -> V.Vector Int -> Tree a -> Tree a
optimizeNode :: forall a.
(Columnable a, Ord a) =>
TaoEnv -> Vector Int -> Tree a -> Tree a
optimizeNode TaoEnv
env Vector Int
indices Tree a
tree
    | Vector Int -> Bool
forall a. Vector a -> Bool
V.null Vector Int
indices = Tree a
tree
    | Bool
otherwise = case Tree a
tree of
        Leaf a
_ -> Tree a
leaf
        Branch Expr Bool
oldCond Tree a
left Tree a
right -> TaoEnv
-> Vector Int -> Expr Bool -> Tree a -> Tree a -> Tree a -> Tree a
forall a.
(Columnable a, Ord a) =>
TaoEnv
-> Vector Int -> Expr Bool -> Tree a -> Tree a -> Tree a -> Tree a
rebuiltBranch TaoEnv
env Vector Int
indices Expr Bool
oldCond Tree a
left Tree a
right Tree a
leaf
  where
    leaf :: Tree a
leaf = a -> Tree a
forall a. a -> Tree a
Leaf (forall a.
(Columnable a, Ord a) =>
Text -> DataFrame -> Vector Int -> a
majorityValueFromIndices @a (TaoEnv -> Text
teTarget TaoEnv
env) (TaoEnv -> DataFrame
teDf TaoEnv
env) Vector Int
indices)

rebuiltBranch ::
    forall a.
    (Columnable a, Ord a) =>
    TaoEnv -> V.Vector Int -> Expr Bool -> Tree a -> Tree a -> Tree a -> Tree a
rebuiltBranch :: forall a.
(Columnable a, Ord a) =>
TaoEnv
-> Vector Int -> Expr Bool -> Tree a -> Tree a -> Tree a -> Tree a
rebuiltBranch TaoEnv
env Vector Int
indices Expr Bool
oldCond Tree a
left Tree a
right Tree a
leaf
    | Bool
underflows = Tree a
leaf
    | Bool
otherwise = Expr Bool -> Tree a -> Tree a -> Tree a
forall a. Expr Bool -> Tree a -> Tree a -> Tree a
Branch Expr Bool
newCond Tree a
left Tree a
right
  where
    newCond :: Expr Bool
newCond = forall a.
Columnable a =>
TaoEnv -> Vector Int -> Tree a -> Tree a -> Expr Bool -> Expr Bool
findBestSplitTAO @a TaoEnv
env Vector Int
indices Tree a
left Tree a
right Expr Bool
oldCond
    (Vector Int
l, Vector Int
r) = CondCache
-> Expr Bool -> DataFrame -> Vector Int -> (Vector Int, Vector Int)
partitionIndicesCached (TaoEnv -> CondCache
teCache TaoEnv
env) Expr Bool
newCond (TaoEnv -> DataFrame
teDf TaoEnv
env) Vector Int
indices
    underflows :: Bool
underflows = Vector Int -> Int
forall a. Vector a -> Int
V.length Vector Int
l Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< TreeConfig -> Int
minLeafSize (TaoEnv -> TreeConfig
teCfg TaoEnv
env) Bool -> Bool -> Bool
|| Vector Int -> Int
forall a. Vector a -> Int
V.length Vector Int
r Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< TreeConfig -> Int
minLeafSize (TaoEnv -> TreeConfig
teCfg TaoEnv
env)

{- | The lowest-penalty replacement condition for a node, falling back to the
current condition when no valid candidate beats it.
-}
findBestSplitTAO ::
    forall a.
    (Columnable a) =>
    TaoEnv -> V.Vector Int -> Tree a -> Tree a -> Expr Bool -> Expr Bool
findBestSplitTAO :: forall a.
Columnable a =>
TaoEnv -> Vector Int -> Tree a -> Tree a -> Expr Bool -> Expr Bool
findBestSplitTAO TaoEnv
env Vector Int
indices Tree a
leftTree Tree a
rightTree Expr Bool
currentCond
    | Vector Int -> Bool
forall a. Vector a -> Bool
V.null Vector Int
indices Bool -> Bool -> Bool
|| [CarePoint] -> Bool
forall a. [a] -> Bool
forall (t :: * -> *) a. Foldable t => t a -> Bool
null [CarePoint]
carePoints = Expr Bool
currentCond
    | TreeConfig -> Bool
pureReplacementLinear TreeConfig
cfg
    , Just Expr Bool
c <- Maybe (Expr Bool)
linearCandidate
    , TreeConfig -> DataFrame -> Vector Int -> Expr Bool -> Bool
isValidAtNode TreeConfig
cfg (TaoEnv -> DataFrame
teDf TaoEnv
env) Vector Int
indices Expr Bool
c =
        Expr Bool
c
    | Bool
otherwise = (CondVec -> (Int, Int)) -> Expr Bool -> [CondVec] -> Expr Bool
bestOfPool CondVec -> (Int, Int)
penaltyCV Expr Bool
currentCond [CondVec]
pool
  where
    cfg :: TreeConfig
cfg = TaoEnv -> TreeConfig
teCfg TaoEnv
env
    carePoints :: [CarePoint]
carePoints =
        forall a.
Columnable a =>
CondCache
-> Text
-> DataFrame
-> Vector Int
-> Tree a
-> Tree a
-> [CarePoint]
identifyCarePointsCached @a
            (TaoEnv -> CondCache
teCache TaoEnv
env)
            (TaoEnv -> Text
teTarget TaoEnv
env)
            (TaoEnv -> DataFrame
teDf TaoEnv
env)
            Vector Int
indices
            Tree a
leftTree
            Tree a
rightTree
    penaltyCV :: CondVec -> (Int, Int)
penaltyCV = TreeConfig -> [CarePoint] -> CondVec -> (Int, Int)
evalWithPenaltyVec TreeConfig
cfg [CarePoint]
carePoints
    linearCandidate :: Maybe (Expr Bool)
linearCandidate = TreeConfig -> Text -> DataFrame -> [CarePoint] -> Maybe (Expr Bool)
bestLinearCandidate TreeConfig
cfg (TaoEnv -> Text
teTarget TaoEnv
env) (TaoEnv -> DataFrame
teDf TaoEnv
env) [CarePoint]
carePoints
    valid :: [CondVec]
valid = TreeConfig -> Vector Int -> [CondVec] -> [CondVec]
filterValidCandidates TreeConfig
cfg Vector Int
indices (TaoEnv -> [CondVec]
teConds TaoEnv
env)
    pool :: [CondVec]
pool =
        TaoEnv
-> Vector Int
-> Expr Bool
-> Maybe CondVec
-> Maybe (Expr Bool)
-> [CondVec]
candidatePool
            TaoEnv
env
            Vector Int
indices
            Expr Bool
currentCond
            (TreeConfig -> (CondVec -> (Int, Int)) -> [CondVec] -> Maybe CondVec
bestDiscreteCandidate TreeConfig
cfg CondVec -> (Int, Int)
penaltyCV [CondVec]
valid)
            Maybe (Expr Bool)
linearCandidate

bestOfPool :: (CondVec -> (Int, Int)) -> Expr Bool -> [CondVec] -> Expr Bool
bestOfPool :: (CondVec -> (Int, Int)) -> Expr Bool -> [CondVec] -> Expr Bool
bestOfPool CondVec -> (Int, Int)
_ Expr Bool
currentCond [] = Expr Bool
currentCond
bestOfPool CondVec -> (Int, Int)
penaltyCV Expr Bool
_ [CondVec]
pool = CondVec -> Expr Bool
cvExpr ((CondVec -> CondVec -> Ordering) -> [CondVec] -> CondVec
forall (t :: * -> *) a.
Foldable t =>
(a -> a -> Ordering) -> t a -> a
minimumBy ((Int, Int) -> (Int, Int) -> Ordering
forall a. Ord a => a -> a -> Ordering
compare ((Int, Int) -> (Int, Int) -> Ordering)
-> (CondVec -> (Int, Int)) -> CondVec -> CondVec -> Ordering
forall b c a. (b -> b -> c) -> (a -> b) -> a -> a -> c
`on` CondVec -> (Int, Int)
penaltyCV) [CondVec]
pool)

{- | Validity-filtered candidates the node could split on: both children must
keep at least 'minLeafSize'. Scored in parallel chunks, order preserved.
-}
filterValidCandidates :: TreeConfig -> V.Vector Int -> [CondVec] -> [CondVec]
filterValidCandidates :: TreeConfig -> Vector Int -> [CondVec] -> [CondVec]
filterValidCandidates TreeConfig
cfg Vector Int
indices [CondVec]
condVecs = ((Bool, CondVec) -> CondVec) -> [(Bool, CondVec)] -> [CondVec]
forall a b. (a -> b) -> [a] -> [b]
map (Bool, CondVec) -> CondVec
forall a b. (a, b) -> b
snd (((Bool, CondVec) -> Bool) -> [(Bool, CondVec)] -> [(Bool, CondVec)]
forall a. (a -> Bool) -> [a] -> [a]
filter (Bool, CondVec) -> Bool
forall a b. (a, b) -> a
fst ([Bool] -> [CondVec] -> [(Bool, CondVec)]
forall a b. [a] -> [b] -> [(a, b)]
zip [Bool]
validity [CondVec]
condVecs))
  where
    validity :: [Bool]
validity =
        (CondVec -> Bool) -> [CondVec] -> [Bool]
forall a b. (a -> b) -> [a] -> [b]
map (TreeConfig -> Vector Int -> CondVec -> Bool
validAtNode TreeConfig
cfg Vector Int
indices) [CondVec]
condVecs
            [Bool] -> Strategy [Bool] -> [Bool]
forall a. a -> Strategy a -> a
`using` Int -> Strategy Bool -> Strategy [Bool]
forall a. Int -> Strategy a -> Strategy [a]
parListChunk Int
candidateParChunk Strategy Bool
forall a. NFData a => Strategy a
rdeepseq

validAtNode :: TreeConfig -> V.Vector Int -> CondVec -> Bool
validAtNode :: TreeConfig -> Vector Int -> CondVec -> Bool
validAtNode TreeConfig
cfg Vector Int
indices CondVec
cv = Int
nTrue Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
minLeaf Bool -> Bool -> Bool
&& (Vector Int -> Int
forall a. Vector a -> Int
V.length Vector Int
indices Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
nTrue) Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
minLeaf
  where
    minLeaf :: Int
minLeaf = TreeConfig -> Int
minLeafSize TreeConfig
cfg
    nTrue :: Int
nTrue =
        (Int -> Int -> Int) -> Int -> Vector Int -> Int
forall a b. (a -> b -> a) -> a -> Vector b -> a
V.foldl'
            (\ !Int
acc Int
i -> if CondVec -> Vector Bool
cvVec CondVec
cv Vector Bool -> Int -> Bool
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i then Int
acc Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1 else Int
acc)
            (Int
0 :: Int)
            Vector Int
indices

{- | The candidate pool to minimize over: the current condition, the best
discrete candidate, and the linear candidate, each kept only if valid.
-}
candidatePool ::
    TaoEnv ->
    V.Vector Int ->
    Expr Bool ->
    Maybe CondVec ->
    Maybe (Expr Bool) ->
    [CondVec]
candidatePool :: TaoEnv
-> Vector Int
-> Expr Bool
-> Maybe CondVec
-> Maybe (Expr Bool)
-> [CondVec]
candidatePool TaoEnv
env Vector Int
indices Expr Bool
currentCond Maybe CondVec
discreteCV Maybe (Expr Bool)
linearCandidate =
    (CondVec -> Bool) -> [CondVec] -> [CondVec]
forall a. (a -> Bool) -> [a] -> [a]
filter
        (TreeConfig -> DataFrame -> Vector Int -> Expr Bool -> Bool
isValidAtNode (TaoEnv -> TreeConfig
teCfg TaoEnv
env) (TaoEnv -> DataFrame
teDf TaoEnv
env) Vector Int
indices (Expr Bool -> Bool) -> (CondVec -> Expr Bool) -> CondVec -> Bool
forall b c a. (b -> c) -> (a -> b) -> a -> c
. CondVec -> Expr Bool
cvExpr)
        ([Maybe CondVec] -> [CondVec]
forall a. [Maybe a] -> [a]
catMaybes [Maybe CondVec
currentCV, Maybe CondVec
discreteCV, Maybe CondVec
linearCV])
  where
    currentCV :: Maybe CondVec
currentCV = Expr Bool -> Vector Bool -> CondVec
CondVec Expr Bool
currentCond (Vector Bool -> CondVec) -> Maybe (Vector Bool) -> Maybe CondVec
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> CondCache -> DataFrame -> Expr Bool -> Maybe (Vector Bool)
lookupCondVec (TaoEnv -> CondCache
teCache TaoEnv
env) (TaoEnv -> DataFrame
teDf TaoEnv
env) Expr Bool
currentCond
    linearCV :: Maybe CondVec
linearCV = Maybe (Expr Bool)
linearCandidate Maybe (Expr Bool) -> (Expr Bool -> Maybe CondVec) -> Maybe CondVec
forall a b. Maybe a -> (a -> Maybe b) -> Maybe b
forall (m :: * -> *) a b. Monad m => m a -> (a -> m b) -> m b
>>= DataFrame -> Expr Bool -> Maybe CondVec
materializeCondVec (TaoEnv -> DataFrame
teDf TaoEnv
env)