{-# LANGUAGE RebindableSyntax #-}
module Circuit.Diff.Star
(
traceStarMatrix,
solveStarBody,
traceStarFromD,
traceStarMatrixD,
)
where
import Circuit.Bimonoid (Merge, MergeZero, Zero)
import Circuit.Bimonoid qualified as CB
import Circuit.Body (Body (..))
import Circuit.Diff.Circuit (Diff (..), traceStarFrom)
import Circuit.Diff.Evidence (StarChannel (..))
import Circuit.Mat.Dense (fromLists, matVec)
import Circuit.Mat.Dense qualified as MD
import Circuit.Pullback (Pullback (..))
import NumHask.Algebra.Additive qualified as NHA
import NumHask.Algebra.Multiplicative qualified as NHM
import NumHask.Algebra.Ring qualified as NHR
import NumHask.Free.Carriers (FieldStar (..))
import NumHask.Prelude
traceStarMatrix ::
(NHR.StarSemiring j, MergeZero (->) c) =>
[j] ->
Int ->
Diff p ([j], b) ([j], c) ->
Diff p b c
traceStarMatrix :: forall {k} j c (p :: k) b.
(StarSemiring j, MergeZero (->) c) =>
[j] -> Int -> Diff p ([j], b) ([j], c) -> Diff p b c
traceStarMatrix [j]
x0 Int
n (Diff ([j], b) -> (([j], c), ([j], c) -> ([j], b))
body) = (b -> (c, c -> b)) -> Diff p b c
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((b -> (c, c -> b)) -> Diff p b c)
-> (b -> (c, c -> b)) -> Diff p b c
forall a b. (a -> b) -> a -> b
$ \b
b ->
let dim :: Int
dim = [j] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [j]
x0
stepFwd :: [j] -> [j]
stepFwd [j]
x = let (([j]
x', c
_), ([j], c) -> ([j], b)
_) = ([j], b) -> (([j], c), ([j], c) -> ([j], b))
body ([j]
x, b
b) in [j]
x'
a :: [j]
a = ([j] -> [j]) -> [j] -> [[j]]
forall a. (a -> a) -> a -> [a]
iterate [j] -> [j]
stepFwd [j]
x0 [[j]] -> Int -> [j]
forall a. HasCallStack => [a] -> Int -> a
!! Int
n
(([j]
_, c
c), ([j], c) -> ([j], b)
backward) = ([j], b) -> (([j], c), ([j], c) -> ([j], b))
body ([j]
a, b
b)
zeroV :: [j]
zeroV = Int -> j -> [j]
forall a. Int -> a -> [a]
replicate Int
dim j
forall a. Additive a => a
NHA.zero
basisVec :: Int -> [j]
basisVec Int
i = [if Int
k Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
i then j
forall a. Multiplicative a => a
NHM.one else j
forall a. Additive a => a
NHA.zero | Int
k <- [Int
0 .. Int
dim Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1]]
cols :: [[j]]
cols = [([j], b) -> [j]
forall a b. (a, b) -> a
fst (([j], c) -> ([j], b)
backward (Int -> [j]
basisVec Int
i, () -> c
forall {k} (arr :: * -> k -> *) (a :: k). Zero arr a => arr () a
CB.zero ())) | Int
i <- [Int
0 .. Int
dim Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1]]
aMat :: Matrix j
aMat = [[j]] -> Matrix j
forall a. [[a]] -> Matrix a
fromLists [[[j]
col [j] -> Int -> j
forall a. HasCallStack => [a] -> Int -> a
!! Int
k | [j]
col <- [[j]]
cols] | Int
k <- [Int
0 .. Int
dim Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1]]
aStar :: Matrix j
aStar = Matrix j -> Matrix j
forall a. StarSemiring a => Matrix a -> Matrix a
MD.starMatrix Matrix j
aMat
pullback :: c -> b
pullback c
dc =
let cdc :: [j]
cdc = ([j], b) -> [j]
forall a b. (a, b) -> a
fst (([j], c) -> ([j], b)
backward ([j]
zeroV, c
dc))
in ([j], b) -> b
forall a b. (a, b) -> b
snd (([j], c) -> ([j], b)
backward (Matrix j -> [j] -> [j]
forall a. (Additive a, Multiplicative a) => Matrix a -> [a] -> [a]
matVec Matrix j
aStar [j]
cdc, c
dc))
in (c
c, c -> b
pullback)
traceStarFromD ::
(MergeZero (->) c) =>
Double ->
Int ->
Diff p (Double, b) (Double, c) ->
Diff p b c
traceStarFromD :: forall {k} c (p :: k) b.
MergeZero (->) c =>
Double -> Int -> Diff p (Double, b) (Double, c) -> Diff p b c
traceStarFromD Double
x0 Int
n (Diff (Double, b) -> ((Double, c), (Double, c) -> (Double, b))
body) =
FieldStar
-> Int -> Diff p (FieldStar, b) (FieldStar, c) -> Diff p b c
forall {k} j c (p :: k) b.
(StarSemiring j, MergeZero (->) c) =>
j -> Int -> Diff p (j, b) (j, c) -> Diff p b c
traceStarFrom (Double -> FieldStar
FieldStar Double
x0) Int
n (((FieldStar, b)
-> ((FieldStar, c), (FieldStar, c) -> (FieldStar, b)))
-> Diff p (FieldStar, b) (FieldStar, c)
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff (FieldStar, b)
-> ((FieldStar, c), (FieldStar, c) -> (FieldStar, b))
body')
where
body' :: (FieldStar, b)
-> ((FieldStar, c), (FieldStar, c) -> (FieldStar, b))
body' (FieldStar Double
x, b
b) =
let ((Double
x', c
c), (Double, c) -> (Double, b)
back) = (Double, b) -> ((Double, c), (Double, c) -> (Double, b))
body (Double
x, b
b)
in ( (Double -> FieldStar
FieldStar Double
x', c
c),
\(FieldStar Double
dx', c
dc) ->
let (Double
dx, b
db) = (Double, c) -> (Double, b)
back (Double
dx', c
dc)
in (Double -> FieldStar
FieldStar Double
dx, b
db)
)
traceStarMatrixD ::
(MergeZero (->) c) =>
[Double] ->
Int ->
Diff p ([Double], b) ([Double], c) ->
Diff p b c
traceStarMatrixD :: forall {k} c (p :: k) b.
MergeZero (->) c =>
[Double] -> Int -> Diff p ([Double], b) ([Double], c) -> Diff p b c
traceStarMatrixD [Double]
x0 Int
n (Diff ([Double], b) -> (([Double], c), ([Double], c) -> ([Double], b))
body) =
[FieldStar]
-> Int -> Diff p ([FieldStar], b) ([FieldStar], c) -> Diff p b c
forall {k} j c (p :: k) b.
(StarSemiring j, MergeZero (->) c) =>
[j] -> Int -> Diff p ([j], b) ([j], c) -> Diff p b c
traceStarMatrix ((Double -> FieldStar) -> [Double] -> [FieldStar]
forall a b. (a -> b) -> [a] -> [b]
map Double -> FieldStar
FieldStar [Double]
x0) Int
n ((([FieldStar], b)
-> (([FieldStar], c), ([FieldStar], c) -> ([FieldStar], b)))
-> Diff p ([FieldStar], b) ([FieldStar], c)
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ([FieldStar], b)
-> (([FieldStar], c), ([FieldStar], c) -> ([FieldStar], b))
body')
where
wrap :: [Double] -> [FieldStar]
wrap = (Double -> FieldStar) -> [Double] -> [FieldStar]
forall a b. (a -> b) -> [a] -> [b]
map Double -> FieldStar
FieldStar
unwrap :: [FieldStar] -> [Double]
unwrap = (FieldStar -> Double) -> [FieldStar] -> [Double]
forall a b. (a -> b) -> [a] -> [b]
map (\(FieldStar Double
d) -> Double
d)
body' :: ([FieldStar], b)
-> (([FieldStar], c), ([FieldStar], c) -> ([FieldStar], b))
body' ([FieldStar]
js, b
b) =
let (([Double]
xs, c
c), ([Double], c) -> ([Double], b)
back) = ([Double], b) -> (([Double], c), ([Double], c) -> ([Double], b))
body ([FieldStar] -> [Double]
unwrap [FieldStar]
js, b
b)
in ( ([Double] -> [FieldStar]
wrap [Double]
xs, c
c),
\([FieldStar]
djs, c
dc) ->
let ([Double]
dxs, b
db) = ([Double], c) -> ([Double], b)
back ([FieldStar] -> [Double]
unwrap [FieldStar]
djs, c
dc)
in ([Double] -> [FieldStar]
wrap [Double]
dxs, b
db)
)
instance Circuit.Bimonoid.Merge (->) FieldStar where
plus :: (FieldStar, FieldStar) -> FieldStar
plus (FieldStar Double
x, FieldStar Double
y) = Double -> FieldStar
FieldStar (Double
x Double -> Double -> Double
forall a. Additive a => a -> a -> a
+ Double
y)
instance Circuit.Bimonoid.Zero (->) FieldStar where
zero :: () -> FieldStar
zero ()
_ = Double -> FieldStar
FieldStar Double
0
solveAffine ::
forall j c b.
((StarChannel j, c) -> (StarChannel j, b)) ->
c ->
b
solveAffine :: forall j c b. ((StarChannel j, c) -> (StarChannel j, b)) -> c -> b
solveAffine (StarChannel j, c) -> (StarChannel j, b)
body c
dc =
let
probe :: StarChannel j
probe :: StarChannel j
probe =
StarChannel
{ starDim :: Int
starDim = [Char] -> Int
forall a. HasCallStack => [Char] -> a
error [Char]
"Circuit.Diff.Star.solveAffine: probe dim evaluated",
starData :: j
starData = [Char] -> j
forall a. HasCallStack => [Char] -> a
error [Char]
"Circuit.Diff.Star.solveAffine: probe data evaluated",
starZero :: Int -> j
starZero = [Char] -> Int -> j
forall a. HasCallStack => [Char] -> a
error [Char]
"Circuit.Diff.Star.solveAffine: probe zero evaluated",
starBasis :: Int -> Int -> j
starBasis = [Char] -> Int -> Int -> j
forall a. HasCallStack => [Char] -> a
error [Char]
"Circuit.Diff.Star.solveAffine: probe basis evaluated",
starAdd :: j -> j -> j
starAdd = [Char] -> j -> j -> j
forall a. HasCallStack => [Char] -> a
error [Char]
"Circuit.Diff.Star.solveAffine: probe add evaluated",
starNegate :: j -> j
starNegate = [Char] -> j -> j
forall a. HasCallStack => [Char] -> a
error [Char]
"Circuit.Diff.Star.solveAffine: probe negate evaluated",
starSelfMatrix :: Int -> (j -> j) -> Matrix (Scalar j)
starSelfMatrix = [Char] -> Int -> (j -> j) -> Matrix (Scalar j)
forall a. HasCallStack => [Char] -> a
error [Char]
"Circuit.Diff.Star.solveAffine: probe selfMatrix evaluated",
starApplyMatrix :: Matrix (Scalar j) -> j -> j
starApplyMatrix = [Char] -> Matrix (Scalar j) -> j -> j
forall a. HasCallStack => [Char] -> a
error [Char]
"Circuit.Diff.Star.solveAffine: probe applyMatrix evaluated",
starMatrix :: Matrix (Scalar j) -> Matrix (Scalar j)
starMatrix = [Char] -> Matrix (Scalar j) -> Matrix (Scalar j)
forall a. HasCallStack => [Char] -> a
error [Char]
"Circuit.Diff.Star.solveAffine: probe matrix evaluated"
}
(StarChannel j
scOut, b
_) = (StarChannel j, c) -> (StarChannel j, b)
body (StarChannel j
probe, c
dc)
zeroJ :: j
zeroJ = StarChannel j -> Int -> j
forall s. StarChannel s -> Int -> s
starZero StarChannel j
scOut (StarChannel j -> Int
forall s. StarChannel s -> Int
starDim StarChannel j
scOut)
zeroSC :: StarChannel j
zeroSC = StarChannel j
scOut {starData = zeroJ}
cdc :: j
cdc = StarChannel j -> j
forall s. StarChannel s -> s
starData ((StarChannel j, b) -> StarChannel j
forall a b. (a, b) -> a
fst ((StarChannel j, c) -> (StarChannel j, b)
body (StarChannel j
zeroSC, c
dc)))
negCdc :: j
negCdc = StarChannel j -> j -> j
forall s. StarChannel s -> s -> s
starNegate StarChannel j
scOut j
cdc
aMat :: Matrix (Scalar j)
aMat =
StarChannel j -> Int -> (j -> j) -> Matrix (Scalar j)
forall s. StarChannel s -> Int -> (s -> s) -> Matrix (Scalar s)
starSelfMatrix
StarChannel j
scOut
(StarChannel j -> Int
forall s. StarChannel s -> Int
starDim StarChannel j
scOut)
(\j
dk -> StarChannel j -> j -> j -> j
forall s. StarChannel s -> s -> s -> s
starAdd StarChannel j
scOut (StarChannel j -> j
forall s. StarChannel s -> s
starData ((StarChannel j, b) -> StarChannel j
forall a b. (a, b) -> a
fst ((StarChannel j, c) -> (StarChannel j, b)
body (StarChannel j
zeroSC {starData = dk}, c
dc)))) j
negCdc)
dj :: j
dj = StarChannel j -> Matrix (Scalar j) -> j -> j
forall s. StarChannel s -> Matrix (Scalar s) -> s -> s
starApplyMatrix StarChannel j
scOut (StarChannel j -> Matrix (Scalar j) -> Matrix (Scalar j)
forall s. StarChannel s -> Matrix (Scalar s) -> Matrix (Scalar s)
starMatrix StarChannel j
scOut Matrix (Scalar j)
aMat) j
cdc
in (StarChannel j, b) -> b
forall a b. (a, b) -> b
snd ((StarChannel j, c) -> (StarChannel j, b)
body (StarChannel j
zeroSC {starData = dj}, c
dc))
solveStarBody ::
forall s a b.
Body (,) (StarChannel s) Pullback b a ->
Pullback b a
solveStarBody :: forall s a b. Body (,) (StarChannel s) Pullback b a -> Pullback b a
solveStarBody (Body Pullback (StarChannel s, b) (StarChannel s, a)
f) = (b -> a) -> Pullback b a
forall b a. (b -> a) -> Pullback b a
Pullback (((StarChannel s, b) -> (StarChannel s, a)) -> b -> a
forall j c b. ((StarChannel j, c) -> (StarChannel j, b)) -> c -> b
solveAffine (Pullback (StarChannel s, b) (StarChannel s, a)
-> (StarChannel s, b) -> (StarChannel s, a)
forall b a. Pullback b a -> b -> a
runPullback Pullback (StarChannel s, b) (StarChannel s, a)
f))