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

{- | Approximate RBF-kernel SVM via Random Fourier Features (Rahimi & Recht): map
each row through @z(x) = √(2/D)·cos(W x + b)@ with @W ~ N(0, 2γI)@ (seeded), then
fit a linear SVC in the random-feature space. 'predict' compiles to a closed
@Σ_r β_r·cos(…)@ expression of size @O(D·d)@, independent of the row count.
-}
module DataFrame.SVM.RFF (
    module DataFrame.Model,
    RFFConfig (..),
    defaultRFFConfig,
    RFFSVMModel (..),
) where

import Control.Exception (throw)
import Data.List (sort)
import qualified Data.Text as T
import qualified Data.Vector as V
import qualified Data.Vector.Unboxed as VU
import DataFrame.Errors (DataFrameException (..))

import DataFrame.Featurize.Internal (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.LinearAlgebra (dot)
import DataFrame.LinearSolver (LinearModel (..), SolverConfig (..), fitProx)
import DataFrame.LinearSolver.Loss (sqHingeLoss)
import DataFrame.Model
import DataFrame.Operators ((.*.), (.+.), (.>.))
import DataFrame.Random (Gen, gaussianVector, mkGen, nextDouble)

data RFFConfig = RFFConfig
    { RFFConfig -> Int
rffD :: !Int
    , RFFConfig -> Double
rffGamma :: !Double
    , RFFConfig -> Double
rffC :: !Double
    , RFFConfig -> Int
rffMaxIter :: !Int
    , RFFConfig -> Double
rffTol :: !Double
    , RFFConfig -> Int
rffSeed :: !Int
    }
    deriving (RFFConfig -> RFFConfig -> Bool
(RFFConfig -> RFFConfig -> Bool)
-> (RFFConfig -> RFFConfig -> Bool) -> Eq RFFConfig
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: RFFConfig -> RFFConfig -> Bool
== :: RFFConfig -> RFFConfig -> Bool
$c/= :: RFFConfig -> RFFConfig -> Bool
/= :: RFFConfig -> RFFConfig -> Bool
Eq, Int -> RFFConfig -> ShowS
[RFFConfig] -> ShowS
RFFConfig -> String
(Int -> RFFConfig -> ShowS)
-> (RFFConfig -> String)
-> ([RFFConfig] -> ShowS)
-> Show RFFConfig
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> RFFConfig -> ShowS
showsPrec :: Int -> RFFConfig -> ShowS
$cshow :: RFFConfig -> String
show :: RFFConfig -> String
$cshowList :: [RFFConfig] -> ShowS
showList :: [RFFConfig] -> ShowS
Show)

defaultRFFConfig :: RFFConfig
defaultRFFConfig :: RFFConfig
defaultRFFConfig =
    RFFConfig
        { rffD :: Int
rffD = Int
100
        , rffGamma :: Double
rffGamma = Double
0.1
        , rffC :: Double
rffC = Double
1.0
        , rffMaxIter :: Int
rffMaxIter = Int
1000
        , rffTol :: Double
rffTol = Double
1.0e-4
        , rffSeed :: Int
rffSeed = Int
0
        }

{- | A fitted RFF SVM (binary). 'rffW' / 'rffB' are the random projection;
'rffCoef' / 'rffIntercept' the linear SVC in feature space.
-}
data RFFSVMModel a = RFFSVMModel
    { forall a. RFFSVMModel a -> Vector (Vector Double)
rffW :: !(V.Vector (VU.Vector Double))
    , forall a. RFFSVMModel a -> Vector Double
rffB :: !(VU.Vector Double)
    , forall a. RFFSVMModel a -> Vector Double
rffCoef :: !(VU.Vector Double)
    , forall a. RFFSVMModel a -> Double
rffIntercept :: !Double
    , forall a. RFFSVMModel a -> Double
rffScale :: !Double
    , forall a. RFFSVMModel a -> a
rffNegClass :: !a
    , forall a. RFFSVMModel a -> a
rffPosClass :: !a
    , forall a. RFFSVMModel a -> Vector Text
rffFeatureNames :: !(V.Vector T.Text)
    }
    deriving (Int -> RFFSVMModel a -> ShowS
[RFFSVMModel a] -> ShowS
RFFSVMModel a -> String
(Int -> RFFSVMModel a -> ShowS)
-> (RFFSVMModel a -> String)
-> ([RFFSVMModel a] -> ShowS)
-> Show (RFFSVMModel a)
forall a. Show a => Int -> RFFSVMModel a -> ShowS
forall a. Show a => [RFFSVMModel a] -> ShowS
forall a. Show a => RFFSVMModel a -> String
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: forall a. Show a => Int -> RFFSVMModel a -> ShowS
showsPrec :: Int -> RFFSVMModel a -> ShowS
$cshow :: forall a. Show a => RFFSVMModel a -> String
show :: RFFSVMModel a -> String
$cshowList :: forall a. Show a => [RFFSVMModel a] -> ShowS
showList :: [RFFSVMModel a] -> ShowS
Show)

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

instance (Columnable a) => Predict (RFFSVMModel a) where
    type Prediction (RFFSVMModel a) = Expr a
    predict :: RFFSVMModel a -> Prediction (RFFSVMModel a)
predict RFFSVMModel a
m =
        Expr Bool -> Expr a -> Expr a -> Expr a
forall a. Columnable a => Expr Bool -> Expr a -> Expr a -> Expr a
If (Expr Double
margin Expr Double -> Expr Double -> Expr Bool
forall a. (Columnable a, Ord a) => Expr a -> Expr a -> Expr Bool
.>. Double -> Expr Double
forall a. Columnable a => a -> Expr a
F.lit Double
0) (a -> Expr a
forall a. Columnable a => a -> Expr a
Lit (RFFSVMModel a -> a
forall a. RFFSVMModel a -> a
rffPosClass RFFSVMModel a
m)) (a -> Expr a
forall a. Columnable a => a -> Expr a
Lit (RFFSVMModel a -> a
forall a. RFFSVMModel a -> a
rffNegClass RFFSVMModel a
m))
      where
        names :: [Text]
names = Vector Text -> [Text]
forall a. Vector a -> [a]
V.toList (RFFSVMModel a -> Vector Text
forall a. RFFSVMModel a -> Vector Text
rffFeatureNames RFFSVMModel a
m)
        margin :: Expr Double
margin =
            (Expr Double -> Expr Double -> Expr Double)
-> Expr Double -> [Expr Double] -> Expr Double
forall a b. (a -> b -> b) -> b -> [a] -> b
forall (t :: * -> *) a b.
Foldable t =>
(a -> b -> b) -> b -> t a -> b
foldr 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 (RFFSVMModel a -> Double
forall a. RFFSVMModel a -> Double
rffIntercept RFFSVMModel a
m)) ([Expr Double] -> Expr Double) -> [Expr Double] -> Expr Double
forall a b. (a -> b) -> a -> b
$
                [ Double -> Expr Double
forall a. Columnable a => a -> Expr a
F.lit (RFFSVMModel a -> Vector Double
forall a. RFFSVMModel a -> Vector Double
rffCoef RFFSVMModel a
m Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
r Double -> Double -> Double
forall a. Num a => a -> a -> a
* RFFSVMModel a -> Double
forall a. RFFSVMModel a -> Double
rffScale RFFSVMModel a
m) Expr Double -> Expr Double -> Expr Double
forall a. (Columnable a, Num a) => Expr a -> Expr a -> Expr a
.*. Int -> Expr Double
cosTerm Int
r
                | Int
r <- [Int
0 .. Vector (Vector Double) -> Int
forall a. Vector a -> Int
V.length (RFFSVMModel a -> Vector (Vector Double)
forall a. RFFSVMModel a -> Vector (Vector Double)
rffW RFFSVMModel a
m) Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]
                , RFFSVMModel a -> Vector Double
forall a. RFFSVMModel a -> Vector Double
rffCoef RFFSVMModel a
m Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
r Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
/= Double
0
                ]
        cosTerm :: Int -> Expr Double
