{-# LANGUAGE RankNTypes #-}
{-# LANGUAGE ScopedTypeVariables #-}
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, (*), (+), (.))
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]
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)