{-# LANGUAGE DataKinds #-}
{-# LANGUAGE DerivingStrategies #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications #-}
{-# LANGUAGE UndecidableInstances #-}

-- | Matrix field calculus for dense, statically-shaped matrices.
--
-- This module provides Cholesky decomposition, triangular-matrix inversion,
-- and square-matrix inversion for 'Harpie.Fixed.Array' matrices.  All
-- operations are expressed in NumHask algebraic classes rather than GHC's
-- 'Num'/'Floating' hierarchy, so they work for any scalar that satisfies the
-- required field structure.
--
-- The matrix-level 'Multiplicative' and 'Divisive' instances are also provided
-- here, in @circuits-mat@, rather than in @harpie@: @harpie@ stays a generic
-- array library with relaxed element-wise constraints, while matrix calculus
-- lives in the categorical/linear-algebra layer.
module Circuit.Mat.Field
  ( -- * Square matrix alias
    MatrixM,

    -- * Field calculus
    cholM,
    invtriM,
    inverseM,
    luM,
    luRank1M,
    householderStep,

    -- * Matrix-as-ring wrapper
    MatField (..),
  )
where

import Data.List (maximumBy)
import Data.Ord (comparing)
import GHC.TypeNats (KnownNat)
import Harpie.Fixed (Array, Matrix)
import Harpie.Fixed qualified as F
import Harpie.Shape (KnownNats)
import Harpie.Shape qualified as S
import NumHask.Algebra.Additive (Additive (..), sum)
import NumHask.Algebra.Field (ExpField (..))
import NumHask.Algebra.Metric (Absolute, abs, signum)
import NumHask.Algebra.Multiplicative (Divisive (..), Multiplicative (..))
import NumHask.Prelude hiding (sum)

-- $setup
--
-- >>> :set -XDataKinds
-- >>> :set -XTypeApplications
-- >>> :m -Prelude
-- >>> :set -XRebindableSyntax
-- >>> import NumHask.Prelude
-- >>> import Harpie.Fixed qualified as F
-- >>> import Prettyprinter hiding (dot, fill)
-- >>> import Circuit.Mat.Field

-- | A square matrix of statically-known size.
type MatrixM n a = Matrix n n a

-- | Cholesky decomposition using the
-- <https://en.wikipedia.org/wiki/Cholesky_decomposition#The_Cholesky_algorithm Cholesky-Crout>
-- algorithm.
--
-- >>> e = F.array @[3,3] @Double [4,12,-16,12,37,-43,-16,-43,98]
-- >>> pretty (cholM e)
-- [[2.0,0.0,0.0],
--  [6.0,1.0,0.0],
--  [-8.0,5.0,3.0]]
-- >>> F.mult (cholM e) (F.transpose (cholM e)) == e
-- True
cholM ::
  forall n a.
  ( KnownNat n,
    KnownNats '[n, n],
    ExpField a
  ) =>
  MatrixM n a ->
  MatrixM n a
cholM :: forall (n :: Nat) a.
(KnownNat n, KnownNats '[n, n], ExpField a) =>
MatrixM n a -> MatrixM n a
cholM MatrixM n a
a = MatrixM n a
l
  where
    l :: MatrixM n a
l = (Rep (Array '[n, n]) -> a) -> MatrixM n a
forall a. (Rep (Array '[n, n]) -> a) -> Array '[n, n] a
forall (f :: * -> *) a. Representable f => (Rep f -> a) -> f a
F.tabulate (\Rep (Array '[n, n])
s -> Int -> MatrixM n a -> Fins '[n, n] -> a -> a
forall (n :: Nat) a.
(KnownNat n, ExpField a) =>
Int -> MatrixM n a -> Fins '[n, n] -> a -> a
norm_ Int
1 MatrixM n a
l Rep (Array '[n, n])
Fins '[n, n]
s (MatrixM n a -> Rep (Array '[n, n]) -> a
forall a. Array '[n, n] a -> Rep (Array '[n, n]) -> a
forall (f :: * -> *) a. Representable f => f a -> Rep f -> a
F.index MatrixM n a
a Rep (Array '[n, n])
s a -> a -> a
forall a. Subtractive a => a -> a -> a
- MatrixM n a -> Fins '[n, n] -> a
forall (n :: Nat) a.
(KnownNat n, Additive a, Multiplicative a) =>
MatrixM n a -> Fins '[n, n] -> a
cross_ MatrixM n a
l Rep (Array '[n, n])
Fins '[n, n]
s))

norm_ ::
  forall n a.
  ( KnownNat n,
    ExpField a
  ) =>
  Int ->
  MatrixM n a ->
  S.Fins '[n, n] ->
  a ->
  a
norm_ :: forall (n :: Nat) a.
(KnownNat n, ExpField a) =>
Int -> MatrixM n a -> Fins '[n, n] -> a -> a
norm_ Int
d MatrixM n a
l (S.UnsafeFins [Int]
s) = (a -> a) -> (a -> a) -> Bool -> a -> a
forall a. a -> a -> Bool -> a
bool (a -> a
forall a. Divisive a => a -> a
recip (Array '[n] a
diagL Array '[n] a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int -> [Int] -> Int
S.getDimL Int
d [Int]
s]) a -> a -> a
forall a. Multiplicative a => a -> a -> a
*) a -> a
forall a. ExpField a => a -> a
sqrt ([Int] -> Bool
forall a. Eq a => [a] -> Bool
S.isDiagL [Int]
s)
  where
    diagL :: Array '[n] a
diagL = MatrixM n a -> Array '[n] a
forall (s' :: [Nat]) a (s :: [Nat]).
(KnownNats s, KnownNats s', s' ~ Eval (MinDim s)) =>
Array s a -> Array s' a
F.diag MatrixM n a
l

-- | Cross term @sum_k l[i,k] * l[j,k]@ used by Cholesky-Crout.
cross_ ::
  forall n a.
  ( KnownNat n,
    Additive a,
    Multiplicative a
  ) =>
  MatrixM n a ->
  S.Fins '[n, n] ->
  a
cross_ :: forall (n :: Nat) a.
(KnownNat n, Additive a, Multiplicative a) =>
MatrixM n a -> Fins '[n, n] -> a
cross_ MatrixM n a
l Fins '[n, n]
s = [a] -> a
forall a (f :: * -> *). (Additive a, Foldable f) => f a -> a
sum [MatrixM n a
l MatrixM n a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
i, Int
k] a -> a -> a
forall a. Multiplicative a => a -> a -> a
* MatrixM n a
l MatrixM n a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
j, Int
k] | Int
k <- [Int
0 .. Int
j Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1]]
  where
    ij :: [Int]
ij = Fins '[n, n] -> [Int]
forall {k} (s :: k). Fins s -> [Int]
S.fromFins Fins '[n, n]
s
    (Int
i, Int
j) = case [Int]
ij of
      [Int
x, Int
y] -> (Int
x, Int
y)
      [Int]
_ -> [Char] -> (Int, Int)
forall a. HasCallStack => [Char] -> a
error [Char]
"cross_: invalid Fins dimension (expected 2D index)"

-- | Inverse of a square matrix via Cholesky decomposition.
--
-- >>> e = F.array @[3,3] @Double [4,12,-16,12,37,-43,-16,-43,98]
-- >>> pretty (inverseM e)
-- [[49.36111111111111,-13.555555555555554,2.1111111111111107],
--  [-13.555555555555554,3.7777777777777772,-0.5555555555555555],
--  [2.1111111111111107,-0.5555555555555555,0.1111111111111111]]
inverseM ::
  forall n a.
  ( KnownNat n,
    KnownNats '[n, n],
    ExpField a
  ) =>
  MatrixM n a ->
  MatrixM n a
inverseM :: forall (n :: Nat) a.
(KnownNat n, KnownNats '[n, n], ExpField a) =>
MatrixM n a -> MatrixM n a
inverseM MatrixM n a
a = MatrixM n a -> MatrixM n a -> MatrixM n a
forall a (ds0 :: [Nat]) (ds1 :: [Nat]) (s0 :: [Nat]) (s1 :: [Nat])
       (so0 :: [Nat]) (so1 :: [Nat]) (st :: [Nat]) (si :: [Nat]).
(Additive a, Multiplicative a, KnownNats s0, KnownNats s1,
 KnownNats ds0, KnownNats ds1, KnownNats so0, KnownNats so1,
 KnownNats st, KnownNats si, so0 ~ Eval (DeleteDims ds0 s0),
 so1 ~ Eval (DeleteDims ds1 s1), si ~ Eval (GetDims ds0 s0),
 si ~ Eval (GetDims ds1 s1), st ~ Eval (so0 ++ so1),
 ds0 ~ '[Eval (Eval (Rank s0) - 1)], ds1 ~ '[0]) =>
Array s0 a -> Array s1 a -> Array st a
F.mult (MatrixM n a -> MatrixM n a
forall (n :: Nat) a.
(KnownNat n, KnownNats '[n, n], Subtractive a, Divisive a) =>
MatrixM n a -> MatrixM n a
invtriM (MatrixM n a -> MatrixM n a
forall a (s :: [Nat]) (s' :: [Nat]).
(KnownNats s, KnownNats s', s' ~ Eval (Reverse s)) =>
Array s a -> Array s' a
F.transpose (MatrixM n a -> MatrixM n a
forall (n :: Nat) a.
(KnownNat n, KnownNats '[n, n], ExpField a) =>
MatrixM n a -> MatrixM n a
cholM MatrixM n a
a))) (MatrixM n a -> MatrixM n a
forall (n :: Nat) a.
(KnownNat n, KnownNats '[n, n], Subtractive a, Divisive a) =>
MatrixM n a -> MatrixM n a
invtriM (MatrixM n a -> MatrixM n a
forall (n :: Nat) a.
(KnownNat n, KnownNats '[n, n], ExpField a) =>
MatrixM n a -> MatrixM n a
cholM MatrixM n a
a))

-- | LU decomposition with partial pivoting.
--
-- Returns @(p, l, u)@ such that @a = p^T . l . u@, where @p@ is the
-- accumulated row-permutation matrix, @l@ is unit lower triangular, and @u@ is
-- upper triangular.  The algorithm is the standard iterative Schur-complement
-- update: at each pivot @k@, swap the pivot row into place, then eliminate the
-- rows below by a rank-1 update.
--
-- The Schur step @a' = a - c . r@ is an outer product followed by subtraction,
-- which is the matrix-level reading of the trace / feedback pattern used in
-- circuit constructions.
luM ::
  forall n a.
  ( KnownNat n,
    KnownNats '[n, n],
    Subtractive a,
    Divisive a,
    Absolute a,
    Ord a
  ) =>
  MatrixM n a ->
  (MatrixM n a, MatrixM n a, MatrixM n a)
luM :: forall (n :: Nat) a.
(KnownNat n, KnownNats '[n, n], Subtractive a, Divisive a,
 Absolute a, Ord a) =>
MatrixM n a -> (MatrixM n a, MatrixM n a, MatrixM n a)
luM MatrixM n a
a = (MatrixM n a
p, MatrixM n a
l, MatrixM n a
u)
  where
    n :: Int
n = forall (n :: Nat). KnownNat n => Int
S.valueOf @n
    (MatrixM n a
p, MatrixM n a
m) = ((MatrixM n a, MatrixM n a) -> Int -> (MatrixM n a, MatrixM n a))
-> (MatrixM n a, MatrixM n a)
-> [Int]
-> (MatrixM n a, MatrixM n a)
forall b a. (b -> a -> b) -> b -> [a] -> b
forall (t :: * -> *) b a.
Foldable t =>
(b -> a -> b) -> b -> t a -> b
foldl (MatrixM n a, MatrixM n a) -> Int -> (MatrixM n a, MatrixM n a)
step (forall (s :: [Nat]) a.
(KnownNats s, Additive a, Multiplicative a) =>
Array s a
F.ident @[n, n], MatrixM n a
a) [Int
0 .. Int
n Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
2]
    step :: (MatrixM n a, MatrixM n a) -> Int -> (MatrixM n a, MatrixM n a)
step (MatrixM n a
p0, MatrixM n a
m0) Int
k =
      let pivotRow :: Int
pivotRow = (Int -> Int -> Ordering) -> [Int] -> Int
forall (t :: * -> *) a.
Foldable t =>
(a -> a -> Ordering) -> t a -> a
maximumBy ((Int -> a) -> Int -> Int -> Ordering
forall a b. Ord a => (b -> a) -> b -> b -> Ordering
comparing (\Int
i -> a -> a
forall a. Absolute a => a -> a
abs (MatrixM n a
m0 MatrixM n a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
i, Int
k]))) [Int
k .. Int
n Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1]
          p1 :: MatrixM n a
p1 = Int -> Int -> MatrixM n a -> MatrixM n a
forall (n :: Nat) a.
KnownNats '[n, n] =>
Int -> Int -> MatrixM n a -> MatrixM n a
swapRows Int
k Int
pivotRow MatrixM n a
p0
          m1 :: MatrixM n a
m1 = Int -> Int -> MatrixM n a -> MatrixM n a
forall (n :: Nat) a.
KnownNats '[n, n] =>
Int -> Int -> MatrixM n a -> MatrixM n a
swapRows Int
k Int
pivotRow MatrixM n a
m0
          m2 :: MatrixM n a
m2 =
            (Rep (Array '[n, n]) -> a) -> MatrixM n a
forall a. (Rep (Array '[n, n]) -> a) -> Array '[n, n] a
forall (f :: * -> *) a. Representable f => (Rep f -> a) -> f a
F.tabulate ((Rep (Array '[n, n]) -> a) -> MatrixM n a)
-> (Rep (Array '[n, n]) -> a) -> MatrixM n a
forall a b. (a -> b) -> a -> b
$ \Rep (Array '[n, n])
s -> case Fins '[n, n] -> [Int]
forall {k} (s :: k). Fins s -> [Int]
S.fromFins Rep (Array '[n, n])
Fins '[n, n]
s of
              [Int
i, Int
j]
                | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Int
k Bool -> Bool -> Bool
&& Int
j Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Int
k ->
                    let mult :: a
mult = MatrixM n a
m1 MatrixM n a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
i, Int
k] a -> a -> a
forall a. Divisive a => a -> a -> a
/ MatrixM n a
m1 MatrixM n a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
k, Int
k]
                     in MatrixM n a
m1 MatrixM n a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
i, Int
j] a -> a -> a
forall a. Subtractive a => a -> a -> a
- a
mult a -> a -> a
forall a. Multiplicative a => a -> a -> a
* MatrixM n a
m1 MatrixM n a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
k, Int
j]
                | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Int
k Bool -> Bool -> Bool
&& Int
j Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
k -> MatrixM n a
m1 MatrixM n a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
i, Int
k] a -> a -> a
forall a. Divisive a => a -> a -> a
/ MatrixM n a
m1 MatrixM n a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
k, Int
k]
                | Bool
otherwise -> MatrixM n a
m1 MatrixM n a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
i, Int
j]
              [Int]
_ -> [Char] -> a
forall a. HasCallStack => [Char] -> a
error [Char]
"luM: expected rank-2 index"
       in (MatrixM n a
p1, MatrixM n a
m2)
    l :: MatrixM n a
l =
      (Rep (Array '[n, n]) -> a) -> MatrixM n a
forall a. (Rep (Array '[n, n]) -> a) -> Array '[n, n] a
forall (f :: * -> *) a. Representable f => (Rep f -> a) -> f a
F.tabulate ((Rep (Array '[n, n]) -> a) -> MatrixM n a)
-> (Rep (Array '[n, n]) -> a) -> MatrixM n a
forall a b. (a -> b) -> a -> b
$ \Rep (Array '[n, n])
s -> case Fins '[n, n] -> [Int]
forall {k} (s :: k). Fins s -> [Int]
S.fromFins Rep (Array '[n, n])
Fins '[n, n]
s of
        [Int
i, Int
j]
          | Int
i Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
j -> a
forall a. Multiplicative a => a
one
          | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Int
j -> MatrixM n a
m MatrixM n a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
i, Int
j]
          | Bool
otherwise -> a
forall a. Additive a => a
zero
        [Int]
_ -> [Char] -> a
forall a. HasCallStack => [Char] -> a
error [Char]
"luM: expected rank-2 index"
    u :: MatrixM n a
u =
      (Rep (Array '[n, n]) -> a) -> MatrixM n a
forall a. (Rep (Array '[n, n]) -> a) -> Array '[n, n] a
forall (f :: * -> *) a. Representable f => (Rep f -> a) -> f a
F.tabulate ((Rep (Array '[n, n]) -> a) -> MatrixM n a)
-> (Rep (Array '[n, n]) -> a) -> MatrixM n a
forall a b. (a -> b) -> a -> b
$ \Rep (Array '[n, n])
s -> case Fins '[n, n] -> [Int]
forall {k} (s :: k). Fins s -> [Int]
S.fromFins Rep (Array '[n, n])
Fins '[n, n]
s of
        [Int
i, Int
j]
          | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
<= Int
j -> MatrixM n a
m MatrixM n a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
i, Int
j]
          | Bool
otherwise -> a
forall a. Additive a => a
zero
        [Int]
_ -> [Char] -> a
forall a. HasCallStack => [Char] -> a
error [Char]
"luM: expected rank-2 index"

-- | LU decomposition via iterated rank-1 updates.
--
-- This is the same factorisation as 'luM', but the Schur-complement step is
-- written explicitly as an outer-product contraction:
--
-- > A' = A - c ⊗ r
--
-- where @c@ is the column of multipliers below the pivot and @r@ is the pivot
-- row.  The outer product uses 'F.expand' and the subtraction is elementwise;
-- together they are the matrix-level reading of the trace / feedback pattern.
--
-- The result agrees with 'luM' on any invertible square matrix.
luRank1M ::
  forall n a.
  ( KnownNat n,
    KnownNats '[n, n],
    Subtractive a,
    Divisive a,
    Absolute a,
    Ord a
  ) =>
  MatrixM n a ->
  (MatrixM n a, MatrixM n a, MatrixM n a)
luRank1M :: forall (n :: Nat) a.
(KnownNat n, KnownNats '[n, n], Subtractive a, Divisive a,
 Absolute a, Ord a) =>
MatrixM n a -> (MatrixM n a, MatrixM n a, MatrixM n a)
luRank1M MatrixM n a
a = (MatrixM n a
p, MatrixM n a
l, MatrixM n a
u)
  where
    n :: Int
n = forall (n :: Nat). KnownNat n => Int
S.valueOf @n
    ((MatrixM n a
p, MatrixM n a
m), MatrixM n a
lAcc) = (((MatrixM n a, MatrixM n a), MatrixM n a)
 -> Int -> ((MatrixM n a, MatrixM n a), MatrixM n a))
-> ((MatrixM n a, MatrixM n a), MatrixM n a)
-> [Int]
-> ((MatrixM n a, MatrixM n a), MatrixM n a)
forall b a. (b -> a -> b) -> b -> [a] -> b
forall (t :: * -> *) b a.
Foldable t =>
(b -> a -> b) -> b -> t a -> b
foldl ((MatrixM n a, MatrixM n a), MatrixM n a)
-> Int -> ((MatrixM n a, MatrixM n a), MatrixM n a)
step ((forall (s :: [Nat]) a.
(KnownNats s, Additive a, Multiplicative a) =>
Array s a
F.ident @[n, n], MatrixM n a
a), forall (s :: [Nat]) a. KnownNats s => a -> Array s a
F.konst @[n, n] a
forall a. Additive a => a
zero) [Int
0 .. Int
n Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
2]
    step :: ((MatrixM n a, MatrixM n a), MatrixM n a)
-> Int -> ((MatrixM n a, MatrixM n a), MatrixM n a)
step ((MatrixM n a
p0, MatrixM n a
m0), MatrixM n a
l0) Int
k =
      let pivotRow :: Int
pivotRow = (Int -> Int -> Ordering) -> [Int] -> Int
forall (t :: * -> *) a.
Foldable t =>
(a -> a -> Ordering) -> t a -> a
maximumBy ((Int -> a) -> Int -> Int -> Ordering
forall a b. Ord a => (b -> a) -> b -> b -> Ordering
comparing (\Int
i -> a -> a
forall a. Absolute a => a -> a
abs (MatrixM n a
m0 MatrixM n a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
i, Int
k]))) [Int
k .. Int
n Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1]
          p1 :: MatrixM n a
p1 = Int -> Int -> MatrixM n a -> MatrixM n a
forall (n :: Nat) a.
KnownNats '[n, n] =>
Int -> Int -> MatrixM n a -> MatrixM n a
swapRows Int
k Int
pivotRow MatrixM n a
p0
          m1 :: MatrixM n a
m1 = Int -> Int -> MatrixM n a -> MatrixM n a
forall (n :: Nat) a.
KnownNats '[n, n] =>
Int -> Int -> MatrixM n a -> MatrixM n a
swapRows Int
k Int
pivotRow MatrixM n a
m0
          l1 :: MatrixM n a
l1 = Int -> Int -> MatrixM n a -> MatrixM n a
forall (n :: Nat) a.
KnownNats '[n, n] =>
Int -> Int -> MatrixM n a -> MatrixM n a
swapRows Int
k Int
pivotRow MatrixM n a
l0
          pivot :: a
pivot = MatrixM n a
m1 MatrixM n a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
k, Int
k]
          mult :: Int -> a
mult Int
i = MatrixM n a
m1 MatrixM n a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
i, Int
k] a -> a -> a
forall a. Divisive a => a -> a -> a
/ a
pivot
          c :: F.Array '[n] a
          c :: Array '[n] a
c =
            (Rep (Array '[n]) -> a) -> Array '[n] a
forall a. (Rep (Array '[n]) -> a) -> Array '[n] a
forall (f :: * -> *) a. Representable f => (Rep f -> a) -> f a
F.tabulate ((Rep (Array '[n]) -> a) -> Array '[n] a)
-> (Rep (Array '[n]) -> a) -> Array '[n] a
forall a b. (a -> b) -> a -> b
$ \Rep (Array '[n])
s -> case Fins '[n] -> [Int]
forall {k} (s :: k). Fins s -> [Int]
S.fromFins Rep (Array '[n])
Fins '[n]
s of
              [Int
i]
                | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Int
k -> Int -> a
mult Int
i
                | Bool
otherwise -> a
forall a. Additive a => a
zero
              [Int]
_ -> [Char] -> a
forall a. HasCallStack => [Char] -> a
error [Char]
"luRank1M: expected rank-1 index"
          r :: F.Array '[n] a
          r :: Array '[n] a
r =
            (Rep (Array '[n]) -> a) -> Array '[n] a
forall a. (Rep (Array '[n]) -> a) -> Array '[n] a
forall (f :: * -> *) a. Representable f => (Rep f -> a) -> f a
F.tabulate ((Rep (Array '[n]) -> a) -> Array '[n] a)
-> (Rep (Array '[n]) -> a) -> Array '[n] a
forall a b. (a -> b) -> a -> b
$ \Rep (Array '[n])
s -> case Fins '[n] -> [Int]
forall {k} (s :: k). Fins s -> [Int]
S.fromFins Rep (Array '[n])
Fins '[n]
s of
              [Int
j]
                | Int
j Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
k -> MatrixM n a
m1 MatrixM n a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
k, Int
j]
                | Bool
otherwise -> a
forall a. Additive a => a
zero
              [Int]
_ -> [Char] -> a
forall a. HasCallStack => [Char] -> a
error [Char]
"luRank1M: expected rank-1 index"
          m2 :: MatrixM n a
m2 = (a -> a -> a) -> MatrixM n a -> MatrixM n a -> MatrixM n a
forall (s :: [Nat]) a b c.
KnownNats s =>
(a -> b -> c) -> Array s a -> Array s b -> Array s c
F.zipWith (-) MatrixM n a
m1 ((a -> a -> a) -> Array '[n] a -> Array '[n] a -> MatrixM n a
forall (sc :: [Nat]) (sa :: [Nat]) (sb :: [Nat]) a b c.
(KnownNats sa, KnownNats sb, KnownNats sc, sc ~ Eval (sa ++ sb)) =>
(a -> b -> c) -> Array sa a -> Array sb b -> Array sc c
F.expand a -> a -> a
forall a. Multiplicative a => a -> a -> a
(*) Array '[n] a
c Array '[n] a
r)
          l2 :: MatrixM n a
l2 =
            (Rep (Array '[n, n]) -> a) -> MatrixM n a
forall a. (Rep (Array '[n, n]) -> a) -> Array '[n, n] a
forall (f :: * -> *) a. Representable f => (Rep f -> a) -> f a
F.tabulate ((Rep (Array '[n, n]) -> a) -> MatrixM n a)
-> (Rep (Array '[n, n]) -> a) -> MatrixM n a
forall a b. (a -> b) -> a -> b
$ \Rep (Array '[n, n])
s -> case Fins '[n, n] -> [Int]
forall {k} (s :: k). Fins s -> [Int]
S.fromFins Rep (Array '[n, n])
Fins '[n, n]
s of
              [Int
i, Int
j]
                | Int
j Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
k Bool -> Bool -> Bool
&& Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Int
k -> Int -> a
mult Int
i
                | Bool
otherwise -> MatrixM n a
l1 MatrixM n a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
i, Int
j]
              [Int]
_ -> [Char] -> a
forall a. HasCallStack => [Char] -> a
error [Char]
"luRank1M: expected rank-2 index"
       in ((MatrixM n a
p1, MatrixM n a
m2), MatrixM n a
l2)
    l :: MatrixM n a
l =
      (Rep (Array '[n, n]) -> a) -> MatrixM n a
forall a. (Rep (Array '[n, n]) -> a) -> Array '[n, n] a
forall (f :: * -> *) a. Representable f => (Rep f -> a) -> f a
F.tabulate ((Rep (Array '[n, n]) -> a) -> MatrixM n a)
-> (Rep (Array '[n, n]) -> a) -> MatrixM n a
forall a b. (a -> b) -> a -> b
$ \Rep (Array '[n, n])
s -> case Fins '[n, n] -> [Int]
forall {k} (s :: k). Fins s -> [Int]
S.fromFins Rep (Array '[n, n])
Fins '[n, n]
s of
        [Int
i, Int
j]
          | Int
i Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
j -> a
forall a. Multiplicative a => a
one
          | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Int
j -> MatrixM n a
lAcc MatrixM n a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
i, Int
j]
          | Bool
otherwise -> a
forall a. Additive a => a
zero
        [Int]
_ -> [Char] -> a
forall a. HasCallStack => [Char] -> a
error [Char]
"luRank1M: expected rank-2 index"
    u :: MatrixM n a
u =
      (Rep (Array '[n, n]) -> a) -> MatrixM n a
forall a. (Rep (Array '[n, n]) -> a) -> Array '[n, n] a
forall (f :: * -> *) a. Representable f => (Rep f -> a) -> f a
F.tabulate ((Rep (Array '[n, n]) -> a) -> MatrixM n a)
-> (Rep (Array '[n, n]) -> a) -> MatrixM n a
forall a b. (a -> b) -> a -> b
$ \Rep (Array '[n, n])
s -> case Fins '[n, n] -> [Int]
forall {k} (s :: k). Fins s -> [Int]
S.fromFins Rep (Array '[n, n])
Fins '[n, n]
s of
        [Int
i, Int
j]
          | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
<= Int
j -> MatrixM n a
m MatrixM n a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
i, Int
j]
          | Bool
otherwise -> a
forall a. Additive a => a
zero
        [Int]
_ -> [Char] -> a
forall a. HasCallStack => [Char] -> a
error [Char]
"luRank1M: expected rank-2 index"

-- | Apply a Householder reflection to zero the subdiagonal of column @k@.
--
-- For column @x = a[:,k]@, compute a reflector @H = I - 2 v v^T / (v^T v)@
-- such that @H x = α e_k@ with @α = -sign(x[k]) · ||x||@.  The action on the
-- whole matrix is the rank-1 update
--
-- > a' = a - (2 / (v^T v)) · v ⊗ (v^T a)
--
-- where @v^T a@ is a contraction over the row axis and @v ⊗ (...)@ is an
-- outer product.  This is the circuit-native QR step: a scheduled reflection
-- implemented as expand / contract / subtract.
householderStep ::
  forall n a.
  ( KnownNat n,
    KnownNats '[n, n],
    ExpField a,
    Ord a
  ) =>
  Int ->
  MatrixM n a ->
  MatrixM n a
householderStep :: forall (n :: Nat) a.
(KnownNat n, KnownNats '[n, n], ExpField a, Ord a) =>
Int -> MatrixM n a -> MatrixM n a
householderStep Int
k MatrixM n a
a = (a -> a -> a) -> MatrixM n a -> MatrixM n a -> MatrixM n a
forall (s :: [Nat]) a b c.
KnownNats s =>
(a -> b -> c) -> Array s a -> Array s b -> Array s c
F.zipWith (-) MatrixM n a
a ((a -> a) -> MatrixM n a -> MatrixM n a
forall a b. (a -> b) -> Array '[n, n] a -> Array '[n, n] 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
*) MatrixM n a
outer)
  where
    x :: F.Array '[n] a
    x :: Array '[n] a
x = (Rep (Array '[n]) -> a) -> Array '[n] a
forall a. (Rep (Array '[n]) -> a) -> Array '[n] a
forall (f :: * -> *) a. Representable f => (Rep f -> a) -> f a
F.tabulate ((Rep (Array '[n]) -> a) -> Array '[n] a)
-> (Rep (Array '[n]) -> a) -> Array '[n] a
forall a b. (a -> b) -> a -> b
$ \Rep (Array '[n])
s -> case Fins '[n] -> [Int]
forall {k} (s :: k). Fins s -> [Int]
S.fromFins Rep (Array '[n])
Fins '[n]
s of
      [Int
i] -> MatrixM n a
a MatrixM n a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
i, Int
k]
      [Int]
_ -> [Char] -> a
forall a. HasCallStack => [Char] -> a
error [Char]
"householderStep: expected rank-1 index"
    xk :: a
xk = Array '[n] a
x Array '[n] a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
k]
    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 '[n] a
x Array '[n] a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
i]) a -> a -> a
forall a. Multiplicative a => a -> a -> a
* (Array '[n] a
x Array '[n] a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
i]) | Int
i <- [Int
0 .. forall (n :: Nat). KnownNat n => Int
S.valueOf @n 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 :: F.Array '[n] a
    v :: Array '[n] a
v = (Rep (Array '[n]) -> a) -> Array '[n] a
forall a. (Rep (Array '[n]) -> a) -> Array '[n] a
forall (f :: * -> *) a. Representable f => (Rep f -> a) -> f a
F.tabulate ((Rep (Array '[n]) -> a) -> Array '[n] a)
-> (Rep (Array '[n]) -> a) -> Array '[n] a
forall a b. (a -> b) -> a -> b
$ \Rep (Array '[n])
s -> case Fins '[n] -> [Int]
forall {k} (s :: k). Fins s -> [Int]
S.fromFins Rep (Array '[n])
Fins '[n]
s of
      [Int
i]
        | Int
i Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
k -> a
forall a. Additive a => a
zero
        | Int
i Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
k -> a
xk a -> a -> a
forall a. Subtractive a => a -> a -> a
- a
alpha
        | Bool
otherwise -> Array '[n] a
x Array '[n] a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
i]
      [Int]
_ -> [Char] -> a
forall a. HasCallStack => [Char] -> a
error [Char]
"householderStep: expected rank-1 index"
    vtv :: a
vtv = [a] -> a
forall a (f :: * -> *). (Additive a, Foldable f) => f a -> a
sum [(Array '[n] a
v Array '[n] a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
i]) a -> a -> a
forall a. Multiplicative a => a -> a -> a
* (Array '[n] a
v Array '[n] a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
i]) | Int
i <- [Int
0 .. forall (n :: Nat). KnownNat n => Int
S.valueOf @n Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1]]
    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
    -- v^T a: contract v's only axis with a's row axis, leaving columns.
    vta :: F.Array '[n] a
    vta :: Array '[n] a
vta = Dims '[0]
-> Dims '[0]
-> (Array '[n] a -> a)
-> (a -> a -> a)
-> Array '[n] a
-> MatrixM n a
-> Array '[n] a
forall a b c d (s0 :: [Nat]) (s1 :: [Nat]) (so0 :: [Nat])
       (so1 :: [Nat]) (si :: [Nat]) (st :: [Nat]) (ds0 :: [Nat])
       (ds1 :: [Nat]).
(KnownNats so0, KnownNats so1, KnownNats si, KnownNats s0,
 KnownNats s1, KnownNats st, KnownNats ds0, KnownNats ds1,
 so0 ~ Eval (DeleteDims ds0 s0), so1 ~ Eval (DeleteDims ds1 s1),
 si ~ Eval (GetDims ds0 s0), si ~ Eval (GetDims ds1 s1),
 st ~ Eval (so0 ++ so1)) =>
Dims ds0
-> Dims ds1
-> (Array si c -> d)
-> (a -> b -> c)
-> Array s0 a
-> Array s1 b
-> Array st d
F.prod (forall (ns :: [Nat]). KnownNats ns => SNats ns
F.Dims @'[0]) (forall (ns :: [Nat]). KnownNats ns => SNats ns
F.Dims @'[0]) Array '[n] a -> a
forall a (f :: * -> *). (Additive a, Foldable f) => f a -> a
sum a -> a -> a
forall a. Multiplicative a => a -> a -> a
(*) Array '[n] a
v MatrixM n a
a
    -- outer product v ⊗ (v^T a), shape [n, n].
    outer :: MatrixM n a
outer = (a -> a -> a) -> Array '[n] a -> Array '[n] a -> MatrixM n a
forall (sc :: [Nat]) (sa :: [Nat]) (sb :: [Nat]) a b c.
(KnownNats sa, KnownNats sb, KnownNats sc, sc ~ Eval (sa ++ sb)) =>
(a -> b -> c) -> Array sa a -> Array sb b -> Array sc c
F.expand a -> a -> a
forall a. Multiplicative a => a -> a -> a
(*) Array '[n] a
v Array '[n] a
vta

-- | Swap two rows of a matrix.
swapRows ::
  forall n a.
  (KnownNats '[n, n]) =>
  Int ->
  Int ->
  MatrixM n a ->
  MatrixM n a
swapRows :: forall (n :: Nat) a.
KnownNats '[n, n] =>
Int -> Int -> MatrixM n a -> MatrixM n a
swapRows Int
k Int
p = ([Int] -> [Int]) -> Array '[n, n] a -> Array '[n, n] a
forall (s' :: [Nat]) (s :: [Nat]) a.
(KnownNats s, KnownNats s') =>
([Int] -> [Int]) -> Array s a -> Array s' a
F.unsafeBackpermute (([Int] -> [Int]) -> Array '[n, n] a -> Array '[n, n] a)
-> ([Int] -> [Int]) -> Array '[n, n] a -> Array '[n, n] a
forall a b. (a -> b) -> a -> b
$ \case
  [Int
i, Int
j]
    | Int
i Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
k -> [Int
p, Int
j]
    | Int
i Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
p -> [Int
k, Int
j]
    | Bool
otherwise -> [Int
i, Int
j]
  [Int]
_ -> [Char] -> [Int]
forall a. HasCallStack => [Char] -> a
error [Char]
"swapRows: expected rank-2 index"

-- | Inversion of a triangular matrix.
--
-- >>> t = F.array @[3,3] @Double [1,0,1,0,1,2,0,0,1]
-- >>> pretty (invtriM t)
-- [[1.0,0.0,-1.0],
--  [0.0,1.0,-2.0],
--  [0.0,0.0,1.0]]
-- >>> F.ident @[3,3] == F.mult t (invtriM t)
-- True
invtriM ::
  forall n a.
  ( KnownNat n,
    KnownNats '[n, n],
    Subtractive a,
    Divisive a
  ) =>
  MatrixM n a ->
  MatrixM n a
invtriM :: forall (n :: Nat) a.
(KnownNat n, KnownNats '[n, n], Subtractive a, Divisive a) =>
MatrixM n a -> MatrixM n a
invtriM MatrixM n a
a = MatrixM n a
i
  where
    ti :: Array ('[n] <> '[n]) a
ti = Array '[n] a -> Array ('[n] <> '[n]) a
forall (s' :: [Nat]) a (s :: [Nat]).
(KnownNats s, KnownNats s', s' ~ Eval (s ++ s), Additive a) =>
Array s a -> Array s' a
F.undiag ((a -> a) -> Array '[n] a -> Array '[n] a
forall a b. (a -> b) -> Array '[n] a -> Array '[n] b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap a -> a
forall a. Divisive a => a -> a
recip (MatrixM n a -> Array '[n] a
forall (s' :: [Nat]) a (s :: [Nat]).
(KnownNats s, KnownNats s', s' ~ Eval (MinDim s)) =>
Array s a -> Array s' a
F.diag MatrixM n a
a))
    tl :: MatrixM n a
tl = (a -> a -> a) -> MatrixM n a -> MatrixM n a -> MatrixM n a
forall (s :: [Nat]) a b c.
KnownNats s =>
(a -> b -> c) -> Array s a -> Array s b -> Array s c
F.zipWith (-) MatrixM n a
a (Array '[n] a -> MatrixM n a
forall (s' :: [Nat]) a (s :: [Nat]).
(KnownNats s, KnownNats s', s' ~ Eval (s ++ s), Additive a) =>
Array s a -> Array s' a
F.undiag (MatrixM n a -> Array '[n] a
forall (s' :: [Nat]) a (s :: [Nat]).
(KnownNats s, KnownNats s', s' ~ Eval (MinDim s)) =>
Array s a -> Array s' a
F.diag MatrixM n a
a))
    l :: Array
  (Eval
     (Foldl'
        (Flip DeleteDim)
        '[n, n]
        (Eval (Rev (Eval (PreDeletePositionsGo '[1] '[])) '[])))
   <> Eval
        (Foldl'
           (Flip DeleteDim)
           '[n, n]
           (Eval (Rev (Eval (PreDeletePositionsGo '[0] '[])) '[]))))
  a
l = (a -> a)
-> Array
     (Eval
        (Foldl'
           (Flip DeleteDim)
           '[n, n]
           (Eval (Rev (Eval (PreDeletePositionsGo '[1] '[])) '[])))
      <> Eval
           (Foldl'
              (Flip DeleteDim)
              '[n, n]
              (Eval (Rev (Eval (PreDeletePositionsGo '[0] '[])) '[]))))
     a
-> Array
     (Eval
        (Foldl'
           (Flip DeleteDim)
           '[n, n]
           (Eval (Rev (Eval (PreDeletePositionsGo '[1] '[])) '[])))
      <> Eval
           (Foldl'
              (Flip DeleteDim)
              '[n, n]
              (Eval (Rev (Eval (PreDeletePositionsGo '[0] '[])) '[]))))
     a
forall a b.
(a -> b)
-> Array
     (Eval
        (Foldl'
           (Flip DeleteDim)
           '[n, n]
           (Eval (Rev (Eval (PreDeletePositionsGo '[1] '[])) '[])))
      <> Eval
           (Foldl'
              (Flip DeleteDim)
              '[n, n]
              (Eval (Rev (Eval (PreDeletePositionsGo '[0] '[])) '[]))))
     a
-> Array
     (Eval
        (Foldl'
           (Flip DeleteDim)
           '[n, n]
           (Eval (Rev (Eval (PreDeletePositionsGo '[1] '[])) '[])))
      <> Eval
           (Foldl'
              (Flip DeleteDim)
              '[n, n]
              (Eval (Rev (Eval (PreDeletePositionsGo '[0] '[])) '[]))))
     b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap a -> a
forall a. Subtractive a => a -> a
negate (Array ('[n] <> '[n]) a
-> MatrixM n a
-> Array
     (Eval
        (Foldl'
           (Flip DeleteDim)
           '[n, n]
           (Eval (Rev (Eval (PreDeletePositionsGo '[1] '[])) '[])))
      <> Eval
           (Foldl'
              (Flip DeleteDim)
              '[n, n]
              (Eval (Rev (Eval (PreDeletePositionsGo '[0] '[])) '[]))))
     a
forall a (ds0 :: [Nat]) (ds1 :: [Nat]) (s0 :: [Nat]) (s1 :: [Nat])
       (so0 :: [Nat]) (so1 :: [Nat]) (st :: [Nat]) (si :: [Nat]).
(Additive a, Multiplicative a, KnownNats s0, KnownNats s1,
 KnownNats ds0, KnownNats ds1, KnownNats so0, KnownNats so1,
 KnownNats st, KnownNats si, so0 ~ Eval (DeleteDims ds0 s0),
 so1 ~ Eval (DeleteDims ds1 s1), si ~ Eval (GetDims ds0 s0),
 si ~ Eval (GetDims ds1 s1), st ~ Eval (so0 ++ so1),
 ds0 ~ '[Eval (Eval (Rank s0) - 1)], ds1 ~ '[0]) =>
Array s0 a -> Array s1 a -> Array st a
F.mult Array ('[n] <> '[n]) a
ti MatrixM n a
tl)
    pow :: Array
  (Eval
     (Foldl'
        (Flip DeleteDim)
        '[n, n]
        (Eval (Rev (Eval (PreDeletePositionsGo '[1] '[])) '[])))
   <> Eval
        (Foldl'
           (Flip DeleteDim)
           '[n, n]
           (Eval (Rev (Eval (PreDeletePositionsGo '[0] '[])) '[]))))
  a
-> Int -> MatrixM n a
pow Array
  (Eval
     (Foldl'
        (Flip DeleteDim)
        '[n, n]
        (Eval (Rev (Eval (PreDeletePositionsGo '[1] '[])) '[])))
   <> Eval
        (Foldl'
           (Flip DeleteDim)
           '[n, n]
           (Eval (Rev (Eval (PreDeletePositionsGo '[0] '[])) '[]))))
  a
xs Int
x = ((MatrixM n a -> MatrixM n a) -> MatrixM n a -> MatrixM n a)
-> MatrixM n a -> [MatrixM n a -> MatrixM n a] -> MatrixM n a
forall a b. (a -> b -> b) -> b -> [a] -> b
forall (t :: * -> *) a b.
Foldable t =>
(a -> b -> b) -> b -> t a -> b
foldr (MatrixM n a -> MatrixM n a) -> MatrixM n a -> MatrixM n a
forall a b. (a -> b) -> a -> b
($) (forall (s :: [Nat]) a.
(KnownNats s, Additive a, Multiplicative a) =>
Array s a
F.ident @[n, n]) (Int -> (MatrixM n a -> MatrixM n a) -> [MatrixM n a -> MatrixM n a]
forall a. Int -> a -> [a]
replicate Int
x (Array
  (Eval
     (Foldl'
        (Flip DeleteDim)
        '[n, n]
        (Eval (Rev (Eval (PreDeletePositionsGo '[1] '[])) '[])))
   <> Eval
        (Foldl'
           (Flip DeleteDim)
           '[n, n]
           (Eval (Rev (Eval (PreDeletePositionsGo '[0] '[])) '[]))))
  a
-> MatrixM n a -> MatrixM n a
forall a (ds0 :: [Nat]) (ds1 :: [Nat]) (s0 :: [Nat]) (s1 :: [Nat])
       (so0 :: [Nat]) (so1 :: [Nat]) (st :: [Nat]) (si :: [Nat]).
(Additive a, Multiplicative a, KnownNats s0, KnownNats s1,
 KnownNats ds0, KnownNats ds1, KnownNats so0, KnownNats so1,
 KnownNats st, KnownNats si, so0 ~ Eval (DeleteDims ds0 s0),
 so1 ~ Eval (DeleteDims ds1 s1), si ~ Eval (GetDims ds0 s0),
 si ~ Eval (GetDims ds1 s1), st ~ Eval (so0 ++ so1),
 ds0 ~ '[Eval (Eval (Rank s0) - 1)], ds1 ~ '[0]) =>
Array s0 a -> Array s1 a -> Array st a
F.mult Array
  (Eval
     (Foldl'
        (Flip DeleteDim)
        '[n, n]
        (Eval (Rev (Eval (PreDeletePositionsGo '[1] '[])) '[])))
   <> Eval
        (Foldl'
           (Flip DeleteDim)
           '[n, n]
           (Eval (Rev (Eval (PreDeletePositionsGo '[0] '[])) '[]))))
  a
xs))
    zero' :: MatrixM n a
zero' = forall (s :: [Nat]) a. KnownNats s => a -> Array s a
F.konst @[n, n] a
forall a. Additive a => a
zero
    add :: MatrixM n a -> MatrixM n a -> MatrixM n a
add = (a -> a -> a) -> MatrixM n a -> MatrixM n a -> MatrixM n a
forall (s :: [Nat]) a b c.
KnownNats s =>
(a -> b -> c) -> Array s a -> Array s b -> Array s c
F.zipWith a -> a -> a
forall a. Additive a => a -> a -> a
(+)
    sum' :: Array '[n] (MatrixM n a) -> MatrixM n a
sum' = (MatrixM n a -> MatrixM n a -> MatrixM n a)
-> MatrixM n a -> Array '[n] (MatrixM n a) -> MatrixM n a
forall b a. (b -> a -> b) -> b -> Array '[n] a -> b
forall (t :: * -> *) b a.
Foldable t =>
(b -> a -> b) -> b -> t a -> b
foldl' MatrixM n a -> MatrixM n a -> MatrixM n a
add MatrixM n a
zero'
    i :: MatrixM n a
i = MatrixM n a -> Array ('[n] <> '[n]) a -> MatrixM n a
forall a (ds0 :: [Nat]) (ds1 :: [Nat]) (s0 :: [Nat]) (s1 :: [Nat])
       (so0 :: [Nat]) (so1 :: [Nat]) (st :: [Nat]) (si :: [Nat]).
(Additive a, Multiplicative a, KnownNats s0, KnownNats s1,
 KnownNats ds0, KnownNats ds1, KnownNats so0, KnownNats so1,
 KnownNats st, KnownNats si, so0 ~ Eval (DeleteDims ds0 s0),
 so1 ~ Eval (DeleteDims ds1 s1), si ~ Eval (GetDims ds0 s0),
 si ~ Eval (GetDims ds1 s1), st ~ Eval (so0 ++ so1),
 ds0 ~ '[Eval (Eval (Rank s0) - 1)], ds1 ~ '[0]) =>
Array s0 a -> Array s1 a -> Array st a
F.mult (Array '[n] (MatrixM n a) -> MatrixM n a
sum' ((Int -> MatrixM n a) -> Array '[n] Int -> Array '[n] (MatrixM n a)
forall a b. (a -> b) -> Array '[n] a -> Array '[n] b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap (Array
  (Eval
     (Foldl'
        (Flip DeleteDim)
        '[n, n]
        (Eval (Rev (Eval (PreDeletePositionsGo '[1] '[])) '[])))
   <> Eval
        (Foldl'
           (Flip DeleteDim)
           '[n, n]
           (Eval (Rev (Eval (PreDeletePositionsGo '[0] '[])) '[]))))
  a
-> Int -> MatrixM n a
pow Array
  (Eval
     (Foldl'
        (Flip DeleteDim)
        '[n, n]
        (Eval (Rev (Eval (PreDeletePositionsGo '[1] '[])) '[])))
   <> Eval
        (Foldl'
           (Flip DeleteDim)
           '[n, n]
           (Eval (Rev (Eval (PreDeletePositionsGo '[0] '[])) '[]))))
  a
l) (forall (s :: [Nat]). KnownNats s => Array s Int
F.range @'[n]))) Array ('[n] <> '[n]) a
ti

-- | Newtype wrapper that gives square matrices a multiplicative (and, over a
-- field, divisible) structure.
--
-- The unit is the identity matrix; multiplication is the usual matrix product.
-- Division uses Cholesky-based inversion.
newtype MatField n a = MatField {forall (n :: Nat) a. MatField n a -> MatrixM n a
unMatField :: MatrixM n a}
  deriving stock (MatField n a -> MatField n a -> Bool
(MatField n a -> MatField n a -> Bool)
-> (MatField n a -> MatField n a -> Bool) -> Eq (MatField n a)
forall (n :: Nat) a. Eq a => MatField n a -> MatField n a -> Bool
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: forall (n :: Nat) a. Eq a => MatField n a -> MatField n a -> Bool
== :: MatField n a -> MatField n a -> Bool
$c/= :: forall (n :: Nat) a. Eq a => MatField n a -> MatField n a -> Bool
/= :: MatField n a -> MatField n a -> Bool
Eq, Int -> MatField n a -> ShowS
[MatField n a] -> ShowS
MatField n a -> [Char]
(Int -> MatField n a -> ShowS)
-> (MatField n a -> [Char])
-> ([MatField n a] -> ShowS)
-> Show (MatField n a)
forall (n :: Nat) a. Show a => Int -> MatField n a -> ShowS
forall (n :: Nat) a. Show a => [MatField n a] -> ShowS
forall (n :: Nat) a. Show a => MatField n a -> [Char]
forall a.
(Int -> a -> ShowS) -> (a -> [Char]) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: forall (n :: Nat) a. Show a => Int -> MatField n a -> ShowS
showsPrec :: Int -> MatField n a -> ShowS
$cshow :: forall (n :: Nat) a. Show a => MatField n a -> [Char]
show :: MatField n a -> [Char]
$cshowList :: forall (n :: Nat) a. Show a => [MatField n a] -> ShowS
showList :: [MatField n a] -> ShowS
Show)

instance
  ( KnownNat n,
    KnownNats '[n, n],
    Additive a,
    Multiplicative a
  ) =>
  Multiplicative (MatField n a)
  where
  one :: MatField n a
one = MatrixM n a -> MatField n a
forall (n :: Nat) a. MatrixM n a -> MatField n a
MatField (forall (s :: [Nat]) a.
(KnownNats s, Additive a, Multiplicative a) =>
Array s a
F.ident @[n, n])
  MatField MatrixM n a
x * :: MatField n a -> MatField n a -> MatField n a
* MatField MatrixM n a
y = MatrixM n a -> MatField n a
forall (n :: Nat) a. MatrixM n a -> MatField n a
MatField (MatrixM n a -> MatrixM n a -> MatrixM n a
forall a (ds0 :: [Nat]) (ds1 :: [Nat]) (s0 :: [Nat]) (s1 :: [Nat])
       (so0 :: [Nat]) (so1 :: [Nat]) (st :: [Nat]) (si :: [Nat]).
(Additive a, Multiplicative a, KnownNats s0, KnownNats s1,
 KnownNats ds0, KnownNats ds1, KnownNats so0, KnownNats so1,
 KnownNats st, KnownNats si, so0 ~ Eval (DeleteDims ds0 s0),
 so1 ~ Eval (DeleteDims ds1 s1), si ~ Eval (GetDims ds0 s0),
 si ~ Eval (GetDims ds1 s1), st ~ Eval (so0 ++ so1),
 ds0 ~ '[Eval (Eval (Rank s0) - 1)], ds1 ~ '[0]) =>
Array s0 a -> Array s1 a -> Array st a
F.mult MatrixM n a
x MatrixM n a
y)

instance
  ( KnownNat n,
    KnownNats '[n, n],
    Additive a,
    Multiplicative a,
    ExpField a
  ) =>
  Divisive (MatField n a)
  where
  recip :: MatField n a -> MatField n a
recip (MatField MatrixM n a
x) = MatrixM n a -> MatField n a
forall (n :: Nat) a. MatrixM n a -> MatField n a
MatField (MatrixM n a -> MatrixM n a
forall (n :: Nat) a.
(KnownNat n, KnownNats '[n, n], ExpField a) =>
MatrixM n a -> MatrixM n a
inverseM MatrixM n a
x)
  MatField MatrixM n a
x / :: MatField n a -> MatField n a -> MatField n a
/ MatField MatrixM n a
y = MatrixM n a -> MatField n a
forall (n :: Nat) a. MatrixM n a -> MatField n a
MatField (MatrixM n a -> MatrixM n a -> MatrixM n a
forall a (ds0 :: [Nat]) (ds1 :: [Nat]) (s0 :: [Nat]) (s1 :: [Nat])
       (so0 :: [Nat]) (so1 :: [Nat]) (st :: [Nat]) (si :: [Nat]).
(Additive a, Multiplicative a, KnownNats s0, KnownNats s1,
 KnownNats ds0, KnownNats ds1, KnownNats so0, KnownNats so1,
 KnownNats st, KnownNats si, so0 ~ Eval (DeleteDims ds0 s0),
 so1 ~ Eval (DeleteDims ds1 s1), si ~ Eval (GetDims ds0 s0),
 si ~ Eval (GetDims ds1 s1), st ~ Eval (so0 ++ so1),
 ds0 ~ '[Eval (Eval (Rank s0) - 1)], ds1 ~ '[0]) =>
Array s0 a -> Array s1 a -> Array st a
F.mult MatrixM n a
x (MatrixM n a -> MatrixM n a
forall (n :: Nat) a.
(KnownNat n, KnownNats '[n, n], ExpField a) =>
MatrixM n a -> MatrixM n a
inverseM MatrixM n a
y))