{-# LANGUAGE LambdaCase #-}
{-# LANGUAGE RebindableSyntax #-}
{-# OPTIONS_GHC -Wno-incomplete-uni-patterns #-}

-- | Value-sized dense matrices backed by 'Harpie.Array'.
--
-- This module provides the canonical dense-matrix carrier for numhask-based
-- computation: Kleene star (reflexive-transitive closure), matrix
-- multiplication, and matrix-vector products.  It is intentionally rank-2 and
-- value-sized so it can serve categories whose object constraints live at
-- runtime (e.g. 'Circuit.Mat.Finite' enumeration or plain list channels) as
-- well as typed settings.
--
-- Moved from @harpie-numhask@ as part of the matrix-calculus extraction into
-- @circuits-mat@.
module Circuit.Mat.Dense
  ( Matrix (..),
    fromLists,
    toLists,
    matPlus,
    matTimes,
    matVec,
    starMatrix,
    qrM,
    forwardSubstStream,
    solve2,
  )
where

import Circuit.Mat.Array.Stream qualified as Stream
import Data.Bool (bool)
import Data.Foldable hiding (sum)
import Data.List (foldl')
import Data.Vector.Unboxed qualified as VU
import Harpie.Array as A
import NumHask.Algebra.Additive (Additive (..), Subtractive (..), sum)
import NumHask.Algebra.Field (ExpField (..))
import NumHask.Algebra.Metric (Absolute, abs)
import NumHask.Algebra.Multiplicative (Divisive (..), Multiplicative (..))
import NumHask.Algebra.Ring (StarSemiring (..))
import Prelude hiding (drop, foldl', length, negate, repeat, sqrt, sum, take, zipWith, (*), (+), (-), (/))
import Prelude qualified as P

-- | Square matrix stored as a rank-2 'Harpie.Array' in row-major order.
newtype Matrix a = Matrix {forall a. Matrix a -> Array a
unMatrix :: A.Array a}
  deriving (Matrix a -> Matrix a -> Bool
(Matrix a -> Matrix a -> Bool)
-> (Matrix a -> Matrix a -> Bool) -> Eq (Matrix a)
forall a. Eq a => Matrix a -> Matrix a -> Bool
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: forall a. Eq a => Matrix a -> Matrix a -> Bool
== :: Matrix a -> Matrix a -> Bool
$c/= :: forall a. Eq a => Matrix a -> Matrix a -> Bool
/= :: Matrix a -> Matrix a -> Bool
Eq, Int -> Matrix a -> ShowS
[Matrix a] -> ShowS
Matrix a -> String
(Int -> Matrix a -> ShowS)
-> (Matrix a -> String) -> ([Matrix a] -> ShowS) -> Show (Matrix a)
forall a. Show a => Int -> Matrix a -> ShowS
forall a. Show a => [Matrix a] -> ShowS
forall a. Show a => Matrix a -> String
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: forall a. Show a => Int -> Matrix a -> ShowS
showsPrec :: Int -> Matrix a -> ShowS
$cshow :: forall a. Show a => Matrix a -> String
show :: Matrix a -> String
$cshowList :: forall a. Show a => [Matrix a] -> ShowS
showList :: [Matrix a] -> ShowS
Show)

-- | Row count.
rows :: Matrix a -> Int
rows :: forall a. Matrix a -> Int
rows (Matrix Array a
a) = case Vector Int -> [Int]
forall a. Unbox a => Vector a -> [a]
VU.toList (Array a -> Vector Int
forall a. Array a -> Vector Int
A.shape Array a
a) of (Int
r : [Int]
_) -> Int
r; [Int]
_ -> Int
0

-- | Column count.
cols :: Matrix a -> Int
cols :: forall a. Matrix a -> Int
cols (Matrix Array a
a) = case Vector Int -> [Int]
forall a. Unbox a => Vector a -> [a]
VU.toList (Array a -> Vector Int
forall a. Array a -> Vector Int
A.shape Array a
a) of (Int
_ : Int
c : [Int]
_) -> Int
c; [Int]
_ -> Int
0

-- | Build a matrix from nested rows.
--
-- An empty list becomes a 0×0 matrix.
fromLists :: [[a]] -> Matrix a
fromLists :: forall a. [[a]] -> Matrix a
fromLists [] = Array a -> Matrix a
forall a. Array a -> Matrix a
Matrix ([Int] -> [a] -> Array a
forall t a. FromVector t a => [Int] -> t -> Array a
A.array [Int
0, Int
0] [])
fromLists xss :: [[a]]
xss@([a]
r : [[a]]
_) = Array a -> Matrix a
forall a. Array a -> Matrix a
Matrix ([Int] -> [a] -> Array a
forall t a. FromVector t a => [Int] -> t -> Array a
A.array [[[a]] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
P.length [[a]]
xss, [a] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
P.length [a]
r] ([[a]] -> [a]
forall (t :: * -> *) a. Foldable t => t [a] -> [a]
P.concat [[a]]
xss))

-- | Convert a matrix to nested rows.
toLists :: Matrix a -> [[a]]
toLists :: forall a. Matrix a -> [[a]]
toLists (Matrix Array a
a) =
  let r :: Int
r = Matrix a -> Int
forall a. Matrix a -> Int
rows (Array a -> Matrix a
forall a. Array a -> Matrix a
Matrix Array a
a)
      c :: Int
c = Matrix a -> Int
forall a. Matrix a -> Int
cols (Array a -> Matrix a
forall a. Array a -> Matrix a
Matrix Array a
a)
   in [[Array a
a Array a -> [Int] -> a
forall a. Array a -> [Int] -> a
A.! [Int
i, Int
j] | Int
j <- [Int
0 .. Int
c Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1]] | Int
i <- [Int
0 .. Int
r Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1]]

-- | Elementwise addition.
matPlus :: (Additive a) => Matrix a -> Matrix a -> Matrix a
matPlus :: forall a. Additive a => Matrix a -> Matrix a -> Matrix a
matPlus (Matrix Array a
a) (Matrix Array a
b) = Array a -> Matrix a
forall a. Array a -> Matrix a
Matrix ((a -> a -> a) -> Array a -> Array a -> Array a
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
A.zipWith a -> a -> a
forall a. Additive a => a -> a -> a
(+) Array a
a Array a
b)

-- | Matrix multiplication.
matTimes ::
  (Additive a, Multiplicative a) =>
  Matrix a ->
  Matrix a ->
  Matrix a
matTimes :: forall a.
(Additive a, Multiplicative a) =>
Matrix a -> Matrix a -> Matrix a
matTimes (Matrix Array a
a) (Matrix Array a
b) =
  case (Vector Int -> [Int]
forall a. Unbox a => Vector a -> [a]
VU.toList (Array a -> Vector Int
forall a. Array a -> Vector Int
A.shape Array a
a), Vector Int -> [Int]
forall a. Unbox a => Vector a -> [a]
VU.toList (Array a -> Vector Int
forall a. Array a -> Vector Int
A.shape Array a
b)) of
    ([Int
ra, Int
ca], [Int
rb, Int
cb]) ->
      case Int
ca Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
rb of
        Bool
False -> String -> Matrix a
forall a. HasCallStack => String -> a
error String
"Circuit.Mat.Dense.matTimes: inner dimension mismatch"
        Bool
True ->
          Array a -> Matrix a
forall a. Array a -> Matrix a
Matrix (Array a -> Matrix a) -> Array a -> Matrix a
forall a b. (a -> b) -> a -> b
$
            [Int] -> ([Int] -> a) -> Array a
forall a. [Int] -> ([Int] -> a) -> Array a
A.tabulate
              [Int
ra, Int
cb]
              ( \case
                  [Int
i, Int
j] -> [a] -> a
forall a (f :: * -> *). (Additive a, Foldable f) => f a -> a
sum [Array a
a Array a -> [Int] -> a
forall a. Array a -> [Int] -> a
A.! [Int
i, Int
k] a -> a -> a
forall a. Multiplicative a => a -> a -> a
* Array a
b Array a -> [Int] -> a
forall a. Array a -> [Int] -> a
A.! [Int
k, Int
j] | Int
k <- [Int
0 .. Int
ca Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1]]
                  [Int]
_ -> String -> a
forall a. HasCallStack => String -> a
error String
"Circuit.Mat.Dense.matTimes: expected rank-2 index"
              )
    ([Int], [Int])
_ -> String -> Matrix a
forall a. HasCallStack => String -> a
error String
"Circuit.Mat.Dense.matTimes: expected rank-2 matrices"

-- | Matrix–vector product.
matVec ::
  (Additive a, Multiplicative a) =>
  Matrix a ->
  [a] ->
  [a]
matVec :: forall a. (Additive a, Multiplicative a) => Matrix a -> [a] -> [a]
matVec (Matrix Array a
a) [a]
v =
  case Vector Int -> [Int]
forall a. Unbox a => Vector a -> [a]
VU.toList (Array a -> Vector Int
forall a. Array a -> Vector Int
A.shape Array a
a) of
    [Int
r, Int
c] ->
      let n :: Int
n = [a] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
P.length [a]
v
       in case Int
c Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
n of
            Bool
False -> String -> [a]
forall a. HasCallStack => String -> a
error String
"Circuit.Mat.Dense.matVec: dimension mismatch"
            Bool
True ->
              [ [a] -> a
forall a (f :: * -> *). (Additive a, Foldable f) => f a -> a
sum [Array a
a Array a -> [Int] -> a
forall a. Array a -> [Int] -> a
A.! [Int
i, Int
k] a -> a -> a
forall a. Multiplicative a => a -> a -> a
* ([a]
v [a] -> Int -> a
forall a. HasCallStack => [a] -> Int -> a
P.!! Int
k) | Int
k <- [Int
0 .. Int
c Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1]]
              | Int
i <- [Int
0 .. Int
r Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1]
              ]
    [Int]
_ -> String -> [a]
forall a. HasCallStack => String -> a
error String
"Circuit.Mat.Dense.matVec: expected a rank-2 matrix"

-- | Kleene star of a square matrix by the standard state-elimination
-- (Warshall / Floyd-Kleene) algorithm.
--
-- For a matrix @A@, computes @A* = I + A + A² + ...@ as the least fixed
-- point of @X ↦ I + A·X@. Requires a 'StarSemiring' element type so that
-- @star x@ is available for the pivot updates.
--
-- When the element type is a 'NumHask.Algebra.Quantale.Quantale', this is the
-- join of the geometric series @I + A + A² + …@; the algorithm uses the
-- 'StarSemiring' fragment (iterative joins) to compute it finitely.
starMatrix ::
  (StarSemiring a) =>
  Matrix a ->
  Matrix a
starMatrix :: forall a. StarSemiring a => Matrix a -> Matrix a
starMatrix (Matrix Array a
a) =
  let sh :: [Int]
sh = Vector Int -> [Int]
forall a. Unbox a => Vector a -> [a]
VU.toList (Array a -> Vector Int
forall a. Array a -> Vector Int
A.shape Array a
a)
   in case [Int]
sh of
        [Int
0, Int
0] -> Array a -> Matrix a
forall a. Array a -> Matrix a
Matrix Array a
a
        [Int
n, Int
m]
          | Int
n Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
m ->
              let step :: Array a -> Int -> Array a
step Array a
arr Int
k =
                    [Int] -> ([Int] -> a) -> Array a
forall a. [Int] -> ([Int] -> a) -> Array a
A.tabulate
                      [Int
n, Int
n]
                      ( \case
                          [Int
i, Int
j] ->
                            let aik :: a
aik = Array a -> [Int] -> a
forall a. Array a -> [Int] -> a
A.index Array a
arr [Int
i, Int
k]
                                akk :: a
akk = Array a -> [Int] -> a
forall a. Array a -> [Int] -> a
A.index Array a
arr [Int
k, Int
k]
                                akj :: a
akj = Array a -> [Int] -> a
forall a. Array a -> [Int] -> a
A.index Array a
arr [Int
k, Int
j]
                                aij :: a
aij = Array a -> [Int] -> a
forall a. Array a -> [Int] -> a
A.index Array a
arr [Int
i, Int
j]
                             in a
aij a -> a -> a
forall a. Additive a => a -> a -> a
+ a
aik a -> a -> a
forall a. Multiplicative a => a -> a -> a
* a -> a
forall a. StarSemiring a => a -> a
star a
akk a -> a -> a
forall a. Multiplicative a => a -> a -> a
* a
akj
                          [Int]
_ -> String -> a
forall a. HasCallStack => String -> a
error String
"Circuit.Mat.Dense.starMatrix: expected rank-2 index"
                      )
                  closed :: Array a
closed = (Array a -> Int -> Array a) -> Array a -> [Int] -> Array a
forall b a. (b -> a -> b) -> b -> [a] -> b
forall (t :: * -> *) b a.
Foldable t =>
(b -> a -> b) -> b -> t a -> b
foldl' Array a -> Int -> Array a
step Array a
a [Int
0 .. Int
n Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1]
               in -- Floyd-Kleene on A is A⁺. Seed I afterwards, matching 'starM':
                  -- A* = I + A⁺.
                  Array a -> Matrix a
forall a. Array a -> Matrix a
Matrix (Array a -> Matrix a) -> Array a -> Matrix a
forall a b. (a -> b) -> a -> b
$
                    [Int] -> ([Int] -> a) -> Array a
forall a. [Int] -> ([Int] -> a) -> Array a
A.tabulate [Int
n, Int
n] (([Int] -> a) -> Array a) -> ([Int] -> a) -> Array a
forall a b. (a -> b) -> a -> b
$ \case
                      [Int
i, Int
j] -> Array a -> [Int] -> a
forall a. Array a -> [Int] -> a
A.index Array a
closed [Int
i, Int
j] a -> a -> a
forall a. Additive a => a -> a -> a
+ a -> a -> Bool -> a
forall a. a -> a -> Bool -> a
bool a
forall a. Additive a => a
zero a
forall a. Multiplicative a => a
one (Int
i Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
j)
                      [Int]
_ -> String -> a
forall a. HasCallStack => String -> a
error String
"Circuit.Mat.Dense.starMatrix: expected rank-2 index"
        [Int]
_ -> String -> Matrix a
forall a. HasCallStack => String -> a
error String
"Circuit.Mat.Dense.starMatrix: expected a square matrix"

-- | QR decomposition via Householder reflections.
--
-- Returns @(q, r)@ with @q@ orthogonal and @r@ upper triangular such that
-- @a = q r@.  The algorithm is the standard column-by-column Householder
-- reduction on value-sized 'Harpie.Array' matrices.
qrM ::
  ( ExpField a,
    Ord a
  ) =>
  Matrix a ->
  (Matrix a, Matrix a)
qrM :: forall a. (ExpField a, Ord a) => Matrix a -> (Matrix a, Matrix a)
qrM (Matrix Array a
a) =
  let sh :: [Int]
sh = Vector Int -> [Int]
forall a. Unbox a => Vector a -> [a]
VU.toList (Array a -> Vector Int
forall a. Array a -> Vector Int
A.shape Array a
a)
   in case [Int]
sh of
        [Int
m, Int
n] ->
          let q0 :: Matrix a
q0 = Array a -> Matrix a
forall a. Array a -> Matrix a
Matrix ([Int] -> Array a
forall a. (Additive a, Multiplicative a) => [Int] -> Array a
A.ident [Int
m, Int
m])
              r0 :: Matrix a
r0 = Array a -> Matrix a
forall a. Array a -> Matrix a
Matrix Array a
a
              steps :: [Int]
steps = [Int
0 .. Int -> Int -> Int
forall a. Ord a => a -> a -> a
P.min Int
m Int
n Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1]
              (Matrix a
q, Matrix a
r) = ((Matrix a, Matrix a) -> Int -> (Matrix a, Matrix a))
-> (Matrix a, Matrix a) -> [Int] -> (Matrix a, Matrix a)
forall b a. (b -> a -> b) -> b -> [a] -> b
forall (t :: * -> *) b a.
Foldable t =>
(b -> a -> b) -> b -> t a -> b
foldl' (Int -> Int -> (Matrix a, Matrix a) -> Int -> (Matrix a, Matrix a)
forall a.
(ExpField a, Ord a) =>
Int -> Int -> (Matrix a, Matrix a) -> Int -> (Matrix a, Matrix a)
householderQRStep Int
m Int
n) (Matrix a
q0, Matrix a
r0) [Int]
steps
           in (Matrix a
q, Matrix a
r)
        [Int]
_ -> String -> (Matrix a, Matrix a)
forall a. HasCallStack => String -> a
error String
"Circuit.Mat.Dense.qrM: expected a rank-2 matrix"

-- | One Householder QR step for column @k@.
householderQRStep ::
  ( ExpField a,
    Ord a
  ) =>
  Int ->
  Int ->
  (Matrix a, Matrix a) ->
  Int ->
  (Matrix a, Matrix a)
householderQRStep :: forall a.
(ExpField a, Ord a) =>
Int -> Int -> (Matrix a, Matrix a) -> Int -> (Matrix a, Matrix a)
householderQRStep Int
m Int
n (Matrix Array a
q, Matrix Array a
r) Int
k =
  let mk :: Int
mk = Int
m Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
k
      nk :: Int
nk = Int
n Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
k
      -- subcolumn x = r[k..m-1, k]
      x :: Array a
x = [Int] -> ([Int] -> a) -> Array a
forall a. [Int] -> ([Int] -> a) -> Array a
A.tabulate [Int
mk] (\[Int
i] -> Array a
r Array a -> [Int] -> a
forall a. Array a -> [Int] -> a
A.! [Int
k Int -> Int -> Int
forall a. Additive a => a -> a -> a
+ Int
i, Int
k])
      xk :: a
xk = Array a
x Array a -> [Int] -> a
forall a. Array a -> [Int] -> a
A.! [Int
0]
      norm :: a
norm = a -> a
forall a. ExpField a => a -> a
sqrt ([a] -> a
forall a (f :: * -> *). (Additive a, Foldable f) => f a -> a
sum [(Array a
x Array a -> [Int] -> a
forall a. Array a -> [Int] -> a
A.! [Int
i]) a -> a -> a
forall a. Multiplicative a => a -> a -> a
* (Array a
x Array a -> [Int] -> a
forall a. Array a -> [Int] -> a
A.! [Int
i]) | Int
i <- [Int
0 .. Int
mk Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1]])
      alpha :: a
alpha = a -> a -> Bool -> a
forall a. a -> a -> Bool -> a
bool (a -> a
forall a. Subtractive a => a -> a
negate a
norm) a
norm (a
xk a -> a -> Bool
forall a. Ord a => a -> a -> Bool
< a
forall a. Additive a => a
zero)
      v :: Array a
v = [Int] -> ([Int] -> a) -> Array a
forall a. [Int] -> ([Int] -> a) -> Array a
A.tabulate [Int
mk] (\[Int
i] -> a -> a -> Bool -> a
forall a. a -> a -> Bool -> a
bool (Array a
x Array a -> [Int] -> a
forall a. Array a -> [Int] -> a
A.! [Int
i]) (a
xk a -> a -> a
forall a. Subtractive a => a -> a -> a
- a
alpha) (Int
i Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
0))
      vtv :: a
vtv = [a] -> a
forall a (f :: * -> *). (Additive a, Foldable f) => f a -> a
sum [(Array a
v Array a -> [Int] -> a
forall a. Array a -> [Int] -> a
A.! [Int
i]) a -> a -> a
forall a. Multiplicative a => a -> a -> a
* (Array a
v Array a -> [Int] -> a
forall a. Array a -> [Int] -> a
A.! [Int
i]) | Int
i <- [Int
0 .. Int
mk Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1]]
   in (Matrix a, Matrix a)
-> (Matrix a, Matrix a) -> Bool -> (Matrix a, Matrix a)
forall a. a -> a -> Bool -> a
bool
        (Array a -> Matrix a
forall a. Array a -> Matrix a
Matrix Array a
q, Array a -> Matrix a
forall a. Array a -> Matrix a
Matrix Array a
r)
        ( let scale :: a
scale = (a
forall a. Multiplicative a => a
one a -> a -> a
forall a. Additive a => a -> a -> a
+ a
forall a. Multiplicative a => a
one) a -> a -> a
forall a. Divisive a => a -> a -> a
/ a
vtv
              -- update r[k..m-1, k..n-1]
              subR :: Array a
subR = [Int] -> ([Int] -> a) -> Array a
forall a. [Int] -> ([Int] -> a) -> Array a
A.tabulate [Int
mk, Int
nk] (\[Int
i, Int
j] -> Array a
r Array a -> [Int] -> a
forall a. Array a -> [Int] -> a
A.! [Int
k Int -> Int -> Int
forall a. Additive a => a -> a -> a
+ Int
i, Int
k Int -> Int -> Int
forall a. Additive a => a -> a -> a
+ Int
j])
              vta :: Array a
vta = [Int]
-> [Int]
-> (Array a -> a)
-> (a -> a -> a)
-> Array a
-> Array a
-> Array a
forall c d a b.
[Int]
-> [Int]
-> (Array c -> d)
-> (a -> b -> c)
-> Array a
-> Array b
-> Array d
A.prod [Int
0] [Int
0] Array a -> a
forall a (f :: * -> *). (Additive a, Foldable f) => f a -> a
sum a -> a -> a
forall a. Multiplicative a => a -> a -> a
(*) Array a
v Array a
subR
              outerR :: Array a
outerR = (a -> a -> a) -> Array a -> Array a -> Array a
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
A.expand a -> a -> a
forall a. Multiplicative a => a -> a -> a
(*) Array a
v Array a
vta
              subR' :: Array a
subR' = (a -> a -> a) -> Array a -> Array a -> Array a
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
A.zipWith (-) Array a
subR ((a -> a) -> Array a -> Array a
forall a b. (a -> b) -> Array a -> Array b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap (a
scale a -> a -> a
forall a. Multiplicative a => a -> a -> a
*) Array a
outerR)
              r' :: Array a
r' = Array a -> Int -> Int -> Array a -> Array a
forall a. Array a -> Int -> Int -> Array a -> Array a
updateSubmatrix Array a
r Int
k Int
k Array a
subR'
              -- update q[:, k..m-1] on the right by H_k
              subQ :: Array a
subQ = [Int] -> ([Int] -> a) -> Array a
forall a. [Int] -> ([Int] -> a) -> Array a
A.tabulate [Int
m, Int
mk] (\[Int
i, Int
j] -> Array a
q Array a -> [Int] -> a
forall a. Array a -> [Int] -> a
A.! [Int
i, Int
k Int -> Int -> Int
forall a. Additive a => a -> a -> a
+ Int
j])
              qv :: Array a
qv = [Int]
-> [Int]
-> (Array a -> a)
-> (a -> a -> a)
-> Array a
-> Array a
-> Array a
forall c d a b.
[Int]
-> [Int]
-> (Array c -> d)
-> (a -> b -> c)
-> Array a
-> Array b
-> Array d
A.prod [Int
1] [Int
0] Array a -> a
forall a (f :: * -> *). (Additive a, Foldable f) => f a -> a
sum a -> a -> a
forall a. Multiplicative a => a -> a -> a
(*) Array a
subQ Array a
v
              outerQ :: Array a
outerQ = (a -> a -> a) -> Array a -> Array a -> Array a
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
A.expand a -> a -> a
forall a. Multiplicative a => a -> a -> a
(*) Array a
qv Array a
v
              subQ' :: Array a
subQ' = (a -> a -> a) -> Array a -> Array a -> Array a
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
A.zipWith (-) Array a
subQ ((a -> a) -> Array a -> Array a
forall a b. (a -> b) -> Array a -> Array b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap (a
scale a -> a -> a
forall a. Multiplicative a => a -> a -> a
*) Array a
outerQ)
              q' :: Array a
q' = Array a -> Int -> Int -> Array a -> Array a
forall a. Array a -> Int -> Int -> Array a -> Array a
updateSubmatrix Array a
q Int
0 Int
k Array a
subQ'
           in (Array a -> Matrix a
forall a. Array a -> Matrix a
Matrix Array a
q', Array a -> Matrix a
forall a. Array a -> Matrix a
Matrix Array a
r')
        )
        (a
vtv a -> a -> Bool
forall a. Eq a => a -> a -> Bool
== a
forall a. Additive a => a
zero)

-- | Overwrite a rectangular region of a matrix with a smaller array.
updateSubmatrix :: A.Array a -> Int -> Int -> A.Array a -> A.Array a
updateSubmatrix :: forall a. Array a -> Int -> Int -> Array a -> Array a
updateSubmatrix Array a
m Int
rowOff Int
colOff Array a
sub =
  let sh :: [Int]
sh = Vector Int -> [Int]
forall a. Unbox a => Vector a -> [a]
VU.toList (Array a -> Vector Int
forall a. Array a -> Vector Int
A.shape Array a
m)
      subSh :: [Int]
subSh = Vector Int -> [Int]
forall a. Unbox a => Vector a -> [a]
VU.toList (Array a -> Vector Int
forall a. Array a -> Vector Int
A.shape Array a
sub)
      rowsSub :: Int
rowsSub = [Int]
subSh [Int] -> Int -> Int
forall a. HasCallStack => [a] -> Int -> a
!! Int
0
      colsSub :: Int
colsSub = [Int]
subSh [Int] -> Int -> Int
forall a. HasCallStack => [a] -> Int -> a
!! Int
1
   in [Int] -> ([Int] -> a) -> Array a
forall a. [Int] -> ([Int] -> a) -> Array a
A.tabulate [Int]
sh (([Int] -> a) -> Array a) -> ([Int] -> a) -> Array a
forall a b. (a -> b) -> a -> b
$ \case
        [Int
i, Int
j]
          | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
rowOff Bool -> Bool -> Bool
&& Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
rowOff Int -> Int -> Int
forall a. Additive a => a -> a -> a
+ Int
rowsSub Bool -> Bool -> Bool
&& Int
j Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
colOff Bool -> Bool -> Bool
&& Int
j Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
colOff Int -> Int -> Int
forall a. Additive a => a -> a -> a
+ Int
colsSub ->
              Array a
sub Array a -> [Int] -> a
forall a. Array a -> [Int] -> a
A.! [Int
i Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
rowOff, Int
j Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
colOff]
          | Bool
otherwise -> Array a
m Array a -> [Int] -> a
forall a. Array a -> [Int] -> a
A.! [Int
i, Int
j]
        [Int]
_ -> String -> a
forall a. HasCallStack => String -> a
error String
"Circuit.Mat.Dense.updateSubmatrix: expected rank-2 index"

-- | Forward substitution as a stream morphism.
--
-- Solves @L y = b@ for unit lower-triangular @L@ by streaming rows of @L@ and
-- components of @b@, accumulating @y@ one component at a time.  Each step is a
-- dot product of the already-computed prefix of @y@ with the active row prefix
-- of @L@, followed by subtraction from the current @b@ component.
forwardSubstStream ::
  (Subtractive a, Multiplicative a) =>
  A.Array a ->
  A.Array a ->
  A.Array a
forwardSubstStream :: forall a.
(Subtractive a, Multiplicative a) =>
Array a -> Array a -> Array a
forwardSubstStream Array a
l Array a
b = These (Array a) (Array a)
-> These (Array a) (Array a) -> Array a -> Array a
forall {c} {f} {f}.
(Subtractive c, Multiplicative c, Uncons f (Array c),
 Uncons f (Array c)) =>
These (Array c) f -> These (Array c) f -> Array c -> Array c
go (Array a -> These (Array a) (Array a)
forall f s. Uncons f s => f -> These s f
Stream.uncons Array a
l) (Array a -> These (Array a) (Array a)
forall f s. Uncons f s => f -> These s f
Stream.uncons Array a
b) Array a
forall a. Array a
A.empty
  where
    go :: These (Array c) f -> These (Array c) f -> Array c -> Array c
go (Stream.These Array c
rowL f
restL) (Stream.These Array c
bi f
restB) Array c
ys =
      let yi :: Array c
yi = Array c -> Array c -> Array c -> Array c
forall {c}.
(Subtractive c, Multiplicative c) =>
Array c -> Array c -> Array c -> Array c
solveRow Array c
rowL Array c
bi Array c
ys
       in These (Array c) f -> These (Array c) f -> Array c -> Array c
go (f -> These (Array c) f
forall f s. Uncons f s => f -> These s f
Stream.uncons f
restL) (f -> These (Array c) f
forall f s. Uncons f s => f -> These s f
Stream.uncons f
restB) (Array c -> Array c -> Array c
forall f s. Snoc f s => f -> s -> f
Stream.snoc Array c
ys Array c
yi)
    go (Stream.This Array c
rowL) (Stream.This Array c
bi) Array c
ys = Array c -> Array c -> Array c
forall f s. Snoc f s => f -> s -> f
Stream.snoc Array c
ys (Array c -> Array c -> Array c -> Array c
forall {c}.
(Subtractive c, Multiplicative c) =>
Array c -> Array c -> Array c -> Array c
solveRow Array c
rowL Array c
bi Array c
ys)
    go These (Array c) f
_ These (Array c) f
_ Array c
ys = Array c
ys
    solveRow :: Array c -> Array c -> Array c -> Array c
solveRow Array c
rowL Array c
bi Array c
ys =
      let k :: Int
k = Array c -> Int
forall a. Array a -> Int
A.length Array c
ys
          rowPrefix :: Array c
rowPrefix = Int -> Int -> Array c -> Array c
forall a. Int -> Int -> Array a -> Array a
A.take Int
0 Int
k Array c
rowL
          rowDot :: c
rowDot = [c] -> c
forall a (f :: * -> *). (Additive a, Foldable f) => f a -> a
sum [Array c
rowPrefix Array c -> [Int] -> c
forall a. Array a -> [Int] -> a
A.! [Int
j] c -> c -> c
forall a. Multiplicative a => a -> a -> a
* Array c
ys Array c -> [Int] -> c
forall a. Array a -> [Int] -> a
A.! [Int
j] | Int
j <- [Int
0 .. Int
k Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1]]
       in (c -> c -> c) -> Array c -> Array c -> Array c
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
A.zipWith (-) Array c
bi ([Int] -> [c] -> Array c
forall t a. FromVector t a => [Int] -> t -> Array a
A.array [] [c
rowDot])

-- | Solve a 2×2 linear system.
--
-- For a system @A x = b@ with @A = [[a,b],[c,d]]@ and @b = [r,s]@,
-- returns @x@ using Cramer's rule.  If the determinant is near zero,
-- returns a zero vector as a soft failure.
--
-- >>> solve2 [[1.0, 2.0], [3.0, 4.0]] [5.0, 6.0]
-- [-4.0,4.5]
solve2 :: [[Double]] -> [Double] -> [Double]
solve2 :: [[Double]] -> [Double] -> [Double]
solve2 [[Double
a, Double
b], [Double
c, Double
d]] [Double
r, Double
s] =
  let det :: Double
det = Double
a Double -> Double -> Double
forall a. Num a => a -> a -> a
P.* Double
d Double -> Double -> Double
forall a. Num a => a -> a -> a
P.- Double
b Double -> Double -> Double
forall a. Num a => a -> a -> a
P.* Double
c
   in [Double] -> [Double] -> Bool -> [Double]
forall a. a -> a -> Bool -> a
bool
        [Double
0, Double
0]
        [(Double
d Double -> Double -> Double
forall a. Num a => a -> a -> a
P.* Double
r Double -> Double -> Double
forall a. Num a => a -> a -> a
P.- Double
b Double -> Double -> Double
forall a. Num a => a -> a -> a
P.* Double
s) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
P./ Double
det, (Double
a Double -> Double -> Double
forall a. Num a => a -> a -> a
P.* Double
s Double -> Double -> Double
forall a. Num a => a -> a -> a
P.- Double
c Double -> Double -> Double
forall a. Num a => a -> a -> a
P.* Double
r) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
P./ Double
det]
        (Double -> Double
forall a. Num a => a -> a
P.abs Double
det Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
P.< Double
1e-14)
solve2 [[Double]]
_ [Double]
_ = String -> [Double]
forall a. HasCallStack => String -> a
P.error String
"solve2: expected 2×2"