{-# LANGUAGE BangPatterns #-}

{- | Householder QR (for ordinary least squares) and Cholesky factorisation (for
ridge normal equations and Gaussian log-densities). Pure, deterministic, no
LAPACK; sound at the @d@ ≤ low-hundreds scales this library targets.
-}
module DataFrame.LinearAlgebra.Solve (
    qrLeastSquares,
    cholesky,
    choleskySolve,
    logDetFromChol,
    forwardSubst,
    backSubst,
) where

import Control.Monad (forM_)
import Control.Monad.ST (ST, runST)
import qualified Data.Vector as V
import qualified Data.Vector.Unboxed as VU
import qualified Data.Vector.Unboxed.Mutable as VUM
import DataFrame.LinearAlgebra (Matrix)

{- | Solve @min ‖A x − b‖₂@ for an @n×d@ matrix @A@ (@n ≥ d@) via Householder QR.
@Left cols@ reports rank deficiency (near-zero @R@ diagonal) with the offending
column indices; @Right x@ is the least-squares solution.
-}
qrLeastSquares :: Matrix -> VU.Vector Double -> Either [Int] (VU.Vector Double)
qrLeastSquares :: Matrix -> Vector Double -> Either [Int] (Vector Double)
qrLeastSquares Matrix
a Vector Double
b
    | Matrix -> Bool
forall a. Vector a -> Bool
V.null Matrix
a = Vector Double -> Either [Int] (Vector Double)
forall a b. b -> Either a b
Right Vector Double
forall a. Unbox a => Vector a
VU.empty
    | Int
n Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
d = [Int] -> Either [Int] (Vector Double)
forall a b. a -> Either a b
Left [Int
0 .. Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]
    | Bool
otherwise = (forall s. ST s (Either [Int] (Vector Double)))
-> Either [Int] (Vector Double)
forall a. (forall s. ST s a) -> a
runST ((forall s. ST s (Either [Int] (Vector Double)))
 -> Either [Int] (Vector Double))
-> (forall s. ST s (Either [Int] (Vector Double)))
-> Either [Int] (Vector Double)
forall a b. (a -> b) -> a -> b
$ do
        STVector s Double
mat <- Int -> ST s (MVector (PrimState (ST s)) Double)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new (Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
d)
        [Int] -> (Int -> ST s ()) -> ST s ()
forall (t :: * -> *) (m :: * -> *) a b.
(Foldable t, Monad m) =>
t a -> (a -> m b) -> m ()
forM_ [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1] ((Int -> ST s ()) -> ST s ()) -> (Int -> ST s ()) -> ST s ()
forall a b. (a -> b) -> a -> b
$ \Int
i ->
            [Int] -> (Int -> ST s ()) -> ST s ()
forall (t :: * -> *) (m :: * -> *) a b.
(Foldable t, Monad m) =>
t a -> (a -> m b) -> m ()
forM_ [Int
0 .. Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1] ((Int -> ST s ()) -> ST s ()) -> (Int -> ST s ()) -> ST s ()
forall a b. (a -> b) -> a -> b
$ \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.write STVector s Double
MVector (PrimState (ST s)) Double
mat (Int
j Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
i) ((Matrix
a Matrix -> 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)
        STVector s Double
rhs <- Int -> ST s (MVector (PrimState (ST s)) Double)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
n
        [Int] -> (Int -> ST s ()) -> ST s ()
forall (t :: * -> *) (m :: * -> *) a b.
(Foldable t, Monad m) =>
t a -> (a -> m b) -> m ()
forM_ [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1] ((Int -> ST s ()) -> ST s ()) -> (Int -> ST s ()) -> ST s ()
forall a b. (a -> b) -> a -> b
$ \Int
i -> MVector (PrimState (ST s)) Double -> Int -> Double -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.write STVector s Double
MVector (PrimState (ST s)) Double
rhs Int
i (Vector Double
b Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i)
        [Int]
deficient <- STVector s Double -> STVector s Double -> Int -> Int -> ST s [Int]
forall s.
STVector s Double -> STVector s Double -> Int -> Int -> ST s [Int]
householder STVector s Double
mat STVector s Double
rhs Int
n Int
d
        if Bool -> Bool
not ([Int] -> Bool
forall a. [a] -> Bool
forall (t :: * -> *) a. Foldable t => t a -> Bool
null [Int]
deficient)
            then Either [Int] (Vector Double) -> ST s (Either [Int] (Vector Double))
forall a. a -> ST s a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ([Int] -> Either [Int] (Vector Double)
forall a b. a -> Either a b
Left [Int]
deficient)
            else do
                STVector s Double
x <- Int -> ST s (MVector (PrimState (ST s)) Double)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
d
                STVector s Double
-> STVector s Double -> Int -> Int -> STVector s Double -> ST s ()
forall s.
STVector s Double
-> STVector s Double -> Int -> Int -> STVector s Double -> ST s ()
backSubstQR STVector s Double
mat STVector s Double
rhs Int
n Int
d STVector s Double
x
                Vector Double -> Either [Int] (Vector Double)
forall a b. b -> Either a b
Right (Vector Double -> Either [Int] (Vector Double))
-> ST s (Vector Double) -> ST s (Either [Int] (Vector Double))
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> MVector (PrimState (ST s)) Double -> ST s (Vector Double)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.freeze STVector s Double
MVector (PrimState (ST s)) Double
x
  where
    n :: Int
n = Matrix -> Int
forall a. Vector a -> Int
V.length Matrix
a
    d :: Int
d = Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length (Matrix -> Vector Double
forall a. Vector a -> a
V.head Matrix
a)

householder ::
    VUM.STVector s Double -> VUM.STVector s Double -> Int -> Int -> ST s [Int]
householder :: forall s.
STVector s Double -> STVector s Double -> Int -> Int -> ST s [Int]
householder STVector s Double
mat STVector s Double
rhs Int
n Int
d = Int -> [Int] -> ST s [Int]
forall {f :: * -> *}.
(PrimState f ~ s, PrimMonad f) =>
Int -> [Int] -> f [Int]
go Int
0 []
  where
    tol :: Double
tol = Double
1e-10
    go :: Int -> [Int] -> f [Int]
go Int
k [Int]
acc
        | Int
k Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
d = [Int] -> f [Int]
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ([Int] -> [Int]
forall a. [a] -> [a]
reverse [Int]
acc)
        | Bool
otherwise = do
            Double
normSq <- Int -> f Double
forall {f :: * -> *}.
(PrimState f ~ s, PrimMonad f) =>
Int -> f Double
sumSq Int
k
            let alphaMag :: Double
alphaMag = Double -> Double
forall a. Floating a => a -> a
sqrt Double
normSq
            Double
akk <- MVector (PrimState f) Double -> Int -> f Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.read STVector s Double
MVector (PrimState f) Double
mat (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
k)
            let alpha :: Double
alpha = if Double
akk Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
> Double
0 then Double -> Double
forall a. Num a => a -> a
negate Double
alphaMag else Double
alphaMag
            if Double
alphaMag Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
< Double
tol
                then Int -> [Int] -> f [Int]
go (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) (Int
k Int -> [Int] -> [Int]
forall a. a -> [a] -> [a]
: [Int]
acc)
                else do
                    MVector (PrimState f) Double -> Int -> Double -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.write STVector s Double
MVector (PrimState f) Double
mat (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
k) (Double
akk Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
alpha)
                    Double
vNormSq <- Int -> f Double
forall {f :: * -> *}.
(PrimState f ~ s, PrimMonad f) =>
Int -> f Double
sumSq Int
k
                    if Double
vNormSq Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
< Double
tol Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
tol
                        then Int -> [Int] -> f [Int]
go (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) [Int]
acc
                        else do
                            [Int] -> (Int -> f ()) -> f ()
forall (t :: * -> *) (m :: * -> *) a b.
(Foldable t, Monad m) =>
t a -> (a -> m b) -> m ()
forM_ [Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1 .. Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1] ((Int -> f ()) -> f ()) -> (Int -> f ()) -> f ()
forall a b. (a -> b) -> a -> b
$ \Int
j -> Int -> Int -> Double -> f ()
forall {m :: * -> *}.
(PrimState m ~ s, PrimMonad m) =>
Int -> Int -> Double -> m ()
reflectColumn Int
k Int
j Double
vNormSq
                            Int -> Double -> f ()
forall {m :: * -> *}.
(PrimState m ~ s, PrimMonad m) =>
Int -> Double -> m ()
reflectRhs Int
k Double
vNormSq
                            MVector (PrimState f) Double -> Int -> Double -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.write STVector s Double
MVector (PrimState f) Double
mat (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
k) Double
alpha
                            Int -> [Int] -> f [Int]
go (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) [Int]
acc
    sumSq :: Int -> f Double
sumSq Int
k = Int -> f Double
forall {f :: * -> *}.
(PrimState f ~ s, PrimMonad f) =>
Int -> f Double
foldRows Int
k
      where
        foldRows :: Int -> f Double
foldRows Int
i
            | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
n = Double -> f Double
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure Double
0
            | Bool
otherwise = do
                Double
x <- MVector (PrimState f) Double -> Int -> f Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.read STVector s Double
MVector (PrimState f) Double
mat (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
i)
                Double
rest <- Int -> f Double
foldRows (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
                Double -> f Double
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
rest)
    reflectColumn :: Int -> Int -> Double -> m ()
reflectColumn Int
k Int
j Double
vNormSq = do
        Double
dotv <- Int -> Int -> Int -> m Double
forall {f :: * -> *}.
(PrimState f ~ s, PrimMonad f) =>
Int -> Int -> Int -> f Double
dotV Int
k Int
j Int
k
        let beta :: Double
beta = Double
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
dotv Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
vNormSq
        [Int] -> (Int -> m ()) -> m ()
forall (t :: * -> *) (m :: * -> *) a b.
(Foldable t, Monad m) =>
t a -> (a -> m b) -> m ()
forM_ [Int
k .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1] ((Int -> m ()) -> m ()) -> (Int -> m ()) -> m ()
forall a b. (a -> b) -> a -> b
$ \Int
i -> do
            Double
vi <- MVector (PrimState m) Double -> Int -> m Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.read STVector s Double
MVector (PrimState m) Double
mat (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
i)
            Double
aij <- MVector (PrimState m) Double -> Int -> m Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.read STVector s Double
MVector (PrimState m) Double
mat (Int
j Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
i)
            MVector (PrimState m) Double -> Int -> Double -> m ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.write STVector s Double
MVector (PrimState m) Double
mat (Int
j Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
i) (Double
aij Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
beta Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
vi)
    reflectRhs :: Int -> Double -> m ()
reflectRhs Int
k Double
vNormSq = do
        Double
dotv <- Int -> Int -> m Double
forall {f :: * -> *}.
(PrimState f ~ s, PrimMonad f) =>
Int -> Int -> f Double
dotRhs Int
k Int
k
        let beta :: Double
beta = Double
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
dotv Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
vNormSq
        [Int] -> (Int -> m ()) -> m ()
forall (t :: * -> *) (m :: * -> *) a b.
(Foldable t, Monad m) =>
t a -> (a -> m b) -> m ()
forM_ [Int
k .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1] ((Int -> m ()) -> m ()) -> (Int -> m ()) -> m ()
forall a b. (a -> b) -> a -> b
$ \Int
i -> do
            Double
vi <- MVector (PrimState m) Double -> Int -> m Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.read STVector s Double
MVector (PrimState m) Double
mat (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
i)
            Double
bi <- MVector (PrimState m) Double -> Int -> m Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.read STVector s Double
MVector (PrimState m) Double
rhs Int
i
            MVector (PrimState m) Double -> Int -> Double -> m ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.write STVector s Double
MVector (PrimState m) Double
rhs Int
i (Double
bi Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
beta Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
vi)
    dotV :: Int -> Int -> Int -> f Double
dotV Int
k Int
j Int
i
        | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
n = Double -> f Double
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure Double
0
        | Bool
otherwise = do
            Double
vi <- MVector (PrimState f) Double -> Int -> f Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.read STVector s Double
MVector (PrimState f) Double
mat (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
i)
            Double
aij <- MVector (PrimState f) Double -> Int -> f Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.read STVector s Double
MVector (PrimState f) Double
mat (Int
j Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
i)
            Double
rest <- Int -> Int -> Int -> f Double
dotV Int
k Int
j (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
            Double -> f Double
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (Double
vi Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
aij Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
rest)
    dotRhs :: Int -> Int -> f Double
dotRhs Int
k Int
i
        | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
n = Double -> f Double
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure Double
0
        | Bool
otherwise = do
            Double
vi <- MVector (PrimState f) Double -> Int -> f Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.read STVector s Double
MVector (PrimState f) Double
mat (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
i)
            Double
bi <- MVector (PrimState f) Double -> Int -> f Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.read STVector s Double
MVector (PrimState f) Double
rhs Int
i
            Double
rest <- Int -> Int -> f Double
dotRhs Int
k (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
            Double -> f Double
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (Double
vi Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
bi Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
rest)

backSubstQR ::
    VUM.STVector s Double ->
    VUM.STVector s Double ->
    Int ->
    Int ->
    VUM.STVector s Double ->
    ST s ()
backSubstQR :: forall s.
STVector s Double
-> STVector s Double -> Int -> Int -> STVector s Double -> ST s ()
backSubstQR STVector s Double
mat STVector s Double
rhs Int
n Int
d STVector s Double
x = [Int] -> (Int -> ST s ()) -> ST s ()
forall (t :: * -> *) (m :: * -> *) a b.
(Foldable t, Monad m) =>
t a -> (a -> m b) -> m ()
forM_ [Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1, Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
2 .. Int
0] ((Int -> ST s ()) -> ST s ()) -> (Int -> ST s ()) -> ST s ()
forall a b. (a -> b) -> a -> b
$ \Int
i -> do
    Double
bi <- MVector (PrimState (ST s)) Double -> Int -> ST s Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.read STVector s Double
MVector (PrimState (ST s)) Double
rhs Int
i
    Double
s <- Int -> Int -> Double -> ST s Double
forall {f :: * -> *}.
(PrimState f ~ s, PrimMonad f) =>
Int -> Int -> Double -> f Double
sumAbove Int
i (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Double
0
    Double
rii <- MVector (PrimState (ST s)) Double -> Int -> ST s Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.read STVector s Double
MVector (PrimState (ST s)) Double
mat (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
i)
    MVector (PrimState (ST s)) Double -> Int -> Double -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.write STVector s Double
MVector (PrimState (ST s)) Double
x Int
i ((Double
bi Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
s) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
rii)
  where
    sumAbove :: Int -> Int -> Double -> f Double
sumAbove Int
i Int
j !Double
acc
        | Int
j Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
d = Double -> f Double
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure Double
acc
        | Bool
otherwise = do
            Double
rij <- MVector (PrimState f) Double -> Int -> f Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.read STVector s Double
MVector (PrimState f) Double
mat (Int
j Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
i)
            Double
xj <- MVector (PrimState f) Double -> Int -> f Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.read STVector s Double
MVector (PrimState f) Double
x Int
j
            Int -> Int -> Double -> f Double
sumAbove Int
i (Int
j Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) (Double
acc Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
rij Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
xj)

{- | Cholesky factor @L@ (lower-triangular, @A = L Lᵀ@) of a symmetric
positive-definite matrix, or 'Nothing' if a non-positive pivot is hit.
-}
cholesky :: Matrix -> Maybe Matrix
cholesky :: Matrix -> Maybe Matrix
cholesky Matrix
a
    | Matrix -> Bool
forall a. Vector a -> Bool
V.null Matrix
a = Matrix -> Maybe Matrix
forall a. a -> Maybe a
Just Matrix
forall a. Vector a
V.empty
    | Bool
otherwise = (forall s. ST s (Maybe Matrix)) -> Maybe Matrix
forall a. (forall s. ST s a) -> a
runST ((forall s. ST s (Maybe Matrix)) -> Maybe Matrix)
-> (forall s. ST s (Maybe Matrix)) -> Maybe Matrix
forall a b. (a -> b) -> a -> b
$ do
        MVector s Double
l <- 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 Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
d) Double
0
        Bool
ok <- MVector (PrimState (ST s)) Double -> ST s Bool
forall {f :: * -> *}.
PrimMonad f =>
MVector (PrimState f) Double -> f Bool
buildL MVector s Double
MVector (PrimState (ST s)) Double
l
        if Bool
ok then Matrix -> Maybe Matrix
forall a. a -> Maybe a
Just (Matrix -> Maybe Matrix) -> ST s Matrix -> ST s (Maybe Matrix)
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> MVector (PrimState (ST s)) Double -> ST s Matrix
forall {m :: * -> *} {a}.
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector (Vector a))
freezeLower MVector s Double
MVector (PrimState (ST s)) Double
l else Maybe Matrix -> ST s (Maybe Matrix)
forall a. a -> ST s a
forall (f :: * -> *) a. Applicative f => a -> f a
pure Maybe Matrix
forall a. Maybe a
Nothing
  where
    d :: Int
d = Matrix -> Int
forall a. Vector a -> Int
V.length Matrix
a
    buildL :: MVector (PrimState f) Double -> f Bool
buildL MVector (PrimState f) Double
l = Int -> f Bool
forall {f :: * -> *}.
(PrimState f ~ PrimState f, PrimMonad f) =>
Int -> f Bool
go Int
0
      where
        go :: Int -> f Bool
go Int
j
            | Int
j Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
d = Bool -> f Bool
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure Bool
True
            | Bool
otherwise = do
                Double
s <- MVector (PrimState f) Double
-> Int -> Int -> Int -> Double -> f Double
forall {f :: * -> *} {t}.
(PrimMonad f, Unbox t, Num t) =>
MVector (PrimState f) t -> Int -> Int -> Int -> t -> f t
sumLk MVector (PrimState f) Double
MVector (PrimState f) Double
l Int
j Int
j (Int
j Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1) Double
0
                let ajj :: Double
ajj = (Matrix
a Matrix -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
j) Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
j
                    diag :: Double
diag = Double
ajj Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
s
                if Double
diag Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
<= Double
0
                    then Bool -> f Bool
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure Bool
False
                    else do
                        let ljj :: Double
ljj = Double -> Double
forall a. Floating a => a -> a
sqrt Double
diag
                        MVector (PrimState f) Double -> Int -> Double -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.write MVector (PrimState f) Double
MVector (PrimState f) Double
l (Int
j Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
j) Double
ljj
                        [Int] -> (Int -> f ()) -> f ()
forall (t :: * -> *) (m :: * -> *) a b.
(Foldable t, Monad m) =>
t a -> (a -> m b) -> m ()
forM_ [Int
j Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1 .. Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1] ((Int -> f ()) -> f ()) -> (Int -> f ()) -> f ()
forall a b. (a -> b) -> a -> b
$ \Int
i -> do
                            Double
sij <- MVector (PrimState f) Double
-> Int -> Int -> Int -> Double -> f Double
forall {f :: * -> *} {t}.
(PrimMonad f, Unbox t, Num t) =>
MVector (PrimState f) t -> Int -> Int -> Int -> t -> f t
sumLk MVector (PrimState f) Double
MVector (PrimState f) Double
l Int
i Int
j (Int
j Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1) Double
0
                            let aij :: Double
aij = (Matrix
a Matrix -> 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
                            MVector (PrimState f) Double -> Int -> Double -> f ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.write MVector (PrimState f) Double
MVector (PrimState f) Double
l (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
j) ((Double
aij Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
sij) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
ljj)
                        Int -> f Bool
go (Int
j Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
    sumLk :: MVector (PrimState f) t -> Int -> Int -> Int -> t -> f t
sumLk MVector (PrimState f) t
l Int
i Int
j Int
k !t
acc
        | Int
k Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
0 = t -> f t
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure t
acc
        | Bool
otherwise = do
            t
lik <- MVector (PrimState f) t -> Int -> f t
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.read MVector (PrimState f) t
l (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
k)
            t
ljk <- MVector (PrimState f) t -> Int -> f t
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.read MVector (PrimState f) t
l (Int
j Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
k)
            MVector (PrimState f) t -> Int -> Int -> Int -> t -> f t
sumLk MVector (PrimState f) t
l Int
i Int
j (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1) (t
acc t -> t -> t
forall a. Num a => a -> a -> a
+ t
lik t -> t -> t
forall a. Num a => a -> a -> a
* t
ljk)
    freezeLower :: MVector (PrimState m) a -> m (Vector (Vector a))
freezeLower MVector (PrimState m) a
l = do
        Vector a
frozen <- MVector (PrimState m) a -> m (Vector a)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.freeze MVector (PrimState m) a
l
        Vector (Vector a) -> m (Vector (Vector a))
forall a. a -> m a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (Vector (Vector a) -> m (Vector (Vector a)))
-> Vector (Vector a) -> m (Vector (Vector a))
forall a b. (a -> b) -> a -> b
$ Int -> (Int -> Vector a) -> Vector (Vector a)
forall a. Int -> (Int -> a) -> Vector a
V.generate Int
d ((Int -> Vector a) -> Vector (Vector a))
-> (Int -> Vector a) -> Vector (Vector a)
forall a b. (a -> b) -> a -> b
$ \Int
i -> Int -> Int -> Vector a -> Vector a
forall a. Unbox a => Int -> Int -> Vector a -> Vector a
VU.slice (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
d) Int
d Vector a
frozen

-- | Solve @L y = b@ for lower-triangular @L@.
forwardSubst :: Matrix -> VU.Vector Double -> VU.Vector Double
forwardSubst :: Matrix -> Vector Double -> Vector Double
forwardSubst Matrix
l Vector Double
b = (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
y <- Int -> ST s (MVector (PrimState (ST s)) Double)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
d
    [Int] -> (Int -> ST s ()) -> ST s ()
forall (t :: * -> *) (m :: * -> *) a b.
(Foldable t, Monad m) =>
t a -> (a -> m b) -> m ()
forM_ [Int
0 .. Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1] ((Int -> ST s ()) -> ST s ()) -> (Int -> ST s ()) -> ST s ()
forall a b. (a -> b) -> a -> b
$ \Int
i -> do
        let row :: Vector Double
row = Matrix
l Matrix -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
i
        Double
s <- MVector (PrimState (ST s)) Double
-> Vector Double -> Int -> Int -> Double -> ST s Double
forall {f :: * -> *} {t}.
(PrimMonad f, Unbox t, Num t) =>
MVector (PrimState f) t -> Vector t -> Int -> Int -> t -> f t
sumKnown MVector s Double
MVector (PrimState (ST s)) Double
y Vector Double
row Int
i Int
0 Double
0
        MVector (PrimState (ST s)) Double -> Int -> Double -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.write MVector s Double
MVector (PrimState (ST s)) Double
y Int
i ((Vector Double
b 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) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ (Vector Double
row Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i))
    MVector (PrimState (ST s)) Double -> ST s (Vector Double)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.freeze MVector s Double
MVector (PrimState (ST s)) Double
y
  where
    d :: Int
d = Matrix -> Int
forall a. Vector a -> Int
V.length Matrix
l
    sumKnown :: MVector (PrimState f) t -> Vector t -> Int -> Int -> t -> f t
sumKnown MVector (PrimState f) t
y Vector t
row Int
i Int
j !t
acc
        | Int
j Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
i = t -> f t
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure t
acc
        | Bool
otherwise = do
            t
yj <- MVector (PrimState f) t -> Int -> f t
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.read MVector (PrimState f) t
y Int
j
            MVector (PrimState f) t -> Vector t -> Int -> Int -> t -> f t
sumKnown MVector (PrimState f) t
y Vector t
row Int
i (Int
j Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) (t
acc t -> t -> t
forall a. Num a => a -> a -> a
+ (Vector t
row Vector t -> Int -> t
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
j) t -> t -> t
forall a. Num a => a -> a -> a
* t
yj)

-- | Solve @Lᵀ x = y@ for lower-triangular @L@.
backSubst :: Matrix -> VU.Vector Double -> VU.Vector Double
backSubst :: Matrix -> Vector Double -> Vector Double
backSubst Matrix
l Vector Double
y = (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
x <- Int -> ST s (MVector (PrimState (ST s)) Double)
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
Int -> m (MVector (PrimState m) a)
VUM.new Int
d
    [Int] -> (Int -> ST s ()) -> ST s ()
forall (t :: * -> *) (m :: * -> *) a b.
(Foldable t, Monad m) =>
t a -> (a -> m b) -> m ()
forM_ [Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1, Int
d Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
2 .. Int
0] ((Int -> ST s ()) -> ST s ()) -> (Int -> ST s ()) -> ST s ()
forall a b. (a -> b) -> a -> b
$ \Int
i -> do
        Double
s <- MVector (PrimState (ST s)) Double
-> Int -> Int -> Double -> ST s Double
forall {f :: * -> *}.
PrimMonad f =>
MVector (PrimState f) Double -> Int -> Int -> Double -> f Double
sumKnown MVector s Double
MVector (PrimState (ST s)) Double
x Int
i (Int
i Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) Double
0
        MVector (PrimState (ST s)) Double -> Int -> Double -> ST s ()
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> a -> m ()
VUM.write MVector s Double
MVector (PrimState (ST s)) Double
x Int
i ((Vector Double
y 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) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ ((Matrix
l Matrix -> 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))
    MVector (PrimState (ST s)) Double -> ST s (Vector Double)
forall a (m :: * -> *).
(Unbox a, PrimMonad m) =>
MVector (PrimState m) a -> m (Vector a)
VU.freeze MVector s Double
MVector (PrimState (ST s)) Double
x
  where
    d :: Int
d = Matrix -> Int
forall a. Vector a -> Int
V.length Matrix
l
    sumKnown :: MVector (PrimState f) Double -> Int -> Int -> Double -> f Double
sumKnown MVector (PrimState f) Double
x Int
i Int
j !Double
acc
        | Int
j Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
d = Double -> f Double
forall a. a -> f a
forall (f :: * -> *) a. Applicative f => a -> f a
pure Double
acc
        | Bool
otherwise = do
            Double
xj <- MVector (PrimState f) Double -> Int -> f Double
forall (m :: * -> *) a.
(PrimMonad m, Unbox a) =>
MVector (PrimState m) a -> Int -> m a
VUM.read MVector (PrimState f) Double
x Int
j
            MVector (PrimState f) Double -> Int -> Int -> Double -> f Double
sumKnown MVector (PrimState f) Double
x Int
i (Int
j Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1) (Double
acc Double -> Double -> Double
forall a. Num a => a -> a -> a
+ ((Matrix
l Matrix -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
j) 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
xj)

{- | Solve the SPD system @A x = b@ via Cholesky; 'Nothing' when @A@ is not
positive-definite.
-}
choleskySolve :: Matrix -> VU.Vector Double -> Maybe (VU.Vector Double)
choleskySolve :: Matrix -> Vector Double -> Maybe (Vector Double)
choleskySolve Matrix
a Vector Double
b = do
    Matrix
l <- Matrix -> Maybe Matrix
cholesky Matrix
a
    Vector Double -> Maybe (Vector Double)
forall a. a -> Maybe a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (Matrix -> Vector Double -> Vector Double
backSubst Matrix
l (Matrix -> Vector Double -> Vector Double
forwardSubst Matrix
l Vector Double
b))

-- | @log det A = 2 Σ log Lᵢᵢ@ from a Cholesky factor @L@.
logDetFromChol :: Matrix -> Double
logDetFromChol :: Matrix -> Double
logDetFromChol Matrix
l = Double
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Vector Double -> Double
forall a. Num a => Vector a -> a
V.sum ((Int -> Vector Double -> Double) -> Matrix -> Vector Double
forall a b. (Int -> a -> b) -> Vector a -> Vector b
V.imap (\Int
i Vector Double
row -> Double -> Double
forall a. Floating a => a -> a
log (Vector Double
row Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
i)) Matrix
l)