{-# LANGUAGE BangPatterns #-}
{-# LANGUAGE ScopedTypeVariables #-}
module DataFrame.LinearSolver (
LinearModel (..),
SolverConfig (..),
defaultSolverConfig,
fitL1Logistic,
fitProx,
modelToExpr,
standardize,
columnStats,
softThreshold,
sigmoid,
dotProduct,
) where
import qualified DataFrame.Functions as F
import DataFrame.Internal.Expression (Expr (..))
import DataFrame.LinearAlgebra (Matrix, gram, scaleV)
import DataFrame.LinearAlgebra.Eigen (powerIterTop)
import DataFrame.LinearSolver.Loss (
SmoothLoss (..),
logisticLoss,
sigmoid,
)
import DataFrame.Operators ((.*.), (.+.), (.>.))
import Control.Monad.ST (ST, runST)
import qualified Data.Text as T
import qualified Data.Vector as V
import qualified Data.Vector.Unboxed as VU
import qualified Data.Vector.Unboxed.Mutable as VUM
data LinearModel = LinearModel
{ LinearModel -> Vector Double
lmWeights :: !(VU.Vector Double)
, LinearModel -> Double
lmIntercept :: !Double
, LinearModel -> Vector Text
lmFeatureNames :: !(V.Vector T.Text)
}
deriving (LinearModel -> LinearModel -> Bool
(LinearModel -> LinearModel -> Bool)
-> (LinearModel -> LinearModel -> Bool) -> Eq LinearModel
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: LinearModel -> LinearModel -> Bool
== :: LinearModel -> LinearModel -> Bool
$c/= :: LinearModel -> LinearModel -> Bool
/= :: LinearModel -> LinearModel -> Bool
Eq, Int -> LinearModel -> ShowS
[LinearModel] -> ShowS
LinearModel -> String
(Int -> LinearModel -> ShowS)
-> (LinearModel -> String)
-> ([LinearModel] -> ShowS)
-> Show LinearModel
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> LinearModel -> ShowS
showsPrec :: Int -> LinearModel -> ShowS
$cshow :: LinearModel -> String
show :: LinearModel -> String
$cshowList :: [LinearModel] -> ShowS
showList :: [LinearModel] -> ShowS
Show)
data SolverConfig = SolverConfig
{ SolverConfig -> Double
scL1Lambda :: !Double
, SolverConfig -> Double
scL2Lambda :: !Double
, SolverConfig -> Int
scMaxIter :: !Int
, SolverConfig -> Double
scTol :: !Double
, SolverConfig -> Maybe (Vector Double)
scSampleWeights :: !(Maybe (VU.Vector Double))
}
deriving (SolverConfig -> SolverConfig -> Bool
(SolverConfig -> SolverConfig -> Bool)
-> (SolverConfig -> SolverConfig -> Bool) -> Eq SolverConfig
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: SolverConfig -> SolverConfig -> Bool
== :: SolverConfig -> SolverConfig -> Bool
$c/= :: SolverConfig -> SolverConfig -> Bool
/= :: SolverConfig -> SolverConfig -> Bool
Eq, Int -> SolverConfig -> ShowS
[SolverConfig] -> ShowS
SolverConfig -> String
(Int -> SolverConfig -> ShowS)
-> (SolverConfig -> String)
-> ([SolverConfig] -> ShowS)
-> Show SolverConfig
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> SolverConfig -> ShowS
showsPrec :: Int -> SolverConfig -> ShowS
$cshow :: SolverConfig -> String
show :: SolverConfig -> String
$cshowList :: [SolverConfig] -> ShowS
showList :: [SolverConfig] -> ShowS
Show)
defaultSolverConfig :: SolverConfig
defaultSolverConfig :: SolverConfig
defaultSolverConfig =
SolverConfig
{ scL1Lambda :: Double
scL1Lambda = Double
0.005
, scL2Lambda :: Double
scL2Lambda = Double
0.005
, scMaxIter :: Int
scMaxIter = Int
200
, scTol :: Double
scTol = Double
1.0e-4
, scSampleWeights :: Maybe (Vector Double)
scSampleWeights = Maybe (Vector Double)
forall a. Maybe a
Nothing
}
fitL1Logistic ::
SolverConfig ->
V.Vector (VU.Vector Double) ->
VU.Vector Double ->
V.Vector T.Text ->
LinearModel
{-# INLINEABLE fitL1Logistic #-}
fitL1Logistic :: SolverConfig
-> Vector (Vector Double)
-> Vector Double
-> Vector Text
-> LinearModel
fitL1Logistic = SmoothLoss
-> (Vector (Vector Double) -> Int -> Double)
-> SolverConfig
-> Vector (Vector Double)
-> Vector Double
-> Vector Text
-> LinearModel
runFista SmoothLoss
logisticLoss (SmoothLoss -> Vector (Vector Double) -> Int -> Double
specNormLipschitz SmoothLoss
logisticLoss)
fitProx ::
SmoothLoss ->
SolverConfig ->
V.Vector (VU.Vector Double) ->
VU.Vector Double ->
V.Vector T.Text ->
LinearModel
fitProx :: SmoothLoss
-> SolverConfig
-> Vector (Vector Double)
-> Vector Double
-> Vector Text
-> LinearModel
fitProx SmoothLoss
loss = SmoothLoss
-> (Vector (Vector Double) -> Int -> Double)
-> SolverConfig
-> Vector (Vector Double)
-> Vector Double
-> Vector Text
-> LinearModel
runFista SmoothLoss
loss (SmoothLoss -> Vector (Vector Double) -> Int -> Double
specNormLipschitz SmoothLoss
loss)
specNormLipschitz :: SmoothLoss -> Matrix -> Int -> Double
specNormLipschitz :: SmoothLoss -> Vector (Vector Double) -> Int -> Double
specNormLipschitz SmoothLoss
loss Vector (Vector Double)
xKept Int
_ =
let n :: Int
n = Vector (Vector Double) -> Int
forall a. Vector a -> Int
V.length Vector (Vector Double)
xKept
gramN :: Vector (Vector Double)
gramN = (Vector Double -> Vector Double)
-> Vector (Vector Double) -> Vector (Vector Double)
forall a b. (a -> b) -> Vector a -> Vector b
V.map (Double -> Vector Double -> Vector Double
scaleV (Double
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)) (Vector (Vector Double) -> Vector (Vector Double)
gram Vector (Vector Double)
xKept)
(Double
specNorm, Vector Double
_) = Int -> Vector (Vector Double) -> (Double, Vector Double)
powerIterTop Int
50 Vector (Vector Double)
gramN
in SmoothLoss -> Double
slCurvBound SmoothLoss
loss Double -> Double -> Double
forall a. Num a => a -> a -> a
* (Double
specNorm Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
1)
runFista ::
SmoothLoss ->
(Matrix -> Int -> Double) ->
SolverConfig ->
V.Vector (VU.Vector Double) ->
VU.Vector Double ->
V.Vector T.Text ->
LinearModel
runFista :: SmoothLoss
-> (Vector (Vector Double) -> Int -> Double)
-> SolverConfig
-> Vector (Vector Double)
-> Vector Double
-> Vector Text
-> LinearModel
runFista SmoothLoss
loss Vector (Vector Double) -> Int -> Double
lipschitzOf SolverConfig
cfg Vector (Vector Double)
rows Vector Double
labels Vector Text
featureNames
| Int
n Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
0 Bool -> Bool -> Bool
|| Int
d Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
0 = LinearModel
zeroModel
| Bool
otherwise =
let (!Vector Double
means, !Vector Double
stds, !Vector Double
variances) = Vector (Vector Double)
-> (Vector Double, Vector Double, Vector Double)
columnStats Vector (Vector Double)
rows
!keep :: Vector Int
keep = Vector Double -> Vector Int
keptIndices Vector Double
variances
in if Vector Int -> Bool
forall a. Unbox a => Vector a -> Bool
VU.null Vector Int
keep
then LinearModel
zeroModel
else
let !meansKept :: Vector Double
meansKept = Vector Int -> Vector Double -> Vector Double
gatherBy Vector Int
keep Vector Double
means
!stdsKept :: Vector Double
stdsKept = Vector Int -> Vector Double -> Vector Double
gatherBy Vector Int
keep Vector Double
stds
!xKept :: Vector (Vector Double)
xKept = (Vector Double -> Vector Double)
-> Vector (Vector Double) -> Vector (Vector Double)
forall a b. (a -> b) -> Vector a -> Vector b
V.map (Vector Int
-> Vector Double -> Vector Double -> Vector Double -> Vector Double
standardizeRowKept Vector Int
keep Vector Double
means Vector Double
stds) Vector (Vector Double)
rows
!lipschitz :: Double
lipschitz =
Vector (Vector Double) -> Int -> Double
lipschitzOf Vector (Vector Double)
xKept (Vector Int -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Int
keep) Double -> Double -> Double
forall a. Num a => a -> a -> a
+ SolverConfig -> Double
scL2Lambda SolverConfig
cfg
(!Vector Double
wStdKept, !Double
bStd) =
SmoothLoss
-> Double
-> Double
-> Double
-> Int
-> Double
-> Maybe (Vector Double)
-> Vector (Vector Double)
-> Vector Double
-> Vector Double
-> Double
-> (Vector Double, Double)
fistaLoop
SmoothLoss
loss
(SolverConfig -> Double
scL1Lambda SolverConfig
cfg)
(SolverConfig -> Double
scL2Lambda SolverConfig
cfg)
Double
lipschitz
(SolverConfig -> Int
scMaxIter SolverConfig
cfg)
(SolverConfig -> Double
scTol SolverConfig
cfg)
(SolverConfig -> Maybe (Vector Double)
scSampleWeights SolverConfig
cfg)
Vector (Vector Double)
xKept
Vector Double
labels
(Int -> Double -> Vector Double
forall a. Unbox a => Int -> a -> Vector a
VU.replicate (Vector Int -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Int
keep) Double
0)
Double
0
!wRawKept :: Vector Double
wRawKept = (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 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
(/) Vector Double
wStdKept Vector Double
stdsKept
!bRaw :: Double
bRaw = Double
bStd Double -> Double -> Double
forall a. Num a => a -> a -> a
- Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum ((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 Double -> Double -> Double
forall a. Num a => a -> a -> a
(*) Vector Double
wRawKept Vector Double
meansKept)
in Vector Double -> Double -> Vector Text -> LinearModel
LinearModel (Int -> Vector Int -> Vector Double -> Vector Double
expandWeights Int
d Vector Int
keep Vector Double
wRawKept) Double
bRaw Vector Text
featureNames
where
!n :: Int
n = Vector (Vector Double) -> Int
forall a. Vector a -> Int
V.length Vector (Vector Double)
rows
!d :: Int
d = Vector Text -> Int
forall a. Vector a -> Int
V.length Vector Text
featureNames
zeroModel :: LinearModel
zeroModel = Vector Double -> Double -> Vector Text -> LinearModel
LinearModel (Int -> Double -> Vector Double
forall a. Unbox a => Int -> a -> Vector a
VU.replicate Int
d Double
0) Double
0 Vector Text
featureNames
keptIndices :: VU.Vector Double -> VU.Vector Int
keptIndices :: Vector Double -> Vector Int
keptIndices Vector Double
variances =
[Int] -> Vector Int
forall a. Unbox a => [a] -> Vector a
VU.fromList
[ Int
j
| Int
j <- [Int
0 .. Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Double
variances Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]
, Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Double
variances Int
j Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
>= Double
1.0e-12
]
gatherBy :: VU.Vector Int -> VU.Vector Double -> VU.Vector Double
gatherBy :: Vector Int -> Vector Double -> Vector Double
gatherBy Vector Int
idxs Vector Double
v = (Int -> Double) -> Vector Int -> Vector Double
forall a b. (Unbox a, Unbox b) => (a -> b) -> Vector a -> Vector b
VU.map (Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Double
v) Vector Int
idxs
standardizeRowKept ::
VU.Vector Int ->
VU.Vector Double ->
VU.Vector Double ->
VU.Vector Double ->
VU.Vector Double
standardizeRowKept :: Vector Int
-> Vector Double -> Vector Double -> Vector Double -> Vector Double
standardizeRowKept Vector Int
keep Vector Double
means Vector Double
stds Vector Double
row = (Int -> Double) -> Vector Int -> Vector Double
forall a b. (Unbox a, Unbox b) => (a -> b) -> Vector a -> Vector b
VU.map Int -> Double
standardizeAt Vector Int
keep
where
standardizeAt :: Int -> Double
standardizeAt Int
j =
(Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Double
row Int
j Double -> Double -> Double
forall a. Num a => a -> a -> a
- Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Double
means Int
j) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Double
stds Int
j
expandWeights :: Int -> VU.Vector Int -> VU.Vector Double -> VU.Vector Double
expandWeights :: Int -> Vector Int -> Vector Double -> Vector Double
expandWeights Int
d Vector Int
keep Vector Double
wKept = (forall s. ST s (MVector s Double)) -> Vector Double
forall a. Unbox a => (forall s. ST s (MVector s a)) -> Vector a
VU.create ((forall s. ST s (MVector s Double)) -> Vector Double)
-> (forall s. ST s (MVector s Double)) -> Vector Double
forall a b. (a -> b) -> a -> b
$ do
MVector s Double
mv <- Int -> Double -> ST s (MVector (PrimState (ST s)) Double)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
d Double
0
Vector Int -> (Int -> Int -> ST s ()) -> ST s ()
forall (m :: * -> *) a b.
(Monad m, Unbox a) =>
Vector a -> (Int -> a -> m b) -> m ()
VU.iforM_ Vector Int
keep ((Int -> Int -> ST s ()) -> ST s ())
-> (Int -> Int -> ST s ()) -> ST s ()
forall a b. (a -> b) -> a -> b
$ \Int
k Int
j -> MVector (PrimState (ST s)) Double -> Int -> Double -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.unsafeWrite MVector s Double
MVector (PrimState (ST s)) Double
mv Int
j (Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Double
wKept Int
k)
MVector s Double -> ST s (MVector s Double)
forall a. a -> ST s a
forall (f :: * -> *) a. Applicative f => a -> f a
pure MVector s Double
mv
modelToExpr :: LinearModel -> Expr Bool
modelToExpr :: LinearModel -> Expr Bool
modelToExpr LinearModel
m =
case [(Double, Text)]
nonZero of
[] -> Bool -> Expr Bool
forall a. Columnable a => a -> Expr a
F.lit (Double
b Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
> Double
0)
(Double
w0, Text
n0) : [(Double, Text)]
rest -> [(Double, Text)] -> Expr Double -> Expr Double
forall {t :: * -> *}.
Foldable t =>
t (Double, Text) -> Expr Double -> Expr Double
score [(Double, Text)]
rest (Double -> Text -> Expr Double
term Double
w0 Text
n0) 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 :: Double)
where
b :: Double
b = LinearModel -> Double
lmIntercept LinearModel
m
nonZero :: [(Double, Text)]
nonZero =
[ (Double
w, Text
n)
| (Double
w, Text
n) <- [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))
, Double
w Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
/= Double
0
]
term :: Double -> Text -> Expr Double
term Double
w Text
n = Double -> Expr Double
forall a. Columnable a => a -> Expr a
F.lit Double
w 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)
score :: t (Double, Text) -> Expr Double -> Expr Double
score t (Double, Text)
rest Expr Double
first = (Expr Double -> (Double, Text) -> Expr Double)
-> Expr Double -> t (Double, Text) -> Expr Double
forall b a. (b -> a -> b) -> b -> t a -> b
forall (t :: * -> *) b a.
Foldable t =>
(b -> a -> b) -> b -> t a -> b
foldl (\Expr Double
acc (Double
w, Text
n) -> Expr Double
acc Expr Double -> Expr Double -> Expr Double
forall a. (Columnable a, Num a) => Expr a -> Expr a -> Expr a
.+. Double -> Text -> Expr Double
term Double
w Text
n) Expr Double
first t (Double, Text)
rest 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
columnStats ::
V.Vector (VU.Vector Double) ->
(VU.Vector Double, VU.Vector Double, VU.Vector Double)
columnStats :: Vector (Vector Double)
-> (Vector Double, Vector Double, Vector Double)
columnStats Vector (Vector Double)
x
| Vector (Vector Double) -> Bool
forall a. Vector a -> Bool
V.null Vector (Vector Double)
x = (Vector Double
forall a. Unbox a => Vector a
VU.empty, Vector Double
forall a. Unbox a => Vector a
VU.empty, Vector Double
forall a. Unbox a => Vector a
VU.empty)
| Bool
otherwise =
let !d :: Int
d = Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length (Vector (Vector Double) -> Vector Double
forall a. Vector a -> a
V.unsafeHead Vector (Vector Double)
x)
!invN :: Double
invN = Double
1 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Vector (Vector Double) -> Int
forall a. Vector a -> Int
V.length Vector (Vector Double)
x)
!means :: Vector Double
means = Int -> Double -> Vector (Vector Double) -> Vector Double
columnMeans Int
d Double
invN Vector (Vector Double)
x
!variances :: Vector Double
variances = Int
-> Double
-> Vector Double
-> Vector (Vector Double)
-> Vector Double
columnVariances Int
d Double
invN Vector Double
means Vector (Vector Double)
x
!stds :: Vector Double
stds = (Double -> Double) -> Vector Double -> Vector Double
forall a b. (Unbox a, Unbox b) => (a -> b) -> Vector a -> Vector b
VU.map (\Double
v -> if Double
v Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
< Double
1e-12 then Double
1 else Double -> Double
forall a. Floating a => a -> a
sqrt Double
v) Vector Double
variances
in (Vector Double
means, Vector Double
stds, Vector Double
variances)
columnMeans :: Int -> Double -> V.Vector (VU.Vector Double) -> VU.Vector Double
columnMeans :: Int -> Double -> Vector (Vector Double) -> Vector Double
columnMeans Int
d Double
invN Vector (Vector Double)
x = (forall s. ST s (Vector Double)) -> Vector Double
forall a. (forall s. ST s a) -> a
runST ((forall s. ST s (Vector Double)) -> Vector Double)
-> (forall s. ST s (Vector Double)) -> Vector Double
forall a b. (a -> b) -> a -> b
$ do
MVector s Double
acc <- Int -> Double -> ST s (MVector (PrimState (ST s)) Double)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
d Double
0
Vector (Vector Double) -> (Vector Double -> ST s ()) -> ST s ()
forall (m :: * -> *) a b. Monad m => Vector a -> (a -> m b) -> m ()
V.forM_ Vector (Vector Double)
x ((Vector Double -> ST s ()) -> ST s ())
-> (Vector Double -> ST s ()) -> ST s ()
forall a b. (a -> b) -> a -> b
$ \Vector Double
row ->
Vector Double -> (Int -> Double -> ST s ()) -> ST s ()
forall (m :: * -> *) a b.
(Monad m, Unbox a) =>
Vector a -> (Int -> a -> m b) -> m ()
VU.iforM_ Vector Double
row ((Int -> Double -> ST s ()) -> ST s ())
-> (Int -> Double -> ST s ()) -> ST s ()
forall a b. (a -> b) -> a -> b
$ \Int
j Double
v -> MVector (PrimState (ST s)) Double
-> (Double -> Double) -> Int -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> (a -> a) -> Int -> m ()
VUM.unsafeModify MVector s Double
MVector (PrimState (ST s)) Double
acc (Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
v) Int
j
Double -> MVector s Double -> ST s ()
forall s. Double -> MVector s Double -> ST s ()
scaleInPlace Double
invN MVector s Double
acc
MVector (PrimState (ST s)) Double -> ST s (Vector Double)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze MVector s Double
MVector (PrimState (ST s)) Double
acc
columnVariances ::
Int ->
Double ->
VU.Vector Double ->
V.Vector (VU.Vector Double) ->
VU.Vector Double
columnVariances :: Int
-> Double
-> Vector Double
-> Vector (Vector Double)
-> Vector Double
columnVariances Int
d Double
invN Vector Double
means Vector (Vector Double)
x = (forall s. ST s (Vector Double)) -> Vector Double
forall a. (forall s. ST s a) -> a
runST ((forall s. ST s (Vector Double)) -> Vector Double)
-> (forall s. ST s (Vector Double)) -> Vector Double
forall a b. (a -> b) -> a -> b
$ do
MVector s Double
acc <- Int -> Double -> ST s (MVector (PrimState (ST s)) Double)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
d Double
0
Vector (Vector Double) -> (Vector Double -> ST s ()) -> ST s ()
forall (m :: * -> *) a b. Monad m => Vector a -> (a -> m b) -> m ()
V.forM_ Vector (Vector Double)
x ((Vector Double -> ST s ()) -> ST s ())
-> (Vector Double -> ST s ()) -> ST s ()
forall a b. (a -> b) -> a -> b
$ \Vector Double
row ->
Vector Double -> (Int -> Double -> ST s ()) -> ST s ()
forall (m :: * -> *) a b.
(Monad m, Unbox a) =>
Vector a -> (Int -> a -> m b) -> m ()
VU.iforM_ Vector Double
row ((Int -> Double -> ST s ()) -> ST s ())
-> (Int -> Double -> ST s ()) -> ST s ()
forall a b. (a -> b) -> a -> b
$ \Int
j Double
v ->
let !c :: Double
c = Double
v Double -> Double -> Double
forall a. Num a => a -> a -> a
- Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Double
means Int
j in MVector (PrimState (ST s)) Double
-> (Double -> Double) -> Int -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> (a -> a) -> Int -> m ()
VUM.unsafeModify MVector s Double
MVector (PrimState (ST s)) Double
acc (Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
c Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
c) Int
j
Double -> MVector s Double -> ST s ()
forall s. Double -> MVector s Double -> ST s ()
scaleInPlace Double
invN MVector s Double
acc
MVector (PrimState (ST s)) Double -> ST s (Vector Double)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze MVector s Double
MVector (PrimState (ST s)) Double
acc
scaleInPlace :: Double -> VUM.MVector s Double -> ST s ()
scaleInPlace :: forall s. Double -> MVector s Double -> ST s ()
scaleInPlace Double
factor MVector s Double
mv = Int -> ST s ()
forall {f :: * -> *}. (PrimState f ~ s, PrimMonad f) => Int -> f ()
go Int
0
where
go :: Int -> f ()
go !Int
j
| Int
j Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= MVector s Double -> Int
forall a s. Unbox a => MVector s a -> Int
VUM.length MVector s Double
mv = () -> f ()
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ()
| Bool
otherwise = MVector (PrimState f) Double -> (Double -> Double) -> Int -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> (a -> a) -> Int -> m ()
VUM.unsafeModify MVector s Double
MVector (PrimState f) Double
mv (Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
factor) Int
j f () -> f () -> f ()
forall a b. f a -> f b -> f b
forall (m :: * -> *) a b. Monad m => m a -> m b -> m b
>> Int -> f ()
go (Int
j Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
standardize ::
V.Vector (VU.Vector Double) ->
( V.Vector (VU.Vector Double)
, VU.Vector Double
, VU.Vector Double
, VU.Vector Double
)
standardize :: Vector (Vector Double)
-> (Vector (Vector Double), Vector Double, Vector Double,
Vector Double)
standardize Vector (Vector Double)
x
| Vector (Vector Double) -> Bool
forall a. Vector a -> Bool
V.null Vector (Vector Double)
x = (Vector (Vector Double)
x, Vector Double
forall a. Unbox a => Vector a
VU.empty, Vector Double
forall a. Unbox a => Vector a
VU.empty, Vector Double
forall a. Unbox a => Vector a
VU.empty)
| Bool
otherwise =
let (!Vector Double
means, !Vector Double
stds, !Vector Double
variances) = Vector (Vector Double)
-> (Vector Double, Vector Double, Vector Double)
columnStats Vector (Vector Double)
x
!d :: Int
d = Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length (Vector (Vector Double) -> Vector Double
forall a. Vector a -> a
V.unsafeHead Vector (Vector Double)
x)
standardizeRow :: Vector Double -> Vector Double
standardizeRow Vector Double
row =
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 ->
(Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Double
row Int
j Double -> Double -> Double
forall a. Num a => a -> a -> a
- Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Double
means Int
j) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Double
stds Int
j
in ((Vector Double -> Vector Double)
-> Vector (Vector Double) -> Vector (Vector Double)
forall a b. (a -> b) -> Vector a -> Vector b
V.map Vector Double -> Vector Double
standardizeRow Vector (Vector Double)
x, Vector Double
means, Vector Double
stds, Vector Double
variances)
softThreshold :: Double -> Double -> Double
softThreshold :: Double -> Double -> Double
softThreshold Double
lambda Double
v
| Double
v Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
> Double
lambda = Double
v Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
lambda
| Double
v Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
< -Double
lambda = Double
v Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
lambda
| Bool
otherwise = Double
0
dotProduct :: VU.Vector Double -> VU.Vector Double -> Double
dotProduct :: Vector Double -> Vector Double -> Double
dotProduct Vector Double
u Vector Double
v = Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum ((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 Double -> Double -> Double
forall a. Num a => a -> a -> a
(*) Vector Double
u Vector Double
v)
lossGradient ::
SmoothLoss ->
Maybe (VU.Vector Double) ->
V.Vector (VU.Vector Double) ->
VU.Vector Double ->
VU.Vector Double ->
Double ->
(VU.Vector Double, Double)
lossGradient :: SmoothLoss
-> Maybe (Vector Double)
-> Vector (Vector Double)
-> Vector Double
-> Vector Double
-> Double
-> (Vector Double, Double)
lossGradient SmoothLoss
loss Maybe (Vector Double)
sampleWeights Vector (Vector Double)
features Vector Double
labels Vector Double
w Double
b = (Vector Double
gradW, Double
gradB)
where
!invN :: Double
invN = Double
1 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Vector (Vector Double) -> Int
forall a. Vector a -> Int
V.length Vector (Vector Double)
features)
!coeffs :: Vector Double
coeffs = SmoothLoss
-> Maybe (Vector Double)
-> Vector (Vector Double)
-> Vector Double
-> Vector Double
-> Double
-> Double
-> Vector Double
rowCoeffs SmoothLoss
loss Maybe (Vector Double)
sampleWeights Vector (Vector Double)
features Vector Double
labels Vector Double
w Double
b Double
invN
!gradW :: Vector Double
gradW = Int -> Vector (Vector Double) -> Vector Double -> Vector Double
accumulateGradW (Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Double
w) Vector (Vector Double)
features Vector Double
coeffs
!gradB :: Double
gradB = Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum Vector Double
coeffs
rowCoeffs ::
SmoothLoss ->
Maybe (VU.Vector Double) ->
V.Vector (VU.Vector Double) ->
VU.Vector Double ->
VU.Vector Double ->
Double ->
Double ->
VU.Vector Double
rowCoeffs :: SmoothLoss
-> Maybe (Vector Double)
-> Vector (Vector Double)
-> Vector Double
-> Vector Double
-> Double
-> Double
-> Vector Double
rowCoeffs SmoothLoss
loss Maybe (Vector Double)
sampleWeights Vector (Vector Double)
features Vector Double
labels Vector Double
w Double
b Double
invN =
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)
features) ((Int -> Double) -> Vector Double)
-> (Int -> Double) -> Vector Double
forall a b. (a -> b) -> a -> b
$ \Int
i ->
let !yi :: Double
yi = Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Double
labels Int
i
!row :: Vector Double
row = Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.unsafeIndex Vector (Vector Double)
features Int
i
!z :: Double
z = Vector Double -> Vector Double -> Double
dotProduct Vector Double
w Vector Double
row Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
b
!base :: Double
base = SmoothLoss -> Double -> Double -> Double
slGradZ SmoothLoss
loss Double
yi Double
z Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
invN
in case Maybe (Vector Double)
sampleWeights of
Maybe (Vector Double)
Nothing -> Double
base
Just Vector Double
ws -> Double
base Double -> Double -> Double
forall a. Num a => a -> a -> a
* Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Double
ws Int
i
accumulateGradW ::
Int -> V.Vector (VU.Vector Double) -> VU.Vector Double -> VU.Vector Double
accumulateGradW :: Int -> Vector (Vector Double) -> Vector Double -> Vector Double
accumulateGradW Int
d Vector (Vector Double)
features Vector Double
coeffs = (forall s. ST s (Vector Double)) -> Vector Double
forall a. (forall s. ST s a) -> a
runST ((forall s. ST s (Vector Double)) -> Vector Double)
-> (forall s. ST s (Vector Double)) -> Vector Double
forall a b. (a -> b) -> a -> b
$ do
MVector s Double
mv <- Int -> Double -> ST s (MVector (PrimState (ST s)) Double)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> a -> m (MVector (PrimState m) a)
VUM.replicate Int
d Double
0
Vector (Vector Double)
-> (Int -> Vector Double -> ST s ()) -> ST s ()
forall (m :: * -> *) a b.
Monad m =>
Vector a -> (Int -> a -> m b) -> m ()
V.iforM_ Vector (Vector Double)
features ((Int -> Vector Double -> ST s ()) -> ST s ())
-> (Int -> Vector Double -> ST s ()) -> ST s ()
forall a b. (a -> b) -> a -> b
$ \Int
i Vector Double
row ->
let !c :: Double
c = Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Double
coeffs Int
i
in Vector Double -> (Int -> Double -> ST s ()) -> ST s ()
forall (m :: * -> *) a b.
(Monad m, Unbox a) =>
Vector a -> (Int -> a -> m b) -> m ()
VU.iforM_ Vector Double
row ((Int -> Double -> ST s ()) -> ST s ())
-> (Int -> Double -> ST s ()) -> ST s ()
forall a b. (a -> b) -> a -> b
$ \Int
j Double
v -> MVector (PrimState (ST s)) Double
-> (Double -> Double) -> Int -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> (a -> a) -> Int -> m ()
VUM.unsafeModify MVector s Double
MVector (PrimState (ST s)) Double
mv (Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
c Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
v) Int
j
MVector (PrimState (ST s)) Double -> ST s (Vector Double)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.unsafeFreeze MVector s Double
MVector (PrimState (ST s)) Double
mv
fistaLoop ::
SmoothLoss ->
Double ->
Double ->
Double ->
Int ->
Double ->
Maybe (VU.Vector Double) ->
V.Vector (VU.Vector Double) ->
VU.Vector Double ->
VU.Vector Double ->
Double ->
(VU.Vector Double, Double)
fistaLoop :: SmoothLoss
-> Double
-> Double
-> Double
-> Int
-> Double
-> Maybe (Vector Double)
-> Vector (Vector Double)
-> Vector Double
-> Vector Double
-> Double
-> (Vector Double, Double)
fistaLoop SmoothLoss
loss Double
lambda1 Double
lambda2 Double
lp Int
maxIter Double
tol Maybe (Vector Double)
sampleWeights Vector (Vector Double)
features Vector Double
labels Vector Double
w0 Double
b0 =
Int
-> Vector Double
-> Double
-> Vector Double
-> Double
-> Double
-> (Vector Double, Double)
go Int
0 Vector Double
w0 Double
b0 Vector Double
w0 Double
b0 Double
1.0
where
!shrink :: Double
shrink = Double
lambda1 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
lp
!ridgeDenom :: Double
ridgeDenom = Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
lambda2 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
lp
!stepInv :: Double
stepInv = Double
1 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
lp
proxStep :: Vector Double -> Double -> (Vector Double, Double)
proxStep = SmoothLoss
-> Maybe (Vector Double)
-> Vector (Vector Double)
-> Vector Double
-> Double
-> Double
-> Double
-> Vector Double
-> Double
-> (Vector Double, Double)
fistaProxStep SmoothLoss
loss Maybe (Vector Double)
sampleWeights Vector (Vector Double)
features Vector Double
labels Double
shrink Double
ridgeDenom Double
stepInv
go :: Int
-> Vector Double
-> Double
-> Vector Double
-> Double
-> Double
-> (Vector Double, Double)
go !Int
iter !Vector Double
xWPrev !Double
xBPrev !Vector Double
yW !Double
yB !Double
t
| Int
iter Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
maxIter = (Vector Double
xWPrev, Double
xBPrev)
| Int
iter Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Int
0 Bool -> Bool -> Bool
&& Double
delta Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
< Double
tol = (Vector Double
xW, Double
xB)
| Bool
otherwise = Int
-> Vector Double
-> Double
-> Vector Double
-> Double
-> Double
-> (Vector Double, Double)
go (Int
iter Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Vector Double
xW Double
xB Vector Double
yWNew Double
yBNew Double
tNew
where
(!Vector Double
xW, !Double
xB) = Vector Double -> Double -> (Vector Double, Double)
proxStep Vector Double
yW Double
yB
!delta :: Double
delta = if Vector Double -> Bool
forall a. Unbox a => Vector a -> Bool
VU.null Vector Double
xW then Double
0 else Vector Double -> Vector Double -> Double
deltaInf Vector Double
xWPrev Vector Double
xW
(!Vector Double
yWNew, !Double
yBNew, !Double
tNew) = Double
-> Vector Double
-> Double
-> Vector Double
-> Double
-> (Vector Double, Double, Double)
fistaMomentum Double
t Vector Double
xWPrev Double
xBPrev Vector Double
xW Double
xB
fistaProxStep ::
SmoothLoss ->
Maybe (VU.Vector Double) ->
V.Vector (VU.Vector Double) ->
VU.Vector Double ->
Double ->
Double ->
Double ->
VU.Vector Double ->
Double ->
(VU.Vector Double, Double)
fistaProxStep :: SmoothLoss
-> Maybe (Vector Double)
-> Vector (Vector Double)
-> Vector Double
-> Double
-> Double
-> Double
-> Vector Double
-> Double
-> (Vector Double, Double)
fistaProxStep SmoothLoss
loss Maybe (Vector Double)
sampleWeights Vector (Vector Double)
features Vector Double
labels Double
shrink Double
ridgeDenom Double
stepInv Vector Double
yW Double
yB =
let (Vector Double
gW, Double
gB) = SmoothLoss
-> Maybe (Vector Double)
-> Vector (Vector Double)
-> Vector Double
-> Vector Double
-> Double
-> (Vector Double, Double)
lossGradient SmoothLoss
loss Maybe (Vector Double)
sampleWeights Vector (Vector Double)
features Vector Double
labels Vector Double
yW Double
yB
!wNew :: Vector Double
wNew =
(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
(\Double
yi Double
gi -> Double -> Double -> Double
softThreshold Double
shrink (Double
yi Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
gi Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
stepInv) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
ridgeDenom)
Vector Double
yW
Vector Double
gW
!bNew :: Double
bNew = Double
yB Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
gB Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
stepInv
in (Vector Double
wNew, Double
bNew)
fistaMomentum ::
Double ->
VU.Vector Double ->
Double ->
VU.Vector Double ->
Double ->
(VU.Vector Double, Double, Double)
fistaMomentum :: Double
-> Vector Double
-> Double
-> Vector Double
-> Double
-> (Vector Double, Double, Double)
fistaMomentum Double
t Vector Double
xWPrev Double
xBPrev Vector Double
xW Double
xB =
let !tNew :: Double
tNew = (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double -> Double
forall a. Floating a => a -> a
sqrt (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
4 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
t Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
t)) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
2
!mom :: Double
mom = (Double
t Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
1) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
tNew
!yW :: Vector Double
yW = (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 (\Double
new Double
old -> Double
new Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
mom Double -> Double -> Double
forall a. Num a => a -> a -> a
* (Double
new Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
old)) Vector Double
xW Vector Double
xWPrev
!yB :: Double
yB = Double
xB Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
mom Double -> Double -> Double
forall a. Num a => a -> a -> a
* (Double
xB Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
xBPrev)
in (Vector Double
yW, Double
yB, Double
tNew)
{-# INLINE deltaInf #-}
deltaInf :: VU.Vector Double -> VU.Vector Double -> Double
deltaInf :: Vector Double -> Vector Double -> Double
deltaInf Vector Double
xWPrev = (Double -> Int -> Double -> Double)
-> Double -> Vector Double -> Double
forall b a. Unbox b => (a -> Int -> b -> a) -> a -> Vector b -> a
VU.ifoldl' (\Double
acc Int
i Double
x -> Double -> Double -> Double
forall a. Ord a => a -> a -> a
max Double
acc (Double -> Double
forall a. Num a => a -> a
abs (Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
- Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Double
xWPrev Int
i))) Double
0