{-# LANGUAGE LambdaCase #-}
{-# LANGUAGE RebindableSyntax #-}
{-# OPTIONS_GHC -Wno-incomplete-uni-patterns #-}
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
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)
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
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
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))
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]]
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)
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"
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"
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
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"
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"
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
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
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'
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)
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"
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])
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"