{-# LANGUAGE FlexibleContexts #-}

{- | Typed sampling and splitting. All operations are schema-preserving: they
change which rows are present, never the columns, so every result reuses the
input schema @cols@.
-}
module DataFrame.Typed.Sampling (
    randomSplit,
    kFolds,
    selectRows,
    stratifiedSample,
    stratifiedSplit,
) where

import System.Random (RandomGen)

import DataFrame.Internal.Column (Columnable)
import DataFrame.Operations.Subset (SplittableGen)
import qualified DataFrame.Operations.Subset as D
import DataFrame.Typed.Types (TExpr (..), TypedDataFrame (..))

-- | Split rows into two DataFrames by a fraction.
randomSplit ::
    (RandomGen g) =>
    g -> Double -> TypedDataFrame cols -> (TypedDataFrame cols, TypedDataFrame cols)
randomSplit :: forall g (cols :: [(Symbol, *)]).
RandomGen g =>
g
-> Double
-> TypedDataFrame cols
-> (TypedDataFrame cols, TypedDataFrame cols)
randomSplit g
g Double
p (TDF DataFrame
df) = let (DataFrame
a, DataFrame
b) = g -> Double -> DataFrame -> (DataFrame, DataFrame)
forall g.
RandomGen g =>
g -> Double -> DataFrame -> (DataFrame, DataFrame)
D.randomSplit g
g Double
p DataFrame
df in (DataFrame -> TypedDataFrame cols
forall (cols :: [(Symbol, *)]). DataFrame -> TypedDataFrame cols
TDF DataFrame
a, DataFrame -> TypedDataFrame cols
forall (cols :: [(Symbol, *)]). DataFrame -> TypedDataFrame cols
TDF DataFrame
b)

-- | Partition rows into @k@ folds.
kFolds ::
    (RandomGen g) => g -> Int -> TypedDataFrame cols -> [TypedDataFrame cols]
kFolds :: forall g (cols :: [(Symbol, *)]).
RandomGen g =>
g -> Int -> TypedDataFrame cols -> [TypedDataFrame cols]
kFolds g
g Int
k (TDF DataFrame
df) = (DataFrame -> TypedDataFrame cols)
-> [DataFrame] -> [TypedDataFrame cols]
forall a b. (a -> b) -> [a] -> [b]
map DataFrame -> TypedDataFrame cols
forall (cols :: [(Symbol, *)]). DataFrame -> TypedDataFrame cols
TDF (g -> Int -> DataFrame -> [DataFrame]
forall g. RandomGen g => g -> Int -> DataFrame -> [DataFrame]
D.kFolds g
g Int
k DataFrame
df)

{- | Select rows by index.
| This may fail if the indices are out of bounds;
| use with caution or use 'filter' to select rows by a predicate instead.
-}
selectRows :: [Int] -> TypedDataFrame cols -> TypedDataFrame cols
selectRows :: forall (cols :: [(Symbol, *)]).
[Int] -> TypedDataFrame cols -> TypedDataFrame cols
selectRows [Int]
ixs (TDF DataFrame
df) = DataFrame -> TypedDataFrame cols
forall (cols :: [(Symbol, *)]). DataFrame -> TypedDataFrame cols
TDF ([Int] -> DataFrame -> DataFrame
D.selectRows [Int]
ixs DataFrame
df)

-- | Sample a fraction of rows, preserving the distribution of a strata column.
stratifiedSample ::
    (SplittableGen g, Columnable a) =>
    g -> Double -> TExpr cols a -> TypedDataFrame cols -> TypedDataFrame cols
stratifiedSample :: forall g a (cols :: [(Symbol, *)]).
(SplittableGen g, Columnable a) =>
g
-> Double
-> TExpr cols a
-> TypedDataFrame cols
-> TypedDataFrame cols
stratifiedSample g
g Double
p (TExpr Expr a
e) (TDF DataFrame
df) = DataFrame -> TypedDataFrame cols
forall (cols :: [(Symbol, *)]). DataFrame -> TypedDataFrame cols
TDF (g -> Double -> Expr a -> DataFrame -> DataFrame
forall a g.
(SplittableGen g, Columnable a) =>
g -> Double -> Expr a -> DataFrame -> DataFrame
D.stratifiedSample g
g Double
p Expr a
e DataFrame
df)

-- | Split rows by a fraction, preserving the distribution of a strata column.
stratifiedSplit ::
    (SplittableGen g, Columnable a) =>
    g ->
    Double ->
    TExpr cols a ->
    TypedDataFrame cols ->
    (TypedDataFrame cols, TypedDataFrame cols)
stratifiedSplit :: forall g a (cols :: [(Symbol, *)]).
(SplittableGen g, Columnable a) =>
g
-> Double
-> TExpr cols a
-> TypedDataFrame cols
-> (TypedDataFrame cols, TypedDataFrame cols)
stratifiedSplit g
g Double
p (TExpr Expr a
e) (TDF DataFrame
df) =
    let (DataFrame
a, DataFrame
b) = g -> Double -> Expr a -> DataFrame -> (DataFrame, DataFrame)
forall a g.
(SplittableGen g, Columnable a) =>
g -> Double -> Expr a -> DataFrame -> (DataFrame, DataFrame)
D.stratifiedSplit g
g Double
p Expr a
e DataFrame
df in (DataFrame -> TypedDataFrame cols
forall (cols :: [(Symbol, *)]). DataFrame -> TypedDataFrame cols
TDF DataFrame
a, DataFrame -> TypedDataFrame cols
forall (cols :: [(Symbol, *)]). DataFrame -> TypedDataFrame cols
TDF DataFrame
b)