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

{- | Cached condition truth vectors and the per-fit cache keyed by structural
form. A condition's truth over a fixed DataFrame is invariant for a whole
fit, so it is materialized once and reused.
-}
module DataFrame.DecisionTree.CondVec (
    CondVec (..),
    materializeCondVec,
    CondCache,
    condCacheKey,
    condCacheFromVecs,
    addTreeCondsToCache,
    lookupCondVec,
    partitionByVec,
    countErrorsByVec,
    consolidateThreshold,
    combineAndVec,
    combineOrVec,
) where

import DataFrame.DecisionTree.Types (CarePoint (..), Direction (..), Tree (..))
import qualified DataFrame.Functions as F
import DataFrame.Internal.Column (TypedColumn (..), toVector)
import DataFrame.Internal.DataFrame (DataFrame)
import DataFrame.Internal.Expression (
    BinaryOp (binaryName),
    Expr (..),
    eqExpr,
    normalize,
 )
import DataFrame.Internal.Interpreter (interpret)

import qualified Data.Map.Strict as M
import Data.Maybe (fromMaybe)
import qualified Data.Text as T
import Data.Type.Equality (testEquality, (:~:) (..))
import qualified Data.Vector as V
import qualified Data.Vector.Unboxed as VU
import Type.Reflection (typeRep)

-- | A boolean condition paired with its truth vector over the full DataFrame.
data CondVec = CondVec
    { CondVec -> Expr Bool
cvExpr :: !(Expr Bool)
    , CondVec -> Vector Bool
cvVec :: !(VU.Vector Bool)
    }

{- | Interpret a condition once over the DataFrame; 'Nothing' on a
type/interpret failure so the candidate is silently dropped.
-}
materializeCondVec :: DataFrame -> Expr Bool -> Maybe CondVec
materializeCondVec :: DataFrame -> Expr Bool -> Maybe CondVec
materializeCondVec 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
    Left DataFrameException
_ -> Maybe CondVec
forall a. Maybe a
Nothing
    Right (TColumn Column
column) -> Expr Bool -> Vector Bool -> CondVec
CondVec Expr Bool
cond (Vector Bool -> CondVec) -> Maybe (Vector Bool) -> Maybe CondVec
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Either DataFrameException (Vector Bool) -> Maybe (Vector Bool)
forall e a. Either e a -> Maybe a
eitherToMaybe (forall a (v :: * -> *).
(Vector v a, Columnable a) =>
Column -> Either DataFrameException (v a)
toVector @Bool @VU.Vector Column
column)

eitherToMaybe :: Either e a -> Maybe a
eitherToMaybe :: forall e a. Either e a -> Maybe a
eitherToMaybe = (e -> Maybe a) -> (a -> Maybe a) -> Either e a -> Maybe a
forall a c b. (a -> c) -> (b -> c) -> Either a b -> c
either (Maybe a -> e -> Maybe a
forall a b. a -> b -> a
const Maybe a
forall a. Maybe a
Nothing) a -> Maybe a
forall a. a -> Maybe a
Just

{- | Full-DataFrame truth vectors keyed by structural form, read-only once
built. Seeded for free from the candidate pool plus the initial tree so the
predict/loss passes index a vector instead of re-interpreting per node.
-}
type CondCache = M.Map T.Text (VU.Vector Bool)

{- | Structural key matching the candidate-dedup key, so a tree branch whose
condition came from the pool hits the cache (equal keys ⟹ equal vector).
-}
condCacheKey :: Expr Bool -> T.Text
condCacheKey :: Expr Bool -> Text
condCacheKey = String -> Text
T.pack (String -> Text) -> (Expr Bool -> String) -> Expr Bool -> Text
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Expr Bool -> String
forall a. Show a => a -> String
show (Expr Bool -> String)
-> (Expr Bool -> Expr Bool) -> Expr Bool -> String
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Expr Bool -> Expr Bool
forall a. (Show a, Typeable a) => Expr a -> Expr a
normalize

-- | Seed a cache from already-materialized candidate 'CondVec's (no interpret).
condCacheFromVecs :: [CondVec] -> CondCache
condCacheFromVecs :: [CondVec] -> CondCache
condCacheFromVecs [CondVec]
cvs = [(Text, Vector Bool)] -> CondCache
forall k a. Ord k => [(k, a)] -> Map k a
M.fromList [(Expr Bool -> Text
condCacheKey (CondVec -> Expr Bool
cvExpr CondVec
cv), CondVec -> Vector Bool
cvVec CondVec
cv) | CondVec
cv <- [CondVec]
cvs]

