{-# LANGUAGE FlexibleContexts #-}
{-# LANGUAGE FlexibleInstances #-}
{-# LANGUAGE MultiParamTypeClasses #-}
{-# LANGUAGE TypeFamilies #-}
{-# LANGUAGE UndecidableInstances #-}
module DataFrame.LinearModel.Logistic (
module DataFrame.Model,
LogisticConfig (..),
defaultLogisticConfig,
LogisticModel (..),
logisticMarginExprs,
logisticProbExprs,
) where
import Control.Parallel.Strategies (Strategy, parList, rseq, using)
import Data.List (sort)
import qualified Data.Map.Strict as M
import qualified Data.Vector as V
import qualified Data.Vector.Unboxed as VU
import DataFrame.Featurize.Internal (
affineExpr,
argMaxExpr,
featureNames,
numericMatrix,
targetValues,
)
import qualified DataFrame.Functions as F
import DataFrame.Internal.Column (Columnable)
import DataFrame.Internal.DataFrame (DataFrame)
import DataFrame.Internal.Expression (Expr)
import DataFrame.LinearSolver (
LinearModel (..),
SolverConfig,
defaultSolverConfig,
fitL1Logistic,
)
import DataFrame.Model
newtype LogisticConfig = LogisticConfig {LogisticConfig -> SolverConfig
lgSolver :: SolverConfig}
deriving (LogisticConfig -> LogisticConfig -> Bool
(LogisticConfig -> LogisticConfig -> Bool)
-> (LogisticConfig -> LogisticConfig -> Bool) -> Eq LogisticConfig
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: LogisticConfig -> LogisticConfig -> Bool
== :: LogisticConfig -> LogisticConfig -> Bool
$c/= :: LogisticConfig -> LogisticConfig -> Bool
/= :: LogisticConfig -> LogisticConfig -> Bool
Eq, Int -> LogisticConfig -> ShowS
[LogisticConfig] -> ShowS
LogisticConfig -> String
(Int -> LogisticConfig -> ShowS)
-> (LogisticConfig -> String)
-> ([LogisticConfig] -> ShowS)
-> Show LogisticConfig
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> LogisticConfig -> ShowS
showsPrec :: Int -> LogisticConfig -> ShowS
$cshow :: LogisticConfig -> String
show :: LogisticConfig -> String
$cshowList :: [LogisticConfig] -> ShowS
showList :: [LogisticConfig] -> ShowS
Show)
defaultLogisticConfig :: LogisticConfig
defaultLogisticConfig :: LogisticConfig
defaultLogisticConfig = SolverConfig -> LogisticConfig
LogisticConfig SolverConfig
defaultSolverConfig
data LogisticModel a = LogisticModel
{ forall a. LogisticModel a -> Vector a
lgClasses :: !(V.Vector a)
, forall a. LogisticModel a -> Vector LinearModel
lgModels :: !(V.Vector LinearModel)
}
deriving (LogisticModel a -> LogisticModel a -> Bool
(LogisticModel a -> LogisticModel a -> Bool)
-> (LogisticModel a -> LogisticModel a -> Bool)
-> Eq (LogisticModel a)
forall a. Eq a => LogisticModel a -> LogisticModel a -> Bool
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: forall a. Eq a => LogisticModel a -> LogisticModel a -> Bool
== :: LogisticModel a -> LogisticModel a -> Bool
$c/= :: forall a. Eq a => LogisticModel a -> LogisticModel a -> Bool
/= :: LogisticModel a -> LogisticModel a -> Bool
Eq, Int -> LogisticModel a -> ShowS
[LogisticModel a] -> ShowS
LogisticModel a -> String
(Int -> LogisticModel a -> ShowS)
-> (LogisticModel a -> String)
-> ([LogisticModel a] -> ShowS)
-> Show (LogisticModel a)
forall a. Show a => Int -> LogisticModel a -> ShowS
forall a. Show a => [LogisticModel a] -> ShowS
forall a. Show a => LogisticModel a -> String
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: forall a. Show a => Int -> LogisticModel a -> ShowS
showsPrec :: Int -> LogisticModel a -> ShowS
$cshow :: forall a. Show a => LogisticModel a -> String
show :: LogisticModel a -> String
$cshowList :: forall a. Show a => [LogisticModel a] -> ShowS
showList :: [LogisticModel a] -> ShowS
Show)
instance (Columnable a, Ord a) => Fit LogisticConfig (Expr a) where
type ModelOf LogisticConfig (Expr a) = (LogisticModel a)
fit :: CheckFrame
(FrameReq LogisticConfig (Expr a)) (FrameFor (Expr a)) =>
LogisticConfig
-> Expr a
-> FrameFor (Expr a)
-> FitResult (FrameFor (Expr a)) (ModelOf LogisticConfig (Expr a))
fit = LogisticConfig -> Expr a -> DataFrame -> LogisticModel a
LogisticConfig
-> Expr a
-> FrameFor (Expr a)
-> FitResult (FrameFor (Expr a)) (ModelOf LogisticConfig (Expr a))
forall a.
(Columnable a, Ord a) =>
LogisticConfig -> Expr a -> DataFrame -> LogisticModel a
fitLogistic
instance (Columnable a, Ord a) => Predict (LogisticModel a) where
type Prediction (LogisticModel a) = Expr a
predict :: LogisticModel a -> Prediction (LogisticModel a)
predict LogisticModel a
m = [(a, Expr Double)] -> Expr a
forall a. Columnable a => [(a, Expr Double)] -> Expr a
argMaxExpr (LogisticModel a -> [(a, Expr Double)]
forall a. LogisticModel a -> [(a, Expr Double)]
labelledMargins LogisticModel a
m)
fitLogistic ::
(Columnable a, Ord a) =>
LogisticConfig -> Expr a -> DataFrame -> LogisticModel a
fitLogistic :: forall a.
(Columnable a, Ord a) =>
LogisticConfig -> Expr a -> DataFrame -> LogisticModel a
fitLogistic (LogisticConfig SolverConfig
cfg) Expr a
target DataFrame
df =
Vector a -> Vector LinearModel -> LogisticModel a
forall a. Vector a -> Vector LinearModel -> LogisticModel a
LogisticModel ([a] -> Vector a
forall a. [a] -> Vector a
V.fromList [a]
classes) ([LinearModel] -> Vector LinearModel
forall a. [a] -> Vector a
V.fromList [LinearModel]
models)
where
names :: [Text]
names = Expr a -> DataFrame -> [Text]
forall a. Expr a -> DataFrame -> [Text]
featureNames Expr a
target DataFrame
df
(Vector Text
nameVec, Matrix
mat) = [Text] -> DataFrame -> (Vector Text, Matrix)
numericMatrix [Text]
names DataFrame
df
ys :: Vector a
ys = Expr a -> DataFrame -> Vector a
forall a. Columnable a => Expr a -> DataFrame -> Vector a
targetValues Expr a
target DataFrame
df
classes :: [a]
classes = [a] -> [a]
forall a. Ord a => [a] -> [a]
sort ((a -> [a] -> [a]) -> [a] -> [a] -> [a]
forall a b. (a -> b -> b) -> b -> [a] -> b
forall (t :: * -> *) a b.
Foldable t =>
(a -> b -> b) -> b -> t a -> b
foldr a -> [a] -> [a]
forall {a}. Eq a => a -> [a] -> [a]
dedup [] (Vector a -> [a]
forall a. Vector a -> [a]
V.toList Vector a
ys))
dedup :: a -> [a] -> [a]
dedup a
x [a]
acc = if a
x a -> [a] -> Bool
forall a. Eq a => a -> [a] -> Bool
forall (t :: * -> *) a. (Foldable t, Eq a) => a -> t a -> Bool
`elem` [a]
acc then [a]
acc else a
x a -> [a] -> [a]
forall a. a -> [a] -> [a]
: [a]
acc
fitOne :: a -> LinearModel
fitOne a
c =
let labels :: Vector Double
labels =
Int -> (Int -> Double) -> Vector Double
forall a. Unbox a => Int -> (Int -> a) -> Vector a
VU.generate (Vector a -> Int
forall a. Vector a -> Int
V.length Vector a
ys) (\Int
i -> if Vector a
ys Vector a -> Int -> a
forall a. Vector a -> Int -> a
V.! Int
i a -> a -> Bool
forall a. Eq a => a -> a -> Bool
== a
c then Double
1 else -Double
1)
in SolverConfig
-> Matrix -> Vector Double -> Vector Text -> LinearModel
fitL1Logistic SolverConfig
cfg Matrix
mat Vector Double
labels Vector Text
nameVec
models :: [LinearModel]
models = case [a]
classes of
[a
c0, a
_c1] -> let m0 :: LinearModel
m0 = a -> LinearModel
fitOne a
c0 in [LinearModel
m0, LinearModel -> LinearModel
negateModel LinearModel
m0]
[a]
_ -> (a -> LinearModel) -> [a] -> [LinearModel]
forall a b. (a -> b) -> [a] -> [b]
map a -> LinearModel
fitOne [a]
classes [LinearModel] -> Strategy [LinearModel] -> [LinearModel]
forall a. a -> Strategy a -> a
`using` Strategy LinearModel -> Strategy [LinearModel]
forall a. Strategy a -> Strategy [a]
parList Strategy LinearModel
forceModel
negateModel :: LinearModel -> LinearModel
negateModel :: LinearModel -> LinearModel
negateModel LinearModel
m =
LinearModel
m{lmWeights = VU.map negate (lmWeights m), lmIntercept = negate (lmIntercept m)}
forceModel :: Strategy LinearModel
forceModel :: Strategy LinearModel
forceModel LinearModel
m = LinearModel -> Vector Double
lmWeights LinearModel
m Vector Double -> Eval LinearModel -> Eval LinearModel
forall a b. a -> b -> b
`seq` Strategy LinearModel
forall a. Strategy a
rseq LinearModel
m
logisticMarginExprs ::
(Columnable a, Ord a) => LogisticModel a -> M.Map a (Expr Double)
logisticMarginExprs :: forall a.
(Columnable a, Ord a) =>
LogisticModel a -> Map a (Expr Double)
logisticMarginExprs LogisticModel a
m = [(a, Expr Double)] -> Map a (Expr Double)
forall k a. Ord k => [(k, a)] -> Map k a
M.fromList (LogisticModel a -> [(a, Expr Double)]
forall a. LogisticModel a -> [(a, Expr Double)]
labelledMargins LogisticModel a
m)
logisticProbExprs ::
(Columnable a, Ord a) => LogisticModel a -> M.Map a (Expr Double)
logisticProbExprs :: forall a.
(Columnable a, Ord a) =>
LogisticModel a -> Map a (Expr Double)
logisticProbExprs = (Expr Double -> Expr Double)
-> Map a (Expr Double) -> Map a (Expr Double)
forall a b k. (a -> b) -> Map k a -> Map k b
M.map Expr Double -> Expr Double
sigmoidExpr (Map a (Expr Double) -> Map a (Expr Double))
-> (LogisticModel a -> Map a (Expr Double))
-> LogisticModel a
-> Map a (Expr Double)
forall b c a. (b -> c) -> (a -> b) -> a -> c
. LogisticModel a -> Map a (Expr Double)
forall a.
(Columnable a, Ord a) =>
LogisticModel a -> Map a (Expr Double)
logisticMarginExprs
labelledMargins :: LogisticModel a -> [(a, Expr Double)]
labelledMargins :: forall a. LogisticModel a -> [(a, Expr Double)]
labelledMargins LogisticModel a
m =
[ (LogisticModel a -> Vector a
forall a. LogisticModel a -> Vector a
lgClasses LogisticModel a
m Vector a -> Int -> a
forall a. Vector a -> Int -> a
V.! Int
i, LinearModel -> Expr Double
marginOf (LogisticModel a -> Vector LinearModel
forall a. LogisticModel a -> Vector LinearModel
lgModels LogisticModel a
m Vector LinearModel -> Int -> LinearModel
forall a. Vector a -> Int -> a
V.! Int
i))
| Int
i <- [Int
0 .. Vector a -> Int
forall a. Vector a -> Int
V.length (LogisticModel a -> Vector a
forall a. LogisticModel a -> Vector a
lgClasses LogisticModel a
m) Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]
]
marginOf :: LinearModel -> Expr Double
marginOf :: LinearModel -> Expr Double
marginOf LinearModel
m =
Double -> [(Double, Text)] -> Expr Double
affineExpr
(LinearModel -> Double
lmIntercept LinearModel
m)
([Double] -> [Text] -> [(Double, Text)]
forall a b. [a] -> [b] -> [(a, b)]
zip (Vector Double -> [Double]
forall a. Unbox a => Vector a -> [a]
VU.toList (LinearModel -> Vector Double
lmWeights LinearModel
m)) (Vector Text -> [Text]
forall a. Vector a -> [a]
V.toList (LinearModel -> Vector Text
lmFeatureNames LinearModel
m)))
sigmoidExpr :: Expr Double -> Expr Double
sigmoidExpr :: Expr Double -> Expr Double
sigmoidExpr Expr Double
z = Double -> Expr Double
forall a. Columnable a => a -> Expr a
F.lit Double
1 Expr Double -> Expr Double -> Expr Double
forall a. Fractional a => a -> a -> a
/ (Double -> Expr Double
forall a. Columnable a => a -> Expr a
F.lit Double
1 Expr Double -> Expr Double -> Expr Double
forall a. Num a => a -> a -> a
+ Expr Double -> Expr Double
forall a. Floating a => a -> a
exp (Expr Double -> Expr Double
forall a. Num a => a -> a
negate Expr Double
z))