{-# LANGUAGE BangPatterns #-}
{-# LANGUAGE FlexibleContexts #-}
{-# LANGUAGE FlexibleInstances #-}
{-# LANGUAGE MultiParamTypeClasses #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeFamilies #-}

{- | Gaussian mixture models fitted by EM. Full covariance by default (with a
diagonal option and an automatic fall-back when a covariance is not positive
definite), log-space responsibilities, and Cholesky-based densities for
stability. 'predict' is the hard (arg-max) component assignment; per-component
log-densities are available via 'gmmLogDensityExprs'.
-}
module DataFrame.GMM (
    module DataFrame.Model,
    CovType (..),
    GMMConfig (..),
    defaultGMMConfig,
    GMMModel (..),
    gmmLogDensityExprs,
    gmmBIC,
    gmmAIC,
) where

import Data.List (sortBy)
import qualified Data.Map.Strict as M
import Data.Ord (comparing)
import qualified Data.Text as T
import qualified Data.Vector as V
import qualified Data.Vector.Unboxed as VU

import DataFrame.Featurize.Internal (Features (..), argMaxExpr, extractFeatures)
import qualified DataFrame.Functions as F
import DataFrame.Internal.DataFrame (DataFrame)
import DataFrame.Internal.Expression (Expr (..))
import DataFrame.LinearAlgebra (Matrix, logSumExp)
import DataFrame.LinearAlgebra.Solve (backSubst, cholesky, forwardSubst)
import DataFrame.Model
import DataFrame.Operators ((.*.), (.+.), (.-.))
import DataFrame.Random (mkGen, sampleIndices)

data CovType = FullCov | DiagCov
    deriving (CovType -> CovType -> Bool
(CovType -> CovType -> Bool)
-> (CovType -> CovType -> Bool) -> Eq CovType
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: CovType -> CovType -> Bool
== :: CovType -> CovType -> Bool
$c/= :: CovType -> CovType -> Bool
/= :: CovType -> CovType -> Bool
Eq, Int -> CovType -> ShowS
[CovType] -> ShowS
CovType -> String
(Int -> CovType -> ShowS)
-> (CovType -> String) -> ([CovType] -> ShowS) -> Show CovType
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> CovType -> ShowS
showsPrec :: Int -> CovType -> ShowS
$cshow :: CovType -> String
show :: CovType -> String
$cshowList :: [CovType] -> ShowS
showList :: [CovType] -> ShowS
Show)

data GMMConfig = GMMConfig
    { GMMConfig -> Int
gmmK :: !Int
    , GMMConfig -> CovType
gmmCovType :: !CovType
    , GMMConfig -> Int
gmmMaxIter :: !Int
    , GMMConfig -> Double
gmmTol :: !Double
    , GMMConfig -> Double
gmmRegCovar :: !Double
    , GMMConfig -> Int
gmmSeed :: !Int
    }
    deriving (GMMConfig -> GMMConfig -> Bool
(GMMConfig -> GMMConfig -> Bool)
-> (GMMConfig -> GMMConfig -> Bool) -> Eq GMMConfig
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: GMMConfig -> GMMConfig -> Bool
== :: GMMConfig -> GMMConfig -> Bool
$c/= :: GMMConfig -> GMMConfig -> Bool
/= :: GMMConfig -> GMMConfig -> Bool
Eq, Int -> GMMConfig -> ShowS
[GMMConfig] -> ShowS
GMMConfig -> String
(Int -> GMMConfig -> ShowS)
-> (GMMConfig -> String)
-> ([GMMConfig] -> ShowS)
-> Show GMMConfig
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> GMMConfig -> ShowS
showsPrec :: Int -> GMMConfig -> ShowS
$cshow :: GMMConfig -> String
show :: GMMConfig -> String
$cshowList :: [GMMConfig] -> ShowS
showList :: [GMMConfig] -> ShowS
Show)

defaultGMMConfig :: GMMConfig
defaultGMMConfig :: GMMConfig
defaultGMMConfig =
    GMMConfig
        { gmmK :: Int
gmmK = Int
2
        , gmmCovType :: CovType
gmmCovType = CovType
FullCov
        , gmmMaxIter :: Int
gmmMaxIter = Int
100
        , gmmTol :: Double
gmmTol = Double
1.0e-3
        , gmmRegCovar :: Double
gmmRegCovar = Double
1.0e-6
        , gmmSeed :: Int
gmmSeed = Int
0
        }

-- | A fitted mixture. 'gmmCovariances' are the per-component covariance matrices.
data GMMModel = GMMModel
    { GMMModel -> Vector Double
gmmWeights :: !(VU.Vector Double)
    , GMMModel -> Vector (Vector Double)
gmmMeans :: !(V.Vector (VU.Vector Double))
    , GMMModel -> Vector (Vector (Vector Double))
gmmCovariances :: !(V.Vector Matrix)
    , GMMModel -> Bool
gmmConverged :: !Bool
    , GMMModel -> Int
gmmNIter :: !Int
    , GMMModel -> Double
gmmLogLikelihood :: !Double
    , GMMModel -> Int
gmmNObs :: !Int
    , GMMModel -> Vector Text
gmmFeatureNames :: !(V.Vector T.Text)
    }
    deriving (GMMModel -> GMMModel -> Bool
(GMMModel -> GMMModel -> Bool)
-> (GMMModel -> GMMModel -> Bool) -> Eq GMMModel
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: GMMModel -> GMMModel -> Bool
== :: GMMModel -> GMMModel -> Bool
$c/= :: GMMModel -> GMMModel -> Bool
/= :: GMMModel -> GMMModel -> Bool
Eq, Int -> GMMModel -> ShowS
[GMMModel] -> ShowS
GMMModel -> String
(Int -> GMMModel -> ShowS)
-> (GMMModel -> String) -> ([GMMModel] -> ShowS) -> Show GMMModel
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> GMMModel -> ShowS
showsPrec :: Int -> GMMModel -> ShowS
$cshow :: GMMModel -> String
show :: GMMModel -> String
$cshowList :: [GMMModel] -> ShowS
showList :: [GMMModel] -> ShowS
Show)

instance Fit GMMConfig [Expr Double] where
    type ModelOf GMMConfig [Expr Double] = GMMModel
    fit :: CheckFrame
  (FrameReq GMMConfig [Expr Double]) (FrameFor [Expr Double]) =>
GMMConfig
-> [Expr Double]
-> FrameFor [Expr Double]
-> FitResult
     (FrameFor [Expr Double]) (ModelOf GMMConfig [Expr Double])
fit = GMMConfig -> [Expr Double] -> DataFrame -> GMMModel
GMMConfig
-> [Expr Double]
-> FrameFor [Expr Double]
-> FitResult
     (FrameFor [Expr Double]) (ModelOf GMMConfig [Expr Double])
fitGMM

instance Predict GMMModel where
    type Prediction GMMModel = Expr Int
    predict :: GMMModel -> Prediction GMMModel
predict = GMMModel -> Expr Int
GMMModel -> Prediction GMMModel
gmmAssignExpr

-- | Fit a Gaussian mixture over the given feature columns.
fitGMM :: GMMConfig -> [Expr Double] -> DataFrame -> GMMModel
fitGMM :: GMMConfig -> [Expr Double] -> DataFrame -> GMMModel
fitGMM GMMConfig
cfg [Expr Double]
features DataFrame
df = GMMModel -> GMMModel
canonical GMMModel
finalModel
  where
    Features [Text]
names [Vector Double]
_ Vector (Vector Double)
rows Int
n Int
d = [Expr Double] -> DataFrame -> Features
extractFeatures [Expr Double]
features DataFrame
df
    k :: Int
k = Int -> Int -> Int
forall a. Ord a => a -> a -> a
min (GMMConfig -> Int
gmmK GMMConfig
cfg) (Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
1 Int
n)
    reg :: Double
reg = GMMConfig -> Double
gmmRegCovar GMMConfig
cfg
    (Vector Int
initIdx, Gen
_) = Int -> Int -> Gen -> (Vector Int, Gen)
sampleIndices Int
k Int
n (Int -> Gen
mkGen (GMMConfig -> Int
gmmSeed GMMConfig
cfg))
    means0 :: Vector (Vector Double)
means0 = (Int -> Vector Double) -> Vector Int -> Vector (Vector Double)
forall a b. (a -> b) -> Vector a -> Vector b
V.map (Vector (Vector Double)
rows Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.!) (Vector Int -> Vector Int
forall (v :: * -> *) a (w :: * -> *).
(Vector v a, Vector w a) =>
v a -> w a
V.convert Vector Int
initIdx)
    varDiag :: Vector Double
varDiag =
        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 ->
            let mu :: Double
mu = [Double] -> Double
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [(Vector (Vector Double)
rows Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
i) Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
j | Int
i <- [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]] Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
1 Int
n)
             in ( [Double] -> Double
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [((Vector (Vector Double)
rows Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
i) Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
j Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
mu) Double -> Int -> Double
forall a b. (Num a, Integral b) => a -> b -> a
^ (Int
2 :: Int) | Int
i <- [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]]
                    Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
1 Int
n)
                )
                    Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
reg
    cov0 :: Vector (Vector Double)
cov0 = Vector Double -> Vector (Vector Double)
diagMatrix Vector Double
varDiag
    covs0 :: Vector (Vector (Vector Double))
covs0 = Int -> Vector (Vector Double) -> Vector (Vector (Vector Double))
forall a. Int -> a -> Vector a
V.replicate Int
k Vector (Vector Double)
cov0
    weights0 :: Vector Double
weights0 = Int -> Double -> Vector Double
forall a. Unbox a => Int -> a -> Vector a
VU.replicate Int
k (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
k)
    finalModel :: GMMModel
finalModel = Int
-> Vector Double
-> Vector (Vector Double)
-> Vector (Vector (Vector Double))
-> Double
-> Bool
-> GMMModel
em Int
0 Vector Double
weights0 Vector (Vector Double)
means0 Vector (Vector (Vector Double))
covs0 (-(Double
1 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
0)) Bool
False
    em :: Int
-> Vector Double
-> Vector (Vector Double)
-> Vector (Vector (Vector Double))
-> Double
-> Bool
-> GMMModel
em !Int
iter Vector Double
weights Vector (Vector Double)
means Vector (Vector (Vector Double))
covs Double
prevLL Bool
converged
        | Int
iter Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= GMMConfig -> Int
gmmMaxIter GMMConfig
cfg Bool -> Bool -> Bool
|| Bool
converged =
            Vector Double
-> Vector (Vector Double)
-> Vector (Vector (Vector Double))
-> Bool
-> Int
-> Double
-> Int
-> Vector Text
-> GMMModel
GMMModel Vector Double
weights Vector (Vector Double)
means Vector (Vector (Vector Double))
covs Bool
converged Int
iter Double
prevLL Int
n ([Text] -> Vector Text
forall a. [a] -> Vector a
V.fromList [Text]
names)
        | Bool
otherwise =
            let (Vector (Vector Double)
logResp, Double
ll) = GMMConfig
-> Vector (Vector Double)
-> Vector Double
-> Vector (Vector Double)
-> Vector (Vector (Vector Double))
-> (Vector (Vector Double), Double)
eStep GMMConfig
cfg Vector (Vector Double)
rows Vector Double
weights Vector (Vector Double)
means Vector (Vector (Vector Double))
covs
                (Vector Double
weights', Vector (Vector Double)
means', Vector (Vector (Vector Double))
covs') = GMMConfig
-> Vector (Vector Double)
-> Vector (Vector Double)
-> Int
-> Double
-> (Vector Double, Vector (Vector Double),
    Vector (Vector (Vector Double)))
mStep GMMConfig
cfg Vector (Vector Double)
rows Vector (Vector Double)
logResp Int
d Double
reg
                done :: Bool
done = Double -> Double
forall a. Num a => a -> a
abs (Double
ll Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
prevLL) Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
< GMMConfig -> Double
gmmTol GMMConfig
cfg
             in Int
-> Vector Double
-> Vector (Vector Double)
-> Vector (Vector (Vector Double))
-> Double
-> Bool
-> GMMModel
em (Int
iter Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Vector Double
weights' Vector (Vector Double)
means' Vector (Vector (Vector Double))
covs' Double
ll Bool
done

-- | Per-component log-density expressions (log weight + Gaussian log pdf).
gmmLogDensityExprs :: GMMModel -> M.Map Int (Expr Double)
gmmLogDensityExprs :: GMMModel -> Map Int (Expr Double)
gmmLogDensityExprs GMMModel
m =
    [(Int, Expr Double)] -> Map Int (Expr Double)
forall k a. Ord k => [(k, a)] -> Map k a
M.fromList
        [ ( Int
c
          , Double
-> Vector Double -> Vector (Vector Double) -> [Text] -> Expr Double
logDensityExpr
                (GMMModel -> Vector Double
gmmWeights GMMModel
m Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
c)
                (GMMModel -> Vector (Vector Double)
gmmMeans GMMModel
m Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
c)
                (GMMModel -> Vector (Vector (Vector Double))
gmmCovariances GMMModel
m Vector (Vector (Vector Double)) -> Int -> Vector (Vector Double)
forall a. Vector a -> Int -> a
V.! Int
c)
                [Text]
names
          )
        | Int
c <- [Int
0 .. Vector (Vector Double) -> Int
forall a. Vector a -> Int
V.length (GMMModel -> Vector (Vector Double)
gmmMeans GMMModel
m) Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]
        ]
  where
    names :: [Text]
names = Vector Text -> [Text]
forall a. Vector a -> [a]
V.toList (GMMModel -> Vector Text
gmmFeatureNames GMMModel
m)

-- | The hard-assignment expression: arg-max of component log-densities.
gmmAssignExpr :: GMMModel -> Expr Int
gmmAssignExpr :: GMMModel -> Expr Int
gmmAssignExpr GMMModel
m =
    [(Int, Expr Double)] -> Expr Int
forall a. Columnable a => [(a, Expr Double)] -> Expr a
argMaxExpr (Map Int (Expr Double) -> [(Int, Expr Double)]
forall k a. Map k a -> [(k, a)]
M.toList (GMMModel -> Map Int (Expr Double)
gmmLogDensityExprs GMMModel
m))

-- | Bayesian information criterion (lower is better).
gmmBIC :: GMMModel -> Double
gmmBIC :: GMMModel -> Double
gmmBIC GMMModel
m =
    Double -> Double
forall a. Num a => a -> a
negate (Double
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* GMMModel -> Double
gmmLogLikelihood GMMModel
m)
        Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral (GMMModel -> Int
nParams GMMModel
m) Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double -> Double
forall a. Floating a => a -> a
log (Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
1 (GMMModel -> Int
gmmNObs GMMModel
m)))

-- | Akaike information criterion (lower is better).
gmmAIC :: GMMModel -> Double
gmmAIC :: GMMModel -> Double
gmmAIC GMMModel
m = Double -> Double
forall a. Num a => a -> a
negate (Double
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* GMMModel -> Double
gmmLogLikelihood GMMModel
m) Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral (GMMModel -> Int
nParams GMMModel
m)

nParams :: GMMModel -> Int
nParams :: GMMModel -> Int
nParams GMMModel
m =
    let k :: Int
k = Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length (GMMModel -> Vector Double
gmmWeights GMMModel
m)
        d :: Int
d = if Vector (Vector Double) -> Bool
forall a. Vector a -> Bool
V.null (GMMModel -> Vector (Vector Double)
gmmMeans GMMModel
m) 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 (GMMModel -> Vector (Vector Double)
gmmMeans GMMModel
m))
     in (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1) Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
* (Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
* (Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Int -> Int -> Int
forall a. Integral a => a -> a -> a
`div` Int
2)

eStep ::
    GMMConfig ->
    Matrix ->
    VU.Vector Double ->
    V.Vector (VU.Vector Double) ->
    V.Vector Matrix ->
    (V.Vector (VU.Vector Double), Double)
eStep :: GMMConfig
-> Vector (Vector Double)
-> Vector Double
-> Vector (Vector Double)
-> Vector (Vector (Vector Double))
-> (Vector (Vector Double), Double)
eStep GMMConfig
_ Vector (Vector Double)
rows Vector Double
weights Vector (Vector Double)
means Vector (Vector (Vector Double))
covs = (Vector (Vector Double)
logResp, Double
totalLL)
  where
    k :: Int
k = Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Double
weights
    comps :: Vector (Double, Vector Double, Vector Double -> Double)
comps =
        Int
-> (Int -> (Double, Vector Double, Vector Double -> Double))
-> Vector (Double, Vector Double, Vector Double -> Double)
forall a. Int -> (Int -> a) -> Vector a
V.generate Int
k ((Int -> (Double, Vector Double, Vector Double -> Double))
 -> Vector (Double, Vector Double, Vector Double -> Double))
-> (Int -> (Double, Vector Double, Vector Double -> Double))
-> Vector (Double, Vector Double, Vector Double -> Double)
forall a b. (a -> b) -> a -> b
$ \Int
c ->
            ( Double -> Double
forall a. Floating a => a -> a
log (Double -> Double -> Double
forall a. Ord a => a -> a -> a
max Double
1e-300 (Vector Double
weights Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
c))
            , Vector (Vector Double)
means Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
c
            , Vector (Vector Double) -> Vector Double -> Vector Double -> Double
gaussianLogPdf (Vector (Vector (Vector Double))
covs Vector (Vector (Vector Double)) -> Int -> Vector (Vector Double)
forall a. Vector a -> Int -> a
V.! Int
c) (Vector (Vector Double)
means Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
c)
            )
    perRow :: Vector Double -> (Vector Double, Double)
perRow Vector Double
x =
        let lps :: Vector Double
lps = Int -> (Int -> Double) -> Vector Double
forall a. Unbox a => Int -> (Int -> a) -> Vector a
VU.generate Int
k (\Int
c -> let (Double
lw, Vector Double
_, Vector Double -> Double
f) = Vector (Double, Vector Double, Vector Double -> Double)
comps Vector (Double, Vector Double, Vector Double -> Double)
-> Int -> (Double, Vector Double, Vector Double -> Double)
forall a. Vector a -> Int -> a
V.! Int
c in Double
lw Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Vector Double -> Double
f Vector Double
x)
            lse :: Double
lse = Vector Double -> Double
logSumExp Vector Double
lps
         in ((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
subtract Double
lse) Vector Double
lps, Double
lse)
    results :: Vector (Vector Double, Double)
results = (Vector Double -> (Vector Double, Double))
-> Vector (Vector Double) -> Vector (Vector Double, Double)
forall a b. (a -> b) -> Vector a -> Vector b
V.map Vector Double -> (Vector Double, Double)
perRow Vector (Vector Double)
rows
    logResp :: Vector (Vector Double)
logResp = ((Vector Double, Double) -> Vector Double)
-> Vector (Vector Double, Double) -> Vector (Vector Double)
forall a b. (a -> b) -> Vector a -> Vector b
V.map (Vector Double, Double) -> Vector Double
forall a b. (a, b) -> a
fst Vector (Vector Double, Double)
results
    totalLL :: Double
totalLL = Vector Double -> Double
forall a. Num a => Vector a -> a
V.sum (((Vector Double, Double) -> Double)
-> Vector (Vector Double, Double) -> Vector Double
forall a b. (a -> b) -> Vector a -> Vector b
V.map (Vector Double, Double) -> Double
forall a b. (a, b) -> b
snd Vector (Vector Double, Double)
results)

mStep ::
    GMMConfig ->
    Matrix ->
    V.Vector (VU.Vector Double) ->
    Int ->
    Double ->
    (VU.Vector Double, V.Vector (VU.Vector Double), V.Vector Matrix)
mStep :: GMMConfig
-> Vector (Vector Double)
-> Vector (Vector Double)
-> Int
-> Double
-> (Vector Double, Vector (Vector Double),
    Vector (Vector (Vector Double)))
mStep GMMConfig
cfg Vector (Vector Double)
rows Vector (Vector Double)
logResp Int
d Double
reg = (Vector Double
weights, Vector (Vector Double)
means, Vector (Vector (Vector Double))
covs)
  where
    n :: Int
n = Vector (Vector Double) -> Int
forall a. Vector a -> Int
V.length Vector (Vector Double)
rows
    k :: Int
k = if Vector (Vector Double) -> Bool
forall a. Vector a -> Bool
V.null Vector (Vector Double)
logResp 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)
logResp)
    resp :: Vector (Vector Double)
resp = (Vector Double -> Vector Double)
-> Vector (Vector Double) -> Vector (Vector Double)
forall a b. (a -> b) -> Vector a -> Vector b
V.map ((Double -> Double) -> Vector Double -> Vector Double
forall a b. (Unbox a, Unbox b) => (a -> b) -> Vector a -> Vector b
VU.map Double -> Double
forall a. Floating a => a -> a
exp) Vector (Vector Double)
logResp
    nk :: Vector Double
nk = Int -> (Int -> Double) -> Vector Double
forall a. Unbox a => Int -> (Int -> a) -> Vector a
VU.generate Int
k (\Int
c -> [Double] -> Double
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [Vector (Vector Double)
resp Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
i Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
c | Int
i <- [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]])
    weights :: Vector Double
weights = (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. Fractional a => a -> a -> a
/ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
1 Int
n)) Vector Double
nk
    means :: Vector (Vector Double)