{- | Add a tree's branch-condition vectors to a cache (one interpret per
distinct, not-yet-cached condition).
-}
addTreeCondsToCache :: DataFrame -> Tree a -> CondCache -> CondCache
addTreeCondsToCache :: forall a. DataFrame -> Tree a -> CondCache -> CondCache
addTreeCondsToCache DataFrame
df = Tree a -> CondCache -> CondCache
go
  where
    go :: Tree a -> CondCache -> CondCache
go (Leaf a
_) CondCache
c = CondCache
c
    go (Branch Expr Bool
cond Tree a
l Tree a
r) CondCache
c = Tree a -> CondCache -> CondCache
go Tree a
r (Tree a -> CondCache -> CondCache
go Tree a
l (DataFrame -> Expr Bool -> CondCache -> CondCache
insertCond DataFrame
df Expr Bool
cond CondCache
c))

insertCond :: DataFrame -> Expr Bool -> CondCache -> CondCache
insertCond :: DataFrame -> Expr Bool -> CondCache -> CondCache
insertCond DataFrame
df Expr Bool
cond CondCache
c
    | Text -> CondCache -> Bool
forall k a. Ord k => k -> Map k a -> Bool
M.member Text
k CondCache
c = CondCache
c
    | Bool
otherwise =
        CondCache -> (CondVec -> CondCache) -> Maybe CondVec -> CondCache
forall b a. b -> (a -> b) -> Maybe a -> b
maybe CondCache
c (\CondVec
cv -> Text -> Vector Bool -> CondCache -> CondCache
forall k a. Ord k => k -> a -> Map k a -> Map k a
M.insert Text
k (CondVec -> Vector Bool
cvVec CondVec
cv) CondCache
c) (DataFrame -> Expr Bool -> Maybe CondVec
materializeCondVec DataFrame
df Expr Bool
cond)
  where
    k :: Text
k = Expr Bool -> Text
condCacheKey Expr Bool
cond

{- | A condition's truth vector: a cache hit, else interpret over the
DataFrame. 'Nothing' mirrors the interpret-failure fallback (route left).
-}
lookupCondVec :: CondCache -> DataFrame -> Expr Bool -> Maybe (VU.Vector Bool)
lookupCondVec :: CondCache -> DataFrame -> Expr Bool -> Maybe (Vector Bool)
lookupCondVec CondCache
cache DataFrame
df Expr Bool
cond = case Text -> CondCache -> Maybe (Vector Bool)
forall k a. Ord k => k -> Map k a -> Maybe a
M.lookup (Expr Bool -> Text
condCacheKey Expr Bool
cond) CondCache
cache of
    hit :: Maybe (Vector Bool)
hit@(Just Vector Bool
_) -> Maybe (Vector Bool)
hit
    Maybe (Vector Bool)
Nothing -> CondVec -> Vector Bool
cvVec (CondVec -> Vector Bool) -> Maybe CondVec -> Maybe (Vector Bool)
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> DataFrame -> Expr Bool -> Maybe CondVec
materializeCondVec DataFrame
df Expr Bool
cond

-- | Partition row indices by a truth vector: @True@ → left, @False@ → right.
partitionByVec :: VU.Vector Bool -> V.Vector Int -> (V.Vector Int, V.Vector Int)
partitionByVec :: Vector Bool -> Vector Int -> (Vector Int, Vector Int)
partitionByVec 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.!)

-- | Count care points the truth vector routes to the wrong child.
countErrorsByVec :: VU.Vector Bool -> [CarePoint] -> Int
countErrorsByVec :: Vector Bool -> [CarePoint] -> Int
countErrorsByVec Vector Bool
boolVals = [CarePoint] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length ([CarePoint] -> Int)
-> ([CarePoint] -> [CarePoint]) -> [CarePoint] -> Int
forall b c a. (b -> c) -> (a -> b) -> a -> c
. (CarePoint -> Bool) -> [CarePoint] -> [CarePoint]
forall a. (a -> Bool) -> [a] -> [a]
filter CarePoint -> Bool
misrouted
  where
    misrouted :: CarePoint -> Bool
misrouted CarePoint
cp = (Vector Bool
boolVals Vector Bool -> Int -> Bool
forall a. Unbox a => Vector a -> Int -> a
VU.! CarePoint -> Int
cpIndex CarePoint
cp) Bool -> Bool -> Bool
forall a. Eq a => a -> a -> Bool
/= (CarePoint -> Direction
cpCorrectDir CarePoint
cp Direction -> Direction -> Bool
forall a. Eq a => a -> a -> Bool
== Direction
GoLeft)

