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

{- | Post-convergence simplification of a fitted tree and its expression form:
drop branches forced by path-condition entailment, collapse identical
siblings, and fold redundant nested conditionals.
-}
module DataFrame.DecisionTree.Prune (
    pruneDead,
    treeEq,
    pruneExpr,
) where

import DataFrame.DecisionTree.Types (Tree (..))
import DataFrame.Internal.Column (Columnable)
import DataFrame.Internal.Expression (Expr (..), eqExpr)
import DataFrame.Internal.Simplify (PredFact, entails, factFalse, factTrue)

{- | Drop branches whose test is forced by the path conditions reaching them,
and collapse @Branch c t t@ to @t@. Sound for the decidable threshold subset;
other tests are left untouched.
-}
pruneDead :: forall a. (Columnable a) => Tree a -> Tree a
pruneDead :: forall a. Columnable a => Tree a -> Tree a
pruneDead = [PredFact] -> Tree a -> Tree a
go []
  where
    go :: [PredFact] -> Tree a -> Tree a
    go :: [PredFact] -> Tree a -> Tree a
go [PredFact]
_ (Leaf a
v) = a -> Tree a
forall a. a -> Tree a
Leaf a
v
    go [PredFact]
facts (Branch Expr Bool
cond Tree a
left Tree a
right) = case [PredFact] -> Expr Bool -> Maybe Bool
entails [PredFact]
facts Expr Bool
cond of
        Just Bool
True -> [PredFact] -> Tree a -> Tree a
go [PredFact]
facts Tree a
left
        Just Bool
False -> [PredFact] -> Tree a -> Tree a
go [PredFact]
facts Tree a
right
        Maybe Bool
Nothing ->
            Expr Bool -> Tree a -> Tree a -> Tree a
forall a. Columnable a => Expr Bool -> Tree a -> Tree a -> Tree a
reconcile
                Expr Bool
cond
                ([PredFact] -> Tree a -> Tree a
go (Maybe PredFact -> [PredFact] -> [PredFact]
addFact (Expr Bool -> Maybe PredFact
factTrue Expr Bool
cond) [PredFact]
facts) Tree a
left)
                ([PredFact] -> Tree a -> Tree a
go (Maybe PredFact -> [PredFact] -> [PredFact]
addFact (Expr Bool -> Maybe PredFact
factFalse Expr Bool
cond) [PredFact]
facts) Tree a
right)

reconcile :: (Columnable a) => Expr Bool -> Tree a -> Tree a -> Tree a
reconcile :: forall a. Columnable a => Expr Bool -> Tree a -> Tree a -> Tree a
reconcile Expr Bool
cond Tree a
left Tree a
right
    | Tree a -> Tree a -> Bool
forall a. Columnable a => Tree a -> Tree a -> Bool
treeEq Tree a
left Tree a
right = Tree a
left
    | Bool
otherwise = 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

addFact :: Maybe PredFact -> [PredFact] -> [PredFact]
addFact :: Maybe PredFact -> [PredFact] -> [PredFact]
addFact (Just PredFact
f) [PredFact]
fs = PredFact
f PredFact -> [PredFact] -> [PredFact]
forall a. a -> [a] -> [a]
: [PredFact]
fs
addFact Maybe PredFact
Nothing [PredFact]
fs = [PredFact]
fs

treeEq :: (Columnable a) => Tree a -> Tree a -> Bool
treeEq :: forall a. Columnable a => Tree a -> Tree a -> Bool
treeEq (Leaf a
x) (Leaf a
y) = a
x a -> a -> Bool
forall a. Eq a => a -> a -> Bool
== a
y
treeEq (Branch Expr Bool
c1 Tree a
l1 Tree a
r1) (Branch Expr Bool
c2 Tree a
l2 Tree a
r2) = Expr Bool -> Expr Bool -> Bool
forall a. Columnable a => Expr a -> Expr a -> Bool
eqExpr Expr Bool
c1 Expr Bool
c2 Bool -> Bool -> Bool
&& Tree a -> Tree a -> Bool
forall a. Columnable a => Tree a -> Tree a -> Bool
treeEq Tree a
l1 Tree a
l2 Bool -> Bool -> Bool
&& Tree a -> Tree a -> Bool
forall a. Columnable a => Tree a -> Tree a -> Bool
treeEq Tree a
r1 Tree a
r2
treeEq Tree a
_ Tree a
_ = Bool
False

{- | Recursively fold @If@ expressions whose branches coincide or nest the same
condition; leave other expressions structurally unchanged.
-}
pruneExpr :: forall a. (Columnable a) => Expr a -> Expr a
pruneExpr :: forall a. Columnable a => Expr a -> Expr a
pruneExpr (If Expr Bool
cond Expr a
t0 Expr a
f0) = Expr Bool -> Expr a -> Expr a -> Expr a
forall a. Columnable a => Expr Bool -> Expr a -> Expr a -> Expr a
collapseIf Expr Bool
cond (Expr a -> Expr a
forall a. Columnable a => Expr a -> Expr a
pruneExpr Expr a
t0) (Expr a -> Expr a
forall a. Columnable a => Expr a -> Expr a
pruneExpr Expr a
f0)
pruneExpr (Unary op b a
op Expr b
e) = op b a -> Expr b -> Expr a
forall (op :: * -> * -> *) a b.
(UnaryOp op, Columnable a, Columnable b) =>
op b a -> Expr b -> Expr a
Unary op b a
op (Expr b -> Expr b
forall a. Columnable a => Expr a -> Expr a
pruneExpr Expr b
e)
pruneExpr (Binary op c b a
op Expr c
l Expr b
r) = op c b a -> Expr c -> Expr b -> Expr a
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 a
op (Expr c -> Expr c
forall a. Columnable a => Expr a -> Expr a
pruneExpr Expr c
l) (Expr b -> Expr b
forall a. Columnable a => Expr a -> Expr a
pruneExpr Expr b
r)
pruneExpr Expr a
e = Expr a
e

collapseIf :: (Columnable a) => Expr Bool -> Expr a -> Expr a -> Expr a
collapseIf :: forall a. Columnable a => Expr Bool -> Expr a -> Expr a -> Expr a
collapseIf Expr Bool
cond Expr a
t Expr a
f
    | Expr a -> Expr a -> Bool
forall a. Columnable a => Expr a -> Expr a -> Bool
eqExpr Expr a
t Expr a
f = Expr a
t
    | If Expr Bool
ci Expr a
ti Expr a
_ <- Expr a
t, Expr Bool -> Expr Bool -> Bool
forall a. Columnable a => Expr a -> Expr a -> Bool
eqExpr Expr Bool
cond Expr Bool
ci = Expr Bool -> Expr a -> Expr a -> Expr a
forall a. Columnable a => Expr Bool -> Expr a -> Expr a -> Expr a
If Expr Bool
cond Expr a
ti Expr a
f
    | If Expr Bool
ci Expr a
_ Expr a
fi <- Expr a
f, Expr Bool -> Expr Bool -> Bool
forall a. Columnable a => Expr a -> Expr a -> Bool
eqExpr Expr Bool
cond Expr Bool
ci = Expr Bool -> Expr a -> Expr a -> Expr a
forall a. Columnable a => Expr Bool -> Expr a -> Expr a -> Expr a
If Expr Bool
cond Expr a
t Expr a
fi
    | Bool
otherwise = Expr Bool -> Expr a -> Expr a -> Expr a
forall a. Columnable a => Expr Bool -> Expr a -> Expr a -> Expr a
If Expr Bool
cond Expr a
t Expr a
f