{-# LANGUAGE BangPatterns #-}
{-# LANGUAGE ScopedTypeVariables #-}
module DataFrame.DecisionTree.Regression (
RegTreeConfig (..),
defaultRegTreeConfig,
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 (..))
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
}
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
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
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)