{-# LANGUAGE BangPatterns #-}
{-# LANGUAGE ScopedTypeVariables #-}

{- | Variance-reduction (weighted-SSE) regression trees over the CART feature
machinery; leaves predict the weighted mean of their rows. 'fitRegTreeOn' lets
gradient boosting refit on residuals without re-extracting features.
-}
module DataFrame.DecisionTree.Regression (
    RegTreeConfig (..),
    defaultRegTreeConfig,
    -- | Implementation verb used by the fit\/predict instances and boosting.
    fitRegTreeOn,
) where

import Control.Parallel (par, pseq)
import Data.Maybe (maybeToList)
import qualified Data.Vector as V
import qualified Data.Vector.Unboxed as VU

import DataFrame.DecisionTree.Cart (CartFeature (..), sortIndicesByValue)
import DataFrame.DecisionTree.Types (Tree (..))

-- | Stopping criteria for the regression tree.
data RegTreeConfig = RegTreeConfig
    { RegTreeConfig -> Int
rtMaxDepth :: !Int
    , RegTreeConfig -> Int
rtMinSamplesSplit :: !Int
    , RegTreeConfig -> Int
rtMinLeafSize :: !Int
    , RegTreeConfig -> Double
rtMinImpurityDecrease :: !Double
    }
    deriving (RegTreeConfig -> RegTreeConfig -> Bool
(RegTreeConfig -> RegTreeConfig -> Bool)
-> (RegTreeConfig -> RegTreeConfig -> Bool) -> Eq RegTreeConfig
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: RegTreeConfig -> RegTreeConfig -> Bool
== :: RegTreeConfig -> RegTreeConfig -> Bool
$c/= :: RegTreeConfig -> RegTreeConfig -> Bool
/= :: RegTreeConfig -> RegTreeConfig -> Bool
Eq, Int -> RegTreeConfig -> ShowS
[RegTreeConfig] -> ShowS
RegTreeConfig -> String
(Int -> RegTreeConfig -> ShowS)
-> (RegTreeConfig -> String)
-> ([RegTreeConfig] -> ShowS)
-> Show RegTreeConfig
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> RegTreeConfig -> ShowS
showsPrec :: Int -> RegTreeConfig -> ShowS
$cshow :: RegTreeConfig -> String
show :: RegTreeConfig -> String
$cshowList :: [RegTreeConfig] -> ShowS
showList :: [RegTreeConfig] -> ShowS
Show)

defaultRegTreeConfig :: RegTreeConfig
defaultRegTreeConfig :: RegTreeConfig
defaultRegTreeConfig =
    RegTreeConfig
        { rtMaxDepth :: Int
rtMaxDepth = Int
3
        , rtMinSamplesSplit :: Int
rtMinSamplesSplit = Int
2
        , rtMinLeafSize :: Int
rtMinLeafSize = Int
1
        , rtMinImpurityDecrease :: Double
rtMinImpurityDecrease = Double
0.0
        }

{- | Fit on pre-extracted features, a target vector, and optional per-row
weights (length @n@). Used by gradient boosting on residual targets.
-}
fitRegTreeOn ::
    RegTreeConfig ->
    V.Vector CartFeature ->
    VU.Vector Double ->
    Maybe (VU.Vector Double) ->
    Tree Double
fitRegTreeOn :: RegTreeConfig
-> Vector CartFeature
-> Vector Double
-> Maybe (Vector Double)
-> Tree Double
fitRegTreeOn RegTreeConfig
cfg Vector CartFeature
feats Vector Double
y Maybe (Vector Double)
mw = Int -> Vector Int -> Vector (Vector Int) -> Tree Double
buildNode Int
0 (Int -> Int -> Vector Int
forall a. (Unbox a, Num a) => a -> Int -> Vector a
VU.enumFromN Int
0 Int
n) Vector (Vector Int)
featSorted
  where
    n :: Int
n = Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Double
y
    weightAt :: Int -> Double
weightAt Int
i = Double
-> (Vector Double -> Double) -> Maybe (Vector Double) -> Double
forall b a. b -> (a -> b) -> Maybe a -> b
maybe Double
1 (Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i) Maybe (Vector Double)
mw
    featSorted :: Vector (Vector Int)
featSorted = (CartFeature -> Vector Int)
-> Vector CartFeature -> Vector (Vector Int)
forall a b. (a -> b) -> Vector a -> Vector b
V.map (Vector Double -> Vector Int
sortIndicesByValue (Vector Double -> Vector Int)
-> (CartFeature -> Vector Double) -> CartFeature -> Vector Int
forall b c a. (b -> c) -> (a -> b) -> a -> c
. CartFeature -> Vector Double
cfValues) Vector CartFeature
feats

    buildNode :: Int -> Vector Int -> Vector (Vector Int) -> Tree Double
buildNode Int
depth Vector Int
idxs Vector (Vector Int)
sortedByFeat
        | Int
depth Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= RegTreeConfig -> Int
rtMaxDepth RegTreeConfig
cfg Bool -> Bool -> Bool
|| Vector Int -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Int
idxs Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< RegTreeConfig -> Int
rtMinSamplesSplit RegTreeConfig
cfg = Tree Double
leaf
        | Bool
otherwise =
            Tree Double
-> ((Int, Double) -> Tree Double)
-> Maybe (Int, Double)
-> Tree Double
forall b a. b -> (a -> b) -> Maybe a -> b
maybe Tree Double
leaf (Int
-> Vector Int
-> Vector (Vector Int)
-> (Int, Double)
-> Tree Double
splitNode Int
depth Vector Int
idxs Vector (Vector Int)
sortedByFeat) (Vector Int -> Vector (Vector Int) -> Maybe (Int, Double)
bestSplit Vector Int
idxs Vector (Vector Int)
sortedByFeat)
      where
        leaf :: Tree Double
leaf = Double -> Tree Double
forall a. a -> Tree a
Leaf (Vector Int -> Double
weightedMean Vector Int
idxs)

    splitNode :: Int
-> Vector Int
-> Vector (Vector Int)
-> (Int, Double)
-> Tree Double
splitNode Int
depth Vector Int
idxs Vector (Vector Int)
sortedByFeat (Int
fj, Double
thr)
        | Vector Int -> Bool
forall a. Unbox a => Vector a -> Bool
VU.null Vector Int
lefts Bool -> Bool -> Bool
|| Vector Int -> Bool
forall a. Unbox a => Vector a -> Bool
VU.null Vector Int
rights = Double -> Tree Double
forall a. a -> Tree a
Leaf (Vector Int -> Double
weightedMean Vector Int
idxs)
        | Bool
otherwise =
            Tree Double -> ()
forceTree Tree Double
l () -> Tree Double -> Tree Double
forall a b. a -> b -> b
`par` (Tree Double -> ()
forceTree Tree Double
r () -> Tree Double -> Tree Double
forall a b. a -> b -> b
`pseq` Expr Bool -> Tree Double -> Tree Double -> Tree Double
forall a. Expr Bool -> Tree a -> Tree a -> Tree a
Branch (CartFeature -> Double -> Expr Bool
cfPred (Vector CartFeature
feats Vector CartFeature -> Int -> CartFeature
forall a. Vector a -> Int -> a
V.! Int
fj) Double
thr) Tree Double
l Tree Double
r)
      where
        vals :: Vector Double
vals = CartFeature -> Vector Double
cfValues (Vector CartFeature
feats Vector CartFeature -> Int -> CartFeature
forall a. Vector a -> Int -> a
V.! Int
fj)
        goesLeft :: Int -> Bool
goesLeft Int
i = Vector Double
vals Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
<= Double
thr
        lefts :: Vector Int
lefts = (Int -> Bool) -> Vector Int -> Vector Int
forall a. Unbox a => (a -> Bool) -> Vector a -> Vector a
VU.filter Int -> Bool
goesLeft Vector Int
idxs
        rights :: Vector Int
rights = (Int -> Bool) -> Vector Int -> Vector Int
forall a. Unbox a => (a -> Bool) -> Vector a -> Vector a
VU.filter (Bool -> Bool
not (Bool -> Bool) -> (Int -> Bool) -> Int -> Bool
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Int -> Bool
goesLeft) Vector Int
idxs
        l :: Tree Double
l = Int -> Vector Int -> Vector (Vector Int) -> Tree Double
buildNode (Int
depth Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Vector Int
lefts ((Vector Int -> Vector Int)
-> Vector (Vector Int) -> Vector (Vector Int)
forall a b. (a -> b) -> Vector a -> Vector b
V.map ((Int -> Bool) -> Vector Int -> Vector Int
forall a. Unbox a => (a -> Bool) -> Vector a -> Vector a
VU.filter Int -> Bool
goesLeft) Vector (Vector Int)
sortedByFeat)
        r :: Tree Double
r =
            Int -> Vector Int -> Vector (Vector Int) -> Tree Double
buildNode (Int
depth Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Vector Int
rights ((Vector Int -> Vector Int)
-> Vector (Vector Int) -> Vector (Vector Int)
forall a b. (a -> b) -> Vector a -> Vector b
V.map ((Int -> Bool) -> Vector Int -> Vector Int
forall a. Unbox a => (a -> Bool) -> Vector a -> Vector a
VU.filter (Bool -> Bool
not (Bool -> Bool) -> (Int -> Bool) -> Int -> Bool
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Int -> Bool
goesLeft)) Vector (Vector Int)
sortedByFeat)

    weightedMean :: Vector Int -> Double
weightedMean Vector Int
idxs =
        let (Double
w, Double
sy) = ((Double, Double) -> Int -> (Double, Double))
-> (Double, Double) -> Vector Int -> (Double, Double)
forall b a. Unbox b => (a -> b -> a) -> a -> Vector b -> a
VU.foldl' (Double, Double) -> Int -> (Double, Double)
step (Double
0, Double
0) Vector Int
idxs
            step :: (Double, Double) -> Int -> (Double, Double)
step (!Double
a, !Double
b) Int
i = (Double
a Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Int -> Double
weightAt Int
i, Double
b Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Int -> Double
weightAt Int
i Double -> Double -> Double
forall a. Num a => a -> a -> a
* (Vector Double
y Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i))
         in if Double
w Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
== Double
0 then Double
0 else Double
sy Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
w

    bestSplit :: Vector Int -> Vector (Vector Int) -> Maybe (Int, Double)
bestSplit Vector Int
idxs Vector (Vector Int)
sortedByFeat
        | [(Double, Int, Double)] -> Bool
forall a. [a] -> Bool
forall (t :: * -> *) a. Foldable t => t a -> Bool
null [(Double, Int, Double)]
candidates = Maybe (Int, Double)
forall a. Maybe a
Nothing
        | Double
red Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
> Double
0 Bool -> Bool -> Bool
&& Double
red Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
>= RegTreeConfig -> Double
rtMinImpurityDecrease RegTreeConfig
cfg = (Int, Double) -> Maybe (Int, Double)
forall a. a -> Maybe a
Just (Int
fj, Double
thr)
        | Bool
otherwise = Maybe (Int, Double)
forall a. Maybe a
Nothing
      where
        (Double
totW, Double
totSY, Double
totSY2) = Vector Int -> (Double, Double, Double)
moments Vector Int
idxs
        nodeSSE :: Double
nodeSSE = Double -> Double -> Double -> Double
sse Double
totSY Double
totSY2 Double
totW
        candidates :: [(Double, Int, Double)]
candidates =
            [ (Double
red', Int
fj', Double
thr')
            | Int
fj' <- [Int
0 .. Vector CartFeature -> Int
forall a. Vector a -> Int
V.length Vector CartFeature
feats Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]
            , (Double
thr', Double
red') <-
                Int
-> Vector Int
-> Double
-> Double
-> Double
-> Double
-> [(Double, Double)]
bestThreshold Int
fj' (Vector (Vector Int)
sortedByFeat Vector (Vector Int) -> Int -> Vector Int
forall a. Vector a -> Int -> a
V.! Int
fj') Double
totW Double
totSY Double
totSY2 Double
nodeSSE
            ]
        (Double
red, Int
fj, Double
thr) = [(Double, Int, Double)] -> (Double, Int, Double)
forall a b c. Ord a => [(a, b, c)] -> (a, b, c)
maximumByFst [(Double, Int, Double)]
candidates

    bestThreshold :: Int
-> Vector Int
-> Double
-> Double
-> Double
-> Double
-> [(Double, Double)]
bestThreshold Int
fj Vector Int
sorted Double
totW Double
totSY Double
totSY2 Double
nodeSSE = Maybe (Double, Double) -> [(Double, Double)]
forall a. Maybe a -> [a]
maybeToList (Int
-> Double
-> Double
-> Double
-> Maybe (Double, Double)
-> Maybe (Double, Double)
go Int
0 Double
0 Double
0 Double
0 Maybe (Double, Double)
forall a. Maybe a
Nothing)
      where
        vals :: Vector Double
vals = CartFeature -> Vector Double
cfValues (Vector CartFeature
feats Vector CartFeature -> Int -> CartFeature
forall a. Vector a -> Int -> a
V.! Int
fj)
        m :: Int
m = Vector Int -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Int
sorted
        go :: Int
-> Double
-> Double
-> Double
-> Maybe (Double, Double)
-> Maybe (Double, Double)
go !Int
k !Double
wl !Double
syl !Double
syl2 Maybe (Double, Double)
best
            | Int
k Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
m Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1 = Maybe (Double, Double)
best
            | Bool
otherwise = Int
-> Double
-> Double
-> Double
-> Maybe (Double, Double)
-> Maybe (Double, Double)
go (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Double
wl' Double
syl' Double
syl2' Maybe (Double, Double)
best'
          where
            i :: Int
i = Vector Int
sorted Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
k
            next :: Int
next = Vector Int
sorted Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.! (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
            wi :: Double
wi = Int -> Double
weightAt Int
i
            yi :: Double
yi = Vector Double
y Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i
            wl' :: Double
wl' = Double
wl Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
wi
            syl' :: Double
syl' = Double
syl Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
wi Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
yi
            syl2' :: Double
syl2' = Double
syl2 Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
wi Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
yi Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
yi
            wr :: Double
wr = Double
totW Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
wl'
            leafSizesOk :: Bool
leafSizesOk = Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1 Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= RegTreeConfig -> Int
rtMinLeafSize RegTreeConfig
cfg Bool -> Bool -> Bool
&& Int
m Int -> Int -> Int
forall a. Num a => a -> a -> a
- (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= RegTreeConfig -> Int
rtMinLeafSize RegTreeConfig
cfg
            splittable :: Bool
splittable = Vector Double
vals Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
/= Vector Double
vals Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
next Bool -> Bool -> Bool
&& Bool
leafSizesOk Bool -> Bool -> Bool
&& Double
wl' Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
> Double
0 Bool -> Bool -> Bool
&& Double
wr Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
> Double
0
            reduction :: Double
reduction = Double
nodeSSE Double -> Double -> Double
forall a. Num a => a -> a -> a
- (Double -> Double -> Double -> Double
sse Double
syl' Double
syl2' Double
wl' Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double -> Double -> Double -> Double
sse (Double
totSY Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
syl') (Double
totSY2 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
syl2') Double
wr)
            best' :: Maybe (Double, Double)
best'
                | Bool
splittable Bool -> Bool -> Bool
&& Bool
-> ((Double, Double) -> Bool) -> Maybe (Double, Double) -> Bool
forall b a. b -> (a -> b) -> Maybe a -> b
maybe Bool
True ((Double
reduction Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
>) (Double -> Bool)
-> ((Double, Double) -> Double) -> (Double, Double) -> Bool
forall b c a. (b -> c) -> (a -> b) -> a -> c
. (Double, Double) -> Double
forall a b. (a, b) -> b
snd) Maybe (Double, Double)
best =
                    (Double, Double) -> Maybe (Double, Double)
forall a. a -> Maybe a
Just ((Vector Double
vals Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Vector Double
vals Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
next) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
2, Double
reduction)
                | Bool
otherwise = Maybe (Double, Double)
best

    moments :: Vector Int -> (Double, Double, Double)
moments = ((Double, Double, Double) -> Int -> (Double, Double, Double))
-> (Double, Double, Double)
-> Vector Int
-> (Double, Double, Double)
forall b a. Unbox b => (a -> b -> a) -> a -> Vector b -> a
VU.foldl' (Double, Double, Double) -> Int -> (Double, Double, Double)
step (Double
0, Double
0, Double
0)
      where
        step :: (Double, Double, Double) -> Int -> (Double, Double, Double)
step (!Double
w, !Double
sy, !Double
sy2) Int
i =
            let wi :: Double
wi = Int -> Double
weightAt Int
i; yi :: Double
yi = Vector Double
y Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i
             in (Double
w Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
wi, Double
sy Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
wi Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
yi, Double
sy2 Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
wi Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
yi Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
yi)

safeDiv :: Double -> Double -> Double
safeDiv :: Double -> Double -> Double
safeDiv Double
a Double
b = if Double
b Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
== Double
0 then Double
0 else Double
a Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
b

-- | Weighted SSE of a node from its Σy, Σy², and total weight: @Σy² − (Σy)²/w@.
sse :: Double -> Double -> Double -> Double
sse :: Double -> Double -> Double -> Double
sse Double
sumY Double
sumSq Double
w = Double
sumSq Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double -> Double -> Double
safeDiv (Double
sumY Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
sumY) Double
w

{- | Force a subtree to WHNF throughout so the spark scoring the sibling has
substantial work to evaluate; pure and value-preserving (cf. 'Tao').
-}
forceTree :: Tree Double -> ()
forceTree :: Tree Double -> ()
forceTree (Leaf Double
v) = Double
v Double -> () -> ()
forall a b. a -> b -> b
`seq` ()
forceTree (Branch Expr Bool
_ Tree Double
l Tree Double
r) = Tree Double -> ()
forceTree Tree Double
l () -> () -> ()
forall a b. a -> b -> b
`seq` Tree Double -> ()
forceTree Tree Double
r

maximumByFst :: (Ord a) => [(a, b, c)] -> (a, b, c)
maximumByFst :: forall a b c. Ord a => [(a, b, c)] -> (a, b, c)
maximumByFst = ((a, b, c) -> (a, b, c) -> (a, b, c)) -> [(a, b, c)] -> (a, b, c)
forall a. (a -> a -> a) -> [a] -> a
forall (t :: * -> *) a. Foldable t => (a -> a -> a) -> t a -> a
foldr1 (\x :: (a, b, c)
x@(a
a, b
_, c
_) y :: (a, b, c)
y@(a
b, b
_, c
_) -> if a
a a -> a -> Bool
forall a. Ord a => a -> a -> Bool
>= a
b then (a, b, c)
x else (a, b, c)
y)