{-# LANGUAGE FlexibleContexts #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications #-}

{- | Fitted column transforms as a composable monoid. A 'Transform' is a list of
named output expressions; @s <> t@ means \"apply @s@, then @t@\", fusing @t@'s
references to @s@'s outputs by simultaneous substitution. 'applyTransform' runs
one against a frame; 'compileThrough' folds a transform into a model's
prediction expression so the result is a single expression over the raw inputs.

Every right-hand side must be row-wise (no aggregation/window), and within one
transform each expression reads the original frame.
-}
module DataFrame.Transform (
    Transform (..),
    applyTransform,
    compileThrough,
    ScalerModel (..),
    standardScaler,
    scalerTransform,
) where

import Control.Exception (throw)
import qualified Data.Map.Strict as M
import qualified Data.Text as T
import qualified Data.Vector as V
import qualified Data.Vector.Unboxed as VU

import qualified DataFrame.Functions as F
import DataFrame.Internal.Column (Columnable)
import DataFrame.Internal.DataFrame (DataFrame)
import DataFrame.Internal.Expression (
    Expr (..),
    NamedExpr,
    UExpr (..),
    substituteColumns,
 )
import DataFrame.Operations.Core (columnAsDoubleVector)
import DataFrame.Operations.Transformations (deriveMany)
import DataFrame.Operators ((.-.), (./.))

-- | A fitted transform: named output columns derived from the input frame.
newtype Transform = Transform {Transform -> [NamedExpr]
transformOutputs :: [NamedExpr]}

instance Semigroup Transform where
    Transform [NamedExpr]
s <> :: Transform -> Transform -> Transform
<> Transform [NamedExpr]
t =
        [NamedExpr] -> Transform
Transform ([NamedExpr]
s [NamedExpr] -> [NamedExpr] -> [NamedExpr]
forall a. [a] -> [a] -> [a]
++ (NamedExpr -> NamedExpr) -> [NamedExpr] -> [NamedExpr]
forall a b. (a -> b) -> [a] -> [b]
map (Map Text UExpr -> NamedExpr -> NamedExpr
subst ([NamedExpr] -> Map Text UExpr
forall k a. Ord k => [(k, a)] -> Map k a
M.fromList [NamedExpr]
s)) [NamedExpr]
t)
      where
        subst :: M.Map T.Text UExpr -> NamedExpr -> NamedExpr
        subst :: Map Text UExpr -> NamedExpr -> NamedExpr
subst Map Text UExpr
m (Text
nm, UExpr Expr a
e) = (Text
nm, Expr a -> UExpr
forall a. Columnable a => Expr a -> UExpr
UExpr (Map Text UExpr -> Expr a -> Expr a
forall a. Columnable a => Map Text UExpr -> Expr a -> Expr a
substituteColumns Map Text UExpr
m Expr a
e))

instance Monoid Transform where
    mempty :: Transform
mempty = [NamedExpr] -> Transform
Transform []

-- | Apply a transform to a frame (deriving its outputs in order).
applyTransform :: Transform -> DataFrame -> DataFrame
applyTransform :: Transform -> DataFrame -> DataFrame
applyTransform (Transform [NamedExpr]
os) = [NamedExpr] -> DataFrame -> DataFrame
deriveMany [NamedExpr]
os

{- | Fold a preprocessing transform into a model's prediction expression,
yielding one expression over the transform's input columns.
-}
compileThrough :: (Columnable a) => Transform -> Expr a -> Expr a
compileThrough :: forall a. Columnable a => Transform -> Expr a -> Expr a
compileThrough (Transform [NamedExpr]
os) = Map Text UExpr -> Expr a -> Expr a
forall a. Columnable a => Map Text UExpr -> Expr a -> Expr a
substituteColumns ([NamedExpr] -> Map Text UExpr
forall k a. Ord k => [(k, a)] -> Map k a
M.fromList [NamedExpr]
os)

-- | A fitted standardizer: per-column means and standard deviations.
data ScalerModel = ScalerModel
    { ScalerModel -> Vector Text
smColumns :: !(V.Vector T.Text)
    , ScalerModel -> Vector Double
smMeans :: !(VU.Vector Double)
    , ScalerModel -> Vector Double
smStds :: !(VU.Vector Double)
    }
    deriving (ScalerModel -> ScalerModel -> Bool
(ScalerModel -> ScalerModel -> Bool)
-> (ScalerModel -> ScalerModel -> Bool) -> Eq ScalerModel
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: ScalerModel -> ScalerModel -> Bool
== :: ScalerModel -> ScalerModel -> Bool
$c/= :: ScalerModel -> ScalerModel -> Bool
/= :: ScalerModel -> ScalerModel -> Bool
Eq, Int -> ScalerModel -> ShowS
[ScalerModel] -> ShowS
ScalerModel -> String
(Int -> ScalerModel -> ShowS)
-> (ScalerModel -> String)
-> ([ScalerModel] -> ShowS)
-> Show ScalerModel
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> ScalerModel -> ShowS
showsPrec :: Int -> ScalerModel -> ShowS
$cshow :: ScalerModel -> String
show :: ScalerModel -> String
$cshowList :: [ScalerModel] -> ShowS
showList :: [ScalerModel] -> ShowS
Show)

-- | Fit a standard scaler over the named columns.
standardScaler :: [T.Text] -> DataFrame -> ScalerModel
standardScaler :: [Text] -> DataFrame -> ScalerModel
standardScaler [Text]
names DataFrame
df =
    Vector Text -> Vector Double -> Vector Double -> ScalerModel
ScalerModel ([Text] -> Vector Text
forall a. [a] -> Vector a
V.fromList [Text]
names) ([Double] -> Vector Double
forall a. Unbox a => [a] -> Vector a
VU.fromList [Double]
means) ([Double] -> Vector Double
forall a. Unbox a => [a] -> Vector a
VU.fromList [Double]
stds)
  where
    cols :: [Vector Double]
cols = (Text -> Vector Double) -> [Text] -> [Vector Double]
forall a b. (a -> b) -> [a] -> [b]
map Text -> Vector Double
column [Text]
names
    column :: Text -> Vector Double
column Text
n = case Expr Double
-> DataFrame -> Either DataFrameException (Vector Double)
forall a.
(Columnable a, Num a) =>
Expr a -> DataFrame -> Either DataFrameException (Vector Double)
columnAsDoubleVector (forall a. Columnable a => Text -> Expr a
F.col @Double Text
n) DataFrame
df of
        Right Vector Double
v -> Vector Double
v
        Left DataFrameException
e -> DataFrameException -> Vector Double
forall a e. Exception e => e -> a
throw DataFrameException
e
    means :: [Double]
means = [Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum Vector Double
c 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 (Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Double
c)) | Vector Double
c <- [Vector Double]
cols]
    stds :: [Double]
