{-# LANGUAGE BangPatterns #-}
{-# LANGUAGE ScopedTypeVariables #-}

{- | Proximal-gradient (FISTA) solver for L1/L2-regularized generalized linear
models. 'fitL1Logistic' is the binary logistic split solver; 'fitProx'
generalizes it to any 'SmoothLoss'. Features are standardized internally.
-}
module DataFrame.LinearSolver (
    -- * Model
    LinearModel (..),

    -- * Configuration
    SolverConfig (..),
    defaultSolverConfig,

    -- * Solvers
    fitL1Logistic,
    fitProx,

    -- * Expr conversion
    modelToExpr,

    -- * Internals (exposed for testing)
    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

{- | A fitted linear classifier: predicts the positive class when
@sum (weights .* features) + intercept > 0@. Weights of exactly @0@ mark
features dropped by the L1 penalty (filtered out by 'modelToExpr').
-}
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)

-- | Hyper-parameters for the FISTA solver.
data SolverConfig = SolverConfig
    { SolverConfig -> Double
scL1Lambda :: !Double
    -- ^ Strength of the L1 penalty on weights (intercept is not regularized).
    , SolverConfig -> Double
scL2Lambda :: !Double
    {- ^ Strength of the L2 penalty @(λ₂/2)·|w|²@ (Elastic Net; Zou & Hastie
    2005). Combined with @scL1Lambda@ this is the elastic-net objective;
    @0@ reduces the solver to pure L1.
    -}
    , SolverConfig -> Int
scMaxIter :: !Int
    -- ^ Maximum number of FISTA iterations.
    , SolverConfig -> Double
scTol :: !Double
    -- ^ Convergence tolerance on the weight delta (L-inf norm).
    , SolverConfig -> Maybe (Vector Double)
scSampleWeights :: !(Maybe (VU.Vector Double))
    {- ^ Optional per-row sample weights, length @n@ (@Nothing@ is uniform).
    Weights should have mean 1 (i.e. @Σ w_i = N@) so the Lipschitz bound stays
    valid; see 'fitLinearCandidate' for the class-balanced construction.
    -}
    }
    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
        }

{- | Fit L1-regularized binary logistic regression by FISTA. Rows are feature
vectors of equal length; labels are in @{\-1,+1}@. Features are standardized
internally and weights de-standardized, so the model applies to raw values.
-}
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)

{- | Fit any 'SmoothLoss' with the elastic-net proximal-gradient engine. The
Lipschitz constant uses the spectral norm of the standardized Gram matrix
(power iteration), tight for squared and squared-hinge losses.
-}
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)

{- | FISTA smooth-part Lipschitz bound: the loss curvature bound times the
spectral norm of the standardized Gram (+1 for the intercept), via power
iteration. Tight for every smooth 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)

{- | Shared FISTA scaffolding: standardize, drop near-constant columns, run the
inner loop, de-standardize. @lipschitzOf@ returns the smooth-part Lipschitz
bound from the kept-feature matrix; the L2 contribution @λ₂@ is added here.
-}
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

{- | Indices of columns whose variance clears the near-constant threshold.
Columns below it are dropped before fitting; their weight ends up @0@.
-}
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
        ]

{- | Gather the entries of @v@ at @idxs@, preserving order. unsafeIndex is
safe: every index in @idxs@ is in range by construction.
-}
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

{- | Standardize one row to the kept columns only (subtract column mean, divide
by column std). unsafeIndex is safe: rows share the column layout.
-}
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

{- | Scatter kept-column weights back into a full-width vector, with @0@ for
the dropped (near-constant) columns.
-}
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

{- | Convert a fitted model to an 'Expr Bool' over its feature columns,
dropping zero-weight features. With no non-zero weights it returns the
constant @Lit (intercept > 0)@.
-}
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

{- | Per-column @(means, stds, variances)@ of a feature matrix. Cheaper than
'standardize' when only the statistics are needed. unsafeIndex within is
safe: all rows share width @d@.
-}
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)

-- | Mean of each of the @d@ columns; @invN@ is @1 / nRows@.
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

-- | Variance of each of the @d@ columns about the supplied @means@.
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

-- | Multiply every element of a mutable vector by @factor@ in place.
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 each column to zero mean and unit variance, also returning
@(means, stds, variances)@. Near-constant columns get std @1@; callers use
the raw variances to detect and drop them (see 'fitL1Logistic').
-}
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)

{- | Proximal operator for the L1 norm: shrink @v@ toward zero by @lambda@,
clamping at zero.
-}
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

{- | Dot product of two unboxed vectors. Caller must ensure equal length;
lengths are not checked.
-}
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)

{- | Gradient of the average loss at @(w, b)@, returning @(gradW, gradB)@.
When @sampleWeights@ is @Just ws@ each row is scaled by @ws[i]@; with mean-1
weights the @1/N@ normalisation is preserved exactly.
-}
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

{- | Per-row loss coefficient @c_i = ℓ'(y_i, z_i) / N@ at margin
@z_i = w·x_i + b@, optionally scaled by @ws[i]@.
-}
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

{- | Accumulate the weight gradient in one pass over every (row, feature)
pair, scattering into a length-@d@ mutable vector.
-}
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

{- | Inner FISTA loop over standardized features, returning the final @(w, b)@
(the caller de-standardizes). @lambda1@/@lambda2@ are the L1/L2 strengths and
@lp@ the smooth-part Lipschitz constant driving the elastic-net prox step.
-}
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

{- | One fused FISTA prox step: gradient step plus the Elastic-Net proximal
operator @softThreshold(z, λ₁/lp) / (1 + λ₂/lp)@ (soft-threshold then ridge
shrinkage). The intercept is unregularised.
-}
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)

{- | Nesterov momentum extrapolation: new look-ahead point @(yW, yB)@ and the
updated step size @t@.
-}
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)

{- | L-inf norm of the weight delta. unsafeIndex is safe: both vectors share
the same length by construction.
-}
{-# 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