{-# LANGUAGE BangPatterns #-}

{- | Dependency-free dense linear algebra over row-major matrices, shared by the
models in @dataframe-learn@. Solvers live in "DataFrame.LinearAlgebra.Solve"
and eigenproblems in "DataFrame.LinearAlgebra.Eigen".
-}
module DataFrame.LinearAlgebra (
    Matrix,
    dot,
    axpy,
    scaleV,
    matVec,
    tMatVec,
    gram,
    transposeM,
    identityM,
    logSumExp,
    sqDist,
    nearestCenter,
    epsNeighbors,
) where

import qualified Data.Vector as V
import qualified Data.Vector.Unboxed as VU

{- | Row-major dense matrix: an outer boxed vector of equal-length rows. An
@n×d@ matrix has @n@ rows of length @d@.
-}
type Matrix = V.Vector (VU.Vector Double)

-- | Inner product of two equal-length vectors.
dot :: VU.Vector Double -> VU.Vector Double -> Double
dot :: Vector Double -> Vector Double -> Double
dot Vector Double
a Vector Double
b = 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
a Vector Double
b)
{-# INLINE dot #-}

-- | @axpy a x y = a*x + y@.
axpy :: Double -> VU.Vector Double -> VU.Vector Double -> VU.Vector Double
axpy :: Double -> Vector Double -> Vector Double -> Vector Double
axpy Double
a = (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
xi Double
yi -> Double
a Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
xi Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
yi)
{-# INLINE axpy #-}

-- | Scalar-vector product.
scaleV :: Double -> VU.Vector Double -> VU.Vector Double
scaleV :: Double -> Vector Double -> Vector Double
scaleV Double
a = (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
a)
{-# INLINE scaleV #-}

-- | @matVec A v@ for @A@ of shape @n×d@ and @v@ of length @d@; result length @n@.
matVec :: Matrix -> VU.Vector Double -> VU.Vector Double
matVec :: Matrix -> Vector Double -> Vector Double
matVec Matrix
a Vector Double
v = Vector Double -> Vector Double
forall (v :: * -> *) a (w :: * -> *).
(Vector v a, Vector w a) =>
v a -> w a
VU.convert ((Vector Double -> Double) -> Matrix -> Vector Double
forall a b. (a -> b) -> Vector a -> Vector b
V.map (Vector Double -> Vector Double -> Double
`dot` Vector Double
v) Matrix
a)

-- | @tMatVec A v = Aᵀ v@ for @A@ of shape @n×d@, @v@ of length @n@; result length @d@.
tMatVec :: Matrix -> VU.Vector Double -> VU.Vector Double
tMatVec :: Matrix -> Vector Double -> Vector Double
tMatVec Matrix
a Vector Double
v
    | Matrix -> Bool
forall a. Vector a -> Bool
V.null Matrix
a = Vector Double
forall a. Unbox a => Vector a
VU.empty
    | Bool
otherwise = (Vector Double -> (Double, Vector Double) -> Vector Double)
-> Vector Double -> Vector (Double, Vector Double) -> Vector Double
forall a b. (a -> b -> a) -> a -> Vector b -> a
V.foldl' Vector Double -> (Double, Vector Double) -> Vector Double
step (Int -> Double -> Vector Double
forall a. Unbox a => Int -> a -> Vector a
VU.replicate Int
d Double
0) ((Double -> Vector Double -> (Double, Vector Double))
-> Vector Double -> Matrix -> Vector (Double, Vector Double)
forall a b c. (a -> b -> c) -> Vector a -> Vector b -> Vector c
V.zipWith (,) Vector Double
vBoxed Matrix
a)
  where
    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)
    vBoxed :: Vector Double
vBoxed = Int -> (Int -> Double) -> Vector Double
forall a. Int -> (Int -> a) -> Vector a
V.generate (Matrix -> Int
forall a. Vector a -> Int
V.length Matrix
a) (Vector Double
v Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.!)
    step :: Vector Double -> (Double, Vector Double) -> Vector Double
step !Vector Double
acc (Double
vi, Vector Double
row) = Double -> Vector Double -> Vector Double -> Vector Double
axpy Double
vi Vector Double
row Vector Double
acc

-- | @gram A = Aᵀ A@, the @d×d@ symmetric matrix of column inner products.
gram :: Matrix -> Matrix
gram :: Matrix -> Matrix
gram Matrix
a
    | Matrix -> Bool
forall a. Vector a -> Bool
V.null Matrix
a = Matrix
forall a. Vector a
V.empty
    | Bool
otherwise =
        Int -> (Int -> Vector Double) -> Matrix
forall a. Int -> (Int -> a) -> Vector a
V.generate Int
d ((Int -> Vector Double) -> Matrix)
-> (Int -> Vector Double) -> Matrix
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 ->
                (Double -> Vector Double -> Double) -> Double -> Matrix -> Double
forall a b. (a -> b -> a) -> a -> Vector b -> a
V.foldl' (\ !Double
acc Vector Double
row -> Double
acc Double -> Double -> Double
forall a. Num a => a -> a -> a
+ (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
* (Vector Double
row Vector Double -> Int -> Double
forall a. Unbox a => Vector a -> Int -> a
VU.! Int
j)) Double
0 Matrix
a
  where
    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)

-- | Transpose an @n×d@ matrix to @d×n@.
transposeM :: Matrix -> Matrix
transposeM :: Matrix -> Matrix
transposeM Matrix
a
    | Matrix -> Bool
forall a. Vector a -> Bool
V.null Matrix
a = Matrix
forall a. Vector a
V.empty
    | Bool
otherwise = Int -> (Int -> Vector Double) -> Matrix
forall a. Int -> (Int -> a) -> Vector a
V.generate Int
d ((Int -> Vector Double) -> Matrix)
-> (Int -> Vector Double) -> Matrix
forall a b. (a -> b) -> a -> b
$ \Int
j -> Int -> (Int -> Double) -> Vector Double
forall a. Unbox a => Int -> (Int -> a) -> Vector a
VU.generate Int
n ((Int -> Double) -> Vector Double)
-> (Int -> Double) -> Vector Double
forall a b. (a -> b) -> a -> b
$ \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
  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)

