{-# LANGUAGE RankNTypes #-}
{-# LANGUAGE ScopedTypeVariables #-}

-- | @Prob@ carrier for @Mat@.
--
-- A @Mat r i j@ is a linear map from input indices @i@ to output indices @j@
-- with weights in @r@.  That is exactly the linear fragment of the double-dual
-- probability arrow 'Circuit.Prob.Prob'.  This module gives the embedding that
-- lets a @Mat@ construction flow through the @Prob@ carrier with no rewrite.
module Circuit.Mat.Prob
  ( matToProb,
    runProbAt,
  )
where

import Circuit.Mat (Finite (..), Mat (..), runMat)
import Circuit.Prob (Prob (..))
import Data.Bool (bool)
import NumHask.Algebra.Additive (Additive (..), sum, zero)
import NumHask.Algebra.Multiplicative (Multiplicative (..), one)
import Prelude hiding (id, sum, (*), (+), (.))

-- | Embed a semiring matrix into the @Prob@ continuation carrier.
--
-- The resulting @Prob (->) r i j@ is the expectation transformer:
--
-- @
--   runProb (matToProb m) k (x, i) = Σⱼ m(i,j) · k(x,j)
-- @
matToProb ::
  forall r i j.
  (Additive r, Multiplicative r, Eq j, Finite i, Finite j) =>
  Mat r i j ->
  Prob (->) r i j
matToProb :: forall r i j.
(Additive r, Multiplicative r, Eq j, Finite i, Finite j) =>
Mat r i j -> Prob (->) r i j
matToProb Mat r i j
m = (forall x. ((x, j) -> r) -> (x, i) -> r) -> Prob (->) r i j
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. ((x, j) -> r) -> (x, i) -> r) -> Prob (->) r i j)
-> (forall x. ((x, j) -> r) -> (x, i) -> r) -> Prob (->) r i j
forall a b. (a -> b) -> a -> b
$ \(x, j) -> r
k (x
x, i
i) -> [r] -> r
forall a (f :: * -> *). (Additive a, Foldable f) => f a -> a
sum [Mat r i j -> i -> j -> r
forall s j i.
(Additive s, Multiplicative s, Eq j) =>
Mat s i j -> i -> j -> s
runMat Mat r i j
m i
i j
j r -> r -> r
forall a. Multiplicative a => a -> a -> a
* (x, j) -> r
k (x
x, j
j) | j
j <- [j]
forall a. Finite a => [a]
universe]

-- | Run a @Prob@ carrier at a single input/output pair.
--
-- This is the analogue of 'runMat' for the @Prob@ carrier: it selects one
-- output index by feeding a continuation that is @one@ at that index and
-- 'zero' elsewhere.
runProbAt ::
  forall r i j.
  (Additive r, Multiplicative r, Eq j) =>
  Prob (->) r i j ->
  i ->
  j ->
  r
runProbAt :: forall r i j.
(Additive r, Multiplicative r, Eq j) =>
Prob (->) r i j -> i -> j -> r
runProbAt (Prob forall x. ((x, j) -> r) -> (x, i) -> r
f) i
i j
j =
  (((), j) -> r) -> ((), i) -> r
forall x. ((x, j) -> r) -> (x, i) -> r
f (\(()
_, j
j') -> r -> r -> Bool -> r
forall a. a -> a -> Bool -> a
bool r
forall a. Additive a => a
zero r
forall a. Multiplicative a => a
one (j
j' j -> j -> Bool
forall a. Eq a => a -> a -> Bool
== j
j)) ((), i
i)