{-# OPTIONS_GHC -Wno-orphans #-}

-- | Effectful probability row: @Prob (K m) r@.
--
-- The function-arrow instances in "Circuit.Prob" are the reference semantics.
-- This module ports the cartesian and cocartesian structural instances to the
-- K base arrow, which is the substrate for sampling-based inference.
-- The tensor remains premonoidal: the two nestings 'parFGK' and 'parGFK' agree
-- on distribution but differ operationally when effects do not commute.
module Circuit.Inference.Prob
  ( -- * Re-export base type
    Prob (..),

    -- * K primitives
    embedK,
    fromWeightedK,
    scoreK,
    massK,
    copyPK,
    discardPK,
    choiceByK,
    orPK,

    -- * Parallel nestings (premonoidal)
    parFGK,
    parGFK,

    -- * Traced Either over effectful base arrows
    traceEK,
    traceENK,
  )
where

import Circuit.Category (Category (..), K (..))
import Circuit.Channel (Channel (..), Strength (..))
import Circuit.Prob (Prob (..))
import Prelude hiding (id, (.))

-- | Lift a pure function into an effectful probability morphism.
embedK :: (a -> b) -> Prob (K m) r a b
embedK :: forall {k} a b (m :: k -> *) (r :: k). (a -> b) -> Prob (K m) r a b
embedK a -> b
h = (forall x. K m (x, b) r -> K m (x, a) r) -> Prob (K m) r a b
forall {k} (arr :: * -> k -> *) (r :: k) a b.
(forall x. arr (x, b) r -> arr (x, a) r) -> Prob arr r a b
Prob ((forall x. K m (x, b) r -> K m (x, a) r) -> Prob (K m) r a b)
-> (forall x. K m (x, b) r -> K m (x, a) r) -> Prob (K m) r a b
forall a b. (a -> b) -> a -> b
$ \(K (x, b) -> m r
k) -> ((x, a) -> m r) -> K m (x, a) r
forall {k} (m :: k -> *) a (b :: k). (a -> m b) -> K m a b
K (((x, a) -> m r) -> K m (x, a) r)
-> ((x, a) -> m r) -> K m (x, a) r
forall a b. (a -> b) -> a -> b
$ \(x
x, a
a) -> (x, b) -> m r
k (x
x, a -> b
h a
a)

-- | Build a probability morphism from a finite weighted table.
fromWeightedK :: (Num r, Monad m) => [(b, r)] -> Prob (K m) r () b
fromWeightedK :: forall r (m :: * -> *) b.
(Num r, Monad m) =>
[(b, r)] -> Prob (K m) r () b
fromWeightedK [(b, r)]
xs = (forall x. K m (x, b) r -> K m (x, ()) r) -> Prob (K m) r () b
forall {k} (arr :: * -> k -> *) (r :: k) a b.
(forall x. arr (x, b) r -> arr (x, a) r) -> Prob arr r a b
Prob ((forall x. K m (x, b) r -> K m (x, ()) r) -> Prob (K m) r () b)
-> (forall x. K m (x, b) r -> K m (x, ()) r) -> Prob (K m) r () b
forall a b. (a -> b) -> a -> b
$ \(K (x, b) -> m r
k) -> ((x, ()) -> m r) -> K m (x, ()) r
forall {k} (m :: k -> *) a (b :: k). (a -> m b) -> K m a b
K (((x, ()) -> m r) -> K m (x, ()) r)
-> ((x, ()) -> m r) -> K m (x, ()) r
forall a b. (a -> b) -> a -> b
$ \(x
x, ()) -> do
  rs <- ((b, r) -> m r) -> [(b, r)] -> m [r]
forall (t :: * -> *) (f :: * -> *) a b.
(Traversable t, Applicative f) =>
(a -> f b) -> t a -> f (t b)
forall (f :: * -> *) a b.
Applicative f =>
(a -> f b) -> [a] -> f [b]
traverse (\(b
b, r
w) -> (r -> r) -> m r -> m r
forall a b. (a -> b) -> m a -> m b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap (r
w r -> r -> r
forall a. Num a => a -> a -> a
*) ((x, b) -> m r
k (x
x, b
b))) [(b, r)]
xs
  pure (sum rs)

-- | Scale the result of a continuation.
scoreK :: (Monad m) => (r -> r) -> Prob (K m) r a a
scoreK :: forall (m :: * -> *) r a. Monad m => (r -> r) -> Prob (K m) r a a
scoreK r -> r
scale = (forall x. K m (x, a) r -> K m (x, a) r) -> Prob (K m) r a a
forall {k} (arr :: * -> k -> *) (r :: k) a b.
(forall x. arr (x, b) r -> arr (x, a) r) -> Prob arr r a b
Prob ((forall x. K m (x, a) r -> K m (x, a) r) -> Prob (K m) r a a)
-> (forall x. K m (x, a) r -> K m (x, a) r) -> Prob (K m) r a a
forall a b. (a -> b) -> a -> b
$ \(K (x, a) -> m r
k) -> ((x, a) -> m r) -> K m (x, a) r
forall {k} (m :: k -> *) a (b :: k). (a -> m b) -> K m a b
K (((x, a) -> m r) -> K m (x, a) r)
-> ((x, a) -> m r) -> K m (x, a) r
forall a b. (a -> b) -> a -> b
$ \(x
x, a
a) -> (r -> r) -> m r -> m r
forall a b. (a -> b) -> m a -> m b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap r -> r
scale ((x, a) -> m r
k (x
x, a
a))

-- | Compute the total mass of an unnormalised morphism against the unit
-- continuation.
massK :: (Monad m) => r -> Prob (K m) r a b -> a -> m r
massK :: forall (m :: * -> *) r a b.
Monad m =>
r -> Prob (K m) r a b -> a -> m r
massK r
one (Prob forall x. K m (x, b) r -> K m (x, a) r
f) a
a = K m ((), a) r -> ((), a) -> m r
forall {k} (m :: k -> *) a (b :: k). K m a b -> a -> m b
runK (K m ((), b) r -> K m ((), a) r
forall x. K m (x, b) r -> K m (x, a) r
f ((((), b) -> m r) -> K m ((), b) r
forall {k} (m :: k -> *) a (b :: k). (a -> m b) -> K m a b
K (\((), b)
_ -> r -> m r
forall a. a -> m a
forall (f :: * -> *) a. Applicative f => a -> f a
pure r
one))) ((), a
a)

-- | Deterministic copy.
copyPK :: Prob (K m) r a (a, a)
copyPK :: forall {k} (m :: k -> *) (r :: k) a. Prob (K m) r a (a, a)
copyPK = (a -> (a, a)) -> Prob (K m) r a (a, a)
forall {k} a b (m :: k -> *) (r :: k). (a -> b) -> Prob (K m) r a b
embedK (\a
a -> (a
a, a
a))

-- | Deterministic discard.
discardPK :: Prob (K m) r a ()
discardPK :: forall {k} (m :: k -> *) (r :: k) a. Prob (K m) r a ()
discardPK = (a -> ()) -> Prob (K m) r a ()
forall {k} a b (m :: k -> *) (r :: k). (a -> b) -> Prob (K m) r a b
embedK (() -> a -> ()
forall a b. a -> b -> a
const ())

-- | Binary choice combined by a scalar operation.
choiceByK ::
  (Monad m) =>
  (r -> r -> r) ->
  Prob (K m) r a b ->
  Prob (K m) r a b ->
  Prob (K m) r a b
choiceByK :: forall (m :: * -> *) r a b.
Monad m =>
(r -> r -> r)
-> Prob (K m) r a b -> Prob (K m) r a b -> Prob (K m) r a b
choiceByK r -> r -> r
(<+>) (Prob forall x. K m (x, b) r -> K m (x, a) r
f) (Prob forall x. K m (x, b) r -> K m (x, a) r
g) = (forall x. K m (x, b) r -> K m (x, a) r) -> Prob (K m) r a b
forall {k} (arr :: * -> k -> *) (r :: k) a b.
(forall x. arr (x, b) r -> arr (x, a) r) -> Prob arr r a b
Prob ((forall x. K m (x, b) r -> K m (x, a) r) -> Prob (K m) r a b)
-> (forall x. K m (x, b) r -> K m (x, a) r) -> Prob (K m) r a b
forall a b. (a -> b) -> a -> b
$ \K m (x, b) r
k -> ((x, a) -> m r) -> K m (x, a) r
forall {k} (m :: k -> *) a (b :: k). (a -> m b) -> K m a b
K (((x, a) -> m r) -> K m (x, a) r)
-> ((x, a) -> m r) -> K m (x, a) r
forall a b. (a -> b) -> a -> b
$ \(x, a)
p -> do
  rf <- K m (x, a) r -> (x, a) -> m r
forall {k} (m :: k -> *) a (b :: k). K m a b -> a -> m b
runK (K m (x, b) r -> K m (x, a) r
forall x. K m (x, b) r -> K m (x, a) r
f K m (x, b) r
k) (x, a)
p
  rg <- runK (g k) p
  pure (rf <+> rg)

-- | Angelic choice for @r = Bool@.
orPK ::
  (Monad m) =>
  Prob (K m) Bool a b ->
  Prob (K m) Bool a b ->
  Prob (K m) Bool a b
orPK :: forall (m :: * -> *) a b.
Monad m =>
Prob (K m) Bool a b -> Prob (K m) Bool a b -> Prob (K m) Bool a b
orPK = (Bool -> Bool -> Bool)
-> Prob (K m) Bool a b
-> Prob (K m) Bool a b
-> Prob (K m) Bool a b
forall (m :: * -> *) r a b.
Monad m =>
(r -> r -> r)
-> Prob (K m) r a b -> Prob (K m) r a b -> Prob (K m) r a b
choiceByK Bool -> Bool -> Bool
(||)

-- ---------------------------------------------------------------------------
-- Cartesian structural instances
-- ---------------------------------------------------------------------------

instance (Monad m) => Channel (,) (Prob (K m) r) where
  assoc :: forall a b c. Prob (K m) r ((a, b), c) (a, (b, c))
assoc = (((a, b), c) -> (a, (b, c)))
-> Prob (K m) r ((a, b), c) (a, (b, c))
forall {k} a b (m :: k -> *) (r :: k). (a -> b) -> Prob (K m) r a b
embedK ((a, b), c) -> (a, (b, c))
forall a b c. ((a, b), c) -> (a, (b, c))
forall {k} (t :: k -> k -> k) (arr :: k -> k -> *) (a :: k)
       (b :: k) (c :: k).
Channel t arr =>
arr (t (t a b) c) (t a (t b c))
assoc
  assoc' :: forall a b c. Prob (K m) r (a, (b, c)) ((a, b), c)
assoc' = ((a, (b, c)) -> ((a, b), c))
-> Prob (K m) r (a, (b, c)) ((a, b), c)
forall {k} a b (m :: k -> *) (r :: k). (a -> b) -> Prob (K m) r a b
embedK (a, (b, c)) -> ((a, b), c)
forall a b c. (a, (b, c)) -> ((a, b), c)
forall {k} (t :: k -> k -> k) (arr :: k -> k -> *) (a :: k)
       (b :: k) (c :: k).
Channel t arr =>
arr (t a (t b c)) (t (t a b) c)
assoc'
  slide :: forall a b c. Prob (K m) r (a, (b, c)) (b, (a, c))
slide = ((a, (b, c)) -> (b, (a, c)))
-> Prob (K m) r (a, (b, c)) (b, (a, c))
forall {k} a b (m :: k -> *) (r :: k). (a -> b) -> Prob (K m) r a b
embedK (a, (b, c)) -> (b, (a, c))
forall a b c. (a, (b, c)) -> (b, (a, c))
forall {k} (t :: k -> k -> k) (arr :: k -> k -> *) (a :: k)
       (b :: k) (c :: k).
Channel t arr =>
arr (t a (t b c)) (t b (t a c))
slide

instance (Monad m) => Strength (,) (Prob (K m) r) where
  strength :: forall b c a. Prob (K m) r b c -> Prob (K m) r (a, b) (a, c)
strength (Prob forall x. K m (x, c) r -> K m (x, b) r
f) = (forall x. K m (x, (a, c)) r -> K m (x, (a, b)) r)
-> Prob (K m) r (a, b) (a, c)
forall {k} (arr :: * -> k -> *) (r :: k) a b.
(forall x. arr (x, b) r -> arr (x, a) r) -> Prob arr r a b
Prob ((forall x. K m (x, (a, c)) r -> K m (x, (a, b)) r)
 -> Prob (K m) r (a, b) (a, c))
-> (forall x. K m (x, (a, c)) r -> K m (x, (a, b)) r)
-> Prob (K m) r (a, b) (a, c)
forall a b. (a -> b) -> a -> b
$ \K m (x, (a, c)) r
k -> K m ((x, a), c) r -> K m ((x, a), b) r
forall x. K m (x, c) r -> K m (x, b) r
f (K m (x, (a, c)) r
k K m (x, (a, c)) r
-> K m ((x, a), c) (x, (a, c)) -> K m ((x, a), c) r
forall b c a. K m b c -> K m a b -> K m a c
forall k (arr :: k -> k -> *) (b :: k) (c :: k) (a :: k).
Category arr =>
arr b c -> arr a b -> arr a c
. K m ((x, a), c) (x, (a, c))
forall a b c. K m ((a, b), c) (a, (b, c))
forall {k} (t :: k -> k -> k) (arr :: k -> k -> *) (a :: k)
       (b :: k) (c :: k).
Channel t arr =>
arr (t (t a b) c) (t a (t b c))
assoc) K m ((x, a), b) r
-> K m (x, (a, b)) ((x, a), b) -> K m (x, (a, b)) r
forall b c a. K m b c -> K m a b -> K m a c
forall k (arr :: k -> k -> *) (b :: k) (c :: k) (a :: k).
Category arr =>
arr b c -> arr a b -> arr a c
. K m (x, (a, b)) ((x, a), b)
forall a b c. K m (a, (b, c)) ((a, b), c)
forall {k} (t :: k -> k -> k) (arr :: k -> k -> *) (a :: k)
       (b :: k) (c :: k).
Channel t arr =>
arr (t a (t b c)) (t (t a b) c)
assoc'

-- ---------------------------------------------------------------------------
-- Cocartesian structural instances
-- ---------------------------------------------------------------------------

instance (Monad m) => Channel Either (Prob (K m) r) where
  assoc :: forall a b c.
Prob (K m) r (Either (Either a b) c) (Either a (Either b c))
assoc = (Either (Either a b) c -> Either a (Either b c))
-> Prob (K m) r (Either (Either a b) c) (Either a (Either b c))
forall {k} a b (m :: k -> *) (r :: k). (a -> b) -> Prob (K m) r a b
embedK Either (Either a b) c -> Either a (Either b c)
forall a b c. Either (Either a b) c -> Either a (Either b c)
forall {k} (t :: k -> k -> k) (arr :: k -> k -> *) (a :: k)
       (b :: k) (c :: k).
Channel t arr =>
arr (t (t a b) c) (t a (t b c))
assoc
  assoc' :: forall a b c.
Prob (K m) r (Either a (Either b c)) (Either (Either a b) c)
assoc' = (Either a (Either b c) -> Either (Either a b) c)
-> Prob (K m) r (Either a (Either b c)) (Either (Either a b) c)
forall {k} a b (m :: k -> *) (r :: k). (a -> b) -> Prob (K m) r a b
embedK Either a (Either b c) -> Either (Either a b) c
forall a b c. Either a (Either b c) -> Either (Either a b) c
forall {k} (t :: k -> k -> k) (arr :: k -> k -> *) (a :: k)
       (b :: k) (c :: k).
Channel t arr =>
arr (t a (t b c)) (t (t a b) c)
assoc'
  slide :: forall a b c.
Prob (K m) r (Either a (Either b c)) (Either b (Either a c))
slide = (Either a (Either b c) -> Either b (Either a c))
-> Prob (K m) r (Either a (Either b c)) (Either b (Either a c))
forall {k} a b (m :: k -> *) (r :: k). (a -> b) -> Prob (K m) r a b
embedK Either a (Either b c) -> Either b (Either a c)
forall a b c. Either a (Either b c) -> Either b (Either a c)
forall {k} (t :: k -> k -> k) (arr :: k -> k -> *) (a :: k)
       (b :: k) (c :: k).
Channel t arr =>
arr (t a (t b c)) (t b (t a c))
slide

instance (Monad m) => Strength Either (Prob (K m) r) where
  strength :: forall b c a.
Prob (K m) r b c -> Prob (K m) r (Either a b) (Either a c)
strength (Prob forall x. K m (x, c) r -> K m (x, b) r
f) = (forall x. K m (x, Either a c) r -> K m (x, Either a b) r)
-> Prob (K m) r (Either a b) (Either a c)
forall {k} (arr :: * -> k -> *) (r :: k) a b.
(forall x. arr (x, b) r -> arr (x, a) r) -> Prob arr r a b
Prob ((forall x. K m (x, Either a c) r -> K m (x, Either a b) r)
 -> Prob (K m) r (Either a b) (Either a c))
-> (forall x. K m (x, Either a c) r -> K m (x, Either a b) r)
-> Prob (K m) r (Either a b) (Either a c)
forall a b. (a -> b) -> a -> b
$ \(K (x, Either a c) -> m r
k) -> ((x, Either a b) -> m r) -> K m (x, Either a b) r
forall {k} (m :: k -> *) a (b :: k). (a -> m b) -> K m a b
K (((x, Either a b) -> m r) -> K m (x, Either a b) r)
-> ((x, Either a b) -> m r) -> K m (x, Either a b) r
forall a b. (a -> b) -> a -> b
$ \(x
x, Either a b
e) -> case Either a b
e of
    Left a
a -> (x, Either a c) -> m r
k (x
x, a -> Either a c
forall a b. a -> Either a b
Left a
a)
    Right b
b -> K m (x, b) r -> (x, b) -> m r
forall {k} (m :: k -> *) a (b :: k). K m a b -> a -> m b
runK (K m (x, c) r -> K m (x, b) r
forall x. K m (x, c) r -> K m (x, b) r
f (((x, c) -> m r) -> K m (x, c) r
forall {k} (m :: k -> *) a (b :: k). (a -> m b) -> K m a b
K (\(x
x', c
c) -> (x, Either a c) -> m r
k (x
x', c -> Either a c
forall a b. b -> Either a b
Right c
c)))) (x
x, b
b)

-- ---------------------------------------------------------------------------
-- Parallel nestings (Fubini on the linear/commutative fragment)
-- ---------------------------------------------------------------------------

-- | Parallel composition: @g@ runs at context @(x, b)@, @f@ runs at context
-- @(x, c)@.  One of two lawful nestings; operationally distinct from 'parGFK'
-- when the underlying monad has ordered effects.
parFGK ::
  Prob (K m) r a b ->
  Prob (K m) r c d ->
  Prob (K m) r (a, c) (b, d)
parFGK :: forall {k} (m :: k -> *) (r :: k) a b c d.
Prob (K m) r a b -> Prob (K m) r c d -> Prob (K m) r (a, c) (b, d)
parFGK (Prob forall x. K m (x, b) r -> K m (x, a) r
f) (Prob forall x. K m (x, d) r -> K m (x, c) r
g) = (forall x. K m (x, (b, d)) r -> K m (x, (a, c)) r)
-> Prob (K m) r (a, c) (b, d)
forall {k} (arr :: * -> k -> *) (r :: k) a b.
(forall x. arr (x, b) r -> arr (x, a) r) -> Prob arr r a b
Prob ((forall x. K m (x, (b, d)) r -> K m (x, (a, c)) r)
 -> Prob (K m) r (a, c) (b, d))
-> (forall x. K m (x, (b, d)) r -> K m (x, (a, c)) r)
-> Prob (K m) r (a, c) (b, d)
forall a b. (a -> b) -> a -> b
$ \(K (x, (b, d)) -> m r
k) ->
  let kg :: K m ((x, b), d) r
kg = (((x, b), d) -> m r) -> K m ((x, b), d) r
forall {k} (m :: k -> *) a (b :: k). (a -> m b) -> K m a b
K ((((x, b), d) -> m r) -> K m ((x, b), d) r)
-> (((x, b), d) -> m r) -> K m ((x, b), d) r
forall a b. (a -> b) -> a -> b
$ \((x
ctx, b
b), d
d) -> (x, (b, d)) -> m r
k (x
ctx, (b
b, d
d))
      K ((x, b), c) -> m r
gc = K m ((x, b), d) r -> K m ((x, b), c) r
forall x. K m (x, d) r -> K m (x, c) r
g K m ((x, b), d) r
kg
      kf :: K m ((x, c), b) r
kf = (((x, c), b) -> m r) -> K m ((x, c), b) r
forall {k} (m :: k -> *) a (b :: k). (a -> m b) -> K m a b
K ((((x, c), b) -> m r) -> K m ((x, c), b) r)
-> (((x, c), b) -> m r) -> K m ((x, c), b) r
forall a b. (a -> b) -> a -> b
$ \((x
ctx, c
c), b
b) -> ((x, b), c) -> m r
gc ((x
ctx, b
b), c
c)
      K ((x, c), a) -> m r
fa = K m ((x, c), b) r -> K m ((x, c), a) r
forall x. K m (x, b) r -> K m (x, a) r
f K m ((x, c), b) r
kf
   in ((x, (a, c)) -> m r) -> K m (x, (a, c)) r
forall {k} (m :: k -> *) a (b :: k). (a -> m b) -> K m a b
K (((x, (a, c)) -> m r) -> K m (x, (a, c)) r)
-> ((x, (a, c)) -> m r) -> K m (x, (a, c)) r
forall a b. (a -> b) -> a -> b
$ \(x
ctx, (a
a, c
c)) -> ((x, c), a) -> m r
fa ((x
ctx, c
c), a
a)

-- | Parallel composition: @f@ runs at context @(x, d)@, @g@ runs at context
-- @(x, a)@.  The other nesting; agrees with 'parFGK' on distribution for the
-- linear fragment, but differs operationally on ordered effects.
parGFK ::
  Prob (K m) r a b ->
  Prob (K m) r c d ->
  Prob (K m) r (a, c) (b, d)
parGFK :: forall {k} (m :: k -> *) (r :: k) a b c d.
Prob (K m) r a b -> Prob (K m) r c d -> Prob (K m) r (a, c) (b, d)
parGFK (Prob forall x. K m (x, b) r -> K m (x, a) r
f) (Prob forall x. K m (x, d) r -> K m (x, c) r
g) = (forall x. K m (x, (b, d)) r -> K m (x, (a, c)) r)
-> Prob (K m) r (a, c) (b, d)
forall {k} (arr :: * -> k -> *) (r :: k) a b.
(forall x. arr (x, b) r -> arr (x, a) r) -> Prob arr r a b
Prob ((forall x. K m (x, (b, d)) r -> K m (x, (a, c)) r)
 -> Prob (K m) r (a, c) (b, d))
-> (forall x. K m (x, (b, d)) r -> K m (x, (a, c)) r)
-> Prob (K m) r (a, c) (b, d)
forall a b. (a -> b) -> a -> b
$ \(K (x, (b, d)) -> m r
k) ->
  let kf :: K m ((x, d), b) r
kf = (((x, d), b) -> m r) -> K m ((x, d), b) r
forall {k} (m :: k -> *) a (b :: k). (a -> m b) -> K m a b
K ((((x, d), b) -> m r) -> K m ((x, d), b) r)
-> (((x, d), b) -> m r) -> K m ((x, d), b) r
forall a b. (a -> b) -> a -> b
$ \((x
ctx, d
d), b
b) -> (x, (b, d)) -> m r
k (x
ctx, (b
b, d
d))
      K ((x, d), a) -> m r
fa = K m ((x, d), b) r -> K m ((x, d), a) r
forall x. K m (x, b) r -> K m (x, a) r
f K m ((x, d), b) r
kf
      kg :: K m ((x, a), d) r
kg = (((x, a), d) -> m r) -> K m ((x, a), d) r
forall {k} (m :: k -> *) a (b :: k). (a -> m b) -> K m a b
K ((((x, a), d) -> m r) -> K m ((x, a), d) r)
-> (((x, a), d) -> m r) -> K m ((x, a), d) r
forall a b. (a -> b) -> a -> b
$ \((x
ctx, a
a), d
d) -> ((x, d), a) -> m r
fa ((x
ctx, d
d), a
a)
      K ((x, a), c) -> m r
gb = K m ((x, a), d) r -> K m ((x, a), c) r
forall x. K m (x, d) r -> K m (x, c) r
g K m ((x, a), d) r
kg
   in ((x, (a, c)) -> m r) -> K m (x, (a, c)) r
forall {k} (m :: k -> *) a (b :: k). (a -> m b) -> K m a b
K (((x, (a, c)) -> m r) -> K m (x, (a, c)) r)
-> ((x, (a, c)) -> m r) -> K m (x, (a, c)) r
forall a b. (a -> b) -> a -> b
$ \(x
ctx, (a
a, c
c)) -> ((x, a), c) -> m r
gb ((x
ctx, a
a), c
c)

-- ---------------------------------------------------------------------------
-- Traced Either (productive under effectful base arrows)
-- ---------------------------------------------------------------------------

-- | Least-fixpoint trace over the 'Either' tensor for @K m@.
--
-- Recursive re-entries are productive because each recursive call is an
-- effectful action in @m@.  For almost-surely terminating bodies (e.g. a
-- geometric sampler) the trace returns a sample rather than diverging.
traceEK ::
  Prob (K m) r (Either a s) (Either b s) ->
  Prob (K m) r a b
traceEK :: forall {k} (m :: k -> *) (r :: k) a s b.
Prob (K m) r (Either a s) (Either b s) -> Prob (K m) r a b
traceEK (Prob forall x. K m (x, Either b s) r -> K m (x, Either a s) r
f) = (forall x. K m (x, b) r -> K m (x, a) r) -> Prob (K m) r a b
forall {k} (arr :: * -> k -> *) (r :: k) a b.
(forall x. arr (x, b) r -> arr (x, a) r) -> Prob arr r a b
Prob ((forall x. K m (x, b) r -> K m (x, a) r) -> Prob (K m) r a b)
-> (forall x. K m (x, b) r -> K m (x, a) r) -> Prob (K m) r a b
forall a b. (a -> b) -> a -> b
$ \(K (x, b) -> m r
k) ->
  let step :: (x, Either b s) -> m r
step (x
x, Left b
b) = (x, b) -> m r
k (x
x, b
b)
      step (x
x, Right s
s) = K m (x, Either a s) r -> (x, Either a s) -> m r
forall {k} (m :: k -> *) a (b :: k). K m a b -> a -> m b
runK (K m (x, Either b s) r -> K m (x, Either a s) r
forall x. K m (x, Either b s) r -> K m (x, Either a s) r
f (((x, Either b s) -> m r) -> K m (x, Either b s) r
forall {k} (m :: k -> *) a (b :: k). (a -> m b) -> K m a b
K (x, Either b s) -> m r
step)) (x
x, s -> Either a s
forall a b. b -> Either a b
Right s
s)
   in ((x, a) -> m r) -> K m (x, a) r
forall {k} (m :: k -> *) a (b :: k). (a -> m b) -> K m a b
K (((x, a) -> m r) -> K m (x, a) r)
-> ((x, a) -> m r) -> K m (x, a) r
forall a b. (a -> b) -> a -> b
$ \(x
x, a
a) -> K m (x, Either a s) r -> (x, Either a s) -> m r
forall {k} (m :: k -> *) a (b :: k). K m a b -> a -> m b
runK (K m (x, Either b s) r -> K m (x, Either a s) r
forall x. K m (x, Either b s) r -> K m (x, Either a s) r
f (((x, Either b s) -> m r) -> K m (x, Either b s) r
forall {k} (m :: k -> *) a (b :: k). (a -> m b) -> K m a b
K (x, Either b s) -> m r
step)) (x
x, a -> Either a s
forall a b. a -> Either a b
Left a
a)

-- | Fuel-bounded variant of 'traceEK'.
traceENK ::
  (Monad m) =>
  r ->
  Int ->
  Prob (K m) r (Either a s) (Either b s) ->
  Prob (K m) r a b
traceENK :: forall (m :: * -> *) r a s b.
Monad m =>
r
-> Int
-> Prob (K m) r (Either a s) (Either b s)
-> Prob (K m) r a b
traceENK r
zero Int
n0 (Prob forall x. K m (x, Either b s) r -> K m (x, Either a s) r
f) = (forall x. K m (x, b) r -> K m (x, a) r) -> Prob (K m) r a b
forall {k} (arr :: * -> k -> *) (r :: k) a b.
(forall x. arr (x, b) r -> arr (x, a) r) -> Prob arr r a b
Prob ((forall x. K m (x, b) r -> K m (x, a) r) -> Prob (K m) r a b)
-> (forall x. K m (x, b) r -> K m (x, a) r) -> Prob (K m) r a b
forall a b. (a -> b) -> a -> b
$ \(K (x, b) -> m r
k) ->
  let step :: Int -> (x, Either b s) -> m r
step Int
_ (x
x, Left b
b) = (x, b) -> m r
k (x
x, b
b)
      step Int
n (x
x, Right s
s)
        | Int
n Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
<= Int
0 = r -> m r
forall a. a -> m a
forall (f :: * -> *) a. Applicative f => a -> f a
pure r
zero
        | Bool
otherwise = K m (x, Either a s) r -> (x, Either a s) -> m r
forall {k} (m :: k -> *) a (b :: k). K m a b -> a -> m b
runK (K m (x, Either b s) r -> K m (x, Either a s) r
forall x. K m (x, Either b s) r -> K m (x, Either a s) r
f (((x, Either b s) -> m r) -> K m (x, Either b s) r
forall {k} (m :: k -> *) a (b :: k). (a -> m b) -> K m a b
K (Int -> (x, Either b s) -> m r
step (Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1)))) (x
x, s -> Either a s
forall a b. b -> Either a b
Right s
s)
   in ((x, a) -> m r) -> K m (x, a) r
forall {k} (m :: k -> *) a (b :: k). (a -> m b) -> K m a b
K (((x, a) -> m r) -> K m (x, a) r)
-> ((x, a) -> m r) -> K m (x, a) r
forall a b. (a -> b) -> a -> b
$ \(x
x, a
a) -> K m (x, Either a s) r -> (x, Either a s) -> m r
forall {k} (m :: k -> *) a (b :: k). K m a b -> a -> m b
runK (K m (x, Either b s) r -> K m (x, Either a s) r
forall x. K m (x, Either b s) r -> K m (x, Either a s) r
f (((x, Either b s) -> m r) -> K m (x, Either b s) r
forall {k} (m :: k -> *) a (b :: k). (a -> m b) -> K m a b
K (Int -> (x, Either b s) -> m r
step Int
n0))) (x
x, a -> Either a s
forall a b. a -> Either a b
Left a
a)