{-# LANGUAGE DataKinds #-}
{-# LANGUAGE TypeApplications #-}
{-# OPTIONS_GHC -Wno-orphans #-}
module Circuit.Mat.Square
( Square,
inverse,
)
where
import Data.Bool (bool)
import Data.List (maximumBy)
import Data.Ord (comparing)
import Data.Proxy (Proxy (..))
import GHC.TypeNats (KnownNat, Nat, natVal)
import Harpie.Fixed (Array)
import Harpie.Fixed qualified as F
import NumHask.Algebra.Additive (Additive (..), Subtractive (..))
import NumHask.Algebra.Metric (Absolute, abs)
import NumHask.Algebra.Multiplicative (Divisive (..), Multiplicative (..))
import NumHask.Data.Integral (FromInteger (..))
import NumHask.Data.Rational (FromRational (..))
import Prelude hiding (abs, fromInteger, fromRational, (*), (+), (-), (/))
import Prelude qualified as P
type Square (n :: Nat) a = Array '[n, n] a
instance
( KnownNat n,
Additive a,
Multiplicative a
) =>
Multiplicative (Square n a)
where
* :: Square n a -> Square n a -> Square n a
(*) = Square n a -> Square n a -> Square 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
one :: Square n a
one = Square n a
forall (s :: [Nat]) a.
(KnownNats s, Additive a, Multiplicative a) =>
Array s a
F.ident
instance
( KnownNat n,
Additive a,
Subtractive a,
Multiplicative a,
Divisive a,
Absolute a,
Ord a
) =>
Divisive (Square n a)
where
recip :: Square n a -> Square n a
recip = Square n a -> Square n a
forall (n :: Nat) a.
(KnownNat n, Subtractive a, Divisive a, Absolute a, Ord a) =>
Square n a -> Square n a
inverse
instance
( KnownNat n,
Additive a,
Multiplicative a,
FromInteger a
) =>
FromInteger (Square n a)
where
fromInteger :: Integer -> Square n a
fromInteger Integer
k = (a -> a -> a) -> Square n a -> Square n a -> Square 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. Multiplicative a => a -> a -> a
(*) (a -> Square n a
forall (s :: [Nat]) a. KnownNats s => a -> Array s a
F.konst (Integer -> a
forall a. FromInteger a => Integer -> a
fromInteger Integer
k)) Square n a
forall (s :: [Nat]) a.
(KnownNats s, Additive a, Multiplicative a) =>
Array s a
F.ident
instance
( KnownNat n,
Additive a,
Multiplicative a,
FromRational a
) =>
FromRational (Square n a)
where
fromRational :: Rational -> Square n a
fromRational Rational
q = (a -> a -> a) -> Square n a -> Square n a -> Square 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. Multiplicative a => a -> a -> a
(*) (a -> Square n a
forall (s :: [Nat]) a. KnownNats s => a -> Array s a
F.konst (Rational -> a
forall a. FromRational a => Rational -> a
fromRational Rational
q)) Square n a
forall (s :: [Nat]) a.
(KnownNats s, Additive a, Multiplicative a) =>
Array s a
F.ident
inverse ::
forall n a.
( KnownNat n,
Subtractive a,
Divisive a,
Absolute a,
Ord a
) =>
Square n a ->
Square n a
inverse :: forall (n :: Nat) a.
(KnownNat n, Subtractive a, Divisive a, Absolute a, Ord a) =>
Square n a -> Square n a
inverse Square n a
a =
let n :: Int
n = Nat -> Int
forall a b. (Integral a, Num b) => a -> b
P.fromIntegral (Proxy n -> Nat
forall (n :: Nat) (proxy :: Nat -> *). KnownNat n => proxy n -> Nat
natVal (forall (t :: Nat). Proxy t
forall {k} (t :: k). Proxy t
Proxy @n))
row :: Int -> [a]
row Int
i = [Square n a
a Square n a -> [Int] -> a
forall (s :: [Nat]) a. KnownNats s => Array s a -> [Int] -> a
F.! [Int
i, Int
j] | Int
j <- [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
P.- Int
1]]
eye :: a -> a -> a
eye a
i a
j = 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 (a
i a -> a -> Bool
forall a. Eq a => a -> a -> Bool
P.== a
j)
aug :: [[a]]
aug = [Int -> [a]
row Int
i [a] -> [a] -> [a]
forall a. [a] -> [a] -> [a]
P.++ [Int -> Int -> a
forall {a} {a}. (Additive a, Multiplicative a, Eq a) => a -> a -> a
eye Int
i Int
j | Int
j <- [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
P.- Int
1]] | Int
i <- [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
P.- Int
1]]
reduced :: [[a]]
reduced = [[a]] -> [[a]]
forall a.
(Subtractive a, Divisive a, Absolute a, Ord a) =>
[[a]] -> [[a]]
gaussJordan [[a]]
aug
invRows :: [[a]]
invRows = ([a] -> [a]) -> [[a]] -> [[a]]
forall a b. (a -> b) -> [a] -> [b]
P.map (Int -> [a] -> [a]
forall a. Int -> [a] -> [a]
P.drop Int
n) [[a]]
reduced
in [a] -> Square n a
forall (s :: [Nat]) a t.
(KnownNats s, FromVector t a) =>
t -> Array s a
F.array ([[a]] -> [a]
forall (t :: * -> *) a. Foldable t => t [a] -> [a]
P.concat [[a]]
invRows)
gaussJordan ::
( Subtractive a,
Divisive a,
Absolute a,
Ord a
) =>
[[a]] ->
[[a]]
gaussJordan :: forall a.
(Subtractive a, Divisive a, Absolute a, Ord a) =>
[[a]] -> [[a]]
gaussJordan [] = []
gaussJordan [[a]]
m0 = ([[a]], [[a]]) -> [[a]]
forall a b. (a, b) -> b
P.snd ((([[a]], [[a]]) -> Int -> ([[a]], [[a]]))
-> ([[a]], [[a]]) -> [Int] -> ([[a]], [[a]])
forall b a. (b -> a -> b) -> b -> [a] -> b
forall (t :: * -> *) b a.
Foldable t =>
(b -> a -> b) -> b -> t a -> b
P.foldl' ([[a]], [[a]]) -> Int -> ([[a]], [[a]])
forall {b}.
(Mag b ~ b, Ord b, Basis b, Divisive b, Subtractive b) =>
([[b]], [[b]]) -> Int -> ([[b]], [[b]])
step ([[a]]
m0, []) [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
P.- Int
1])
where
n :: Int
n = [[a]] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
P.length [[a]]
m0
step :: ([[b]], [[b]]) -> Int -> ([[b]], [[b]])
step ([[b]]
rows, [[b]]
doneRows) Int
k =
let pivotOffset :: Int
pivotOffset =
(Int, [b]) -> Int
forall a b. (a, b) -> a
fst ((Int, [b]) -> Int) -> (Int, [b]) -> Int
forall a b. (a -> b) -> a -> b
$
((Int, [b]) -> (Int, [b]) -> Ordering)
-> [(Int, [b])] -> (Int, [b])
forall (t :: * -> *) a.
Foldable t =>
(a -> a -> Ordering) -> t a -> a
maximumBy
(((Int, [b]) -> b) -> (Int, [b]) -> (Int, [b]) -> Ordering
forall a b. Ord a => (b -> a) -> b -> b -> Ordering
comparing (b -> b
forall a. Absolute a => a -> a
abs (b -> b) -> ((Int, [b]) -> b) -> (Int, [b]) -> b
forall b c a. (b -> c) -> (a -> b) -> a -> c
. ([b] -> Int -> b
forall a. HasCallStack => [a] -> Int -> a
!! Int
k) ([b] -> b) -> ((Int, [b]) -> [b]) -> (Int, [b]) -> b
forall b c a. (b -> c) -> (a -> b) -> a -> c
. (Int, [b]) -> [b]
forall a b. (a, b) -> b
snd))
([Int] -> [[b]] -> [(Int, [b])]
forall a b. [a] -> [b] -> [(a, b)]
P.zip [Int
0 ..] [[b]]
rows)
rows' :: [[b]]
rows' = Int -> Int -> [[b]] -> [[b]]
forall a. Int -> Int -> [a] -> [a]
swap Int
k Int
pivotOffset [[b]]
rows
pivot :: b
pivot = [[b]]
rows' [[b]] -> Int -> [b]
forall a. HasCallStack => [a] -> Int -> a
!! Int
k [b] -> Int -> b
forall a. HasCallStack => [a] -> Int -> a
!! Int
k
rowK :: [b]
rowK = (b -> b) -> [b] -> [b]
forall a b. (a -> b) -> [a] -> [b]
P.map (b -> b -> b
forall a. Divisive a => a -> a -> a
/ b
pivot) ([[b]]
rows' [[b]] -> Int -> [b]
forall a. HasCallStack => [a] -> Int -> a
!! Int
k)
eliminate :: Int -> [b] -> [b]
eliminate Int
i [b]
row
| Int
i Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
P.== Int
k = [b]
rowK
| Bool
P.otherwise =
let factor :: b
factor = [b]
row [b] -> Int -> b
forall a. HasCallStack => [a] -> Int -> a
!! Int
k
in (b -> b -> b) -> [b] -> [b] -> [b]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
P.zipWith (-) [b]
row ((b -> b) -> [b] -> [b]
forall a b. (a -> b) -> [a] -> [b]
P.map (b
factor b -> b -> b
forall a. Multiplicative a => a -> a -> a
*) [b]
rowK)
remaining :: [[b]]
remaining = (Int -> [b] -> [b]) -> [Int] -> [[b]] -> [[b]]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
P.zipWith Int -> [b] -> [b]
eliminate [Int
0 ..] [[b]]
rows'
in ([[b]]
remaining, [[b]]
doneRows [[b]] -> [[b]] -> [[b]]
forall a. [a] -> [a] -> [a]
P.++ [[b]
rowK])
swap :: Int -> Int -> [a] -> [a]
swap :: forall a. Int -> Int -> [a] -> [a]
swap Int
i Int
j [a]
xs =
[ case Int
k of
Int
_ | Int
k Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
P.== Int
i -> [a]
xs [a] -> Int -> a
forall a. HasCallStack => [a] -> Int -> a
P.!! Int
j
Int
_ | Int
k Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
P.== Int
j -> [a]
xs [a] -> Int -> a
forall a. HasCallStack => [a] -> Int -> a
P.!! Int
i
Int
_ -> a
x
| (Int
k, a
x) <- [Int] -> [a] -> [(Int, a)]
forall a b. [a] -> [b] -> [(a, b)]
P.zip [Int
0 ..] [a]
xs
]