{- | Cross-validation and grid search for hyperparameter tuning. The model
fitters have heterogeneous types, so these helpers are parameterized by a
user-supplied @train -> test -> score@ closure; the search maximizes the mean
cross-validated score (use a negated error metric to minimize). Splitting reuses
the deterministic 'kFolds' from @dataframe-operations@.
-}
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)

{- | Per-fold scores from k-fold cross-validation. @scoreFn train test@ fits on
the training rows and returns a score on the held-out fold.
-}
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
(<>)

{- | scikit-learn @cross_val_score@: fit a model on each training fold and score
its prediction expression against a truth column on the held-out fold.

@fitPredict train@ fits on the training frame and returns the prediction
expression; @truth@ is the target column. Returns the per-fold metric values.

> crossValidate 5 0 rmse (F.col @Double "target")
>   (\tr -> predict (fit defaultLinearConfig (F.col @Double "target") tr)) df
-}
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

-- | The outcome of a grid search: the best config, its score, and all results.
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)

{- | Search configurations by mean cross-validated score, returning the
maximizer. @scoreFn cfg train test@ fits @cfg@ on @train@ and scores on @test@.
-}
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)