{-# LANGUAGE DataKinds #-}
{-# LANGUAGE FlexibleInstances #-}
{-# LANGUAGE MultiParamTypeClasses #-}
{-# LANGUAGE TypeFamilies #-}

{- | Linear regression with the standard penalties: OLS (QR), ridge (Cholesky),
and lasso\/elastic net (FISTA). 'fit' produces a 'LinearRegressor'; 'predict'
compiles it to an @Expr Double@ over the raw feature columns.
-}
module DataFrame.LinearModel.Regression (
    module DataFrame.Model,
    Penalty (..),
    LinearConfig (..),
    defaultLinearConfig,
    LinearRegressor (..),
) where

import qualified Data.Text as T
import qualified Data.Vector as V
import qualified Data.Vector.Unboxed as VU

import DataFrame.Featurize.Internal (
    affineExpr,
    featureNames,
    numericMatrix,
    targetDoubles,
 )
import DataFrame.Internal.Expression (Expr)
import DataFrame.LinearAlgebra (Matrix, dot, gram, tMatVec)
import DataFrame.LinearAlgebra.Solve (choleskySolve, qrLeastSquares)
import DataFrame.LinearSolver (
    LinearModel (..),
    SolverConfig (..),
    defaultSolverConfig,
    fitProx,
 )
import DataFrame.LinearSolver.Loss (squaredLoss)
import DataFrame.Model

-- | Regularization choice. @alpha@ is the penalty strength; @l1Ratio@ mixes L1/L2.
data Penalty
    = OLS
    | Ridge !Double
    | Lasso !Double
    | ElasticNet !Double !Double
    deriving (Penalty -> Penalty -> Bool
(Penalty -> Penalty -> Bool)
-> (Penalty -> Penalty -> Bool) -> Eq Penalty
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: Penalty -> Penalty -> Bool
== :: Penalty -> Penalty -> Bool
$c/= :: Penalty -> Penalty -> Bool
/= :: Penalty -> Penalty -> Bool
Eq, Int -> Penalty -> ShowS
[Penalty] -> ShowS
Penalty -> String
(Int -> Penalty -> ShowS)
-> (Penalty -> String) -> ([Penalty] -> ShowS) -> Show Penalty
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> Penalty -> ShowS
showsPrec :: Int -> Penalty -> ShowS
$cshow :: Penalty -> String
show :: Penalty -> String
$cshowList :: [Penalty] -> ShowS
showList :: [Penalty] -> ShowS
Show)

-- | Hyperparameters for linear regression: the penalty and the FISTA solver config.
data LinearConfig = LinearConfig
    { LinearConfig -> Penalty
lcPenalty :: !Penalty
    , LinearConfig -> SolverConfig
lcSolver :: !SolverConfig
    }
    deriving (LinearConfig -> LinearConfig -> Bool
(LinearConfig -> LinearConfig -> Bool)
-> (LinearConfig -> LinearConfig -> Bool) -> Eq LinearConfig
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: LinearConfig -> LinearConfig -> Bool
== :: LinearConfig -> LinearConfig -> Bool
$c/= :: LinearConfig -> LinearConfig -> Bool
/= :: LinearConfig -> LinearConfig -> Bool
Eq, Int -> LinearConfig -> ShowS
[LinearConfig] -> ShowS
LinearConfig -> String
(Int -> LinearConfig -> ShowS)
-> (LinearConfig -> String)
-> ([LinearConfig] -> ShowS)
-> Show LinearConfig
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> LinearConfig -> ShowS
showsPrec :: Int -> LinearConfig -> ShowS
$cshow :: LinearConfig -> String
show :: LinearConfig -> String
$cshowList :: [LinearConfig] -> ShowS
showList :: [LinearConfig] -> ShowS
Show)

defaultLinearConfig :: LinearConfig
defaultLinearConfig :: LinearConfig
defaultLinearConfig = LinearConfig{lcPenalty :: Penalty
lcPenalty = Penalty
OLS, lcSolver :: SolverConfig
lcSolver = SolverConfig
defaultSolverConfig}

{- | A fitted linear regressor. @regCoef@ and @regIntercept@ are sklearn's
@coef_@ / @intercept_@ in raw feature space.
-}
data LinearRegressor = LinearRegressor
    { LinearRegressor -> Vector Double
regCoef :: !(VU.Vector Double)
    , LinearRegressor -> Double
regIntercept :: !Double
    , LinearRegressor -> Vector Text
regFeatureNames :: !(V.Vector T.Text)
    , LinearRegressor -> Penalty
regPenalty :: !Penalty
    }
    deriving (LinearRegressor -> LinearRegressor -> Bool
(LinearRegressor -> LinearRegressor -> Bool)
-> (LinearRegressor -> LinearRegressor -> Bool)
-> Eq LinearRegressor
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: LinearRegressor -> LinearRegressor -> Bool
== :: LinearRegressor -> LinearRegressor -> Bool
$c/= :: LinearRegressor -> LinearRegressor -> Bool
/= :: LinearRegressor -> LinearRegressor -> Bool
Eq, Int -> LinearRegressor -> ShowS
[LinearRegressor] -> ShowS
LinearRegressor -> String
(Int -> LinearRegressor -> ShowS)
-> (LinearRegressor -> String)
-> ([LinearRegressor] -> ShowS)
-> Show LinearRegressor
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> LinearRegressor -> ShowS
showsPrec :: Int -> LinearRegressor -> ShowS
$cshow :: LinearRegressor -> String
show :: LinearRegressor -> String
$cshowList :: [LinearRegressor] -> ShowS
showList :: [LinearRegressor] -> ShowS
Show)

instance Fit LinearConfig (Expr Double) where
    type ModelOf LinearConfig (Expr Double) = LinearRegressor
    type FrameReq LinearConfig (Expr Double) = 'AllDoubleFrame
    fit :: CheckFrame
  (FrameReq LinearConfig (Expr Double)) (FrameFor (Expr Double)) =>
LinearConfig
-> Expr Double
-> FrameFor (Expr Double)
-> FitResult
     (FrameFor (Expr Double)) (ModelOf LinearConfig (Expr Double))
fit (LinearConfig Penalty
penalty SolverConfig
cfg) Expr Double
target FrameFor (Expr Double)
df =
        case Penalty
penalty of
            Penalty
OLS -> (Vector Double, Double) -> LinearRegressor
closedForm (Matrix -> Vector Double -> (Vector Double, Double)
olsSolve Matrix
mat Vector Double
y)
            Ridge Double
alpha -> (Vector Double, Double) -> LinearRegressor
closedForm (Double -> Matrix -> Vector Double -> (Vector Double, Double)
ridgeSolve Double
alpha Matrix
mat Vector Double
y)
            Lasso Double
alpha -> Double -> Double -> LinearRegressor
proxFit Double
alpha Double
1.0
            ElasticNet Double
alpha Double
l1r -> Double -> Double -> LinearRegressor
proxFit Double
alpha Double
l1r
      where
        names :: [Text]
names = Expr Double -> DataFrame -> [Text]
forall a. Expr a -> DataFrame -> [Text]
featureNames Expr Double
target DataFrame
FrameFor (Expr Double)
df
        (Vector Text
nameVec, Matrix
mat) = [Text] -> DataFrame -> (Vector Text, Matrix)
numericMatrix [Text]
names DataFrame
FrameFor (Expr Double)
df
        y :: Vector Double
y = Expr Double -> DataFrame -> Vector Double
targetDoubles Expr Double
target DataFrame
FrameFor (Expr Double)
df
        closedForm :: (Vector Double, Double) -> LinearRegressor
closedForm (Vector Double
coef, Double
intercept) =
            Vector Double
-> Double -> Vector Text -> Penalty -> LinearRegressor
LinearRegressor Vector Double
coef Double
intercept Vector Text
nameVec Penalty
penalty
        proxFit :: Double -> Double -> LinearRegressor
proxFit Double
alpha Double
l1r =
            let proxCfg :: SolverConfig
proxCfg =
                    SolverConfig
cfg{scL1Lambda = alpha * l1r, scL2Lambda = alpha * (1 - l1r)}
                m :: LinearModel
m = SmoothLoss
-> SolverConfig
-> Matrix
-> Vector Double
-> Vector Text
-> LinearModel
fitProx SmoothLoss
squaredLoss SolverConfig
proxCfg Matrix
mat Vector Double
y Vector Text
nameVec
             in Vector Double
-> Double -> Vector Text -> Penalty -> LinearRegressor
LinearRegressor (LinearModel -> Vector Double
lmWeights LinearModel
m) (LinearModel -> Double
lmIntercept LinearModel
m) Vector Text
nameVec Penalty
penalty

instance Predict LinearRegressor where
    type Prediction LinearRegressor = Expr Double
    predict :: LinearRegressor -> Prediction LinearRegressor
predict LinearRegressor
m =
        Double -> [(Double, Text)] -> Expr Double
affineExpr
            (LinearRegressor -> Double
regIntercept LinearRegressor
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 (LinearRegressor -> Vector Double
regCoef LinearRegressor
m)) (Vector Text -> [Text]
forall a. Vector a -> [a]
V.toList (LinearRegressor -> Vector Text
regFeatureNames LinearRegressor
m)))

-- | OLS via QR on the intercept-augmented design matrix.
olsSolve :: Matrix -> VU.Vector Double -> (VU.Vector Double, Double)
olsSolve :: Matrix -> Vector Double -> (Vector Double, Double)
olsSolve Matrix
mat Vector Double
y =
    case Matrix -> Vector Double -> Either [Int] (Vector Double)
qrLeastSquares Matrix
augmented Vector Double
y of
        Right Vector Double
sol -> (Int -> Vector Double -> Vector Double
forall a. Unbox a => Int -> Vector a -> Vector a
VU.drop Int
1 Vector Double
sol, Vector Double
sol Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
0)
        Left [Int]
_ -> Double -> Matrix -> Vector Double -> (Vector Double, Double)
ridgeSolve Double
1e-8 Matrix
mat Vector Double
y
  where
    augmented :: Matrix
augmented = (Vector Double -> Vector Double) -> Matrix -> Matrix
forall a b. (a -> b) -> Vector a -> Vector b
V.map (Double -> Vector Double -> Vector Double
forall a. Unbox a => a -> Vector a -> Vector a
VU.cons Double
1) Matrix
mat

{- | Ridge via Cholesky on @(XcᵀXc + αI) w = Xcᵀ yc@ over centred data; the
intercept is recovered from the column/target means.
-}
ridgeSolve :: Double -> Matrix -> VU.Vector Double -> (VU.Vector Double, Double)
ridgeSolve :: Double -> Matrix -> Vector Double -> (Vector Double, Double)
ridgeSolve Double
alpha Matrix
mat Vector Double
y =
    case Matrix -> Vector Double -> Maybe (Vector Double)
choleskySolve Matrix
a Vector Double
rhs of
        Just Vector Double
w -> (Vector Double
w, Double
meanY Double -> Double -> Double
forall a. Num a => a -> a -> a
- Vector Double -> Vector Double -> Double
dot Vector Double
w Vector Double
meansX)
        Maybe (Vector Double)
Nothing -> (Int -> Double -> Vector Double
forall a. Unbox a => Int -> a -> Vector a
VU.replicate Int
d Double
0, Double
meanY)
  where
    n :: Int
n = Matrix -> Int
forall a. Vector a -> Int
V.length Matrix
mat
    d :: Int
d = if Int
n Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
0 then Int
0 else Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length (Matrix -> Vector Double
forall a. Vector a -> a
V.head Matrix
mat)
    meansX :: Vector Double
meansX =
        Int -> (Int -> Double) -> Vector Double
forall a. Unbox a => Int -> (Int -> a) -> Vector a
VU.generate Int
d ((Int -> Double) -> Vector Double)
-> (Int -> Double) -> Vector Double
forall a b. (a -> b) -> a -> b
$ \Int
j ->
            [Double] -> Double
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [(Matrix
mat Matrix -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
i) Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
j | Int
i <- [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]] Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
n
    meanY :: Double
meanY = Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum Vector Double
y Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
n
    centered :: Matrix
centered = (Vector Double -> Vector Double) -> Matrix -> Matrix
forall a b. (a -> b) -> Vector a -> Vector b
V.map (\Vector Double
row -> (Double -> Double -> Double)
-> Vector Double -> Vector Double -> Vector Double
forall a b c.
(Unbox a, Unbox b, Unbox c) =>
(a -> b -> c) -> Vector a -> Vector b -> Vector c
VU.zipWith (-) Vector Double
row Vector Double
meansX) Matrix
mat
    yc :: Vector Double
yc = (Double -> Double) -> Vector Double -> Vector Double
forall a b. (Unbox a, Unbox b) => (a -> b) -> Vector a -> Vector b
VU.map (Double -> Double -> Double
forall a. Num a => a -> a -> a
subtract Double
meanY) Vector Double
y
    a :: Matrix
a = Double -> Matrix -> Matrix
addDiag Double
alpha (Matrix -> Matrix
gram Matrix
centered)
    rhs :: Vector Double
rhs = Matrix -> Vector Double -> Vector Double
tMatVec Matrix
centered Vector Double
yc

-- | Add @alpha@ to the diagonal of a square matrix.
addDiag :: Double -> Matrix -> Matrix
addDiag :: Double -> Matrix -> Matrix
addDiag Double
alpha = (Int -> Vector Double -> Vector Double) -> Matrix -> Matrix
forall a b. (Int -> a -> b) -> Vector a -> Vector b
V.imap (\Int
i Vector Double
row -> Vector Double
row Vector Double -> [(Int, Double)] -> Vector Double
forall a. Unbox a => Vector a -> [(Int, a)] -> Vector a
VU.// [(Int
i, Vector Double
row Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
alpha)])

-- | Predict the target for each row of a feature matrix.
predictLinear :: LinearRegressor -> Matrix -> VU.Vector Double
predictLinear :: LinearRegressor -> Matrix -> Vector Double
predictLinear LinearRegressor
m = Vector Double -> Vector Double
forall (v :: * -> *) a (w :: * -> *).
(Vector v a, Vector w a) =>
v a -> w a
VU.convert (Vector Double -> Vector Double)
-> (Matrix -> Vector Double) -> Matrix -> Vector Double
forall b c a. (b -> c) -> (a -> b) -> a -> c
. (Vector Double -> Double) -> Matrix -> Vector Double
forall a b. (a -> b) -> Vector a -> Vector b
V.map (\Vector Double
x -> LinearRegressor -> Double
regIntercept LinearRegressor
m Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Vector Double -> Vector Double -> Double
dot (LinearRegressor -> Vector Double
regCoef LinearRegressor
m) Vector Double
x)