{-# LANGUAGE CPP #-}
{-# LANGUAGE TupleSections #-}
{-# LANGUAGE UndecidableInstances #-}
module Circuit.Pullback
(
Pullback (..),
evalPullback,
)
where
import Circuit.Bimonoid (Copy (..), Discard (..), Merge (..), Zero (..))
import Circuit.Category (Category (..))
import Circuit.Channel (Channel (..), Strength (..), Traced (..))
import Circuit.Layer (run)
import Circuit.Net (Net)
import Circuit.Tensor (Action (..), Tensor (..), Unital (..))
import Data.Bifunctor
import Prelude hiding (id, (.))
newtype Pullback b a = Pullback
{
forall b a. Pullback b a -> b -> a
runPullback :: b -> a
}
instance Category Pullback where
id :: forall a. Pullback a a
id = (a -> a) -> Pullback a a
forall b a. (b -> a) -> Pullback b a
Pullback a -> a
forall a. a -> a
forall k (arr :: k -> k -> *) (a :: k). Category arr => arr a a
id
Pullback b -> c
g . :: forall b c a. Pullback b c -> Pullback a b -> Pullback a c
. Pullback a -> b
f = (a -> c) -> Pullback a c
forall b a. (b -> a) -> Pullback b a
Pullback (\a
x -> b -> c
g (a -> b
f a
x))
{-# INLINE id #-}
{-# INLINE (.) #-}
instance Unital (,) Pullback where
unitl :: forall a. Pullback (Unit (,), a) a
unitl = (((), a) -> a) -> Pullback ((), a) a
forall b a. (b -> a) -> Pullback b a
Pullback ((), a) -> a
forall a b. (a, b) -> b
snd
{-# INLINE unitl #-}
unitl' :: forall a. Pullback a (Unit (,), a)
unitl' = (a -> ((), a)) -> Pullback a ((), a)
forall b a. (b -> a) -> Pullback b a
Pullback ((),)
{-# INLINE unitl' #-}
unitr :: forall a. Pullback (a, Unit (,)) a
unitr = ((a, ()) -> a) -> Pullback (a, ()) a
forall b a. (b -> a) -> Pullback b a
Pullback (a, ()) -> a
forall a b. (a, b) -> a
fst
{-# INLINE unitr #-}
unitr' :: forall a. Pullback a (a, Unit (,))
unitr' = (a -> (a, ())) -> Pullback a (a, ())
forall b a. (b -> a) -> Pullback b a
Pullback (,())
{-# INLINE unitr' #-}
instance Tensor (,) Pullback where
tensor :: forall a b c d.
Pullback a b -> Pullback c d -> Pullback (a, c) (b, d)
tensor (Pullback a -> b
f) (Pullback c -> d
g) = ((a, c) -> (b, d)) -> Pullback (a, c) (b, d)
forall b a. (b -> a) -> Pullback b a
Pullback ((a -> b) -> (c -> d) -> (a, c) -> (b, d)
forall a b c d. (a -> b) -> (c -> d) -> (a, c) -> (b, d)
forall (p :: * -> * -> *) a b c d.
Bifunctor p =>
(a -> b) -> (c -> d) -> p a c -> p b d
Data.Bifunctor.bimap a -> b
f c -> d
g)
{-# INLINE tensor #-}
instance Action (,) Pullback where
braid :: forall a b. Pullback (a, b) (b, a)
braid = ((a, b) -> (b, a)) -> Pullback (a, b) (b, a)
forall b a. (b -> a) -> Pullback b a
Pullback (\(a
b, b
a) -> (b
a, a
b))
{-# INLINE braid #-}
instance Strength (,) Pullback where
strength :: forall b c a. Pullback b c -> Pullback (a, b) (a, c)
strength (Pullback b -> c
f) = ((a, b) -> (a, c)) -> Pullback (a, b) (a, c)
forall b a. (b -> a) -> Pullback b a
Pullback (\(a
a, b
b) -> (a
a, b -> c
f b
b))
{-# INLINE strength #-}
instance Traced (,) Pullback where
trace :: forall a b c. Pullback (a, b) (a, c) -> Pullback b c
trace (Pullback (a, b) -> (a, c)
f) = (b -> c) -> Pullback b c
forall b a. (b -> a) -> Pullback b a
Pullback ((b -> c) -> Pullback b c) -> (b -> c) -> Pullback b c
forall a b. (a -> b) -> a -> b
$ \b
dc ->
let ~(a
dx, c
db) = (a, b) -> (a, c)
f (a
dx, b
dc)
in c
db
{-# INLINE trace #-}
instance Copy Pullback a where
copy :: Pullback a (a, a)
copy = (a -> (a, a)) -> Pullback a (a, a)
forall b a. (b -> a) -> Pullback b a
Pullback (\a
b -> (a
b, a
b))
{-# INLINE copy #-}
instance Discard Pullback a where
discard :: Pullback a ()
discard = (a -> ()) -> Pullback a ()
forall b a. (b -> a) -> Pullback b a
Pullback (() -> a -> ()
forall a b. a -> b -> a
const ())
{-# INLINE discard #-}
instance (Merge (->) a) => Merge Pullback a where
plus :: Pullback (a, a) a
plus = ((a, a) -> a) -> Pullback (a, a) a
forall b a. (b -> a) -> Pullback b a
Pullback (\(a
b1, a
b2) -> (a, a) -> a
forall (arr :: * -> * -> *) a. Merge arr a => arr (a, a) a
plus (a
b1, a
b2))
{-# INLINE plus #-}
instance (Zero (->) a) => Zero Pullback a where
zero :: Pullback () a
zero = (() -> a) -> Pullback () a
forall b a. (b -> a) -> Pullback b a
Pullback (\() -> () -> a
forall {k} (arr :: * -> k -> *) (a :: k). Zero arr a => arr () a
zero ())
{-# INLINE zero #-}
evalPullback :: Net (,) Pullback b a -> b -> a
evalPullback :: forall b a. Net (,) Pullback b a -> b -> a
evalPullback Net (,) Pullback b a
n = Pullback b a -> b -> a
forall b a. Pullback b a -> b -> a
runPullback (Net (,) Pullback b a -> Pullback b a
forall (arr :: * -> * -> *) a b.
(Run
(Syntax
(SigCompose
:+: (SigPar (,)
:+: (SigSwap (,)
:+: (SigCopy (,)
:+: (SigDiscard (,) :+: (SigPlus (,) :+: SigZero (,))))))))
arr,
Law
(Syntax
(SigCompose
:+: (SigPar (,)
:+: (SigSwap (,)
:+: (SigCopy (,)
:+: (SigDiscard (,) :+: (SigPlus (,) :+: SigZero (,))))))))
arr,
Bind
(Syntax
(SigCompose
:+: (SigPar (,)
:+: (SigSwap (,)
:+: (SigCopy (,)
:+: (SigDiscard (,) :+: (SigPlus (,) :+: SigZero (,))))))))
arr) =>
Syntax
(SigCompose
:+: (SigPar (,)
:+: (SigSwap (,)
:+: (SigCopy (,)
:+: (SigDiscard (,) :+: (SigPlus (,) :+: SigZero (,)))))))
arr
a
b
-> arr a b
forall (f :: (* -> * -> *) -> * -> * -> *) (arr :: * -> * -> *) a
b.
(Layer f, Run f arr, Law f arr, Bind f arr) =>
f arr a b -> arr a b
run Net (,) Pullback b a
n)
{-# INLINE evalPullback #-}
instance Channel (,) Pullback where
assoc :: forall a b c. Pullback ((a, b), c) (a, (b, c))
assoc = (((a, b), c) -> (a, (b, c))) -> Pullback ((a, b), c) (a, (b, c))
forall b a. (b -> a) -> Pullback b a
Pullback (\((a
s, b
s'), c
x) -> (a
s, (b
s', c
x)))
assoc' :: forall a b c. Pullback (a, (b, c)) ((a, b), c)
assoc' = ((a, (b, c)) -> ((a, b), c)) -> Pullback (a, (b, c)) ((a, b), c)
forall b a. (b -> a) -> Pullback b a
Pullback (\(a
s, (b
s', c
x)) -> ((a
s, b
s'), c
x))
slide :: forall a b c. Pullback (a, (b, c)) (b, (a, c))
slide = ((a, (b, c)) -> (b, (a, c))) -> Pullback (a, (b, c)) (b, (a, c))
forall b a. (b -> a) -> Pullback b a
Pullback (\(a
s, (b
s', c
x)) -> (b
s', (a
s, c
x)))