{-# LANGUAGE OverloadedStrings #-}

{- | Smooth losses for the proximal-gradient engine. Each carries its
derivative @∂ℓ/∂z@ at @z = w·x + b@ and a global bound on the curvature
@∂²ℓ/∂z²@ (used for the FISTA step size).
-}
module DataFrame.LinearSolver.Loss (
    SmoothLoss (..),
    sigmoid,
    logisticLoss,
    squaredLoss,
    sqHingeLoss,
) where

import qualified Data.Text as T

{- | A convex, @C¹@ per-sample loss @ℓ(y, z)@. 'slGradZ' is @∂ℓ/∂z@;
'slCurvBound' bounds @∂²ℓ/∂z²@ over all @(y, z)@.
-}
data SmoothLoss = SmoothLoss
    { SmoothLoss -> Text
slName :: !T.Text
    , SmoothLoss -> Double -> Double -> Double
slGradZ :: Double -> Double -> Double
    , SmoothLoss -> Double
slCurvBound :: !Double
    }

-- | Numerically stable logistic sigmoid.
sigmoid :: Double -> Double
sigmoid :: Double -> Double
sigmoid Double
z
    | Double
z Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
>= Double
0 = Double
1 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double -> Double
forall a. Floating a => a -> a
exp (-Double
z))
    | Bool
otherwise = let ez :: Double
ez = Double -> Double
forall a. Floating a => a -> a
exp Double
z in Double
ez Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
ez)

-- | Binary logistic loss for labels in @{\-1,+1}@: @ℓ = log(1 + exp(-y z))@.
logisticLoss :: SmoothLoss
logisticLoss :: SmoothLoss
logisticLoss =
    Text -> (Double -> Double -> Double) -> Double -> SmoothLoss
SmoothLoss Text
"logistic" (\Double
y Double
z -> Double -> Double
forall a. Num a => a -> a
negate (Double
y Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double -> Double
sigmoid (Double -> Double
forall a. Num a => a -> a
negate (Double
y Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
z)))) Double
0.25

-- | Squared error for regression: @ℓ = ½ (z - y)²@.
squaredLoss :: SmoothLoss
squaredLoss :: SmoothLoss
squaredLoss = Text -> (Double -> Double -> Double) -> Double -> SmoothLoss
SmoothLoss Text
"squared" ((Double -> Double -> Double) -> Double -> Double -> Double
forall a b c. (a -> b -> c) -> b -> a -> c
flip (-)) Double
1.0

-- | Squared hinge for classification (LinearSVC default), labels @{\-1,+1}@.
sqHingeLoss :: SmoothLoss
sqHingeLoss :: SmoothLoss
sqHingeLoss =
    Text -> (Double -> Double -> Double) -> Double -> SmoothLoss
SmoothLoss
        Text
"squared_hinge"
        (\Double
y Double
z -> let m :: Double
m = Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
y Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
z in if Double
m Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
> Double
0 then Double -> Double
forall a. Num a => a -> a
negate (Double
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
y Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
m) else Double
0)
        Double
2.0