{-# LANGUAGE FlexibleContexts #-}
{-# LANGUAGE FlexibleInstances #-}
{-# LANGUAGE MultiParamTypeClasses #-}
{-# LANGUAGE TypeFamilies #-}
{-# LANGUAGE UndecidableInstances #-}

{- | Linear support vector classification: L2-regularized squared hinge via
FISTA (sklearn's LinearSVC default). 'fit' trains a one-vs-rest 'LinearSVCModel';
'predict' is the arg-max class margin (no @predict_proba@, as in sklearn).
-}
module DataFrame.SVM (
    module DataFrame.Model,
    LinearSVCModel (..),
    SVCConfig (..),
    defaultSVCConfig,
    svcMarginExprs,
    -- | Surfaced by @LinearSVCModel.svcModels@.
    LinearModel (..),
) where

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 DataFrame.Internal.Column (Columnable)
import DataFrame.Internal.DataFrame (DataFrame)
import DataFrame.Internal.Expression (Expr)
import DataFrame.LinearSolver (LinearModel (..), SolverConfig (..), fitProx)
import DataFrame.LinearSolver.Loss (sqHingeLoss)
import DataFrame.Model

-- | Hyper-parameters. @svcC@ is the inverse regularization strength (sklearn @C@).
data SVCConfig = SVCConfig
    { SVCConfig -> Double
svcC :: !Double
    , SVCConfig -> Int
svcMaxIter :: !Int
    , SVCConfig -> Double
svcTol :: !Double
    }
    deriving (SVCConfig -> SVCConfig -> Bool
(SVCConfig -> SVCConfig -> Bool)
-> (SVCConfig -> SVCConfig -> Bool) -> Eq SVCConfig
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: SVCConfig -> SVCConfig -> Bool
== :: SVCConfig -> SVCConfig -> Bool
$c/= :: SVCConfig -> SVCConfig -> Bool
/= :: SVCConfig -> SVCConfig -> Bool
Eq, Int -> SVCConfig -> ShowS
[SVCConfig] -> ShowS
SVCConfig -> String
(Int -> SVCConfig -> ShowS)
-> (SVCConfig -> String)
-> ([SVCConfig] -> ShowS)
-> Show SVCConfig
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> SVCConfig -> ShowS
showsPrec :: Int -> SVCConfig -> ShowS
$cshow :: SVCConfig -> String
show :: SVCConfig -> String
$cshowList :: [SVCConfig] -> ShowS
showList :: [SVCConfig] -> ShowS
Show)

defaultSVCConfig :: SVCConfig
defaultSVCConfig :: SVCConfig
defaultSVCConfig = SVCConfig{svcC :: Double
svcC = Double
1.0, svcMaxIter :: Int
svcMaxIter = Int
1000, svcTol :: Double
svcTol = Double
1.0e-4}

-- | A fitted one-vs-rest linear SVC: class labels and their margin sub-models.
data LinearSVCModel a = LinearSVCModel
    { forall a. LinearSVCModel a -> Vector a
svcClasses :: !(V.Vector a)
    , forall a. LinearSVCModel a -> Vector LinearModel
svcModels :: !(V.Vector LinearModel)
    }
    deriving (LinearSVCModel a -> LinearSVCModel a -> Bool
(LinearSVCModel a -> LinearSVCModel a -> Bool)
-> (LinearSVCModel a -> LinearSVCModel a -> Bool)
-> Eq (LinearSVCModel a)
forall a. Eq a => LinearSVCModel a -> LinearSVCModel a -> Bool
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: forall a. Eq a => LinearSVCModel a -> LinearSVCModel a -> Bool
== :: LinearSVCModel a -> LinearSVCModel a -> Bool
$c/= :: forall a. Eq a => LinearSVCModel a -> LinearSVCModel a -> Bool
/= :: LinearSVCModel a -> LinearSVCModel a -> Bool
Eq, Int -> LinearSVCModel a -> ShowS
[LinearSVCModel a] -> ShowS
LinearSVCModel a -> String
(Int -> LinearSVCModel a -> ShowS)
-> (LinearSVCModel a -> String)
-> ([LinearSVCModel a] -> ShowS)
-> Show (LinearSVCModel a)
forall a. Show a => Int -> LinearSVCModel a -> ShowS
forall a. Show a => [LinearSVCModel a] -> ShowS
forall a. Show a => LinearSVCModel a -> String
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: forall a. Show a => Int -> LinearSVCModel a -> ShowS
showsPrec :: Int -> LinearSVCModel a -> ShowS
$cshow :: forall a. Show a => LinearSVCModel a -> String
show :: LinearSVCModel a -> String
$cshowList :: forall a. Show a => [LinearSVCModel a] -> ShowS
showList :: [LinearSVCModel a] -> ShowS
Show)

instance (Columnable a, Ord a) => Fit SVCConfig (Expr a) where
    type ModelOf SVCConfig (Expr a) = (LinearSVCModel a)
    fit :: CheckFrame (FrameReq SVCConfig (Expr a)) (FrameFor (Expr a)) =>
SVCConfig
-> Expr a
-> FrameFor (Expr a)
-> FitResult (FrameFor (Expr a)) (ModelOf SVCConfig (Expr a))
fit = SVCConfig -> Expr a -> DataFrame -> LinearSVCModel a
SVCConfig
-> Expr a
-> FrameFor (Expr a)
-> FitResult (FrameFor (Expr a)) (ModelOf SVCConfig (Expr a))
forall a.
(Columnable a, Ord a) =>
SVCConfig -> Expr a -> DataFrame -> LinearSVCModel a
fitLinearSVC

instance (Columnable a, Ord a) => Predict (LinearSVCModel a) where
    type Prediction (LinearSVCModel a) = Expr a
    predict :: LinearSVCModel a -> Prediction (LinearSVCModel a)
predict LinearSVCModel a
m = [(a, Expr Double)] -> Expr a
forall a. Columnable a => [(a, Expr Double)] -> Expr a
argMaxExpr (LinearSVCModel a -> [(a, Expr Double)]
forall a. LinearSVCModel a -> [(a, Expr Double)]
labelledMargins LinearSVCModel a
m)

-- | Fit a one-vs-rest linear SVC.
fitLinearSVC ::
    (Columnable a, Ord a) =>
    SVCConfig -> Expr a -> DataFrame -> LinearSVCModel a
fitLinearSVC :: forall a.
(Columnable a, Ord a) =>
SVCConfig -> Expr a -> DataFrame -> LinearSVCModel a
fitLinearSVC SVCConfig
cfg Expr a
target DataFrame
df =
    Vector a -> Vector LinearModel -> LinearSVCModel a
forall a. Vector a -> Vector LinearModel -> LinearSVCModel a
LinearSVCModel ([a] -> Vector a
forall a. [a] -> Vector a
V.fromList [a]
classes) ([LinearModel] -> Vector LinearModel
forall a. [a] -> Vector a
V.fromList ((a -> LinearModel) -> [a] -> [LinearModel]
forall a b. (a -> b) -> [a] -> [b]
map a -> LinearModel
fitOne [a]
classes))
  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
    solverCfg :: SolverConfig
solverCfg =
        SolverConfig
            { scL1Lambda :: Double
scL1Lambda = Double
0
            , scL2Lambda :: Double
scL2Lambda = Double
1 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ SVCConfig -> Double
svcC SVCConfig
cfg
            , scMaxIter :: Int
scMaxIter = SVCConfig -> Int
svcMaxIter SVCConfig
cfg
            , scTol :: Double
scTol = SVCConfig -> Double
svcTol SVCConfig
cfg
            , scSampleWeights :: Maybe (Vector Double)
scSampleWeights = Maybe (Vector Double)
forall a. Maybe a
Nothing
            }
    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 SmoothLoss
-> SolverConfig
-> Matrix
-> Vector Double
-> Vector Text
-> LinearModel
fitProx SmoothLoss
sqHingeLoss SolverConfig
solverCfg Matrix
mat Vector Double
labels Vector Text
nameVec

-- | The raw margin expression for each class.
svcMarginExprs ::
    (Columnable a, Ord a) => LinearSVCModel a -> M.Map a (Expr Double)
svcMarginExprs :: forall a.
(Columnable a, Ord a) =>
LinearSVCModel a -> Map a (Expr Double)
svcMarginExprs LinearSVCModel a
m = [(a, Expr Double)] -> Map a (Expr Double)
forall k a. Ord k => [(k, a)] -> Map k a
M.fromList (LinearSVCModel a -> [(a, Expr Double)]
forall a. LinearSVCModel a -> [(a, Expr Double)]
labelledMargins LinearSVCModel a
m)

labelledMargins :: LinearSVCModel a -> [(a, Expr Double)]
labelledMargins :: forall a. LinearSVCModel a -> [(a, Expr Double)]
labelledMargins LinearSVCModel a
m =
    [ (LinearSVCModel a -> Vector a
forall a. LinearSVCModel a -> Vector a
svcClasses LinearSVCModel a
m Vector a -> Int -> a
forall a. Vector a -> Int -> a
V.! Int
i, LinearModel -> Expr Double
marginOf (LinearSVCModel a -> Vector LinearModel
forall a. LinearSVCModel a -> Vector LinearModel
svcModels LinearSVCModel 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 (LinearSVCModel a -> Vector a
forall a. LinearSVCModel a -> Vector a
svcClasses LinearSVCModel 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)))