cosTerm Int
r = Expr Double -> Expr Double
forall a. Floating a => a -> a
cos (Vector Double -> Double -> Expr Double
linComb (RFFSVMModel a -> Vector (Vector Double)
forall a. RFFSVMModel a -> Vector (Vector Double)
rffW RFFSVMModel a
m Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
r) (RFFSVMModel a -> Vector Double
forall a. RFFSVMModel a -> Vector Double
rffB RFFSVMModel a
m Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
r))
        linComb :: Vector Double -> Double -> Expr Double
linComb Vector Double
w Double
b =
            (Expr Double -> Expr Double -> Expr Double)
-> Expr Double -> [Expr Double] -> Expr Double
forall a b. (a -> b -> b) -> b -> [a] -> b
forall (t :: * -> *) a b.
Foldable t =>
(a -> b -> b) -> b -> t a -> b
foldr 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
b) ([Expr Double] -> Expr Double) -> [Expr Double] -> Expr Double
forall a b. (a -> b) -> a -> b
$
                [ Double -> Expr Double
forall a. Columnable a => a -> Expr a
F.lit (Vector Double
w Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
j) Expr Double -> Expr Double -> Expr Double
forall a. (Columnable a, Num a) => Expr a -> Expr a -> Expr a
.*. (Text -> Expr Double
forall a. Columnable a => Text -> Expr a
Col Text
n :: Expr Double)
                | (Int
j, Text
n) <- [Int] -> [Text] -> [(Int, Text)]
forall a b. [a] -> [b] -> [(a, b)]
zip [Int
0 ..] [Text]
names
                ]

