{-# LANGUAGE BangPatterns #-}
{-# LANGUAGE FlexibleContexts #-}
{-# LANGUAGE FlexibleInstances #-}
{-# LANGUAGE MultiParamTypeClasses #-}
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications #-}
{-# LANGUAGE TypeFamilies #-}
{-# LANGUAGE UndecidableInstances #-}
module DataFrame.Boosting.AdaBoost (
module DataFrame.Model,
AdaBoostConfig (..),
defaultAdaBoostConfig,
AdaBoostModel (..),
) where
import Control.Exception (throw)
import Data.List (sort)
import Data.Maybe (fromMaybe, maybeToList)
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 (
CartFeature (..),
cartFeatures,
sortIndicesByValue,
)
import DataFrame.DecisionTree.Fit (treeToExpr)
import DataFrame.DecisionTree.Types (Tree (..))
import DataFrame.Featurize.Internal (argMaxExpr, targetValues)
import qualified DataFrame.Functions as F
import DataFrame.Internal.Column (Columnable, TypedColumn (..), toVector)
import DataFrame.Internal.DataFrame (DataFrame)
import DataFrame.Internal.Expression (Expr (..))
import DataFrame.Internal.Interpreter (interpret)
import DataFrame.Model
import DataFrame.Operators ((.*.), (.+.), (.==.))
data AdaBoostConfig = AdaBoostConfig
{ AdaBoostConfig -> Int
abNEstimators :: !Int
, AdaBoostConfig -> Int
abMaxDepth :: !Int
}
deriving (AdaBoostConfig -> AdaBoostConfig -> Bool
(AdaBoostConfig -> AdaBoostConfig -> Bool)
-> (AdaBoostConfig -> AdaBoostConfig -> Bool) -> Eq AdaBoostConfig
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: AdaBoostConfig -> AdaBoostConfig -> Bool
== :: AdaBoostConfig -> AdaBoostConfig -> Bool
$c/= :: AdaBoostConfig -> AdaBoostConfig -> Bool
/= :: AdaBoostConfig -> AdaBoostConfig -> Bool
Eq, Int -> AdaBoostConfig -> ShowS
[AdaBoostConfig] -> ShowS
AdaBoostConfig -> String
(Int -> AdaBoostConfig -> ShowS)
-> (AdaBoostConfig -> String)
-> ([AdaBoostConfig] -> ShowS)
-> Show AdaBoostConfig
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> AdaBoostConfig -> ShowS
showsPrec :: Int -> AdaBoostConfig -> ShowS
$cshow :: AdaBoostConfig -> String
show :: AdaBoostConfig -> String
$cshowList :: [AdaBoostConfig] -> ShowS
showList :: [AdaBoostConfig] -> ShowS
Show)
defaultAdaBoostConfig :: AdaBoostConfig
defaultAdaBoostConfig :: AdaBoostConfig
defaultAdaBoostConfig = AdaBoostConfig{abNEstimators :: Int
abNEstimators = Int
50, abMaxDepth :: Int
abMaxDepth = Int
1}
data AdaBoostModel a = AdaBoostModel
{ forall a. AdaBoostModel a -> Vector Double
abAlphas :: !(VU.Vector Double)
, forall a. AdaBoostModel a -> Vector (Tree a)
abStumps :: !(V.Vector (Tree a))
, forall a. AdaBoostModel a -> Vector a
abClasses :: !(V.Vector a)
}
deriving (Int -> AdaBoostModel a -> ShowS
[AdaBoostModel a] -> ShowS
AdaBoostModel a -> String
(Int -> AdaBoostModel a -> ShowS)
-> (AdaBoostModel a -> String)
-> ([AdaBoostModel a] -> ShowS)
-> Show (AdaBoostModel a)
forall a. Show a => Int -> AdaBoostModel a -> ShowS
forall a. Show a => [AdaBoostModel a] -> ShowS
forall a. Show a => AdaBoostModel a -> String
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: forall a. Show a => Int -> AdaBoostModel a -> ShowS
showsPrec :: Int -> AdaBoostModel a -> ShowS
$cshow :: forall a. Show a => AdaBoostModel a -> String
show :: AdaBoostModel a -> String
$cshowList :: forall a. Show a => [AdaBoostModel a] -> ShowS
showList :: [AdaBoostModel a] -> ShowS
Show)
instance (Columnable a, Ord a) => Fit AdaBoostConfig (Expr a) where
type ModelOf AdaBoostConfig (Expr a) = (AdaBoostModel a)
fit :: CheckFrame
(FrameReq AdaBoostConfig (Expr a)) (FrameFor (Expr a)) =>
AdaBoostConfig
-> Expr a
-> FrameFor (Expr a)
-> FitResult (FrameFor (Expr a)) (ModelOf AdaBoostConfig (Expr a))
fit = AdaBoostConfig -> Expr a -> DataFrame -> AdaBoostModel a
AdaBoostConfig
-> Expr a
-> FrameFor (Expr a)
-> FitResult (FrameFor (Expr a)) (ModelOf AdaBoostConfig (Expr a))
forall a.
(Columnable a, Ord a) =>
AdaBoostConfig -> Expr a -> DataFrame -> AdaBoostModel a
fitAdaBoost
instance (Columnable a, Ord a) => Predict (AdaBoostModel a) where
type Prediction (AdaBoostModel a) = Expr a
predict :: AdaBoostModel a -> Prediction (AdaBoostModel a)
predict = AdaBoostModel a -> Expr a
AdaBoostModel a -> Prediction (AdaBoostModel a)
forall a. (Columnable a, Ord a) => AdaBoostModel a -> Expr a
adaBoostExpr
fitAdaBoost ::
(Columnable a, Ord a) =>
AdaBoostConfig -> Expr a -> DataFrame -> AdaBoostModel a
fitAdaBoost :: forall a.
(Columnable a, Ord a) =>
AdaBoostConfig -> Expr a -> DataFrame -> AdaBoostModel a
fitAdaBoost AdaBoostConfig
cfg target :: Expr a
target@(Col Text
name) DataFrame
df =
Vector Double -> Vector (Tree a) -> Vector a -> AdaBoostModel a
forall a.
Vector Double -> Vector (Tree a) -> Vector a -> AdaBoostModel a
AdaBoostModel
([Double] -> Vector Double
forall a. Unbox a => [a] -> Vector a
VU.fromList ([Double] -> [Double]
forall a. [a] -> [a]
reverse [Double]
alphas))
([Tree a] -> Vector (Tree a)
forall a. [a] -> Vector a
V.fromList ([Tree a] -> [Tree a]
forall a. [a] -> [a]
reverse [Tree a]
stumps))
Vector a
classesV
where
feats :: Vector CartFeature
feats = [CartFeature] -> Vector CartFeature
forall a. [a] -> Vector a
V.fromList (Text -> DataFrame -> [CartFeature]
cartFeatures Text
name DataFrame
df)
ys :: Vector a
ys = Expr a -> DataFrame -> Vector a
forall a. Columnable a => Expr a -> DataFrame -> Vector a
targetValues Expr a
target DataFrame
df
n :: Int
n = Vector a -> Int
forall a. Vector a -> Int
V.length Vector a
ys
classes :: [a]
classes = [a] -> [a]
forall a. Ord a => [a] -> [a]
sort ((a -> [a] -> [a]) -> [a] -> [a] -> [a]
forall a b. (a -> b -> b) -> b -> [a] -> b
forall (t :: * -> *) a b.
Foldable t =>
(a -> b -> b) -> b -> t a -> b
foldr a -> [a] -> [a]
forall {a}. Eq a => a -> [a] -> [a]
dedup [] (Vector a -> [a]
forall a. Vector a -> [a]
V.toList Vector a
ys))
dedup :: a -> [a] -> [a]
dedup a
x [a]
acc = if a
x a -> [a] -> Bool
forall a. Eq a => a -> [a] -> Bool
forall (t :: * -> *) a. (Foldable t, Eq a) => a -> t a -> Bool
`elem` [a]
acc then [a]
acc else a
x a -> [a] -> [a]
forall a. a -> [a] -> [a]
: [a]
acc
classesV :: Vector a
classesV = [a] -> Vector a
forall a. [a] -> Vector a
V.fromList [a]
classes
kClasses :: Int
kClasses = [a] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [a]
classes
codes :: Vector Int
codes = Int -> (Int -> Int) -> Vector Int
forall a. Unbox a => Int -> (Int -> a) -> Vector a
VU.generate Int
n (\Int
i -> a -> Int
classIndex (Vector a
ys Vector a -> Int -> a
forall a. Vector a -> Int -> a
V.! Int
i))
classIndex :: a -> Int
classIndex a
v = [a] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length ((a -> Bool) -> [a] -> [a]
forall a. (a -> Bool) -> [a] -> [a]
takeWhile (a -> a -> Bool
forall a. Ord a => a -> a -> Bool
< a
v) [a]
classes)
([Double]
alphas, [Tree a]
stumps) = Int
-> Vector Double -> [Double] -> [Tree a] -> ([Double], [Tree a])
boost Int
0 (Int -> Double -> Vector Double
forall a. Unbox a => Int -> a -> Vector a
VU.replicate Int
n (Double
1 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))) [] []
boost :: Int
-> Vector Double -> [Double] -> [Tree a] -> ([Double], [Tree a])
boost !Int
m Vector Double
w [Double]
as [Tree a]
ts
| Int
m Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= AdaBoostConfig -> Int
abNEstimators AdaBoostConfig
cfg = ([Double]
as, [Tree a]
ts)
| Bool
otherwise =
let stump :: Tree a
stump = Int
-> Vector CartFeature
-> Vector a
-> Vector Int
-> Int
-> Vector Double
-> Tree a
forall a.
Columnable a =>
Int
-> Vector CartFeature
-> Vector a
-> Vector Int
-> Int
-> Vector Double
-> Tree a
fitWeightedTree (AdaBoostConfig -> Int
abMaxDepth AdaBoostConfig
cfg) Vector CartFeature
feats Vector a
classesV Vector Int
codes Int
kClasses Vector Double
w
pred :: Vector Int
pred = DataFrame -> Vector a -> Tree a -> Vector Int
forall a.
(Columnable a, Ord a) =>
DataFrame -> Vector a -> Tree a -> Vector Int
predictCodes DataFrame
df Vector a
classesV Tree a
stump
wrong :: VU.Vector Int
wrong :: Vector Int
wrong = Int -> (Int -> Int) -> Vector Int
forall a. Unbox a => Int -> (Int -> a) -> Vector a
VU.generate Int
n (\Int
i -> if Vector Int
pred Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
/= Vector Int
codes Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i then Int
1 else Int
0)
err :: Double
err = Double -> Double
forall {a}. (Ord a, Fractional a) => a -> a
clamp (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 -> Double -> Double
forall a. Num a => a -> a -> a
(*) Vector Double
w ((Int -> Double) -> Vector Int -> Vector Double
forall a b. (Unbox a, Unbox b) => (a -> b) -> Vector a -> Vector b
VU.map Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Vector Int
wrong)) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum Vector Double
w)
alpha :: Double
alpha = Double -> Double
forall a. Floating a => a -> a
log ((Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
err) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
err) Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double -> Double
forall a. Floating a => a -> a
log (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
kClasses Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1)))
w' :: Vector Double
w' = Vector Double -> Vector Double
forall {a}. (Unbox a, Eq a, Fractional a) => Vector a -> Vector a
normalize ((Double -> Int -> Double)
-> Vector Double -> Vector Int -> Vector Double
forall a b c.
(Unbox a, Unbox b, Unbox c) =>
(a -> b -> c) -> Vector a -> Vector b -> Vector c
VU.zipWith (\Double
wi Int
e -> Double
wi Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double -> Double
forall a. Floating a => a -> a
exp (Double
alpha Double -> Double -> Double
forall a. Num a => a -> a -> a
* Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
e)) Vector Double
w Vector Int
wrong)
in if Double
err Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
<= Double
0 Bool -> Bool -> Bool
|| Double
err Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
>= Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
1 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
kClasses
then (Double
alpha Double -> [Double] -> [Double]
forall a. a -> [a] -> [a]
: [Double]
as, Tree a
stump Tree a -> [Tree a] -> [Tree a]
forall a. a -> [a] -> [a]
: [Tree a]
ts)
else Int
-> Vector Double -> [Double] -> [Tree a] -> ([Double], [Tree a])
boost (Int
m Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Vector Double
w' (Double
alpha Double -> [Double] -> [Double]
forall a. a -> [a] -> [a]
: [Double]
as) (Tree a
stump Tree a -> [Tree a] -> [Tree a]
forall a. a -> [a] -> [a]
: [Tree a]
ts)
clamp :: a -> a
clamp a
e = a -> a -> a
forall a. Ord a => a -> a -> a
max a
1e-10 (a -> a -> a
forall a. Ord a => a -> a -> a
min (a
1 a -> a -> a
forall a. Num a => a -> a -> a
- a
1e-10) a
e)
normalize :: Vector a -> Vector a
normalize Vector a
v = let s :: a
s = Vector a -> a
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum Vector a
v in if a
s a -> a -> Bool
forall a. Eq a => a -> a -> Bool
== a
0 then Vector a
v else (a -> a) -> Vector a -> Vector a
forall a b. (Unbox a, Unbox b) => (a -> b) -> Vector a -> Vector b
VU.map (a -> a -> a
forall a. Fractional a => a -> a -> a
/ a
s) Vector a
v
fitAdaBoost AdaBoostConfig
_ Expr a
expr DataFrame
_ =
DataFrameException -> AdaBoostModel a
forall a e. Exception e => e -> a
throw (Text -> DataFrameException
NonColumnReferenceException (Text
"fitAdaBoost: " Text -> Text -> Text
forall a. Semigroup a => a -> a -> a
<> String -> Text
T.pack (Expr a -> String
forall a. Show a => a -> String
show Expr a
expr)))
predictCodes ::
forall a.
(Columnable a, Ord a) =>
DataFrame -> V.Vector a -> Tree a -> VU.Vector Int
predictCodes :: forall a.
(Columnable a, Ord a) =>
DataFrame -> Vector a -> Tree a -> Vector Int
predictCodes DataFrame
df Vector a
classesV Tree a
stump =
[Int] -> Vector Int
forall a. Unbox a => [a] -> Vector a
VU.fromList ((a -> Int) -> [a] -> [Int]
forall a b. (a -> b) -> [a] -> [b]
map a -> Int
toCode [a]
preds)
where
preds :: [a]
preds :: [a]
preds = case DataFrame -> Expr a -> Either DataFrameException (TypedColumn a)
forall a.
Columnable a =>
DataFrame -> Expr a -> Either DataFrameException (TypedColumn a)
interpret DataFrame
df (Tree a -> Expr a
forall a. Columnable a => Tree a -> Expr a
treeToExpr Tree a
stump) of
Right (TColumn Column
c) -> (DataFrameException -> [a])
-> (Vector a -> [a]) -> Either DataFrameException (Vector a) -> [a]
forall a c b. (a -> c) -> (b -> c) -> Either a b -> c
either ([a] -> DataFrameException -> [a]
forall a b. a -> b -> a
const []) Vector a -> [a]
forall a. Vector a -> [a]
V.toList (forall a (v :: * -> *).
(Vector v a, Columnable a) =>
Column -> Either DataFrameException (v a)
toVector @a @V.Vector Column
c)
Left DataFrameException
e -> DataFrameException -> [a]
forall a e. Exception e => e -> a
throw DataFrameException
e
toCode :: a -> Int
toCode a
v = Int -> Maybe Int -> Int
forall a. a -> Maybe a -> a
fromMaybe Int
0 ((a -> Bool) -> Vector a -> Maybe Int
forall a. (a -> Bool) -> Vector a -> Maybe Int
V.findIndex (a -> a -> Bool
forall a. Eq a => a -> a -> Bool
== a
v) Vector a
classesV)
fitWeightedTree ::
(Columnable a) =>
Int ->
V.Vector CartFeature ->
V.Vector a ->
VU.Vector Int ->
Int ->
VU.Vector Double ->
Tree a
fitWeightedTree :: forall a.
Columnable a =>
Int
-> Vector CartFeature
-> Vector a
-> Vector Int
-> Int
-> Vector Double
-> Tree a
fitWeightedTree Int
maxDepth Vector CartFeature
feats Vector a
classesV Vector Int
codes Int
kClasses Vector Double
weights =
Int -> Vector Int -> Tree a
go Int
0 (Int -> Int -> Vector Int
forall a. (Unbox a, Num a) => a -> Int -> Vector a
VU.enumFromN Int
0 (Vector Int -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Int
codes))
where
go :: Int -> Vector Int -> Tree a
go Int
depth Vector Int
idxs
| Int
depth Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
maxDepth 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
< Int
2 Bool -> Bool -> Bool
|| Vector Int -> Bool
isPure Vector Int
idxs =
a -> Tree a
forall a. a -> Tree a
Leaf (Vector a
classesV Vector a -> Int -> a
forall a. Vector a -> Int -> a
V.! Vector Int -> Int
majority Vector Int
idxs)
| Bool
otherwise = case Vector Int -> Maybe (Int, Double)
bestSplit Vector Int
idxs of
Maybe (Int, Double)
Nothing -> a -> Tree a
forall a. a -> Tree a
Leaf (Vector a
classesV Vector a -> Int -> a
forall a. Vector a -> Int -> a
V.! Vector Int -> Int
majority Vector Int
idxs)
Just (Int
fj, Double
thr) ->
let vals :: Vector Double
vals = CartFeature -> Vector Double
cfValues (Vector CartFeature
feats Vector CartFeature -> Int -> CartFeature
forall a. Vector a -> Int -> a
V.! Int
fj)
(Vector Int
l, Vector Int
r) = (Int -> Bool) -> Vector Int -> (Vector Int, Vector Int)
forall a.
Unbox a =>
(a -> Bool) -> Vector a -> (Vector a, Vector a)
VU.partition (\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) Vector Int
idxs
in if Vector Int -> Bool
forall a. Unbox a => Vector a -> Bool
VU.null Vector Int
l Bool -> Bool -> Bool
|| Vector Int -> Bool
forall a. Unbox a => Vector a -> Bool
VU.null Vector Int
r
then a -> Tree a
forall a. a -> Tree a
Leaf (Vector a
classesV Vector a -> Int -> a
forall a. Vector a -> Int -> a
V.! Vector Int -> Int
majority Vector Int
idxs)
else
Expr Bool -> Tree a -> Tree a -> Tree a
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)
(Int -> Vector Int -> Tree a
go (Int
depth Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Vector Int
l)
(Int -> Vector Int -> Tree a
go (Int
depth Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Vector Int
r)
classWeights :: Vector Int -> Vector Double
classWeights Vector Int
idxs =
(Double -> Double -> Double)
-> Vector Double -> Vector (Int, Double) -> Vector Double
forall a b.
(Unbox a, Unbox b) =>
(a -> b -> a) -> Vector a -> Vector (Int, b) -> Vector a
VU.accumulate
Double -> Double -> Double
forall a. Num a => a -> a -> a
(+)
(Int -> Double -> Vector Double
forall a. Unbox a => Int -> a -> Vector a
VU.replicate Int
kClasses Double
0)
((Int -> (Int, Double)) -> Vector Int -> Vector (Int, Double)
forall a b. (Unbox a, Unbox b) => (a -> b) -> Vector a -> Vector b
VU.map (\Int
i -> (Vector Int
codes Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i, Vector Double
weights Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i)) Vector Int
idxs)
majority :: Vector Int -> Int
majority Vector Int
idxs = Vector Double -> Int
forall a. (Unbox a, Ord a) => Vector a -> Int
VU.maxIndex (Vector Int -> Vector Double
classWeights Vector Int
idxs)
isPure :: Vector Int -> Bool
isPure Vector Int
idxs = Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length ((Double -> Bool) -> Vector Double -> Vector Double
forall a. Unbox a => (a -> Bool) -> Vector a -> Vector a
VU.filter (Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
> Double
0) (Vector Int -> Vector Double
classWeights Vector Int
idxs)) Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
<= Int
1
bestSplit :: Vector Int -> Maybe (Int, Double)
bestSplit Vector Int
idxs =
let cands :: [(Double, Int, Double)]
cands =
[ (Double
score, 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
score) <- Vector Int -> Int -> [(Double, Double)]
featureSplits Vector Int
idxs Int
fj
]
in case [(Double, Int, Double)]
cands of
[] -> Maybe (Int, Double)
forall a. Maybe a
Nothing
[(Double, Int, Double)]
_ -> let (Double
_, Int
fj, Double
thr) = [(Double, Int, Double)] -> (Double, Int, Double)
forall a b c. Ord a => [(a, b, c)] -> (a, b, c)
minimum3 [(Double, Int, Double)]
cands in (Int, Double) -> Maybe (Int, Double)
forall a. a -> Maybe a
Just (Int
fj, Double
thr)
featureSplits :: Vector Int -> Int -> [(Double, Double)]
featureSplits Vector Int
idxs Int
fj =
let vals :: Vector Double
vals = CartFeature -> Vector Double
cfValues (Vector CartFeature
feats Vector CartFeature -> Int -> CartFeature
forall a. Vector a -> Int -> a
V.! Int
fj)
member :: Vector Bool
member =
Int -> Bool -> Vector Bool
forall a. Unbox a => Int -> a -> Vector a
VU.replicate (Vector Int -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Int
codes) Bool
False
Vector Bool -> [(Int, Bool)] -> Vector Bool
forall a. Unbox a => Vector a -> [(Int, a)] -> Vector a
VU.// [(Int
i, Bool
True) | Int
i <- Vector Int -> [Int]
forall a. Unbox a => Vector a -> [a]
VU.toList Vector Int
idxs]
sorted :: Vector Int
sorted = (Int -> Bool) -> Vector Int -> Vector Int
forall a. Unbox a => (a -> Bool) -> Vector a -> Vector a
VU.filter (Vector Bool
member Vector Bool -> Int -> Bool
forall a. Unbox a => Vector a -> Int -> a
VU.!) (Vector Double -> Vector Int
sortIndicesByValue Vector Double
vals)
in Vector Double -> Vector Int -> Vector Double -> [(Double, Double)]
sweep Vector Double
vals Vector Int
sorted (Vector Int -> Vector Double
classWeights Vector Int
idxs)
sweep :: Vector Double -> Vector Int -> Vector Double -> [(Double, Double)]
sweep Vector Double
vals Vector Int
sorted Vector Double
totW = Int
-> Vector Double -> Maybe (Double, Double) -> [(Double, Double)]
go0 Int
0 (Int -> Double -> Vector Double
forall a. Unbox a => Int -> a -> Vector a
VU.replicate Int
kClasses Double
0) Maybe (Double, Double)
forall a. Maybe a
Nothing
where
m :: Int
m = Vector Int -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Int
sorted
totWsum :: Double
totWsum = Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum Vector Double
totW
go0 :: Int
-> Vector Double -> Maybe (Double, Double) -> [(Double, Double)]
go0 !Int
k Vector Double
leftW 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) -> [(Double, Double)]
forall a. Maybe a -> [a]
maybeToList Maybe (Double, Double)
best
| Bool
otherwise =
let i :: Int
i = Vector Int
sorted Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
k
leftW' :: Vector Double
leftW' = Vector Double
leftW Vector Double -> [(Int, Double)] -> Vector Double
forall a. Unbox a => Vector a -> [(Int, a)] -> Vector a
VU.// [(Vector Int
codes Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i, Vector Double
leftW Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! (Vector Int
codes Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i) Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Vector Double
weights Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i)]
vCur :: Double
vCur = Vector Double
vals Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i
vNext :: Double
vNext = Vector Double
vals Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! (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))
wl :: Double
wl = Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum Vector Double
leftW'
wr :: Double
wr = Double
totWsum Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
wl
score :: Double
score = Double
wl Double -> Double -> Double
forall a. Num a => a -> a -> a
* Vector Double -> Double
gini Vector Double
leftW' Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
wr Double -> Double -> Double
forall a. Num a => a -> a -> a
* Vector Double -> Double
gini ((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
totW Vector Double
leftW')
valid :: Bool
valid = Double
vCur Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
/= Double
vNext 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
best' :: Maybe (Double, Double)
best' =
if Bool
valid Bool -> Bool -> Bool
&& Bool
-> ((Double, Double) -> Bool) -> Maybe (Double, Double) -> Bool
forall b a. b -> (a -> b) -> Maybe a -> b
maybe Bool
True (\(Double
_, Double
s) -> Double
score Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
< Double
s) Maybe (Double, Double)
best
then (Double, Double) -> Maybe (Double, Double)
forall a. a -> Maybe a
Just ((Double
vCur Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
vNext) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
2, Double
score)
else Maybe (Double, Double)
best
in Int
-> Vector Double -> Maybe (Double, Double) -> [(Double, Double)]
go0 (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Vector Double
leftW' Maybe (Double, Double)
best'
gini :: VU.Vector Double -> Double
gini :: Vector Double -> Double
gini Vector Double
cw =
let total :: Double
total = Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum Vector Double
cw
in if Double
total Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
== Double
0
then Double
0
else Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum ((Double -> Double) -> Vector Double -> Vector Double
forall a b. (Unbox a, Unbox b) => (a -> b) -> Vector a -> Vector b
VU.map (\Double
c -> (Double
c Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
total) Double -> Int -> Double
forall a b. (Num a, Integral b) => a -> b -> a
^ (Int
2 :: Int)) Vector Double
cw)
minimum3 :: (Ord a) => [(a, b, c)] -> (a, b, c)
minimum3 :: forall a b c. Ord a => [(a, b, c)] -> (a, b, c)
minimum3 = ((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)
adaBoostExpr :: (Columnable a, Ord a) => AdaBoostModel a -> Expr a
adaBoostExpr :: forall a. (Columnable a, Ord a) => AdaBoostModel a -> Expr a
adaBoostExpr AdaBoostModel a
m = [(a, Expr Double)] -> Expr a
forall a. Columnable a => [(a, Expr Double)] -> Expr a
argMaxExpr ([a] -> [Expr Double] -> [(a, Expr Double)]
forall a b. [a] -> [b] -> [(a, b)]
zip [a]
classes [Expr Double]
scores)
where
classes :: [a]
classes = Vector a -> [a]
forall a. Vector a -> [a]
V.toList (AdaBoostModel a -> Vector a
forall a. AdaBoostModel a -> Vector a
abClasses AdaBoostModel a
m)
stumpExprs :: [Expr a]
stumpExprs = (Tree a -> Expr a) -> [Tree a] -> [Expr a]
forall a b. (a -> b) -> [a] -> [b]
map Tree a -> Expr a
forall a. Columnable a => Tree a -> Expr a
treeToExpr (Vector (Tree a) -> [Tree a]
forall a. Vector a -> [a]
V.toList (AdaBoostModel a -> Vector (Tree a)
forall a. AdaBoostModel a -> Vector (Tree a)
abStumps AdaBoostModel a
m))
alphas :: [Double]
alphas = Vector Double -> [Double]
forall a. Unbox a => Vector a -> [a]
VU.toList (AdaBoostModel a -> Vector Double
forall a. AdaBoostModel a -> Vector Double
abAlphas AdaBoostModel a
m)
scores :: [Expr Double]
scores =
[ (Expr Double -> Expr Double -> Expr Double)
-> Expr Double -> [Expr 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
(.+.) (Double -> Expr Double
forall a. Columnable a => a -> Expr a
F.lit Double
0) ((Double -> Expr a -> Expr Double)
-> [Double] -> [Expr a] -> [Expr Double]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith (a -> Double -> Expr a -> Expr Double
forall {a} {a}.
(When (Unboxable a) (Unbox a), When (IntegralTypes a) (Integral a),
When (FloatingTypes a) (Real a, Fractional a),
When (Unboxable a) (Unbox a), When (IntegralTypes a) (Integral a),
When (FloatingTypes a) (Real a, Fractional a), Fractional a,
Typeable a, Typeable a, Show a, Show a, ColumnifyRep (KindOf a) a,
ColumnifyRep (KindOf a) a, SBoolI (Unboxable a),
SBoolI (Unboxable a), SBoolI (Numeric a), SBoolI (Numeric a),
SBoolI (IntegralTypes a), SBoolI (IntegralTypes a),
SBoolI (FloatingTypes a), SBoolI (FloatingTypes a), Eq a, Eq a) =>
a -> a -> Expr a -> Expr a
vote a
c) [Double]
alphas [Expr a]
stumpExprs)
| a
c <- [a]
classes
]
vote :: a -> a -> Expr a -> Expr a
vote a
c a
a Expr a
se = a -> Expr a
forall a. Columnable a => a -> Expr a
F.lit a
a Expr a -> Expr a -> Expr a
forall a. (Columnable a, Num a) => Expr a -> Expr a -> Expr a
.*. Expr Bool -> Expr a
forall {a}.
(When (Unboxable a) (Unbox a), When (IntegralTypes a) (Integral a),
When (FloatingTypes a) (Real a, Fractional a), Typeable a, Show a,
Eq a, ColumnifyRep (KindOf a) a, SBoolI (Unboxable a),
SBoolI (Numeric a), SBoolI (IntegralTypes a),
SBoolI (FloatingTypes a), Fractional a) =>
Expr Bool -> Expr a
indicator (Expr a
se Expr a -> Expr a -> Expr Bool
forall a. (Columnable a, Eq a) => Expr a -> Expr a -> Expr Bool
.==. a -> Expr a
forall a. Columnable a => a -> Expr a
F.lit a
c)
indicator :: Expr Bool -> Expr a
indicator Expr Bool
cond = Expr Bool -> Expr a -> Expr a -> Expr a
forall a. Columnable a => Expr Bool -> Expr a -> Expr a -> Expr a
F.ifThenElse Expr Bool
cond (a -> Expr a
forall a. Columnable a => a -> Expr a
F.lit a
1.0) (a -> Expr a
forall a. Columnable a => a -> Expr a
F.lit a
0.0)