{-# LANGUAGE BangPatterns #-}
{-# LANGUAGE FlexibleInstances #-}
{-# LANGUAGE MultiParamTypeClasses #-}
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications #-}
{-# LANGUAGE TypeFamilies #-}

{- | Gradient boosting of regression trees (Friedman). Trees are fitted to the
negative gradient of the loss each round and accumulated with a shrinkage
factor; squared error gives regression, logistic deviance gives binary
classification. 'predict' is the additive score; 'gbProbaExpr' /
'gbDecisionExpr' give the classification probability / decision.
-}
module DataFrame.Boosting.GBM (
    module DataFrame.Model,
    GBLoss (..),
    GBConfig (..),
    defaultGBConfig,
    GBModel (..),
    gbExprAtStage,
    gbProbaExpr,
    gbDecisionExpr,
) where

import Control.Exception (throw)
import Data.Either (fromRight)
import qualified Data.Map.Strict as M
import qualified Data.Text as T
import qualified Data.Vector as V
import qualified Data.Vector.Unboxed as VU
import DataFrame.Errors (DataFrameException (..))

import DataFrame.DecisionTree.Cart (cartFeatures)
import DataFrame.DecisionTree.Fit (treeToExpr)
import DataFrame.DecisionTree.Regression (RegTreeConfig (..), fitRegTreeOn)
import DataFrame.DecisionTree.Types (Tree)
import DataFrame.Featurize.Internal (targetDoubles)
import qualified DataFrame.Functions as F
import DataFrame.Internal.Column (TypedColumn (..), toVector)
import DataFrame.Internal.DataFrame (DataFrame)
import DataFrame.Internal.Expression (Expr (..), getColumns)
import DataFrame.Internal.Interpreter (interpret)
import DataFrame.Model
import DataFrame.Operators ((.*.), (.+.), (.>.))

-- | The boosting loss.
data GBLoss = SquaredError | LogisticDeviance
    deriving (GBLoss -> GBLoss -> Bool
(GBLoss -> GBLoss -> Bool)
-> (GBLoss -> GBLoss -> Bool) -> Eq GBLoss
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: GBLoss -> GBLoss -> Bool
== :: GBLoss -> GBLoss -> Bool
$c/= :: GBLoss -> GBLoss -> Bool
/= :: GBLoss -> GBLoss -> Bool
Eq, Int -> GBLoss -> ShowS
[GBLoss] -> ShowS
GBLoss -> String
(Int -> GBLoss -> ShowS)
-> (GBLoss -> String) -> ([GBLoss] -> ShowS) -> Show GBLoss
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> GBLoss -> ShowS
showsPrec :: Int -> GBLoss -> ShowS
$cshow :: GBLoss -> String
show :: GBLoss -> String
$cshowList :: [GBLoss] -> ShowS
showList :: [GBLoss] -> ShowS
Show)

data GBConfig = GBConfig
    { GBConfig -> GBLoss
gbLoss :: !GBLoss
    , GBConfig -> Int
gbNEstimators :: !Int
    , GBConfig -> Double
gbLearningRate :: !Double
    , GBConfig -> Int
gbMaxDepth :: !Int
    , GBConfig -> Int
gbSeed :: !Int
    }
    deriving (GBConfig -> GBConfig -> Bool
(GBConfig -> GBConfig -> Bool)
-> (GBConfig -> GBConfig -> Bool) -> Eq GBConfig
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: GBConfig -> GBConfig -> Bool
== :: GBConfig -> GBConfig -> Bool
$c/= :: GBConfig -> GBConfig -> Bool
/= :: GBConfig -> GBConfig -> Bool
Eq, Int -> GBConfig -> ShowS
[GBConfig] -> ShowS
GBConfig -> String
(Int -> GBConfig -> ShowS)
-> (GBConfig -> String) -> ([GBConfig] -> ShowS) -> Show GBConfig
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> GBConfig -> ShowS
showsPrec :: Int -> GBConfig -> ShowS
$cshow :: GBConfig -> String
show :: GBConfig -> String
$cshowList :: [GBConfig] -> ShowS
showList :: [GBConfig] -> ShowS
Show)

defaultGBConfig :: GBConfig
defaultGBConfig :: GBConfig
defaultGBConfig =
    GBConfig
        { gbLoss :: GBLoss
gbLoss = GBLoss
SquaredError
        , gbNEstimators :: Int
gbNEstimators = Int
100
        , gbLearningRate :: Double
gbLearningRate = Double
0.1
        , gbMaxDepth :: Int
gbMaxDepth = Int
3
        , gbSeed :: Int
gbSeed = Int
0
        }

{- | A fitted gradient-boosting model. 'gbInit' is the constant initial score
(mean, or log-odds for classification); 'gbTrees' are the staged regression
trees.
-}
data GBModel = GBModel
    { GBModel -> Double
gbInit :: !Double
    , GBModel -> Vector (Tree Double)
gbTrees :: !(V.Vector (Tree Double))
    , GBModel -> Double
gbRate :: !Double
    , GBModel -> GBLoss
gbModelLoss :: !GBLoss
    , GBModel -> Vector Double
gbTrainScore :: !(VU.Vector Double)
    , GBModel -> Map Text Int
gbFeatureUsage :: !(M.Map T.Text Int)
    }
    deriving (Int -> GBModel -> ShowS
[GBModel] -> ShowS
GBModel -> String
(Int -> GBModel -> ShowS)
-> (GBModel -> String) -> ([GBModel] -> ShowS) -> Show GBModel
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> GBModel -> ShowS
showsPrec :: Int -> GBModel -> ShowS
$cshow :: GBModel -> String
show :: GBModel -> String
$cshowList :: [GBModel] -> ShowS
showList :: [GBModel] -> ShowS
Show)

instance Fit GBConfig (Expr Double) where
    type ModelOf GBConfig (Expr Double) = GBModel
    fit :: CheckFrame
  (FrameReq GBConfig (Expr Double)) (FrameFor (Expr Double)) =>
GBConfig
-> Expr Double
-> FrameFor (Expr Double)
-> FitResult
     (FrameFor (Expr Double)) (ModelOf GBConfig (Expr Double))
fit = GBConfig -> Expr Double -> DataFrame -> GBModel
GBConfig
-> Expr Double
-> FrameFor (Expr Double)
-> FitResult
     (FrameFor (Expr Double)) (ModelOf GBConfig (Expr Double))
fitGBM

instance Predict GBModel where
    type Prediction GBModel = Expr Double
    predict :: GBModel -> Prediction GBModel
predict = GBModel -> Expr Double
GBModel -> Prediction GBModel
gbExpr

-- | Fit a gradient-boosting ensemble predicting @target@ from the other columns.
fitGBM :: GBConfig -> Expr Double -> DataFrame -> GBModel
fitGBM :: GBConfig -> Expr Double -> DataFrame -> GBModel
fitGBM GBConfig
cfg target :: Expr Double
target@(Col Text
name) DataFrame
df =
    Double
-> Vector (Tree Double)
-> Double
-> GBLoss
-> Vector Double
-> Map Text Int
-> GBModel
GBModel
        Double
f0
        ([Tree Double] -> Vector (Tree Double)
forall a. [a] -> Vector a
V.fromList ([Tree Double] -> [Tree Double]
forall a. [a] -> [a]
reverse [Tree Double]
trees))
        Double
lr
        (GBConfig -> GBLoss
gbLoss GBConfig
cfg)
        ([Double] -> Vector Double
forall a. Unbox a => [a] -> Vector a
VU.fromList ([Double] -> [Double]
forall a. [a] -> [a]
reverse [Double]
scores))
        Map Text Int
usage
  where
    feats :: Vector CartFeature
feats = [CartFeature] -> Vector CartFeature
forall a. [a] -> Vector a
V.fromList (Text -> DataFrame -> [CartFeature]
cartFeatures Text
name DataFrame
df)
    y :: Vector Double
y = Expr Double -> DataFrame -> Vector Double
targetDoubles Expr Double
target DataFrame
df
    n :: Int
n = Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Double
y
    lr :: Double
lr = GBConfig -> Double
gbLearningRate GBConfig
cfg
    rtCfg :: RegTreeConfig
rtCfg =
        RegTreeConfig
            { rtMaxDepth :: Int
rtMaxDepth = GBConfig -> Int
gbMaxDepth GBConfig
cfg
            , rtMinSamplesSplit :: Int
rtMinSamplesSplit = Int
2
            , rtMinLeafSize :: Int
rtMinLeafSize = Int
1
            , rtMinImpurityDecrease :: Double
rtMinImpurityDecrease = Double
0.0
            }
    f0 :: Double
f0 = case GBConfig -> GBLoss
gbLoss GBConfig
cfg of
        GBLoss
SquaredError -> Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum Vector Double
y Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
1 Int
n)
        GBLoss
LogisticDeviance ->
            let p :: Double
p = Double -> Double
clamp01 (Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum Vector Double
y Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
1 Int
n))
             in Double -> Double
forall a. Floating a => a -> a
log (Double
p Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
p))
    ([Tree Double]
trees, [Double]
scores, Map Text Int
usage) = Int
-> Vector Double
-> [Tree Double]
-> [Double]
-> Map Text Int
-> ([Tree Double], [Double], Map Text Int)
boost Int
0 (Int -> Double -> Vector Double
forall a. Unbox a => Int -> a -> Vector a
VU.replicate Int
n Double
f0) [] [] Map Text Int
forall k a. Map k a
M.empty
    boost :: Int
-> Vector Double
-> [Tree Double]
-> [Double]
-> Map Text Int
-> ([Tree Double], [Double], Map Text Int)
boost !Int
m Vector Double
fScores [Tree Double]
ts [Double]
ss Map Text Int
usageAcc
        | Int
m Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= GBConfig -> Int
gbNEstimators GBConfig
cfg = ([Tree Double]
ts, [Double]
ss, Map Text Int
usageAcc)
        | Bool
otherwise =
            let (Vector Double
target', Maybe (Vector Double)
weights) = GBLoss
-> Vector Double
-> Vector Double
-> (Vector Double, Maybe (Vector Double))
newtonStep (GBConfig -> GBLoss
gbLoss GBConfig
cfg) Vector Double
y Vector Double
fScores
                tree :: Tree Double
tree = RegTreeConfig
-> Vector CartFeature
-> Vector Double
-> Maybe (Vector Double)
-> Tree Double
fitRegTreeOn RegTreeConfig
rtCfg Vector CartFeature
feats Vector Double
target' Maybe (Vector Double)
weights
                pred :: Vector Double
pred = DataFrame -> Tree Double -> Vector Double
predictTree DataFrame
df Tree Double
tree
                fScores' :: Vector Double
fScores' = (Double -> Double -> Double)
-> Vector Double -> Vector Double -> Vector Double
forall a b c.
(Unbox a, Unbox b, Unbox c) =>
(a -> b -> c) -> Vector a -> Vector b -> Vector c
VU.zipWith (\Double
f Double
p -> Double
f Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
lr Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
p) Vector Double
fScores Vector Double
pred
                score :: Double
score = GBLoss -> Vector Double -> Vector Double -> Double
lossValue (GBConfig -> GBLoss
gbLoss GBConfig
cfg) Vector Double
y Vector Double
fScores'
                usage' :: Map Text Int
usage' = (Text -> Map Text Int -> Map Text Int)
-> Map Text Int -> [Text] -> Map Text Int
forall a b. (a -> b -> b) -> b -> [a] -> b
forall (t :: * -> *) a b.
Foldable t =>
(a -> b -> b) -> b -> t a -> b
foldr (\Text
c -> (Int -> Int -> Int) -> Text -> Int -> Map Text Int -> Map Text Int
forall k a. Ord k => (a -> a -> a) -> k -> a -> Map k a -> Map k a
M.insertWith Int -> Int -> Int
forall a. Num a => a -> a -> a
(+) Text
c Int
1) Map Text Int
usageAcc (Tree Double -> [Text]
treeColumns Tree Double
tree)
             in Int
-> Vector Double
-> [Tree Double]
-> [Double]
-> Map Text Int
-> ([Tree Double], [Double], Map Text Int)
boost (Int
m Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Vector Double
fScores' (Tree Double
tree Tree Double -> [Tree Double] -> [Tree Double]
forall a. a -> [a] -> [a]
: [Tree Double]
ts) (Double
score Double -> [Double] -> [Double]
forall a. a -> [a] -> [a]
: [Double]
ss) Map Text Int
usage'
fitGBM GBConfig
_ Expr Double
expr DataFrame
_ =
    DataFrameException -> GBModel
forall a e. Exception e => e -> a
throw (Text -> DataFrameException
NonColumnReferenceException (Text
"fitGBM: " Text -> Text -> Text
forall a. Semigroup a => a -> a -> a
<> String -> Text
T.pack (Expr Double -> String
forall a. Show a => a -> String
show Expr Double
expr)))

newtonStep ::
    GBLoss ->
    VU.Vector Double ->
    VU.Vector Double ->
    (VU.Vector Double, Maybe (VU.Vector Double))
newtonStep :: GBLoss
-> Vector Double
-> Vector Double
-> (Vector Double, Maybe (Vector Double))
newtonStep GBLoss
SquaredError Vector Double
y Vector Double
f = (GBLoss -> Vector Double -> Vector Double -> Vector Double
negGradient GBLoss
SquaredError Vector Double
y Vector Double
f, Maybe (Vector Double)
forall a. Maybe a
Nothing)
newtonStep GBLoss
LogisticDeviance Vector Double
y Vector Double
f = (Vector Double
z, Vector Double -> Maybe (Vector Double)
forall a. a -> Maybe a
Just Vector Double
h)
  where
    p :: Vector Double
p = (Double -> Double) -> Vector Double -> Vector Double
forall a b. (Unbox a, Unbox b) => (a -> b) -> Vector a -> Vector b
VU.map Double -> Double
sigmoid Vector Double
f
    h :: Vector Double
h = (Double -> Double) -> Vector Double -> Vector Double
forall a b. (Unbox a, Unbox b) => (a -> b) -> Vector a -> Vector b
VU.map (\Double
pi' -> Double -> Double -> Double
forall a. Ord a => a -> a -> a
max Double
hFloor (Double
pi' Double -> Double -> Double
forall a. Num a => a -> a -> a
* (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
pi'))) Vector Double
p
    z :: Vector Double
z = (Double -> Double -> Double -> Double)
-> Vector Double -> Vector Double -> Vector Double -> Vector Double
forall a b c d.
(Unbox a, Unbox b, Unbox c, Unbox d) =>
(a -> b -> c -> d) -> Vector a -> Vector b -> Vector c -> Vector d
VU.zipWith3 (\Double
yi Double
pi' Double
hi -> (Double
yi Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
pi') Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
hi) Vector Double
y Vector Double
p Vector Double
h

-- | Floor on the Hessian, so a saturated row cannot produce an unbounded step.
hFloor :: Double
hFloor :: Double
hFloor = Double
1e-6

negGradient ::
    GBLoss -> VU.Vector Double -> VU.Vector Double -> VU.Vector Double
negGradient :: GBLoss -> Vector Double -> Vector Double -> Vector Double
negGradient GBLoss
SquaredError Vector Double
y Vector Double
f = (Double -> Double -> Double)
-> Vector Double -> Vector Double -> Vector Double
forall a b c.
(Unbox a, Unbox b, Unbox c) =>
(a -> b -> c) -> Vector a -> Vector b -> Vector c
VU.zipWith (-) Vector Double
y Vector Double
f
negGradient GBLoss
LogisticDeviance Vector Double
y Vector Double
f =
    (Double -> Double -> Double)
-> Vector Double -> Vector Double -> Vector Double
forall a b c.
(Unbox a, Unbox b, Unbox c) =>
(a -> b -> c) -> Vector a -> Vector b -> Vector c
VU.zipWith (\Double
yi Double
fi -> Double
yi Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double -> Double
sigmoid Double
fi) Vector Double
y Vector Double
f

lossValue :: GBLoss -> VU.Vector Double -> VU.Vector Double -> Double
lossValue :: GBLoss -> Vector Double -> Vector Double -> Double
lossValue GBLoss
SquaredError Vector Double
y Vector Double
f =
    Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum ((Double -> Double -> Double)
-> Vector Double -> Vector Double -> Vector Double
forall a b c.
(Unbox a, Unbox b, Unbox c) =>
(a -> b -> c) -> Vector a -> Vector b -> Vector c
VU.zipWith (\Double
yi Double
fi -> (Double
yi Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
fi) Double -> Int -> Double
forall a b. (Num a, Integral b) => a -> b -> a
^ (Int
2 :: Int)) Vector Double
y Vector Double
f)
        Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
1 (Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Double
y))
lossValue GBLoss
LogisticDeviance Vector Double
y Vector Double
f =
    Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum
        ( (Double -> Double -> Double)
-> Vector Double -> Vector Double -> Vector Double
forall a b c.
(Unbox a, Unbox b, Unbox c) =>
(a -> b -> c) -> Vector a -> Vector b -> Vector c
VU.zipWith
            ( \Double
yi Double
fi -> let p :: Double
p = Double -> Double
clamp01 (Double -> Double
sigmoid Double
fi) in Double -> Double
forall a. Num a => a -> a
negate (Double
yi Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double -> Double
forall a. Floating a => a -> a
log Double
p Double -> Double -> Double
forall a. Num a => a -> a -> a
+ (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
yi) Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double -> Double
forall a. Floating a => a -> a
log (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
p))
            )
            Vector Double
y
            Vector Double
f
        )
        Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
1 (Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Double
y))

sigmoid :: Double -> Double
sigmoid :: Double -> Double
sigmoid Double
z
    | Double
z Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
>= Double
0 = Double
1 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double -> Double
forall a. Floating a => a -> a
exp (-Double
z))
    | Bool
otherwise = let e :: Double
e = Double -> Double
forall a. Floating a => a -> a
exp Double
z in Double
e Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
e)

clamp01 :: Double -> Double
clamp01 :: Double -> Double
clamp01 Double
p = Double -> Double -> Double
forall a. Ord a => a -> a -> a
max Double
1e-12 (Double -> Double -> Double
forall a. Ord a => a -> a -> a
min (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
1e-12) Double
p)

predictTree :: DataFrame -> Tree Double -> VU.Vector Double
predictTree :: DataFrame -> Tree Double -> Vector Double
predictTree DataFrame
df Tree Double
t = case forall a.
Columnable a =>
DataFrame -> Expr a -> Either DataFrameException (TypedColumn a)
interpret @Double DataFrame
df (Tree Double -> Expr Double
forall a. Columnable a => Tree a -> Expr a
treeToExpr Tree Double
t) of
    Right (TColumn Column
c) -> Vector Double
-> Either DataFrameException (Vector Double) -> Vector Double
forall b a. b -> Either a b -> b
fromRight Vector Double
forall a. Unbox a => Vector a
VU.empty (forall a (v :: * -> *).
(Vector v a, Columnable a) =>
Column -> Either DataFrameException (v a)
toVector @Double @VU.Vector Column
c)
    Left DataFrameException
e -> DataFrameException -> Vector Double
forall a e. Exception e => e -> a
throw DataFrameException
e

treeColumns :: Tree Double -> [T.Text]
treeColumns :: Tree Double -> [Text]
treeColumns = Expr Double -> [Text]
forall a. Expr a -> [Text]
getColumns (Expr Double -> [Text])
-> (Tree Double -> Expr Double) -> Tree Double -> [Text]
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Tree Double -> Expr Double
forall a. Columnable a => Tree a -> Expr a
treeToExpr

-- | The full additive prediction expression: @f0 + lr · Σ treeᵢ@.
gbExpr :: GBModel -> Expr Double
gbExpr :: GBModel -> Expr Double
gbExpr GBModel
m = Int -> GBModel -> Expr Double
stageExpr (Vector (Tree Double) -> Int
forall a. Vector a -> Int
V.length (GBModel -> Vector (Tree Double)
gbTrees GBModel
m)) GBModel
m

-- | The prediction expression using only the first @k@ trees (staged predict).
gbExprAtStage :: Int -> GBModel -> Maybe (Expr Double)
gbExprAtStage :: Int -> GBModel -> Maybe (Expr Double)
gbExprAtStage Int
k GBModel
m
    | Int
k Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
0 Bool -> Bool -> Bool
|| Int
k Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Vector (Tree Double) -> Int
forall a. Vector a -> Int
V.length (GBModel -> Vector (Tree Double)
gbTrees GBModel
m) = Maybe (Expr Double)
forall a. Maybe a
Nothing
    | Bool
otherwise = Expr Double -> Maybe (Expr Double)
forall a. a -> Maybe a
Just (Int -> GBModel -> Expr Double
stageExpr Int
k GBModel
m)

stageExpr :: Int -> GBModel -> Expr Double
stageExpr :: Int -> GBModel -> Expr Double
stageExpr Int
k GBModel
m =
    (Tree Double -> Expr Double -> Expr Double)
-> Expr Double -> [Tree Double] -> Expr Double
forall a b. (a -> b -> b) -> b -> [a] -> b
forall (t :: * -> *) a b.
Foldable t =>
(a -> b -> b) -> b -> t a -> b
foldr (Expr Double -> Expr Double -> Expr Double
forall a. (Columnable a, Num a) => Expr a -> Expr a -> Expr a
(.+.) (Expr Double -> Expr Double -> Expr Double)
-> (Tree Double -> Expr Double)
-> Tree Double
-> Expr Double
-> Expr Double
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Tree Double -> Expr Double
scaled) (Double -> Expr Double
forall a. Columnable a => a -> Expr a
F.lit (GBModel -> Double
gbInit GBModel
m)) (Int -> [Tree Double] -> [Tree Double]
forall a. Int -> [a] -> [a]
take Int
k (Vector (Tree Double) -> [Tree Double]
forall a. Vector a -> [a]
V.toList (GBModel -> Vector (Tree Double)
gbTrees GBModel
m)))
  where
    scaled :: Tree Double -> Expr Double
scaled Tree Double
t = Double -> Expr Double
forall a. Columnable a => a -> Expr a
F.lit (GBModel -> Double
gbRate GBModel
m) Expr Double -> Expr Double -> Expr Double
forall a. (Columnable a, Num a) => Expr a -> Expr a -> Expr a
.*. Tree Double -> Expr Double
forall a. Columnable a => Tree a -> Expr a
treeToExpr Tree Double
t

-- | Probability expression for classification: @sigmoid(score)@.
gbProbaExpr :: GBModel -> Expr Double
gbProbaExpr :: GBModel -> Expr Double
gbProbaExpr GBModel
m = Double -> Expr Double
forall a. Columnable a => a -> Expr a
F.lit Double
1 Expr Double -> Expr Double -> Expr Double
forall a. Fractional a => a -> a -> a
/ (Double -> Expr Double
forall a. Columnable a => a -> Expr a
F.lit Double
1 Expr Double -> Expr Double -> Expr Double
forall a. Num a => a -> a -> a
+ Expr Double -> Expr Double
forall a. Floating a => a -> a
exp (Expr Double -> Expr Double
forall a. Num a => a -> a
negate (GBModel -> Expr Double
gbExpr GBModel
m)))

-- | Decision expression for classification: positive class when score > 0.
gbDecisionExpr :: GBModel -> Expr Bool
gbDecisionExpr :: GBModel -> Expr Bool
gbDecisionExpr GBModel
m = GBModel -> Expr Double
gbExpr GBModel
m Expr Double -> Expr Double -> Expr Bool
forall a. (Columnable a, Ord a) => Expr a -> Expr a -> Expr Bool
.>. Double -> Expr Double
forall a. Columnable a => a -> Expr a
F.lit Double
0