-- | Fit a binary RFF SVM. Targets with more than two classes are rejected.
fitRFFSVM ::
    (Columnable a, Ord a) => RFFConfig -> Expr a -> DataFrame -> RFFSVMModel a
fitRFFSVM :: forall a.
(Columnable a, Ord a) =>
RFFConfig -> Expr a -> DataFrame -> RFFSVMModel a
fitRFFSVM RFFConfig
cfg Expr a
target DataFrame
df =
    case [a]
classes of
        [a
neg, a
pos] -> a -> a -> RFFSVMModel a
build a
neg a
pos
        [a]
_ ->
            DataFrameException -> RFFSVMModel a
forall a e. Exception e => e -> a
throw
                ( Text -> DataFrameException
InternalException
                    Text
"fitRFFSVM: binary classification only, but the target has /= 2 classes"
                )
  where
    names :: [Text]
names = Expr a -> DataFrame -> [Text]
forall a. Expr a -> DataFrame -> [Text]
featureNames Expr a
target DataFrame
df
    (Vector Text
nameVec, Vector (Vector Double)
mat) = [Text] -> DataFrame -> (Vector Text, Vector (Vector Double))
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
    d :: Int
d = if Vector (Vector Double) -> Bool
forall a. Vector a -> Bool
V.null Vector (Vector Double)
mat then Int
0 else Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length (Vector (Vector Double) -> Vector Double
forall a. Vector a -> a
V.head Vector (Vector Double)
mat)
    bigD :: Int
bigD = Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
1 (RFFConfig -> Int
rffD RFFConfig
cfg)
    (Vector (Vector Double)
ws, Vector Double
bs) = Int
-> Int -> Double -> Gen -> (Vector (Vector Double), Vector Double)
sampleRFF Int
bigD Int
d (RFFConfig -> Double
rffGamma RFFConfig
cfg) (Int -> Gen
mkGen (RFFConfig -> Int
rffSeed RFFConfig
cfg))
    scale :: Double
scale = Double -> Double
forall a. Floating a => a -> a
sqrt (Double
2 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
bigD)
    z :: Vector (Vector Double)
z = (Vector Double -> Vector Double)
-> Vector (Vector Double) -> Vector (Vector Double)
forall a b. (a -> b) -> Vector a -> Vector b
V.map (Vector (Vector Double)
-> Vector Double -> Double -> Vector Double -> Vector Double
featureRow Vector (Vector Double)
ws Vector Double
bs Double
scale) Vector (Vector Double)
mat
    featNames :: Vector Text
featNames = [Text] -> Vector Text
forall a. [a] -> Vector a
V.fromList [Text
"rff" Text -> Text -> Text
forall a. Semigroup a => a -> a -> a
<> String -> Text
T.pack (Int -> String
forall a. Show a => a -> String
show Int
r) | Int
r <- [Int
0 .. Int
bigD Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]]
    build :: a -> a -> RFFSVMModel a
build a
neg a
pos =
        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
pos then Double
1 else -Double
1)
            solverCfg :: SolverConfig
solverCfg =
                SolverConfig
                    { scL1Lambda :: Double
scL1Lambda = Double
0
                    , scL2Lambda :: Double
scL2Lambda = Double
1 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ RFFConfig -> Double
rffC RFFConfig
cfg
                    , scMaxIter :: Int
scMaxIter = RFFConfig -> Int
rffMaxIter RFFConfig
cfg
                    , scTol :: Double
scTol = RFFConfig -> Double
rffTol RFFConfig
cfg
                    , scSampleWeights :: Maybe (Vector Double)
scSampleWeights = Maybe (Vector Double)
forall a. Maybe a
Nothing
                    }
            model :: LinearModel
