{-# LANGUAGE ConstraintKinds #-}
{-# LANGUAGE DataKinds #-}
{-# LANGUAGE FlexibleContexts #-}
{-# LANGUAGE FlexibleInstances #-}
{-# LANGUAGE MultiParamTypeClasses #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications #-}
{-# LANGUAGE TypeFamilies #-}
{-# LANGUAGE UndecidableInstances #-}
module DataFrame.Model (
Fit (..),
ToDataFrame (..),
Fitted (..),
FitResult,
FrameFor,
FrameKind (..),
CheckFrame,
AllDouble,
Predict (..),
AsTExpr,
ToTExpr (..),
) where
import Data.Kind (Constraint, Type)
import DataFrame.Internal.DataFrame (DataFrame)
import DataFrame.Internal.Expression (Expr (..))
import DataFrame.Typed.Freeze (ToDataFrame (..), thaw)
import DataFrame.Typed.Schema (AllDouble)
import DataFrame.Typed.Types (AsTExpr, TExpr (..), ToTExpr (..), TypedDataFrame)
import GHC.TypeLits (Symbol)
newtype Fitted (cols :: [(Symbol, Type)]) model = Fitted {forall (cols :: [(Symbol, *)]) model. Fitted cols model -> model
fittedModel :: model}
type family FitResult (f :: Type) (model :: Type) :: Type where
FitResult DataFrame model = model
FitResult (TypedDataFrame cols) model = Fitted cols model
type family FrameFor (input :: Type) :: Type where
FrameFor (Expr a) = DataFrame
FrameFor (TExpr cols a) = TypedDataFrame cols
FrameFor [Expr Double] = DataFrame
FrameFor [TExpr cols Double] = TypedDataFrame cols
data FrameKind = AnyFrame | AllDoubleFrame
type family CheckFrame (req :: FrameKind) (f :: Type) :: Constraint where
CheckFrame _ DataFrame = ()
CheckFrame 'AnyFrame _ = ()
CheckFrame 'AllDoubleFrame (TypedDataFrame cols) = AllDouble cols
class Fit cfg input where
type ModelOf cfg input :: Type
type FrameReq cfg input :: FrameKind
type FrameReq cfg input = 'AnyFrame
fit ::
(CheckFrame (FrameReq cfg input) (FrameFor input)) =>
cfg ->
input ->
FrameFor input ->
FitResult (FrameFor input) (ModelOf cfg input)
instance (Fit cfg (Expr a)) => Fit cfg (TExpr cols a) where
type ModelOf cfg (TExpr cols a) = ModelOf cfg (Expr a)
type FrameReq cfg (TExpr cols a) = FrameReq cfg (Expr a)
fit :: CheckFrame
(FrameReq cfg (TExpr cols a)) (FrameFor (TExpr cols a)) =>
cfg
-> TExpr cols a
-> FrameFor (TExpr cols a)
-> FitResult (FrameFor (TExpr cols a)) (ModelOf cfg (TExpr cols a))
fit cfg
cfg (TExpr Expr a
e) FrameFor (TExpr cols a)
tdf = ModelOf cfg (Expr a) -> Fitted cols (ModelOf cfg (Expr a))
forall (cols :: [(Symbol, *)]) model. model -> Fitted cols model
Fitted (cfg
-> Expr a
-> FrameFor (Expr a)
-> FitResult (FrameFor (Expr a)) (ModelOf cfg (Expr a))
forall cfg input.
(Fit cfg input,
CheckFrame (FrameReq cfg input) (FrameFor input)) =>
cfg
-> input
-> FrameFor input
-> FitResult (FrameFor input) (ModelOf cfg input)
fit cfg
cfg Expr a
e (TypedDataFrame cols -> DataFrame
forall (cols :: [(Symbol, *)]). TypedDataFrame cols -> DataFrame
thaw TypedDataFrame cols
FrameFor (TExpr cols a)
tdf))
instance (Fit cfg [Expr Double]) => Fit cfg [TExpr cols Double] where
type ModelOf cfg [TExpr cols Double] = ModelOf cfg [Expr Double]
type FrameReq cfg [TExpr cols Double] = FrameReq cfg [Expr Double]
fit :: CheckFrame
(FrameReq cfg [TExpr cols Double])
(FrameFor [TExpr cols Double]) =>
cfg
-> [TExpr cols Double]
-> FrameFor [TExpr cols Double]
-> FitResult
(FrameFor [TExpr cols Double]) (ModelOf cfg [TExpr cols Double])
fit cfg
cfg [TExpr cols Double]
feats FrameFor [TExpr cols Double]
tdf = ModelOf cfg [Expr Double]
-> Fitted cols (ModelOf cfg [Expr Double])
forall (cols :: [(Symbol, *)]) model. model -> Fitted cols model
Fitted (cfg
-> [Expr Double]
-> FrameFor [Expr Double]
-> FitResult (FrameFor [Expr Double]) (ModelOf cfg [Expr Double])
forall cfg input.
(Fit cfg input,
CheckFrame (FrameReq cfg input) (FrameFor input)) =>
cfg
-> input
-> FrameFor input
-> FitResult (FrameFor input) (ModelOf cfg input)
fit cfg
cfg ((TExpr cols Double -> Expr Double)
-> [TExpr cols Double] -> [Expr Double]
forall a b. (a -> b) -> [a] -> [b]
map TExpr cols Double -> Expr Double
forall (cols :: [(Symbol, *)]) a. TExpr cols a -> Expr a
unTExpr [TExpr cols Double]
feats) (TypedDataFrame cols -> DataFrame
forall (cols :: [(Symbol, *)]). TypedDataFrame cols -> DataFrame
thaw TypedDataFrame cols
FrameFor [TExpr cols Double]
tdf))
class Predict model where
type Prediction model :: Type
predict :: model -> Prediction model
instance (Predict model, ToTExpr cols (Prediction model)) => Predict (Fitted cols model) where
type Prediction (Fitted cols model) = AsTExpr cols (Prediction model)
predict :: Fitted cols model -> Prediction (Fitted cols model)
predict (Fitted model
m) = forall (cols :: [(Symbol, *)]) e.
ToTExpr cols e =>
e -> AsTExpr cols e
toTExpr @cols (model -> Prediction model
forall model. Predict model => model -> Prediction model
predict model
m)