{-# LANGUAGE DataKinds #-}
{-# LANGUAGE UndecidableInstances #-}
{-# OPTIONS_GHC -Wno-orphans #-}
module Circuit.Poly.DiffP
(
diffPAt,
diffPAsFamily,
diffPParamGrad,
diffPFromFamily,
traceDiffPFrom,
traceDiffPD,
traceDiffPMatrix,
)
where
import Circuit.Bimonoid (MergeZero, Zero (..))
import Circuit.Category (Category (..))
import Circuit.Channel (Channel (..), Strength (..))
import Circuit.Diff.Param (DiffP (..))
import Circuit.Mat.Dense (Matrix, fromLists, matVec, starMatrix, toLists)
import Circuit.Poly (Mono, Morphism (..), Poly (..), applyLens, lens)
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 Prelude hiding (id, (.))
diffPAt :: DiffP p a b -> p -> Morphism (Mono a a) (Mono b b)
diffPAt :: forall p a b. DiffP p a b -> p -> Morphism (Mono a a) (Mono b b)
diffPAt (DiffP p -> a -> (b, b -> (a, p))
f) p
p = (a -> b) -> (a -> b -> a) -> Morphism (Mono a a) (Mono b b)
forall a b db da.
(a -> b) -> (a -> db -> da) -> Morphism (Mono da a) (Mono db b)
lens a -> b
get a -> b -> a
put
where
get :: a -> b
get a
a = (b, b -> (a, p)) -> b
forall a b. (a, b) -> a
fst (p -> a -> (b, b -> (a, p))
f p
p a
a)
put :: a -> b -> a
put a
a b
db = (a, p) -> a
forall a b. (a, b) -> a
fst ((b, b -> (a, p)) -> b -> (a, p)
forall a b. (a, b) -> b
snd (p -> a -> (b, b -> (a, p))
f p
p a
a) b
db)
diffPParamGrad :: DiffP p a b -> p -> a -> b -> p
diffPParamGrad :: forall p a b. DiffP p a b -> p -> a -> b -> p
diffPParamGrad (DiffP p -> a -> (b, b -> (a, p))
f) p
p a
a b
db = (a, p) -> p
forall a b. (a, b) -> b
snd ((b, b -> (a, p)) -> b -> (a, p)
forall a b. (a, b) -> b
snd (p -> a -> (b, b -> (a, p))
f p
p a
a) b
db)
diffPAsFamily :: DiffP p a b -> Morphism ('Prod ('Const p) (Mono a a)) (Mono b b)
diffPAsFamily :: forall p a b.
DiffP p a b -> Morphism ('Prod ('Const p) (Mono a a)) (Mono b b)
diffPAsFamily DiffP p a b
d = (p -> Morphism (Mono a a) (Mono b b))
-> Morphism ('Prod ('Const p) (Mono a a)) (Mono b b)
forall a (p1 :: Poly) (q :: Poly).
(a -> Morphism p1 q) -> Morphism ('Prod ('Const a) p1) q
Depend (DiffP p a b -> p -> Morphism (Mono a a) (Mono b b)
forall p a b. DiffP p a b -> p -> Morphism (Mono a a) (Mono b b)
diffPAt DiffP p a b
d)
diffPFromFamily ::
(p -> Morphism (Mono a a) (Mono b b)) ->
(p -> a -> b -> p) ->
DiffP p a b
diffPFromFamily :: forall p a b.
(p -> Morphism (Mono a a) (Mono b b))
-> (p -> a -> b -> p) -> DiffP p a b
diffPFromFamily p -> Morphism (Mono a a) (Mono b b)
family p -> a -> b -> p
grad = (p -> a -> (b, b -> (a, p))) -> DiffP p a b
forall p a b. (p -> a -> (b, b -> (a, p))) -> DiffP p a b
DiffP ((p -> a -> (b, b -> (a, p))) -> DiffP p a b)
-> (p -> a -> (b, b -> (a, p))) -> DiffP p a b
forall a b. (a -> b) -> a -> b
$ \p
p a
a ->
let (b
b, b -> a
put) = Morphism (Mono a a) (Mono b b) -> a -> (b, b -> a)
forall da a db b.
Morphism (Mono da a) (Mono db b) -> a -> (b, db -> da)
applyLens (p -> Morphism (Mono a a) (Mono b b)
family p
p) a
a
in (b
b, \b
db -> (b -> a
put b
db, p -> a -> b -> p
grad p
p a
a b
db))
instance (MergeZero (->) p) => Strength (,) (DiffP p) where
strength :: forall b c a. DiffP p b c -> DiffP p (a, b) (a, c)
strength (DiffP p -> b -> (c, c -> (b, p))
f) = (p -> (a, b) -> ((a, c), (a, c) -> ((a, b), p)))
-> DiffP p (a, b) (a, c)
forall p a b. (p -> a -> (b, b -> (a, p))) -> DiffP p a b
DiffP ((p -> (a, b) -> ((a, c), (a, c) -> ((a, b), p)))
-> DiffP p (a, b) (a, c))
-> (p -> (a, b) -> ((a, c), (a, c) -> ((a, b), p)))
-> DiffP p (a, b) (a, c)
forall a b. (a -> b) -> a -> b
$ \p
p0 (a
a, b
b) ->
let (c
c, c -> (b, p)
back) = p -> b -> (c, c -> (b, p))
f p
p0 b
b
in ( (a
a, c
c),
\(a
da, c
dc) ->
let (b
db, p
dp) = c -> (b, p)
back c
dc
in ((a
da, b
db), p
dp)
)
traceDiffPFrom ::
(NHR.StarSemiring j, MergeZero (->) o) =>
j ->
Int ->
DiffP p (j, i) (j, o) ->
DiffP p i o
traceDiffPFrom :: forall j o p i.
(StarSemiring j, MergeZero (->) o) =>
j -> Int -> DiffP p (j, i) (j, o) -> DiffP p i o
traceDiffPFrom j
j0 Int
n (DiffP p -> (j, i) -> ((j, o), (j, o) -> ((j, i), p))
body) = (p -> i -> (o, o -> (i, p))) -> DiffP p i o
forall p a b. (p -> a -> (b, b -> (a, p))) -> DiffP p a b
DiffP ((p -> i -> (o, o -> (i, p))) -> DiffP p i o)
-> (p -> i -> (o, o -> (i, p))) -> DiffP p i o
forall a b. (a -> b) -> a -> b
$ \p
p0 i
i ->
let stepFwd :: j -> j
stepFwd j
j = let ((j
j', o
_), (j, o) -> ((j, i), p)
_) = p -> (j, i) -> ((j, o), (j, o) -> ((j, i), p))
body p
p0 (j
j, i
i) in j
j'
a :: j
a = (j -> j) -> j -> [j]
forall a. (a -> a) -> a -> [a]
iterate j -> j
stepFwd j
j0 [j] -> Int -> j
forall a. HasCallStack => [a] -> Int -> a
!! Int
n
((j
_, o
o), (j, o) -> ((j, i), p)
backward) = p -> (j, i) -> ((j, o), (j, o) -> ((j, i), p))
body p
p0 (j
a, i
i)
((j
aJ, i
_), p
_) = (j, o) -> ((j, i), p)
backward (j
forall a. Multiplicative a => a
NHM.one, () -> o
forall {k} (arr :: * -> k -> *) (a :: k). Zero arr a => arr () a
zero ())
aStar :: j
aStar = j -> j
forall a. StarSemiring a => a -> a
NHR.star j
aJ
pullback :: o -> (i, p)
pullback o
do_ =
let
cdc :: j
cdc = (j, i) -> j
forall a b. (a, b) -> a
fst (((j, i), p) -> (j, i)
forall a b. (a, b) -> a
fst ((j, o) -> ((j, i), p)
backward (j
forall a. Additive a => a
NHA.zero, o
do_)))
dj :: j
dj = j
aStar j -> j -> j
forall a. Multiplicative a => a -> a -> a
NHM.* j
cdc
((j
_, i
di), p
dp) = (j, o) -> ((j, i), p)
backward (j
dj, o
do_)
in (i
di, p
dp)
in (o
o, o -> (i, p)
pullback)
traceDiffPD ::
(MergeZero (->) o) =>
Double ->
Double ->
Int ->
DiffP p (Double, i) (Double, o) ->
DiffP p i o
traceDiffPD :: forall o p i.
MergeZero (->) o =>
Double
-> Double -> Int -> DiffP p (Double, i) (Double, o) -> DiffP p i o
traceDiffPD Double
j0 Double
tol Int
maxIter (DiffP p -> (Double, i) -> ((Double, o), (Double, o) -> ((Double, i), p))
body) = (p -> i -> (o, o -> (i, p))) -> DiffP p i o
forall p a b. (p -> a -> (b, b -> (a, p))) -> DiffP p a b
DiffP ((p -> i -> (o, o -> (i, p))) -> DiffP p i o)
-> (p -> i -> (o, o -> (i, p))) -> DiffP p i o
forall a b. (a -> b) -> a -> b
$ \p
p0 i
i ->
let stepFwd :: Double -> Double
stepFwd Double
j = let ((Double
j', o
_), (Double, o) -> ((Double, i), p)
_) = p -> (Double, i) -> ((Double, o), (Double, o) -> ((Double, i), p))
body p
p0 (Double
j, i
i) in Double
j'
go :: Int -> Double -> Double
go Int
0 Double
j = Double
j
go Int
n Double
j =
let j' :: Double
j' = Double -> Double
stepFwd Double
j
in if Double -> Double
forall a. Num a => a -> a
abs (Double
j' Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
j) Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
<= Double
tol then Double
j' else Int -> Double -> Double
go (Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1) Double
j'
a :: Double
a = Int -> Double -> Double
go Int
maxIter Double
j0
aNext :: Double
aNext = Double -> Double
stepFwd Double
a
((Double
_, o
o), (Double, o) -> ((Double, i), p)
backward) = p -> (Double, i) -> ((Double, o), (Double, o) -> ((Double, i), p))
body p
p0 (Double
a, i
i)
((Double
aJ, i
_), p
_) = (Double, o) -> ((Double, i), p)
backward (Double
1.0, () -> o
forall {k} (arr :: * -> k -> *) (a :: k). Zero arr a => arr () a
zero ())
aStar :: Double
aStar = Double -> Double
forall a. Fractional a => a -> a
recip (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
aJ)
pullback :: o -> (i, p)
pullback o
do_ =
let
cdc :: Double
cdc = (Double, i) -> Double
forall a b. (a, b) -> a
fst (((Double, i), p) -> (Double, i)
forall a b. (a, b) -> a
fst ((Double, o) -> ((Double, i), p)
backward (Double
0.0, o
do_)))
dj :: Double
dj = Double
aStar Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
cdc
((Double
_, i
di), p
dp) = (Double, o) -> ((Double, i), p)
backward (Double
dj, o
do_)
in (i
di, p
dp)
in if Double -> Double
forall a. Num a => a -> a
abs (Double
aNext Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
a) Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
> Double
tol
then [Char] -> (o, o -> (i, p))
forall a. HasCallStack => [Char] -> a
error ([Char]
"traceDiffPD: primal fixed point did not converge within " [Char] -> [Char] -> [Char]
forall a. [a] -> [a] -> [a]
++ Int -> [Char]
forall a. Show a => a -> [Char]
show Int
maxIter [Char] -> [Char] -> [Char]
forall a. [a] -> [a] -> [a]
++ [Char]
" iterations")
else
if Double -> Double
forall a. Num a => a -> a
abs Double
aJ Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
>= Double
1.0
then [Char] -> (o, o -> (i, p))
forall a. HasCallStack => [Char] -> a
error ([Char]
"traceDiffPD: feedback Jacobian |" [Char] -> [Char] -> [Char]
forall a. [a] -> [a] -> [a]
++ Double -> [Char]
forall a. Show a => a -> [Char]
show Double
aJ [Char] -> [Char] -> [Char]
forall a. [a] -> [a] -> [a]
++ [Char]
"| >= 1 is outside the contractive regime")
else (o
o, o -> (i, p)
pullback)
traceDiffPMatrix ::
(MergeZero (->) o) =>
[Double] ->
Double ->
Int ->
DiffP p ([Double], i) ([Double], o) ->
DiffP p i o
traceDiffPMatrix :: forall o p i.
MergeZero (->) o =>
[Double]
-> Double
-> Int
-> DiffP p ([Double], i) ([Double], o)
-> DiffP p i o
traceDiffPMatrix [Double]
x0 Double
tol Int
maxIter (DiffP p
-> ([Double], i)
-> (([Double], o), ([Double], o) -> (([Double], i), p))
body) = (p -> i -> (o, o -> (i, p))) -> DiffP p i o
forall p a b. (p -> a -> (b, b -> (a, p))) -> DiffP p a b
DiffP ((p -> i -> (o, o -> (i, p))) -> DiffP p i o)
-> (p -> i -> (o, o -> (i, p))) -> DiffP p i o
forall a b. (a -> b) -> a -> b
$ \p
p0 i
i ->
let dim :: Int
dim = [Double] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [Double]
x0
stepFwd :: [Double] -> [Double]
stepFwd [Double]
xs = let (([Double]
xs', o
_), ([Double], o) -> (([Double], i), p)
_) = p
-> ([Double], i)
-> (([Double], o), ([Double], o) -> (([Double], i), p))
body p
p0 ([Double]
xs, i
i) in [Double]
xs'
zeroV :: [Double]
zeroV = Int -> Double -> [Double]
forall a. Int -> a -> [a]
replicate Int
dim Double
0.0
oneV :: Int -> [Double]
oneV Int
k = [if Int
j Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
k then Double
1.0 else Double
0.0 | Int
j <- [Int
0 .. Int
dim Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]]
diffV :: [Double] -> [Double] -> [Double]
diffV = (Double -> Double -> Double) -> [Double] -> [Double] -> [Double]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith (-)
normInf :: [a] -> a
normInf [a]
v = [a] -> a
forall a. Ord a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Ord a) => t a -> a
maximum ((a -> a) -> [a] -> [a]
forall a b. (a -> b) -> [a] -> [b]
map a -> a
forall a. Num a => a -> a
abs [a]
v)
go :: Int -> [Double] -> [Double]
go Int
0 [Double]
xs = [Double]
xs
go Int
n [Double]
xs =
let xs' :: [Double]
xs' = [Double] -> [Double]
stepFwd [Double]
xs
in if [Double] -> Double
forall {a}. (Ord a, Num a) => [a] -> a
normInf ([Double] -> [Double] -> [Double]
diffV [Double]
xs' [Double]
xs) Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
<= Double
tol then [Double]
xs' else Int -> [Double] -> [Double]
go (Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1) [Double]
xs'
a :: [Double]
a = Int -> [Double] -> [Double]
go Int
maxIter [Double]
x0
aNext :: [Double]
aNext = [Double] -> [Double]
stepFwd [Double]
a
(([Double]
_, o
o), ([Double], o) -> (([Double], i), p)
backward) = p
-> ([Double], i)
-> (([Double], o), ([Double], o) -> (([Double], i), p))
body p
p0 ([Double]
a, i
i)
cols :: [[Double]]
cols = [([Double], i) -> [Double]
forall a b. (a, b) -> a
fst ((([Double], i), p) -> ([Double], i)
forall a b. (a, b) -> a
fst (([Double], o) -> (([Double], i), p)
backward (Int -> [Double]
oneV Int
k, () -> o
forall {k} (arr :: * -> k -> *) (a :: k). Zero arr a => arr () a
zero ()))) | Int
k <- [Int
0 .. Int
dim Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]]
aMat :: Matrix FieldStar
aMat = [[FieldStar]] -> Matrix FieldStar
forall a. [[a]] -> Matrix a
fromLists [[Double -> FieldStar
FieldStar ([[Double]]
cols [[Double]] -> Int -> [Double]
forall a. HasCallStack => [a] -> Int -> a
!! Int
j [Double] -> Int -> Double
forall a. HasCallStack => [a] -> Int -> a
!! Int
k) | Int
j <- [Int
0 .. Int
dim Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]] | Int
k <- [Int
0 .. Int
dim Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]]
aStar :: Matrix Double
aStar = [[Double]] -> Matrix Double
forall a. [[a]] -> Matrix a
fromLists (([FieldStar] -> [Double]) -> [[FieldStar]] -> [[Double]]
forall a b. (a -> b) -> [a] -> [b]
map ((FieldStar -> Double) -> [FieldStar] -> [Double]
forall a b. (a -> b) -> [a] -> [b]
map FieldStar -> Double
unFieldStar) (Matrix FieldStar -> [[FieldStar]]
forall a. Matrix a -> [[a]]
toLists (Matrix FieldStar -> Matrix FieldStar
forall a. StarSemiring a => Matrix a -> Matrix a
starMatrix Matrix FieldStar
aMat)))
pullback :: o -> (i, p)
pullback o
do_ =
let
cdc :: [Double]
cdc = ([Double], i) -> [Double]
forall a b. (a, b) -> a
fst ((([Double], i), p) -> ([Double], i)
forall a b. (a, b) -> a
fst (([Double], o) -> (([Double], i), p)
backward ([Double]
zeroV, o
do_)))
dj :: [Double]
dj = Matrix Double -> [Double] -> [Double]
forall a. (Additive a, Multiplicative a) => Matrix a -> [a] -> [a]
matVec Matrix Double
aStar [Double]
cdc
(([Double]
_, i
di), p
dp) = ([Double], o) -> (([Double], i), p)
backward ([Double]
dj, o
do_)
in (i
di, p
dp)
in if [Double] -> Double
forall {a}. (Ord a, Num a) => [a] -> a
normInf ([Double] -> [Double] -> [Double]
diffV [Double]
aNext [Double]
a) Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
> Double
tol
then [Char] -> (o, o -> (i, p))
forall a. HasCallStack => [Char] -> a
error ([Char]
"traceDiffPMatrix: primal fixed point did not converge within " [Char] -> [Char] -> [Char]
forall a. [a] -> [a] -> [a]
++ Int -> [Char]
forall a. Show a => a -> [Char]
show Int
maxIter [Char] -> [Char] -> [Char]
forall a. [a] -> [a] -> [a]
++ [Char]
" iterations")
else (o
o, o -> (i, p)
pullback)