model = SmoothLoss
-> SolverConfig
-> Vector (Vector Double)
-> Vector Double
-> Vector Text
-> LinearModel
fitProx SmoothLoss
sqHingeLoss SolverConfig
solverCfg Vector (Vector Double)
z Vector Double
labels Vector Text
featNames
         in Vector (Vector Double)
-> Vector Double
-> Vector Double
-> Double
-> Double
-> a
-> a
-> Vector Text
-> RFFSVMModel a
forall a.
Vector (Vector Double)
-> Vector Double
-> Vector Double
-> Double
-> Double
-> a
-> a
-> Vector Text
-> RFFSVMModel a
RFFSVMModel Vector (Vector Double)
ws Vector Double
bs (LinearModel -> Vector Double
lmWeights LinearModel
model) (LinearModel -> Double
lmIntercept LinearModel
model) Double
scale a
neg a
pos Vector Text
nameVec

sampleRFF ::
    Int -> Int -> Double -> Gen -> (V.Vector (VU.Vector Double), VU.Vector Double)
sampleRFF :: Int
-> Int -> Double -> Gen -> (Vector (Vector Double), Vector Double)
sampleRFF Int
bigD Int
d Double
gamma Gen
g0 = ([Vector Double] -> Vector (Vector Double)
forall a. [a] -> Vector a
V.fromList [Vector Double]
ws, [Double] -> Vector Double
forall a. Unbox a => [a] -> Vector a
VU.fromList [Double]
bs)
  where
    sigma :: Double
sigma = Double -> Double
forall a. Floating a => a -> a
sqrt (Double
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
gamma)
    ([Vector Double]
ws, Gen
g1) = Int -> Gen -> [Vector Double] -> ([Vector Double], Gen)
goW Int
bigD Gen
g0 []
    goW :: Int -> Gen -> [Vector Double] -> ([Vector Double], Gen)
goW Int
0 Gen
g [Vector Double]
acc = ([Vector Double] -> [Vector Double]
forall a. [a] -> [a]
reverse [Vector Double]
acc, Gen
g)
    goW Int
k Gen
g [Vector Double]
acc =
        let (Vector Double
vec, Gen
g') = Int -> Gen -> (Vector Double, Gen)
gaussianVector Int
d Gen
g
         in Int -> Gen -> [Vector Double] -> ([Vector Double], Gen)
goW (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1) Gen
g' ((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
* Double
sigma) Vector Double
vec Vector Double -> [Vector Double] -> [Vector Double]
forall a. a -> [a] -> [a]
: [Vector Double]
acc)
    bs :: [Double]
bs = Int -> [Double] -> [Double]
forall a. Int -> [a] -> [a]
take Int
bigD (Gen -> [Double]
goB Gen
g1)
    goB :: Gen -> [Double]
goB Gen
g = let (Double
u, Gen
g') = Gen -> (Double, Gen)
nextDouble Gen
g in (Double
u Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
forall a. Floating a => a
pi) Double -> [Double] -> [Double]
forall a. a -> [a] -> [a]
: Gen -> [Double]
goB Gen
g'

featureRow ::
    V.Vector (VU.Vector Double) ->
    VU.Vector Double ->
    Double ->
    VU.Vector Double ->
    VU.Vector Double
featureRow :: Vector (Vector Double)
-> Vector Double -> Double -> Vector Double -> Vector Double
featureRow Vector (Vector Double)
ws Vector Double
bs Double
scale Vector Double
x =
    Int -> (Int -> Double) -> Vector Double
forall a. Unbox a => Int -> (Int -> a) -> Vector a
VU.generate (Vector (Vector Double) -> Int
forall a. Vector a -> Int
V.length Vector (Vector Double)
ws) ((Int -> Double) -> Vector Double)
-> (Int -> Double) -> Vector Double
forall a b. (a -> b) -> a -> b
$ \Int
r ->
        Double
scale Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double -> Double
forall a. Floating a => a -> a
cos (Vector Double -> Vector Double -> Double
dot (Vector (Vector Double)
ws Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
r) Vector Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Vector Double
bs Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
r)