stds =
        [ let mu :: Double
mu = Double
mean
              v :: Double
v =
                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
x -> (Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
mu) Double -> Int -> Double
forall a b. (Num a, Integral b) => a -> b -> a
^ (Int
2 :: Int)) Vector Double
c)
                    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 (Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Double
c))
              s :: Double
s = Double -> Double
forall a. Floating a => a -> a
sqrt Double
v
           in if Double
s Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
< Double
1e-12 then Double
1 else Double
s
        | (Double
mean, Vector Double
c) <- [Double] -> [Vector Double] -> [(Double, Vector Double)]
forall a b. [a] -> [b] -> [(a, b)]
zip [Double]
means [Vector Double]
cols
        ]

-- | The scaler as a 'Transform': @(col - μ) / σ@ per column.
scalerTransform :: ScalerModel -> Transform
scalerTransform :: ScalerModel -> Transform
scalerTransform ScalerModel
m =
    [NamedExpr] -> Transform
Transform
        [ (Text
n, Expr Double -> UExpr
forall a. Columnable a => Expr a -> UExpr
UExpr ((forall a. Columnable a => Text -> Expr a
F.col @Double Text
n 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
mu) Expr Double -> Expr Double -> Expr Double
forall a.
(Columnable a, Fractional a) =>
Expr a -> Expr a -> Expr a
./. Double -> Expr Double
forall a. Columnable a => a -> Expr a
F.lit Double
sigma))
        | (Text
n, Double
mu, Double
sigma) <-
            [Text] -> [Double] -> [Double] -> [(Text, Double, Double)]
forall a b c. [a] -> [b] -> [c] -> [(a, b, c)]
zip3
                (Vector Text -> [Text]
forall a. Vector a -> [a]
V.toList (ScalerModel -> Vector Text
smColumns ScalerModel
m))
                (Vector Double -> [Double]
forall a. Unbox a => Vector a -> [a]
VU.toList (ScalerModel -> Vector Double
smMeans ScalerModel
m))
                (Vector Double -> [Double]
forall a. Unbox a => Vector a -> [a]
VU.toList (ScalerModel -> Vector Double
smStds ScalerModel
m))
        ]