{- | A same-column same-direction Double threshold comparison, with a rebuild
function to swap in a new threshold.
-}
data ThreshCmp = ThreshCmp
    { ThreshCmp -> Text
tcCol :: !T.Text
    , ThreshCmp -> Text
tcName :: !T.Text
    , ThreshCmp -> Double
tcThr :: !Double
    , ThreshCmp -> Double -> Expr Bool
tcRebuild :: Double -> Expr Bool
    }

asDoubleThreshold :: Expr Bool -> Maybe ThreshCmp
asDoubleThreshold :: Expr Bool -> Maybe ThreshCmp
asDoubleThreshold (Binary op c b Bool
op (Col Text
c :: Expr cc) (Lit (b
t :: tt))) =
    case ( TypeRep c -> TypeRep Double -> Maybe (c :~: 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 @cc) (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @Double)
         , 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 @tt) (forall a. Typeable a => TypeRep a
forall {k} (a :: k). Typeable a => TypeRep a
typeRep @Double)
         ) of
        (Just c :~: Double
Refl, Just b :~: Double
Refl) -> ThreshCmp -> Maybe ThreshCmp
forall a. a -> Maybe a
Just (Text -> Text -> Double -> (Double -> Expr Bool) -> ThreshCmp
ThreshCmp Text
c (op c b Bool -> Text
forall a b c. op a b c -> Text
forall (op :: * -> * -> * -> *) a b c.
BinaryOp op =>
op a b c -> Text
binaryName op c b Bool
op) b
Double
t (op c b Bool -> Expr c -> Expr b -> Expr Bool
forall (op :: * -> * -> * -> *) c b a.
(BinaryOp op, Columnable c, Columnable b, Columnable a) =>
op c b a -> Expr c -> Expr b -> Expr a
Binary op c b Bool
op (Text -> Expr c
forall a. Columnable a => Text -> Expr a
Col Text
c) (Expr b -> Expr Bool) -> (Double -> Expr b) -> Double -> Expr Bool
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Double -> Expr b
Double -> Expr Double
forall a. Columnable a => a -> Expr a
Lit))
        (Maybe (c :~: Double), Maybe (b :~: Double))
_ -> Maybe ThreshCmp
forall a. Maybe a
Nothing
asDoubleThreshold Expr Bool
_ = Maybe ThreshCmp
forall a. Maybe a
Nothing

directionalNames :: [T.Text]
directionalNames :: [Text]
directionalNames = [Text
"lt", Text
"leq", Text
"gt", Text
"geq"]

{- | Tighter (AND) or looser (OR) of two same-direction thresholds: @<@/@<=@
are left-half-spaces (AND = min), @>@/@>=@ are right-half-spaces (AND = max).
-}
chooseThreshold :: Bool -> T.Text -> Double -> Double -> Double
chooseThreshold :: Bool -> Text -> Double -> Double -> Double
chooseThreshold Bool
isAnd Text
name Double
t1 Double
t2
    | Bool
leftDir = if Bool
isAnd then Double -> Double -> Double
forall a. Ord a => a -> a -> a
min Double
t1 Double
t2 else Double -> Double -> Double
forall a. Ord a => a -> a -> a
max Double
t1 Double
t2
    | Bool
otherwise = if Bool
isAnd then Double -> Double -> Double
forall a. Ord a => a -> a -> a
max Double
t1 Double
t2 else Double -> Double -> Double
forall a. Ord a => a -> a -> a
min Double
t1 Double
t2
  where
    leftDir :: Bool
leftDir = Text
name Text -> Text -> Bool
forall a. Eq a => a -> a -> Bool
== Text
"lt" Bool -> Bool -> Bool
|| Text
name Text -> Text -> Bool
forall a. Eq a => a -> a -> Bool
== Text
"leq"

{- | Collapse two same-column same-direction strict-Double comparisons into one
comparison (the @True@ argument selects AND, @False@ OR); 'Nothing' otherwise.
-}
consolidateThreshold :: Bool -> Expr Bool -> Expr Bool -> Maybe (Expr Bool)
consolidateThreshold :: Bool -> Expr Bool -> Expr Bool -> Maybe (Expr Bool)
consolidateThreshold Bool
isAnd Expr Bool
ea Expr Bool
eb = do
    ThreshCmp
a <- Expr Bool -> Maybe ThreshCmp
asDoubleThreshold Expr Bool
ea
    ThreshCmp
b <- Expr Bool -> Maybe ThreshCmp
asDoubleThreshold Expr Bool
eb
    if ThreshCmp -> Text
