-- | Finite-state stationary distribution via linear solve.
--
-- A 3-state Markov chain transition matrix P. The stationary distribution
-- π satisfies π·P = π with Σπ = 1.
-- Oracle: power iteration converges to exact solution within 1e-12.
module Circuit.Inference.LinearSolve
  ( transition,
    exactStationary,
    powerIteration,
  )
where

-- | 3-state transition matrix: rows sum to 1.
transition :: [[Double]]
transition :: [[Double]]
transition =
  [ [Double
0.5, Double
0.3, Double
0.2],
    [Double
0.1, Double
0.7, Double
0.2],
    [Double
0.3, Double
0.3, Double
0.4]
  ]

-- | Exact stationary distribution: solved from π·(P-I) = 0, Σπ = 1.
-- π = [0.25, 0.5, 0.25]
exactStationary :: [Double]
exactStationary :: [Double]
exactStationary = [Double
0.25, Double
0.5, Double
0.25]

-- | Power iteration converges to stationary distribution.
powerIteration :: [Double]
powerIteration :: [Double]
powerIteration =
  let step :: [Double] -> [Double]
step [Double]
v = [[Double] -> Double
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [[Double]
v [Double] -> Int -> Double
forall a. HasCallStack => [a] -> Int -> a
!! Int
j Double -> Double -> Double
forall a. Num a => a -> a -> a
* [[Double]]
transition [[Double]] -> Int -> [Double]
forall a. HasCallStack => [a] -> Int -> a
!! Int
j [Double] -> Int -> Double
forall a. HasCallStack => [a] -> Int -> a
!! Int
i | Int
j <- [Int
0 .. Int
2]] | Int
i <- [Int
0 .. Int
2]]
   in ([Double] -> [Double]) -> [Double] -> [Double]
forall {a}. (Fractional a, Ord a) => ([a] -> [a]) -> [a] -> [a]
converge [Double] -> [Double]
step [Double
1 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
3, Double
1 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
3, Double
1 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
3]
  where
    converge :: ([a] -> [a]) -> [a] -> [a]
converge [a] -> [a]
step [a]
v =
      let v' :: [a]
v' = (a -> a) -> [a] -> [a]
forall a b. (a -> b) -> [a] -> [b]
map (a -> a -> a
forall a. Fractional a => a -> a -> a
/ [a] -> a
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [a]
v) ([a] -> [a]
step [a]
v)
       in if [a] -> [a] -> a
forall {a}. (Ord a, Num a) => [a] -> [a] -> a
maxDiff [a]
v [a]
v' a -> a -> Bool
forall a. Ord a => a -> a -> Bool
< a
1e-14 then [a]
v' else ([a] -> [a]) -> [a] -> [a]
converge [a] -> [a]
step [a]
v'
    maxDiff :: [a] -> [a] -> a
maxDiff [a]
v [a]
w = [a] -> a
forall a. Ord a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Ord a) => t a -> a
maximum ((a -> a -> a) -> [a] -> [a] -> [a]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith (\a
x a
y -> a -> a
forall a. Num a => a -> a
abs (a
x a -> a -> a
forall a. Num a => a -> a -> a
- a
y)) [a]
v [a]
w)