module DataFrame.ModelSelection (
crossValScore,
crossValidate,
GridSearchResult (..),
gridSearch,
) where
import Data.List (maximumBy)
import Data.Ord (comparing)
import System.Random (mkStdGen)
import DataFrame.Internal.DataFrame (DataFrame)
import DataFrame.Internal.Expression (Expr)
import DataFrame.Metrics (Metric, evaluate)
import DataFrame.Operations.Merge ()
import DataFrame.Operations.Subset (kFolds)
crossValScore ::
Int -> Int -> (DataFrame -> DataFrame -> Double) -> DataFrame -> [Double]
crossValScore :: Int
-> Int
-> (DataFrame -> DataFrame -> Double)
-> DataFrame
-> [Double]
crossValScore Int
folds Int
seed DataFrame -> DataFrame -> Double
scoreFn DataFrame
df =
[ DataFrame -> DataFrame -> Double
scoreFn ([DataFrame] -> DataFrame
combine (Int -> [DataFrame]
forall {a}. (Num a, Enum a, Eq a) => a -> [DataFrame]
others Int
i)) ([DataFrame]
fs [DataFrame] -> Int -> DataFrame
forall a. HasCallStack => [a] -> Int -> a
!! Int
i)
| Int
i <- [Int
0 .. [DataFrame] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [DataFrame]
fs Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]
, Bool -> Bool
not ([DataFrame] -> Bool
forall a. [a] -> Bool
forall (t :: * -> *) a. Foldable t => t a -> Bool
null (Int -> [DataFrame]
forall {a}. (Num a, Enum a, Eq a) => a -> [DataFrame]
others Int
i))
]
where
fs :: [DataFrame]
fs = StdGen -> Int -> DataFrame -> [DataFrame]
forall g. RandomGen g => g -> Int -> DataFrame -> [DataFrame]
kFolds (Int -> StdGen
mkStdGen Int
seed) Int
folds DataFrame
df
others :: a -> [DataFrame]
others a
i = [DataFrame
f | (a
j, DataFrame
f) <- [a] -> [DataFrame] -> [(a, DataFrame)]
forall a b. [a] -> [b] -> [(a, b)]
zip [a
0 ..] [DataFrame]
fs, a
j a -> a -> Bool
forall a. Eq a => a -> a -> Bool
/= a
i]
combine :: [DataFrame] -> DataFrame
combine = (DataFrame -> DataFrame -> DataFrame) -> [DataFrame] -> DataFrame
forall a. (a -> a -> a) -> [a] -> a
forall (t :: * -> *) a. Foldable t => (a -> a -> a) -> t a -> a
foldr1 DataFrame -> DataFrame -> DataFrame
forall a. Semigroup a => a -> a -> a
(<>)
crossValidate ::
Int ->
Int ->
Metric ->
Expr Double ->
(DataFrame -> Expr Double) ->
DataFrame ->
[Double]
crossValidate :: Int
-> Int
-> Metric
-> Expr Double
-> (DataFrame -> Expr Double)
-> DataFrame
-> [Double]
crossValidate Int
folds Int
seed Metric
metric Expr Double
truth DataFrame -> Expr Double
fitPredict =
Int
-> Int
-> (DataFrame -> DataFrame -> Double)
-> DataFrame
-> [Double]
crossValScore Int
folds Int
seed DataFrame -> DataFrame -> Double
score
where
score :: DataFrame -> DataFrame -> Double
score DataFrame
train = Metric -> Expr Double -> Expr Double -> DataFrame -> Double
evaluate Metric
metric (DataFrame -> Expr Double
fitPredict DataFrame
train) Expr Double
truth
data GridSearchResult c = GridSearchResult
{ forall c. GridSearchResult c -> c
gsBest :: !c
, forall c. GridSearchResult c -> Double
gsBestScore :: !Double
, forall c. GridSearchResult c -> [(c, Double)]
gsAll :: ![(c, Double)]
}
deriving (Int -> GridSearchResult c -> ShowS
[GridSearchResult c] -> ShowS
GridSearchResult c -> String
(Int -> GridSearchResult c -> ShowS)
-> (GridSearchResult c -> String)
-> ([GridSearchResult c] -> ShowS)
-> Show (GridSearchResult c)
forall c. Show c => Int -> GridSearchResult c -> ShowS
forall c. Show c => [GridSearchResult c] -> ShowS
forall c. Show c => GridSearchResult c -> String
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: forall c. Show c => Int -> GridSearchResult c -> ShowS
showsPrec :: Int -> GridSearchResult c -> ShowS
$cshow :: forall c. Show c => GridSearchResult c -> String
show :: GridSearchResult c -> String
$cshowList :: forall c. Show c => [GridSearchResult c] -> ShowS
showList :: [GridSearchResult c] -> ShowS
Show)
gridSearch ::
Int ->
Int ->
[c] ->
(c -> DataFrame -> DataFrame -> Double) ->
DataFrame ->
GridSearchResult c
gridSearch :: forall c.
Int
-> Int
-> [c]
-> (c -> DataFrame -> DataFrame -> Double)
-> DataFrame
-> GridSearchResult c
gridSearch Int
folds Int
seed [c]
configs c -> DataFrame -> DataFrame -> Double
scoreFn DataFrame
df =
c -> Double -> [(c, Double)] -> GridSearchResult c
forall c. c -> Double -> [(c, Double)] -> GridSearchResult c
GridSearchResult c
bestC Double
bestS [(c, Double)]
scored
where
scored :: [(c, Double)]
scored = [(c
c, [Double] -> Double
forall {t :: * -> *} {a}. (Foldable t, Fractional a) => t a -> a
mean (Int
-> Int
-> (DataFrame -> DataFrame -> Double)
-> DataFrame
-> [Double]
crossValScore Int
folds Int
seed (c -> DataFrame -> DataFrame -> Double
scoreFn c
c) DataFrame
df)) | c
c <- [c]
configs]
(c
bestC, Double
bestS) = ((c, Double) -> (c, Double) -> Ordering)
-> [(c, Double)] -> (c, Double)
forall (t :: * -> *) a.
Foldable t =>
(a -> a -> Ordering) -> t a -> a
maximumBy (((c, Double) -> Double) -> (c, Double) -> (c, Double) -> Ordering
forall a b. Ord a => (b -> a) -> b -> b -> Ordering
comparing (c, Double) -> Double
forall a b. (a, b) -> b
snd) [(c, Double)]
scored
mean :: t a -> a
mean t a
xs = if t a -> Bool
forall a. t a -> Bool
forall (t :: * -> *) a. Foldable t => t a -> Bool
null t a
xs then -(a
1 a -> a -> a
forall a. Fractional a => a -> a -> a
/ a
0) else t a -> a
forall a. Num a => t a -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum t a
xs a -> a -> a
forall a. Fractional a => a -> a -> a
/ Int -> a
forall a b. (Integral a, Num b) => a -> b
fromIntegral (t a -> Int
forall a. t a -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length t a
xs)