tcCol ThreshCmp
a Text -> Text -> Bool
forall a. Eq a => a -> a -> Bool
== ThreshCmp -> Text
tcCol ThreshCmp
b Bool -> Bool -> Bool
&& ThreshCmp -> Text
tcName ThreshCmp
a Text -> Text -> Bool
forall a. Eq a => a -> a -> Bool
== ThreshCmp -> Text
tcName ThreshCmp
b Bool -> Bool -> Bool
&& ThreshCmp -> Text
tcName ThreshCmp
a Text -> [Text] -> Bool
forall a. Eq a => a -> [a] -> Bool
forall (t :: * -> *) a. (Foldable t, Eq a) => a -> t a -> Bool
`elem` [Text]
directionalNames
        then Expr Bool -> Maybe (Expr Bool)
forall a. a -> Maybe a
Just (ThreshCmp -> Double -> Expr Bool
tcRebuild ThreshCmp
a (Bool -> Text -> Double -> Double -> Double
chooseThreshold Bool
isAnd (ThreshCmp -> Text
tcName ThreshCmp
a) (ThreshCmp -> Double
tcThr ThreshCmp
a) (ThreshCmp -> Double
tcThr ThreshCmp
b)))
        else Maybe (Expr Bool)
forall a. Maybe a
Nothing

{- | AND-combine two cached conditions: idempotence and threshold consolidation
first, else the generic @F.and@; the vector is always the elementwise AND.
-}
combineAndVec :: CondVec -> CondVec -> CondVec
combineAndVec :: CondVec -> CondVec -> CondVec
combineAndVec CondVec
a CondVec
b
    | Expr Bool -> Expr Bool -> Bool
forall a. Columnable a => Expr a -> Expr a -> Bool
eqExpr (CondVec -> Expr Bool
cvExpr CondVec
a) (CondVec -> Expr Bool
cvExpr CondVec
b) = CondVec
a
    | Bool
otherwise = Expr Bool -> Vector Bool -> CondVec
CondVec Expr Bool
expr ((Bool -> Bool -> Bool) -> Vector Bool -> Vector Bool -> Vector Bool
forall a b c.
(Unbox a, Unbox b, Unbox c) =>
(a -> b -> c) -> Vector a -> Vector b -> Vector c
VU.zipWith Bool -> Bool -> Bool
(&&) (CondVec -> Vector Bool
cvVec CondVec
a) (CondVec -> Vector Bool
cvVec CondVec
b))
  where
    expr :: Expr Bool
expr =
        Expr Bool -> Maybe (Expr Bool) -> Expr Bool
forall a. a -> Maybe a -> a
fromMaybe
            (Expr Bool -> Expr Bool -> Expr Bool
F.and (CondVec -> Expr Bool
cvExpr CondVec
a) (CondVec -> Expr Bool
cvExpr CondVec
b))
            (Bool -> Expr Bool -> Expr Bool -> Maybe (Expr Bool)
consolidateThreshold Bool
True (CondVec -> Expr Bool
cvExpr CondVec
a) (CondVec -> Expr Bool
cvExpr CondVec
b))

{- | OR-combine two cached conditions (see 'combineAndVec'; AND/OR direction
differs in 'consolidateThreshold').
-}
combineOrVec :: CondVec -> CondVec -> CondVec
combineOrVec :: CondVec -> CondVec -> CondVec
combineOrVec CondVec
a CondVec
b
    | Expr Bool -> Expr Bool -> Bool
forall a. Columnable a => Expr a -> Expr a -> Bool
eqExpr (CondVec -> Expr Bool
cvExpr CondVec
a) (CondVec -> Expr Bool
cvExpr CondVec
b) = CondVec
a
    | Bool
otherwise = Expr Bool -> Vector Bool -> CondVec
CondVec Expr Bool
expr ((Bool -> Bool -> Bool) -> Vector Bool -> Vector Bool -> Vector Bool
forall a b c.
(Unbox a, Unbox b, Unbox c) =>
(a -> b -> c) -> Vector a -> Vector b -> Vector c
VU.zipWith Bool -> Bool -> Bool
(||) (CondVec -> Vector Bool
cvVec CondVec
a) (CondVec -> Vector Bool
cvVec CondVec
b))
  where
    expr :: Expr Bool
expr =
        Expr Bool -> Maybe (Expr Bool) -> Expr Bool
forall a. a -> Maybe a -> a
fromMaybe
            (Expr Bool -> Expr Bool -> Expr Bool
F.or (CondVec -> Expr Bool
cvExpr CondVec
a) (CondVec -> Expr Bool
cvExpr CondVec
b))
            (Bool -> Expr Bool -> Expr Bool -> Maybe (Expr Bool)
consolidateThreshold Bool
False (CondVec -> Expr Bool
cvExpr CondVec
a) (CondVec -> Expr Bool
cvExpr CondVec
b))