{-# LANGUAGE RankNTypes #-}
{-# LANGUAGE RebindableSyntax #-}
module Circuit.Diff.RDC
(
rdc,
fstD,
sndD,
pairD,
forkD,
terminalD,
rdcAdditive,
rdcLinear,
rdcHomogeneous,
rdcIdentity,
rdcFst,
rdcSnd,
rdcPairing,
rdcTerminal,
rdcChain,
rdcMixedPartials,
)
where
import Circuit.Diff (Diff (..), runDiff)
import Circuit.Diff.Jet (constant, taylor)
import NumHask.Algebra.Additive (Additive (..))
import NumHask.Algebra.Field (ExpField (..), TrigField (..))
import NumHask.Algebra.Multiplicative (Multiplicative (..))
import NumHask.Data.Integral (FromInteger (..))
import NumHask.Prelude
rdc :: Diff p a b -> (a, b) -> a
rdc :: forall {k} (p :: k) a b. Diff p a b -> (a, b) -> a
rdc (Diff a -> (b, b -> a)
f) (a
a, b
db) =
let (b
_, b -> a
pullback) = a -> (b, b -> a)
f a
a
in b -> a
pullback b
db
fstD :: (Additive b) => Diff p (a, b) a
fstD :: forall {k} b (p :: k) a. Additive b => Diff p (a, b) a
fstD = ((a, b) -> (a, a -> (a, b))) -> Diff p (a, b) a
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff (((a, b) -> (a, a -> (a, b))) -> Diff p (a, b) a)
-> ((a, b) -> (a, a -> (a, b))) -> Diff p (a, b) a
forall a b. (a -> b) -> a -> b
$ \(a
a, b
_) -> (a
a, \a
da -> (a
da, b
forall a. Additive a => a
zero))
sndD :: (Additive a) => Diff p (a, b) b
sndD :: forall {k} a (p :: k) b. Additive a => Diff p (a, b) b
sndD = ((a, b) -> (b, b -> (a, b))) -> Diff p (a, b) b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff (((a, b) -> (b, b -> (a, b))) -> Diff p (a, b) b)
-> ((a, b) -> (b, b -> (a, b))) -> Diff p (a, b) b
forall a b. (a -> b) -> a -> b
$ \(a
_, b
b) -> (b
b, \b
db -> (a
forall a. Additive a => a
zero, b
db))
pairD ::
(Additive a) =>
Diff p a b ->
Diff p a c ->
Diff p a (b, c)
pairD :: forall {k} a (p :: k) b c.
Additive a =>
Diff p a b -> Diff p a c -> Diff p a (b, c)
pairD (Diff a -> (b, b -> a)
f) (Diff a -> (c, c -> a)
g) = (a -> ((b, c), (b, c) -> a)) -> Diff p a (b, c)
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((a -> ((b, c), (b, c) -> a)) -> Diff p a (b, c))
-> (a -> ((b, c), (b, c) -> a)) -> Diff p a (b, c)
forall a b. (a -> b) -> a -> b
$ \a
a ->
let (b
b, b -> a
pb) = a -> (b, b -> a)
f a
a
(c
c, c -> a
pc) = a -> (c, c -> a)
g a
a
in ( (b
b, c
c),
\(b
db, c
dc) -> b -> a
pb b
db a -> a -> a
forall a. Additive a => a -> a -> a
+ c -> a
pc c
dc
)
forkD ::
(Additive a) =>
Diff p a b ->
Diff p a (b, b)
forkD :: forall {k} a (p :: k) b.
Additive a =>
Diff p a b -> Diff p a (b, b)
forkD Diff p a b
f = Diff p a b -> Diff p a b -> Diff p a (b, b)
forall {k} a (p :: k) b c.
Additive a =>
Diff p a b -> Diff p a c -> Diff p a (b, c)
pairD Diff p a b
f Diff p a b
f
terminalD :: (Additive a) => Diff p a ()
terminalD :: forall {k} a (p :: k). Additive a => Diff p a ()
terminalD = (a -> ((), () -> a)) -> Diff p a ()
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((a -> ((), () -> a)) -> Diff p a ())
-> (a -> ((), () -> a)) -> Diff p a ()
forall a b. (a -> b) -> a -> b
$ \a
_ -> ((), \()
_ -> a
forall a. Additive a => a
zero)
near :: (Fractional a, Ord a, Absolute a, Subtractive a) => a -> a -> a -> Bool
near :: forall a.
(Fractional a, Ord a, Absolute a, Subtractive a) =>
a -> a -> a -> Bool
near a
tol a
x a
y = a -> a
forall a. Absolute a => a -> a
abs (a
x a -> a -> a
forall a. Subtractive a => a -> a -> a
- a
y) a -> a -> Bool
forall a. Ord a => a -> a -> Bool
< a
tol
rdcAdditive ::
(Additive a, Additive b, Fractional a, Ord a, Absolute a, Subtractive a) =>
a ->
Diff p a b ->
Diff p a b ->
a ->
b ->
b ->
Bool
rdcAdditive :: forall {k} a b (p :: k).
(Additive a, Additive b, Fractional a, Ord a, Absolute a,
Subtractive a) =>
a -> Diff p a b -> Diff p a b -> a -> b -> b -> Bool
rdcAdditive a
tol Diff p a b
f Diff p a b
g a
a b
db1 b
db2 =
a -> a -> a -> Bool
forall a.
(Fractional a, Ord a, Absolute a, Subtractive a) =>
a -> a -> a -> Bool
near
a
tol
(Diff p a b -> (a, b) -> a
forall {k} (p :: k) a b. Diff p a b -> (a, b) -> a
rdc (Diff p a b
f Diff p a b -> Diff p a b -> Diff p a b
forall a. Additive a => a -> a -> a
+ Diff p a b
g) (a
a, b
db1 b -> b -> b
forall a. Additive a => a -> a -> a
+ b
db2))
(Diff p a b -> (a, b) -> a
forall {k} (p :: k) a b. Diff p a b -> (a, b) -> a
rdc Diff p a b
f (a
a, b
db1 b -> b -> b
forall a. Additive a => a -> a -> a
+ b
db2) a -> a -> a
forall a. Additive a => a -> a -> a
+ Diff p a b -> (a, b) -> a
forall {k} (p :: k) a b. Diff p a b -> (a, b) -> a
rdc Diff p a b
g (a
a, b
db1 b -> b -> b
forall a. Additive a => a -> a -> a
+ b
db2))
Bool -> Bool -> Bool
&& a -> a -> a -> Bool
forall a.
(Fractional a, Ord a, Absolute a, Subtractive a) =>
a -> a -> a -> Bool
near a
tol (Diff (ZonkAny 1) a b -> (a, b) -> a
forall {k} (p :: k) a b. Diff p a b -> (a, b) -> a
rdc Diff (ZonkAny 1) a b
forall a. Additive a => a
zero (a
a, b
db1)) a
forall a. Additive a => a
zero
rdcLinear ::
(Additive a, Additive b, Fractional a, Ord a, Absolute a, Subtractive a) =>
a ->
Diff p a b ->
a ->
b ->
b ->
b ->
Bool
rdcLinear :: forall {k} a b (p :: k).
(Additive a, Additive b, Fractional a, Ord a, Absolute a,
Subtractive a) =>
a -> Diff p a b -> a -> b -> b -> b -> Bool
rdcLinear a
tol Diff p a b
f a
a b
db1 b
db2 b
db3 =
a -> a -> a -> Bool
forall a.
(Fractional a, Ord a, Absolute a, Subtractive a) =>
a -> a -> a -> Bool
near
a
tol
(Diff p a b -> (a, b) -> a
forall {k} (p :: k) a b. Diff p a b -> (a, b) -> a
rdc Diff p a b
f (a
a, b
db1 b -> b -> b
forall a. Additive a => a -> a -> a
+ b
db2 b -> b -> b
forall a. Additive a => a -> a -> a
+ b
db3))
(Diff p a b -> (a, b) -> a
forall {k} (p :: k) a b. Diff p a b -> (a, b) -> a
rdc Diff p a b
f (a
a, b
db1) a -> a -> a
forall a. Additive a => a -> a -> a
+ Diff p a b -> (a, b) -> a
forall {k} (p :: k) a b. Diff p a b -> (a, b) -> a
rdc Diff p a b
f (a
a, b
db2) a -> a -> a
forall a. Additive a => a -> a -> a
+ Diff p a b -> (a, b) -> a
forall {k} (p :: k) a b. Diff p a b -> (a, b) -> a
rdc Diff p a b
f (a
a, b
db3))
Bool -> Bool -> Bool
&& a -> a -> a -> Bool
forall a.
(Fractional a, Ord a, Absolute a, Subtractive a) =>
a -> a -> a -> Bool
near a
tol (Diff p a b -> (a, b) -> a
forall {k} (p :: k) a b. Diff p a b -> (a, b) -> a
rdc Diff p a b
f (a
a, b
forall a. Additive a => a
zero)) a
forall a. Additive a => a
zero
rdcIdentity :: (Fractional a, Ord a, Absolute a, Subtractive a) => a -> a -> a -> Bool
rdcIdentity :: forall a.
(Fractional a, Ord a, Absolute a, Subtractive a) =>
a -> a -> a -> Bool
rdcIdentity a
tol a
a a
da = a -> a -> a -> Bool
forall a.
(Fractional a, Ord a, Absolute a, Subtractive a) =>
a -> a -> a -> Bool
near a
tol (Diff (ZonkAny 3) a a -> (a, a) -> a
forall {k} (p :: k) a b. Diff p a b -> (a, b) -> a
rdc Diff (ZonkAny 3) a a
forall a. Diff (ZonkAny 3) a a
forall {k} (cat :: k -> k -> *) (a :: k). Category cat => cat a a
id (a
a, a
da)) a
da
rdcFst ::
(Additive b, Fractional a, Ord a, Absolute a, Subtractive a, Eq b) =>
a ->
(a, b) ->
a ->
Bool
rdcFst :: forall b a.
(Additive b, Fractional a, Ord a, Absolute a, Subtractive a,
Eq b) =>
a -> (a, b) -> a -> Bool
rdcFst a
tol (a, b)
x a
da =
let (a
da', b
zeroB) = Diff (ZonkAny 5) (a, b) a -> ((a, b), a) -> (a, b)
forall {k} (p :: k) a b. Diff p a b -> (a, b) -> a
rdc Diff (ZonkAny 5) (a, b) a
forall {k} b (p :: k) a. Additive b => Diff p (a, b) a
fstD ((a, b)
x, a
da)
in a -> a -> a -> Bool
forall a.
(Fractional a, Ord a, Absolute a, Subtractive a) =>
a -> a -> a -> Bool
near a
tol a
da' a
da Bool -> Bool -> Bool
&& b
zeroB b -> b -> Bool
forall a. Eq a => a -> a -> Bool
== b
forall a. Additive a => a
zero
rdcSnd ::
(Additive a, Fractional a, Ord a, Absolute a, Subtractive a, Eq b) =>
a ->
(a, b) ->
b ->
Bool
rdcSnd :: forall a b.
(Additive a, Fractional a, Ord a, Absolute a, Subtractive a,
Eq b) =>
a -> (a, b) -> b -> Bool
rdcSnd a
tol (a, b)
x b
db =
let (a
zeroA, b
db') = Diff (ZonkAny 7) (a, b) b -> ((a, b), b) -> (a, b)
forall {k} (p :: k) a b. Diff p a b -> (a, b) -> a
rdc Diff (ZonkAny 7) (a, b) b
forall {k} a (p :: k) b. Additive a => Diff p (a, b) b
sndD ((a, b)
x, b
db)
in a -> a -> a -> Bool
forall a.
(Fractional a, Ord a, Absolute a, Subtractive a) =>
a -> a -> a -> Bool
near a
tol a
zeroA a
forall a. Additive a => a
zero Bool -> Bool -> Bool
&& b
db' b -> b -> Bool
forall a. Eq a => a -> a -> Bool
== b
db
rdcPairing ::
(Additive a, Additive b, Additive c, Fractional a, Ord a, Absolute a, Subtractive a) =>
a ->
Diff p a b ->
Diff p a c ->
a ->
(b, c) ->
(b, c) ->
Bool
rdcPairing :: forall {k} a b c (p :: k).
(Additive a, Additive b, Additive c, Fractional a, Ord a,
Absolute a, Subtractive a) =>
a -> Diff p a b -> Diff p a c -> a -> (b, c) -> (b, c) -> Bool
rdcPairing a
tol Diff p a b
f Diff p a c
g a
a (b
db1, c
dc1) (b
db2, c
dc2) =
let lhs :: a
lhs = Diff p a (b, c) -> (a, (b, c)) -> a
forall {k} (p :: k) a b. Diff p a b -> (a, b) -> a
rdc (Diff p a b -> Diff p a c -> Diff p a (b, c)
forall {k} a (p :: k) b c.
Additive a =>
Diff p a b -> Diff p a c -> Diff p a (b, c)
pairD Diff p a b
f Diff p a c
g) (a
a, (b
db1 b -> b -> b
forall a. Additive a => a -> a -> a
+ b
db2, c
dc1 c -> c -> c
forall a. Additive a => a -> a -> a
+ c
dc2))
rhs :: a
rhs = Diff p a b -> (a, b) -> a
forall {k} (p :: k) a b. Diff p a b -> (a, b) -> a
rdc Diff p a b
f (a
a, b
db1 b -> b -> b
forall a. Additive a => a -> a -> a
+ b
db2) a -> a -> a
forall a. Additive a => a -> a -> a
+ Diff p a c -> (a, c) -> a
forall {k} (p :: k) a b. Diff p a b -> (a, b) -> a
rdc Diff p a c
g (a
a, c
dc1 c -> c -> c
forall a. Additive a => a -> a -> a
+ c
dc2)
in a -> a -> a -> Bool
forall a.
(Fractional a, Ord a, Absolute a, Subtractive a) =>
a -> a -> a -> Bool
near a
tol a
lhs a
rhs
rdcTerminal ::
(Additive a, Fractional a, Ord a, Absolute a, Subtractive a) =>
a ->
a ->
Bool
rdcTerminal :: forall a.
(Additive a, Fractional a, Ord a, Absolute a, Subtractive a) =>
a -> a -> Bool
rdcTerminal a
tol a
a = a -> a -> a -> Bool
forall a.
(Fractional a, Ord a, Absolute a, Subtractive a) =>
a -> a -> a -> Bool
near a
tol (Diff (ZonkAny 9) a () -> (a, ()) -> a
forall {k} (p :: k) a b. Diff p a b -> (a, b) -> a
rdc Diff (ZonkAny 9) a ()
forall {k} a (p :: k). Additive a => Diff p a ()
terminalD (a
a, ())) a
forall a. Additive a => a
zero
rdcChain ::
(Fractional a, Ord a, Absolute a, Subtractive a) =>
a ->
Diff p a b ->
Diff p b c ->
a ->
c ->
Bool
rdcChain :: forall {k} a (p :: k) b c.
(Fractional a, Ord a, Absolute a, Subtractive a) =>
a -> Diff p a b -> Diff p b c -> a -> c -> Bool
rdcChain a
tol Diff p a b
f Diff p b c
g a
a c
dc =
let lhs :: a
lhs = Diff p a c -> (a, c) -> a
forall {k} (p :: k) a b. Diff p a b -> (a, b) -> a
rdc (Diff p b c
g Diff p b c -> Diff p a b -> Diff p a c
forall b c a. Diff p b c -> Diff p a b -> Diff p a c
forall {k} (cat :: k -> k -> *) (b :: k) (c :: k) (a :: k).
Category cat =>
cat b c -> cat a b -> cat a c
. Diff p a b
f) (a
a, c
dc)
rhs :: a
rhs = Diff p a b -> (a, b) -> a
forall {k} (p :: k) a b. Diff p a b -> (a, b) -> a
rdc Diff p a b
f (a
a, Diff p b c -> (b, c) -> b
forall {k} (p :: k) a b. Diff p a b -> (a, b) -> a
rdc Diff p b c
g ((b, b -> a) -> b
forall a b. (a, b) -> a
fst (Diff p a b -> a -> (b, b -> a)
forall k (p :: k) a b. Diff p a b -> a -> (b, b -> a)
runDiff Diff p a b
f a
a), c
dc))
in a -> a -> a -> Bool
forall a.
(Fractional a, Ord a, Absolute a, Subtractive a) =>
a -> a -> a -> Bool
near a
tol a
lhs a
rhs
rdcHomogeneous ::
(Multiplicative a, Fractional a, Ord a, Absolute a, Subtractive a) =>
a ->
Diff p a a ->
a ->
a ->
a ->
Bool
rdcHomogeneous :: forall {k} a (p :: k).
(Multiplicative a, Fractional a, Ord a, Absolute a,
Subtractive a) =>
a -> Diff p a a -> a -> a -> a -> Bool
rdcHomogeneous a
tol Diff p a a
f a
a a
db a
c =
a -> a -> a -> Bool
forall a.
(Fractional a, Ord a, Absolute a, Subtractive a) =>
a -> a -> a -> Bool
near a
tol (Diff p a a -> (a, a) -> a
forall {k} (p :: k) a b. Diff p a b -> (a, b) -> a
rdc Diff p a a
f (a
a, a
c a -> a -> a
forall a. Multiplicative a => a -> a -> a
* a
db)) (a
c a -> a -> a
forall a. Multiplicative a => a -> a -> a
* Diff p a a -> (a, a) -> a
forall {k} (p :: k) a b. Diff p a b -> (a, b) -> a
rdc Diff p a a
f (a
a, a
db))
rdcMixedPartials ::
(ExpField a, TrigField a, FromInteger a, Fractional a, Ord a, Absolute a, Subtractive a) =>
a ->
(forall b. (ExpField b, TrigField b, FromInteger b) => (b, b) -> b) ->
(a, a) ->
Bool
rdcMixedPartials :: forall a.
(ExpField a, TrigField a, FromInteger a, Fractional a, Ord a,
Absolute a, Subtractive a) =>
a
-> (forall b.
(ExpField b, TrigField b, FromInteger b) =>
(b, b) -> b)
-> (a, a)
-> Bool
rdcMixedPartials a
tol forall b. (ExpField b, TrigField b, FromInteger b) => (b, b) -> b
f (a
x0, a
y0) =
let x0J :: Jet a
x0J = Int -> a -> Jet a
forall a. Additive a => Int -> a -> Jet a
constant Int
0 a
x0
y0J :: Jet a
y0J = Int -> a -> Jet a
forall a. Additive a => Int -> a -> Jet a
constant Int
0 a
y0
dx :: Jet a -> Jet a
dx Jet a
y = (Jet (Jet a) -> Jet (Jet a)) -> Int -> Jet a -> [Jet a]
forall a.
(ExpField a, FromInteger a) =>
(Jet a -> Jet a) -> Int -> a -> [a]
taylor (\Jet (Jet a)
x -> (Jet (Jet a), Jet (Jet a)) -> Jet (Jet a)
forall b. (ExpField b, TrigField b, FromInteger b) => (b, b) -> b
f (Jet (Jet a)
x, Int -> Jet a -> Jet (Jet a)
forall a. Additive a => Int -> a -> Jet a
constant Int
0 Jet a
y)) Int
1 Jet a
x0J [Jet a] -> Int -> Jet a
forall a. HasCallStack => [a] -> Int -> a
!! Int
1
dy :: Jet a -> Jet a
dy Jet a
x = (Jet (Jet a) -> Jet (Jet a)) -> Int -> Jet a -> [Jet a]
forall a.
(ExpField a, FromInteger a) =>
(Jet a -> Jet a) -> Int -> a -> [a]
taylor (\Jet (Jet a)
y -> (Jet (Jet a), Jet (Jet a)) -> Jet (Jet a)
forall b. (ExpField b, TrigField b, FromInteger b) => (b, b) -> b
f (Int -> Jet a -> Jet (Jet a)
forall a. Additive a => Int -> a -> Jet a
constant Int
0 Jet a
x, Jet (Jet a)
y)) Int
1 Jet a
y0J [Jet a] -> Int -> Jet a
forall a. HasCallStack => [a] -> Int -> a
!! Int
1
dxy :: a
dxy = (Jet a -> Jet a) -> Int -> a -> [a]
forall a.
(ExpField a, FromInteger a) =>
(Jet a -> Jet a) -> Int -> a -> [a]
taylor Jet a -> Jet a
dx Int
1 a
y0 [a] -> Int -> a
forall a. HasCallStack => [a] -> Int -> a
!! Int
1
dyx :: a
dyx = (Jet a -> Jet a) -> Int -> a -> [a]
forall a.
(ExpField a, FromInteger a) =>
(Jet a -> Jet a) -> Int -> a -> [a]
taylor Jet a -> Jet a
dy Int
1 a
x0 [a] -> Int -> a
forall a. HasCallStack => [a] -> Int -> a
!! Int
1
in a -> a -> a -> Bool
forall a.
(Fractional a, Ord a, Absolute a, Subtractive a) =>
a -> a -> a -> Bool
near a
tol a
dxy a
dyx