-- | @d×d@ identity matrix.
identityM :: Int -> Matrix
identityM :: Int -> Matrix
identityM Int
d = Int -> (Int -> Vector Double) -> Matrix
forall a. Int -> (Int -> a) -> Vector a
V.generate Int
d ((Int -> Vector Double) -> Matrix)
-> (Int -> Vector Double) -> Matrix
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 -> if Int
i Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
j then Double
1 else Double
0

-- | Numerically stable @log Σ exp xᵢ@.
logSumExp :: VU.Vector Double -> Double
logSumExp :: Vector Double -> Double
logSumExp Vector Double
xs
    | Vector Double -> Bool
forall a. Unbox a => Vector a -> Bool
VU.null Vector Double
xs = Double -> Double
forall a. Num a => a -> a
negate (Double
1 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
0)
    | Bool
otherwise = Double
m Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double -> Double
forall a. Floating a => a -> a
log (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
x -> Double -> Double
forall a. Floating a => a -> a
exp (Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
m)) Vector Double
xs))
  where
    m :: Double
m = Vector Double -> Double
forall a. (Unbox a, Ord a) => Vector a -> a
VU.maximum Vector Double
xs

-- | Squared Euclidean distance.
sqDist :: VU.Vector Double -> VU.Vector Double -> Double
sqDist :: Vector Double -> Vector Double -> Double
sqDist Vector Double
a Vector Double
b = 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
x Double
y -> let z :: Double
z = Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
y in Double
z Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
z) Vector Double
a Vector Double
b)
{-# INLINE sqDist #-}

-- | Index of and squared distance to the nearest centre.
nearestCenter ::
    V.Vector (VU.Vector Double) -> VU.Vector Double -> (Int, Double)
nearestCenter :: Matrix -> Vector Double -> (Int, Double)
nearestCenter Matrix
centers Vector Double
p =
    ((Int, Double) -> Int -> Vector Double -> (Int, Double))
-> (Int, Double) -> Matrix -> (Int, Double)
forall a b. (a -> Int -> b -> a) -> a -> Vector b -> a
V.ifoldl'
        ( \(!Int
bi, !Double
bd) Int
i Vector Double
c ->
            let dd :: Double
dd = Vector Double -> Vector Double -> Double
sqDist Vector Double
c Vector Double
p in if Double
dd Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
< Double
bd then (Int
i, Double
dd) else (Int
bi, Double
bd)
        )
        (-Int
1, Double
1 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
0)
        Matrix
centers

{- | Indices @j@ (excluding @i@) within squared radius @eps²@ of row @i@, by
brute force; @O(n d)@ per query.
-}
epsNeighbors :: Double -> Matrix -> Int -> VU.Vector Int
epsNeighbors :: Double -> Matrix -> Int -> Vector Int
epsNeighbors Double
eps Matrix
rows Int
i =
    [Int] -> Vector Int
forall a. Unbox a => [a] -> Vector a
VU.fromList
        [ Int
j
        | Int
j <- [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]
        , Int
j Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
/= Int
i
        , Vector Double -> Vector Double -> Double
sqDist (Matrix
rows Matrix -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
i) (Matrix
rows Matrix -> Int -> Vector Double
forall a. Vector a -> Int -> a
V.! Int
j) Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
<= Double
eps2
        ]
  where
    n :: Int
n = Matrix -> Int
forall a. Vector a -> Int
V.length Matrix
rows
    eps2 :: Double
eps2 = Double
eps Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
eps