{-# LANGUAGE BangPatterns #-}
{-# LANGUAGE FlexibleContexts #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications #-}
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
data TaoEnv = TaoEnv
{ TaoEnv -> CondCache
teCache :: !CondCache
, TaoEnv -> TreeConfig
teCfg :: !TreeConfig
, TaoEnv -> Text
teTarget :: !T.Text
, TaoEnv -> [CondVec]
teConds :: ![CondVec]
, TaoEnv -> DataFrame
teDf :: !DataFrame
}
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
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'
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
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
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
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
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)
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)
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
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)