means =
        Int -> (Int -> Vector Double) -> Vector (Vector Double)
forall a. Int -> (Int -> a) -> Vector a
V.generate Int
k ((Int -> Vector Double) -> Vector (Vector Double))
-> (Int -> Vector Double) -> Vector (Vector Double)
forall a b. (a -> b) -> a -> b
$ \Int
c ->
            let s :: Vector Double
s =
                    (Vector Double -> Vector Double -> Vector Double)
-> Vector Double -> [Vector Double] -> Vector Double
forall a b. (a -> b -> b) -> b -> [a] -> b
forall (t :: * -> *) a b.
Foldable t =>
(a -> b -> b) -> b -> t a -> b
foldr
                        ((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
(+))
                        (Int -> Double -> Vector Double
forall a. Unbox a => Int -> a -> Vector a
VU.replicate Int
d Double
0)
                        [(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
* (Vector (Vector Double)
resp Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
i Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
c)) (Vector (Vector Double)
rows Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
i) | Int
i <- [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]]
             in (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. Fractional a => a -> a -> a
/ Double -> Double -> Double
forall a. Ord a => a -> a -> a
max Double
1e-12 (Vector Double
nk Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
c)) Vector Double
s
    covs :: Vector (Vector (Vector Double))
covs =
        Int
-> (Int -> Vector (Vector Double))
-> Vector (Vector (Vector Double))
forall a. Int -> (Int -> a) -> Vector a
V.generate Int
k ((Int -> Vector (Vector Double))
 -> Vector (Vector (Vector Double)))
-> (Int -> Vector (Vector Double))
-> Vector (Vector (Vector Double))
forall a b. (a -> b) -> a -> b
$ \Int
c ->
            let mu :: Vector Double
mu = Vector (Vector Double)
means Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
c
                acc :: Vector (Vector Double)
acc = (Int -> Vector (Vector Double) -> Vector (Vector Double))
-> Vector (Vector Double) -> [Int] -> Vector (Vector Double)
forall a b. (a -> b -> b) -> b -> [a] -> b
forall (t :: * -> *) a b.
Foldable t =>
(a -> b -> b) -> b -> t a -> b
foldr Int -> Vector (Vector Double) -> Vector (Vector Double)
addOuter (Int -> Vector (Vector Double)
zeroMatrix Int
d) [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]
                addOuter :: Int -> Vector (Vector Double) -> Vector (Vector Double)
addOuter Int
i Vector (Vector Double)
m =
                    let diff :: Vector Double
diff = (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 (-) (Vector (Vector Double)
rows Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
i) Vector Double
mu
                        w :: Double
w = Vector (Vector Double)
resp Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
i Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
c
                     in Double
-> Vector Double
-> Vector (Vector Double)
-> Vector (Vector Double)
addScaledOuter Double
w Vector Double
diff Vector (Vector Double)
m
                scaled :: Vector (Vector Double)
scaled = Double -> Vector (Vector Double) -> Vector (Vector Double)
scaleMatrix (Double
1 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double -> Double -> Double
forall a. Ord a => a -> a -> a
max Double
1e-12 (Vector Double
nk Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
c)) Vector (Vector Double)
acc
                regd :: Vector (Vector Double)
regd = Double -> Vector (Vector Double) -> Vector (Vector Double)
addDiagScalar Double
reg Vector (Vector Double)
scaled
             in case GMMConfig -> CovType
gmmCovType GMMConfig
cfg of
                    CovType
FullCov -> Vector (Vector Double)
regd
                    CovType
DiagCov -> Vector (Vector Double) -> Vector (Vector Double)
diagOnly Vector (Vector Double)
regd

gaussianLogPdf :: Matrix -> VU.Vector Double -> VU.Vector Double -> Double
gaussianLogPdf :: Vector (Vector Double) -> Vector Double -> Vector Double -> Double
gaussianLogPdf Vector (Vector Double)
cov Vector Double
mu =
    case Vector (Vector Double) -> Maybe (Vector (Vector Double))
cholesky Vector (Vector Double)
cov of
        Just Vector (Vector Double)
l ->
            let logdet :: Double
logdet = Double
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* [Double] -> Double
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [Double -> Double
forall a. Floating a => a -> a
log ((Vector (Vector Double)
l Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
i) Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i) | Int
i <- [Int
0 .. Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]]
             in \Vector Double
x ->
                    let diff :: Vector Double
diff = (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 (-) Vector Double
x Vector Double
mu
                        z :: Vector Double
z = Vector (Vector Double) -> Vector Double -> Vector Double
forwardSubst Vector (Vector Double)
l Vector Double
diff
                        quad :: Double
quad = Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum ((Double -> Double) -> Vector Double -> Vector Double
forall a b. (Unbox a, Unbox b) => (a -> b) -> Vector a -> Vector b
VU.map (Double -> Int -> Double
forall a b. (Num a, Integral b) => a -> b -> a
^ (Int
2 :: Int)) Vector Double
z)
                     in Double -> Double
forall a. Num a => a -> a
negate Double
0.5 Double -> Double -> Double
forall a. Num a => a -> a -> a
* (Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
d Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double -> Double
forall a. Floating a => a -> a
log (Double
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
forall a. Floating a => a
pi) Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
logdet Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
quad)
        Maybe (Vector (Vector Double))
Nothing ->
            let var :: Vector Double
var = Int -> (Int -> Double) -> Vector Double
forall a. Unbox a => Int -> (Int -> a) -> Vector a
VU.generate Int
d (\Int
i -> Double -> Double -> Double
forall a. Ord a => a -> a -> a
max Double
1e-12 ((Vector (Vector Double)
cov Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
i) Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i))
             in \Vector Double
x ->
                    Double -> Double
forall a. Num a => a -> a
negate Double
0.5
                        Double -> Double -> Double
forall a. Num a => a -> a -> a
* Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum
                            ( 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 ->
                                let diff :: Double
diff = Vector Double
x Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
j Double -> Double -> Double
forall a. Num a => a -> a -> a
- Vector Double
mu Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
j
                                 in Double
diff Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
diff Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Vector Double
var Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
j Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double -> Double
forall a. Floating a => a -> a
log (Double
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
forall a. Floating a => a
pi Double -> Double -> Double
forall a. Num a => a -> a -> a
* Vector Double
var Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
j)
                            )
  where
    d :: Int
d = Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Double
mu

logDensityExpr ::
    Double -> VU.Vector Double -> Matrix -> [T.Text] -> Expr Double
logDensityExpr :: Double
-> Vector Double -> Vector (Vector Double) -> [Text] -> Expr Double
logDensityExpr Double
weight Vector Double
mu Vector (Vector Double)
cov [Text]
names =
    case Vector (Vector Double) -> Maybe (Vector (Vector Double), Double)
precisionAndLogdet Vector (Vector Double)
cov of
        Just (Vector (Vector Double)
prec, Double
logdet) ->
            let constTerm :: Double
constTerm =
                    Double -> Double
forall a. Floating a => a -> a
log (Double -> Double -> Double
forall a. Ord a => a -> a -> a
max Double
1e-300 Double
weight)
                        Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
0.5 Double -> Double -> Double
forall a. Num a => a -> a -> a
* (Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
d Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double -> Double
forall a. Floating a => a -> a
log (Double
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
forall a. Floating a => a
pi) Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
logdet)
                quad :: Expr Double
quad =
                    (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
0) ([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 (-(Double
0.5 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Vector (Vector Double)
prec Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
a Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
b)) Expr Double -> Expr Double -> Expr Double
forall a. (Columnable a, Num a) => Expr a -> Expr a -> Expr a
.*. (Int -> Expr Double
centered Int
a Expr Double -> Expr Double -> Expr Double
forall a. (Columnable a, Num a) => Expr a -> Expr a -> Expr a
.*. Int -> Expr Double
centered Int
b)
                        | Int
a <- [Int
0 .. Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]
                        , Int
b <- [Int
0 .. Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]
                        ]
             in Double -> Expr Double
forall a. Columnable a => a -> Expr a
F.lit Double
constTerm Expr Double -> Expr Double -> Expr Double
forall a. (Columnable a, Num a) => Expr a -> Expr a -> Expr a
.+. Expr Double
quad
        Maybe (Vector (Vector Double), Double)
Nothing -> Double -> Expr Double
forall a. Columnable a => a -> Expr a
F.lit (Double -> Double
forall a. Floating a => a -> a
log (Double -> Double -> Double
forall a. Ord a => a -> a -> a
max Double
1e-300 Double
weight))
  where
    d :: Int
d = Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Double
mu
    centered :: Int -> Expr Double
centered Int
j = (Text -> Expr Double
forall a. Columnable a => Text -> Expr a
Col ([Text]
names [Text] -> Int -> Text
forall a. HasCallStack => [a] -> Int -> a
!! Int
j) :: Expr Double) 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 (Vector Double
mu Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
j)

precisionAndLogdet :: Matrix -> Maybe (Matrix, Double)
precisionAndLogdet :: Vector (Vector Double) -> Maybe (Vector (Vector Double), Double)
precisionAndLogdet Vector (Vector Double)
cov = do
    Vector (Vector Double)
l <- Vector (Vector Double) -> Maybe (Vector (Vector Double))
cholesky Vector (Vector Double)
cov
    let d :: Int
d = Vector (Vector Double) -> Int
forall a. Vector a -> Int
V.length Vector (Vector Double)
cov
        logdet :: Double
logdet = Double
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* [Double] -> Double
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [Double -> Double
forall a. Floating a => a -> a
log ((Vector (Vector Double)
l Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
i) Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i) | Int
i <- [Int
0 .. Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]]
        cols :: [Vector Double]
cols = [Vector (Vector Double) -> Vector Double -> Vector Double
forwardThenBack Vector (Vector Double)
l (Int -> Int -> Vector Double
unitVec Int
d Int
i) | Int
i <- [Int
0 .. Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]]
        prec :: Vector (Vector Double)
prec =
            [Vector Double] -> Vector (Vector Double)
forall a. [a] -> Vector a
V.fromList
                [[Double] -> Vector Double
forall a. Unbox a => [a] -> Vector a
VU.fromList [[Vector Double]
cols [Vector Double] -> Int -> Vector Double
forall a. HasCallStack => [a] -> Int -> a
!! Int
j Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i | Int
j <- [Int
0 .. Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]] | Int
i <- [Int
0 .. Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]]
    (Vector (Vector Double), Double)
-> Maybe (Vector (Vector Double), Double)
forall a. a -> Maybe a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (Vector (Vector Double)
prec, Double
logdet)

forwardThenBack :: Matrix -> VU.Vector Double -> VU.Vector Double
forwardThenBack :: Vector (Vector Double) -> Vector Double -> Vector Double
forwardThenBack Vector (Vector Double)
l Vector Double
b = Vector (Vector Double) -> Vector Double -> Vector Double
backSubst Vector (Vector Double)
l (Vector (Vector Double) -> Vector Double -> Vector Double
forwardSubst Vector (Vector Double)
l Vector Double
b)

canonical :: GMMModel -> GMMModel
canonical :: GMMModel -> GMMModel
canonical GMMModel
m =
    let order :: [Int]
order =
            ((Double, Int) -> Int) -> [(Double, Int)] -> [Int]
forall a b. (a -> b) -> [a] -> [b]
map (Double, Int) -> Int
forall a b. (a, b) -> b
snd ([(Double, Int)] -> [Int]) -> [(Double, Int)] -> [Int]
forall a b. (a -> b) -> a -> b
$
                ((Double, Int) -> (Double, Int) -> Ordering)
-> [(Double, Int)] -> [(Double, Int)]
forall a. (a -> a -> Ordering) -> [a] -> [a]
sortBy
                    (((Double, Int) -> Double)
-> (Double, Int) -> (Double, Int) -> Ordering
forall a b. Ord a => (b -> a) -> b -> b -> Ordering
comparing (Double, Int) -> Double
forall a b. (a, b) -> a
fst)
                    [ (Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
firstCoord (GMMModel -> Vector (Vector Double)
gmmMeans GMMModel
m Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
c), Int
c)
                    | Int
c <- [Int
0 .. Vector (Vector Double) -> Int
forall a. Vector a -> Int
V.length (GMMModel -> Vector (Vector Double)
gmmMeans GMMModel
m) Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]
                    ]
        firstCoord :: Vector a -> a
firstCoord Vector a
v = if Vector a -> Bool
forall a. Unbox a => Vector a -> Bool
VU.null Vector a
v then a
0 else Vector a -> a
forall a. Unbox a => Vector a -> a
VU.head Vector a
v
     in GMMModel
m
            { gmmWeights = VU.fromList [gmmWeights m VU.! c | c <- order]
            , gmmMeans = V.fromList [gmmMeans m V.! c | c <- order]
            , gmmCovariances = V.fromList [gmmCovariances m V.! c | c <- order]
            }

diagMatrix :: VU.Vector Double -> Matrix
diagMatrix :: Vector Double -> Vector (Vector Double)
diagMatrix Vector Double
v =
    let d :: Int
d = Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Double
v
     in Int -> (Int -> Vector Double) -> Vector (Vector Double)
forall a. Int -> (Int -> a) -> Vector a
V.generate Int
d (\Int
i -> Int -> (Int -> Double) -> Vector Double
forall a. Unbox a => Int -> (Int -> a) -> Vector a
VU.generate Int
d (\Int
j -> if Int
i Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
j then Vector Double
v Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i else Double
0))

diagOnly :: Matrix -> Matrix
diagOnly :: Vector (Vector Double) -> Vector (Vector Double)
diagOnly Vector (Vector Double)
m =
    let d :: Int
d = Vector (Vector Double) -> Int
forall a. Vector a -> Int
V.length Vector (Vector Double)
m
     in Int -> (Int -> Vector Double) -> Vector (Vector Double)
forall a. Int -> (Int -> a) -> Vector a
V.generate
            Int
d
            (\Int
i -> Int -> (Int -> Double) -> Vector Double
forall a. Unbox a => Int -> (Int -> a) -> Vector a
VU.generate Int
d (\Int
j -> if Int
i Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
j then (Vector (Vector Double)
m Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
i) Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
j else Double
0))

zeroMatrix :: Int -> Matrix
zeroMatrix :: Int -> Vector (Vector Double)
zeroMatrix Int
d = Int -> Vector Double -> Vector (Vector Double)
forall a. Int -> a -> Vector a
V.replicate Int
d (Int -> Double -> Vector Double
forall a. Unbox a => Int -> a -> Vector a
VU.replicate Int
d Double
0)

scaleMatrix :: Double -> Matrix -> Matrix
scaleMatrix :: Double -> Vector (Vector Double) -> Vector (Vector Double)
scaleMatrix Double
s = (Vector Double -> Vector Double)
-> Vector (Vector Double) -> Vector (Vector Double)
forall a b. (a -> b) -> Vector a -> Vector b
V.map ((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
s))

addDiagScalar :: Double -> Matrix -> Matrix
addDiagScalar :: Double -> Vector (Vector Double) -> Vector (Vector Double)
addDiagScalar Double
s = (Int -> Vector Double -> Vector Double)
-> Vector (Vector Double) -> Vector (Vector Double)
forall a b. (Int -> a -> b) -> Vector a -> Vector b
V.imap (\Int
i Vector Double
row -> Vector Double
row Vector Double -> [(Int, Double)] -> Vector Double
forall a. Unbox a => Vector a -> [(Int, a)] -> Vector a
VU.// [(Int
i, Vector Double
row Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
s)])

addScaledOuter :: Double -> VU.Vector Double -> Matrix -> Matrix
addScaledOuter :: Double
-> Vector Double
-> Vector (Vector Double)
-> Vector (Vector Double)
addScaledOuter Double
w Vector Double
diff Vector (Vector Double)
m =
    let d :: Int
d = Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Double
diff
     in Int -> (Int -> Vector Double) -> Vector (Vector Double)
forall a. Int -> (Int -> a) -> Vector a
V.generate Int
d ((Int -> Vector Double) -> Vector (Vector Double))
-> (Int -> Vector Double) -> Vector (Vector Double)
forall a b. (a -> b) -> a -> b
$ \Int
i ->
            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 (Vector Double)
m Vector (Vector Double) -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
i) Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
j Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
w Double -> Double -> Double
forall a. Num a => a -> a -> a
* (Vector Double
diff Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i) Double -> Double -> Double
forall a. Num a => a -> a -> a
* (Vector Double
diff Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
j)

unitVec :: Int -> Int -> VU.Vector Double
unitVec :: Int -> Int -> Vector Double
unitVec Int
d Int
i = Int -> (Int -> Double) -> Vector Double
forall a. Unbox a => Int -> (Int -> a) -> Vector a
VU.generate Int
d (\Int
j -> if Int
i Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
j then Double
1 else Double
0)