{-# LANGUAGE AllowAmbiguousTypes #-}
{-# LANGUAGE DataKinds #-}
{-# LANGUAGE FlexibleContexts #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications #-}
{-# LANGUAGE TypeFamilies #-}
{-# LANGUAGE TypeOperators #-}

{- | Typed statistical reducers over a 'TypedDataFrame'.

These mirror the untyped reducers in "DataFrame.Operations.Statistics", taking
a schema-checked 'TExpr' instead of a raw @Expr@. The names (@sum@, @mean@,
@median@, …) deliberately collide with the aggregation-expression combinators
in "DataFrame.Typed.Expr", so this module is meant to be imported qualified:

@
import qualified DataFrame.Typed.Statistics as TS

avg = TS.mean (col \@\"salary\") employees
@
-}
module DataFrame.Typed.Statistics (
    mean,
    meanMaybe,
    median,
    medianMaybe,
    percentile,
    genericPercentile,
    standardDeviation,
    skewness,
    variance,
    interQuartileRange,
    sum,
    correlation,
    frequencies,
    imputeWith,
    summarize,
    describeColumns,
) where

import Data.Proxy (Proxy (..))
import qualified Data.Text as T
import qualified Data.Vector.Unboxed as VU
import GHC.TypeLits (KnownSymbol, symbolVal)
import Prelude hiding (sum)

import DataFrame.Internal.Column (Columnable)
import qualified DataFrame.Internal.DataFrame as D
import DataFrame.Internal.Nullable (BaseType)
import qualified DataFrame.Operations.Core as Core
import qualified DataFrame.Operations.Statistics as Stats
import DataFrame.Operations.Transformations (ImputeOp)
import DataFrame.Typed.Schema (AssertPresent, SafeLookup)
import DataFrame.Typed.Types (TExpr (..), TypedDataFrame (..))

-- | Mean of a column.
mean ::
    (Columnable a, Real a, VU.Unbox a) =>
    TExpr cols a -> TypedDataFrame cols -> Double
mean :: forall a (cols :: [(Symbol, *)]).
(Columnable a, Real a, Unbox a) =>
TExpr cols a -> TypedDataFrame cols -> Double
mean (TExpr Expr a
e) (TDF DataFrame
df) = Expr a -> DataFrame -> Double
forall a.
(Columnable a, Real a, Unbox a) =>
Expr a -> DataFrame -> Double
Stats.mean Expr a
e DataFrame
df

-- | Mean of a nullable column, ignoring 'Nothing'.
meanMaybe ::
    (Columnable a, Real a) =>
    TExpr cols (Maybe a) -> TypedDataFrame cols -> Double
meanMaybe :: forall a (cols :: [(Symbol, *)]).
(Columnable a, Real a) =>
TExpr cols (Maybe a) -> TypedDataFrame cols -> Double
meanMaybe (TExpr Expr (Maybe a)
e) (TDF DataFrame
df) = Expr (Maybe a) -> DataFrame -> Double
forall a.
(Columnable a, Real a) =>
Expr (Maybe a) -> DataFrame -> Double
Stats.meanMaybe Expr (Maybe a)
e DataFrame
df

-- | Median of a column.
median ::
    (Columnable a, Real a, VU.Unbox a) =>
    TExpr cols a -> TypedDataFrame cols -> Double
median :: forall a (cols :: [(Symbol, *)]).
(Columnable a, Real a, Unbox a) =>
TExpr cols a -> TypedDataFrame cols -> Double
median (TExpr Expr a
e) (TDF DataFrame
df) = Expr a -> DataFrame -> Double
forall a.
(Columnable a, Real a, Unbox a) =>
Expr a -> DataFrame -> Double
Stats.median Expr a
e DataFrame
df

-- | Median of a nullable column, ignoring 'Nothing'.
medianMaybe ::
    (Columnable a, Real a) =>
    TExpr cols (Maybe a) -> TypedDataFrame cols -> Double
medianMaybe :: forall a (cols :: [(Symbol, *)]).
(Columnable a, Real a) =>
TExpr cols (Maybe a) -> TypedDataFrame cols -> Double
medianMaybe (TExpr Expr (Maybe a)
e) (TDF DataFrame
df) = Expr (Maybe a) -> DataFrame -> Double
forall a.
(Columnable a, Real a) =>
Expr (Maybe a) -> DataFrame -> Double
Stats.medianMaybe Expr (Maybe a)
e DataFrame
df

-- | The @n@-th percentile of a column.
percentile ::
    (Columnable a, Real a, VU.Unbox a) =>
    Int -> TExpr cols a -> TypedDataFrame cols -> Double
percentile :: forall a (cols :: [(Symbol, *)]).
(Columnable a, Real a, Unbox a) =>
Int -> TExpr cols a -> TypedDataFrame cols -> Double
percentile Int
n (TExpr Expr a
e) (TDF DataFrame
df) = Int -> Expr a -> DataFrame -> Double
forall a.
(Columnable a, Real a, Unbox a) =>
Int -> Expr a -> DataFrame -> Double
Stats.percentile Int
n Expr a
e DataFrame
df

-- | The @n@-th percentile of a column of any 'Ord' type.
genericPercentile ::
    (Columnable a, Ord a) =>
    Int -> TExpr cols a -> TypedDataFrame cols -> a
genericPercentile :: forall a (cols :: [(Symbol, *)]).
(Columnable a, Ord a) =>
Int -> TExpr cols a -> TypedDataFrame cols -> a
genericPercentile Int
n (TExpr Expr a
e) (TDF DataFrame
df) = Int -> Expr a -> DataFrame -> a
forall a. (Columnable a, Ord a) => Int -> Expr a -> DataFrame -> a
Stats.genericPercentile Int
n Expr a
e DataFrame
df

-- | Standard deviation of a column.
standardDeviation ::
    (Columnable a, Real a, VU.Unbox a) =>
    TExpr cols a -> TypedDataFrame cols -> Double
standardDeviation :: forall a (cols :: [(Symbol, *)]).
(Columnable a, Real a, Unbox a) =>
TExpr cols a -> TypedDataFrame cols -> Double
standardDeviation (TExpr Expr a
e) (TDF DataFrame
df) = Expr a -> DataFrame -> Double
forall a.
(Columnable a, Real a, Unbox a) =>
Expr a -> DataFrame -> Double
Stats.standardDeviation Expr a
e DataFrame
df

-- | Skewness of a column.
skewness ::
    (Columnable a, Real a, VU.Unbox a) =>
    TExpr cols a -> TypedDataFrame cols -> Double
skewness :: forall a (cols :: [(Symbol, *)]).
(Columnable a, Real a, Unbox a) =>
TExpr cols a -> TypedDataFrame cols -> Double
skewness (TExpr Expr a
e) (TDF DataFrame
df) = Expr a -> DataFrame -> Double
forall a.
(Columnable a, Real a, Unbox a) =>
Expr a -> DataFrame -> Double
Stats.skewness Expr a
e DataFrame
df

-- | Variance of a column.
variance ::
    (Columnable a, Real a, VU.Unbox a) =>
    TExpr cols a -> TypedDataFrame cols -> Double
variance :: forall a (cols :: [(Symbol, *)]).
(Columnable a, Real a, Unbox a) =>
TExpr cols a -> TypedDataFrame cols -> Double
variance (TExpr Expr a
e) (TDF DataFrame
df) = Expr a -> DataFrame -> Double
forall a.
(Columnable a, Real a, Unbox a) =>
Expr a -> DataFrame -> Double
Stats.variance Expr a
e DataFrame
df

-- | Inter-quartile range of a column.
interQuartileRange ::
    (Columnable a, Real a, VU.Unbox a) =>
    TExpr cols a -> TypedDataFrame cols -> Double
interQuartileRange :: forall a (cols :: [(Symbol, *)]).
(Columnable a, Real a, Unbox a) =>
TExpr cols a -> TypedDataFrame cols -> Double
interQuartileRange (TExpr Expr a
e) (TDF DataFrame
df) = Expr a -> DataFrame -> Double
forall a.
(Columnable a, Real a, Unbox a) =>
Expr a -> DataFrame -> Double
Stats.interQuartileRange Expr a
e DataFrame
df

-- | Sum of a column.
sum :: (Columnable a, Num a) => TExpr cols a -> TypedDataFrame cols -> a
sum :: forall a (cols :: [(Symbol, *)]).
(Columnable a, Num a) =>
TExpr cols a -> TypedDataFrame cols -> a
sum (TExpr Expr a
e) (TDF DataFrame
df) = Expr a -> DataFrame -> a
forall a. (Columnable a, Num a) => Expr a -> DataFrame -> a
Stats.sum Expr a
e DataFrame
df

{- | Pearson's correlation coefficient between two columns, named by type
application. Both columns must exist in the schema and be numeric — these are
checked at compile time via 'SafeLookup' on each name.

@
TS.correlation \@\"height\" \@\"weight\" people
@
-}
correlation ::
    forall c1 c2 a b cols.
    ( KnownSymbol c1
    , KnownSymbol c2
    , a ~ SafeLookup c1 cols
    , b ~ SafeLookup c2 cols
    , Columnable a
    , Columnable b
    , Real a
    , Real b
    , VU.Unbox a
    , VU.Unbox b
    , AssertPresent c1 cols
    , AssertPresent c2 cols
    ) =>
    TypedDataFrame cols -> Maybe Double
correlation :: forall (c1 :: Symbol) (c2 :: Symbol) a b (cols :: [(Symbol, *)]).
(KnownSymbol c1, KnownSymbol c2, a ~ SafeLookup c1 cols,
 b ~ SafeLookup c2 cols, Columnable a, Columnable b, Real a, Real b,
 Unbox a, Unbox b, AssertPresent c1 cols, AssertPresent c2 cols) =>
TypedDataFrame cols -> Maybe Double
correlation (TDF DataFrame
df) =
    Text -> Text -> DataFrame -> Maybe Double
Stats.correlation
        (String -> Text
T.pack (Proxy c1 -> String
forall (n :: Symbol) (proxy :: Symbol -> *).
KnownSymbol n =>
proxy n -> String
symbolVal (forall {k} (t :: k). Proxy t
forall (t :: Symbol). Proxy t
Proxy @c1)))
        (String -> Text
T.pack (Proxy c2 -> String
forall (n :: Symbol) (proxy :: Symbol -> *).
KnownSymbol n =>
proxy n -> String
symbolVal (forall {k} (t :: k). Proxy t
forall (t :: Symbol). Proxy t
Proxy @c2)))
        DataFrame
df

{- | Frequency table for a column. The result schema is data-dependent
(one column per distinct value), so an untyped 'D.DataFrame' is returned.
-}
frequencies ::
    (Columnable a, Ord a) => TExpr cols a -> TypedDataFrame cols -> D.DataFrame
frequencies :: forall a (cols :: [(Symbol, *)]).
(Columnable a, Ord a) =>
TExpr cols a -> TypedDataFrame cols -> DataFrame
frequencies (TExpr Expr a
e) (TDF DataFrame
df) = Expr a -> DataFrame -> DataFrame
forall a. (Columnable a, Ord a) => Expr a -> DataFrame -> DataFrame
Stats.frequencies Expr a
e DataFrame
df

{- | Impute missing values in a column using a derived scalar (e.g. the mean).
Schema-preserving: the imputed column keeps its type-level @Maybe@ even though
its runtime values are now fully populated.
-}
imputeWith ::
    (ImputeOp a, Columnable (BaseType a)) =>
    (TExpr cols (BaseType a) -> TExpr cols (BaseType a)) ->
    TExpr cols a ->
    TypedDataFrame cols ->
    TypedDataFrame cols
imputeWith :: forall a (cols :: [(Symbol, *)]).
(ImputeOp a, Columnable (BaseType a)) =>
(TExpr cols (BaseType a) -> TExpr cols (BaseType a))
-> TExpr cols a -> TypedDataFrame cols -> TypedDataFrame cols
imputeWith TExpr cols (BaseType a) -> TExpr cols (BaseType a)
f (TExpr Expr a
e) (TDF DataFrame
df) =
    DataFrame -> TypedDataFrame cols
forall (cols :: [(Symbol, *)]). DataFrame -> TypedDataFrame cols
TDF ((Expr (BaseType a) -> Expr (BaseType a))
-> Expr a -> DataFrame -> DataFrame
forall a.
(ImputeOp a, Columnable (BaseType a)) =>
(Expr (BaseType a) -> Expr (BaseType a))
-> Expr a -> DataFrame -> DataFrame
Stats.imputeWith (TExpr cols (BaseType a) -> Expr (BaseType a)
forall (cols :: [(Symbol, *)]) a. TExpr cols a -> Expr a
unTExpr (TExpr cols (BaseType a) -> Expr (BaseType a))
-> (Expr (BaseType a) -> TExpr cols (BaseType a))
-> Expr (BaseType a)
-> Expr (BaseType a)
forall b c a. (b -> c) -> (a -> b) -> a -> c
. TExpr cols (BaseType a) -> TExpr cols (BaseType a)
f (TExpr cols (BaseType a) -> TExpr cols (BaseType a))
-> (Expr (BaseType a) -> TExpr cols (BaseType a))
-> Expr (BaseType a)
-> TExpr cols (BaseType a)
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Expr (BaseType a) -> TExpr cols (BaseType a)
forall (cols :: [(Symbol, *)]) a. Expr a -> TExpr cols a
TExpr) Expr a
e DataFrame
df)

{- | Descriptive statistics of the numeric columns. Returns an untyped
'D.DataFrame' (the result is a fixed set of statistic rows, not the input schema).
-}
summarize :: TypedDataFrame cols -> D.DataFrame
summarize :: forall (cols :: [(Symbol, *)]). TypedDataFrame cols -> DataFrame
summarize (TDF DataFrame
df) = DataFrame -> DataFrame
Stats.summarize DataFrame
df

{- | Per-column summary (non-null\/null counts, unique values, type). Returns an
untyped 'D.DataFrame'.
-}
describeColumns :: TypedDataFrame cols -> D.DataFrame
describeColumns :: forall (cols :: [(Symbol, *)]). TypedDataFrame cols -> DataFrame
describeColumns (TDF DataFrame
df) = DataFrame -> DataFrame
Core.describeColumns DataFrame
df