-- | Hamiltonian Monte Carlo on a standard Gaussian target.
--
-- HMC with leapfrog integration and Metropolis-Hastings accept/reject.
-- The oracle verifies empirical moments match N(0,1) within tolerance.
--
-- This module also exposes the leapfrog integrator as a 'System' and 'Process'
-- over the circuits polynomial interface, together with exact oracles for
-- reversibility and volume preservation on the Gaussian target.
module Circuit.Inference.HMC
  ( -- * Sampling
    hmcSamples,
    momentTest,

    -- * Leapfrog dynamics
    leapfrogStep,
    leapfrog,
    negateMomentum,
    reverseLeapfrog,

    -- * Circuit integrations
    leapfrogSystem,
    leapfrogProcess,

    -- * Yoshida-4 dynamics
    yoshida4Step,
    yoshida4,
    yoshida4System,
    yoshida4Process,

    -- * Exact oracles
    leapfrogReversible,
    leapfrogJacobianDet,
    yoshida4Reversible,
    yoshida4JacobianDet,

    -- * Order oracles
    leapfrogOrderSlope,
    yoshida4OrderSlope,
  )
where

import Circuit (Mono, Process, System, system)
import Circuit.Process (systemToProcess)
import Data.Void (absurd)
import System.Random (randomRIO)

-- | N(0,1) log-density.
logPTarget :: Double -> Double
logPTarget :: Double -> Double
logPTarget Double
x = -((Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
x) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
2)

-- | Gradient of log-density.
gradLogP :: Double -> Double
gradLogP :: Double -> Double
gradLogP Double
x = -Double
x

-- | One leapfrog step on the Gaussian target.
leapfrogStep :: (Double, Double) -> Double -> (Double, Double)
leapfrogStep :: (Double, Double) -> Double -> (Double, Double)
leapfrogStep (Double
x, Double
p) Double
eps =
  let pHalf :: Double
pHalf = Double
p Double -> Double -> Double
forall a. Num a => a -> a -> a
+ (Double
eps Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
2) Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double -> Double
gradLogP Double
x
      xNext :: Double
xNext = Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
eps Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
pHalf
      pNext :: Double
pNext = Double
pHalf Double -> Double -> Double
forall a. Num a => a -> a -> a
+ (Double
eps Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
2) Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double -> Double
gradLogP Double
xNext
   in (Double
xNext, Double
pNext)

-- | L leapfrog steps.
leapfrog :: Double -> Int -> (Double, Double) -> (Double, Double)
leapfrog :: Double -> Int -> (Double, Double) -> (Double, Double)
leapfrog Double
eps Int
n (Double, Double)
state = ((Double, Double) -> (Double, Double))
-> (Double, Double) -> [(Double, Double)]
forall a. (a -> a) -> a -> [a]
iterate ((Double, Double) -> Double -> (Double, Double)
`leapfrogStep` Double
eps) (Double, Double)
state [(Double, Double)] -> Int -> (Double, Double)
forall a. HasCallStack => [a] -> Int -> a
!! Int
n

-- | Negate the momentum component of a phase-space point.
negateMomentum :: (Double, Double) -> (Double, Double)
negateMomentum :: (Double, Double) -> (Double, Double)
negateMomentum (Double
x, Double
p) = (Double
x, -Double
p)

-- | Run the leapfrog integrator backwards in time.
--
-- Time reversal for symplectic integrators is achieved by negating the momentum,
-- integrating forward, then negating the momentum again.
reverseLeapfrog :: Double -> Int -> (Double, Double) -> (Double, Double)
reverseLeapfrog :: Double -> Int -> (Double, Double) -> (Double, Double)
reverseLeapfrog Double
eps Int
n = (Double, Double) -> (Double, Double)
negateMomentum ((Double, Double) -> (Double, Double))
-> ((Double, Double) -> (Double, Double))
-> (Double, Double)
-> (Double, Double)
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Double -> Int -> (Double, Double) -> (Double, Double)
leapfrog Double
eps Int
n ((Double, Double) -> (Double, Double))
-> ((Double, Double) -> (Double, Double))
-> (Double, Double)
-> (Double, Double)
forall b c a. (b -> c) -> (a -> b) -> a -> c
. (Double, Double) -> (Double, Double)
negateMomentum

-- | The leapfrog integrator as a cartesian 'System' with phase-space state.
--
-- The state is the current @(position, momentum)@ pair.  The monomial direction
-- is ignored because the Gaussian-target dynamics are autonomous; the output
-- position is the next phase-space point.
leapfrogSystem :: Double -> System (->) (Double, Double) (Mono (Double, Double) (Double, Double))
leapfrogSystem :: Double
-> System
     (->) (Double, Double) (Mono (Double, Double) (Double, Double))
leapfrogSystem Double
eps = (((Double, Double), Dir (Mono (Double, Double) (Double, Double)))
 -> ((Double, Double),
     Pos (Mono (Double, Double) (Double, Double))))
-> System
     (->) (Double, Double) (Mono (Double, Double) (Double, Double))
forall (arr :: * -> * -> *) s (p :: Poly).
arr (s, Dir p) (s, Pos p) -> System arr s p
system ((((Double, Double), Dir (Mono (Double, Double) (Double, Double)))
  -> ((Double, Double),
      Pos (Mono (Double, Double) (Double, Double))))
 -> System
      (->) (Double, Double) (Mono (Double, Double) (Double, Double)))
-> (((Double, Double),
     Dir (Mono (Double, Double) (Double, Double)))
    -> ((Double, Double),
        Pos (Mono (Double, Double) (Double, Double))))
-> System
     (->) (Double, Double) (Mono (Double, Double) (Double, Double))
forall a b. (a -> b) -> a -> b
$ \case
  ((Double, Double)
_, Left Void
v) -> Void -> ((Double, Double), ((Double, Double), ()))
forall a. Void -> a
absurd Void
v
  ((Double, Double)
s, Right (Double, Double)
_) ->
    let s' :: (Double, Double)
s' = (Double, Double) -> Double -> (Double, Double)
leapfrogStep (Double, Double)
s Double
eps
     in ((Double, Double)
s', ((Double, Double)
s', ()))

-- | The leapfrog integrator as a first-input-seeded 'Process'.
--
-- The seed is @(0, 1)@; the observation returns the current state as the
-- position, and the step uses 'leapfrogSystem'.
leapfrogProcess :: Double -> Process (Double, Double) (Double, Double)
leapfrogProcess :: Double -> Process (Double, Double) (Double, Double)
leapfrogProcess Double
eps = (Double, Double)
-> ((Double, Double) -> (Double, Double))
-> System
     (->) (Double, Double) (Mono (Double, Double) (Double, Double))
-> Process (Double, Double) (Double, Double)
forall s b a.
s -> (s -> b) -> System (->) s (Mono a b) -> Process a b
systemToProcess (Double
0.0, Double
1.0) (Double, Double) -> (Double, Double)
forall a. a -> a
id (Double
-> System
     (->) (Double, Double) (Mono (Double, Double) (Double, Double))
leapfrogSystem Double
eps)

-- ---------------------------------------------------------------------------
-- Yoshida-4 composition
-- ---------------------------------------------------------------------------

-- | Classic Yoshida coefficients for 4th-order symplectic composition.
--
-- One Yoshida-4 macro-step is the symmetric composition
-- @S(c1 eps) ∘ S(c2 eps) ∘ S(c3 eps)@ of three leapfrog steps.
yoshidaC1 :: Double
yoshidaC1 :: Double
yoshidaC1 = Double
1.0 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ (Double
2.0 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
2.0 Double -> Double -> Double
forall a. Floating a => a -> a -> a
** (Double
1.0 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
3.0))

yoshidaC2 :: Double
yoshidaC2 :: Double
yoshidaC2 = -((Double
2.0 Double -> Double -> Double
forall a. Floating a => a -> a -> a
** (Double
1.0 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
3.0)) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ (Double
2.0 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
2.0 Double -> Double -> Double
forall a. Floating a => a -> a -> a
** (Double
1.0 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
3.0)))

yoshidaC3 :: Double
yoshidaC3 :: Double
yoshidaC3 = Double
yoshidaC1

-- | One Yoshida-4 macro-step on the Gaussian target.
yoshida4Step :: (Double, Double) -> Double -> (Double, Double)
yoshida4Step :: (Double, Double) -> Double -> (Double, Double)
yoshida4Step (Double, Double)
state Double
eps =
  let s1 :: (Double, Double)
s1 = (Double, Double) -> Double -> (Double, Double)
leapfrogStep (Double, Double)
state (Double
yoshidaC1 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
eps)
      s2 :: (Double, Double)
s2 = (Double, Double) -> Double -> (Double, Double)
leapfrogStep (Double, Double)
s1 (Double
yoshidaC2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
eps)
      s3 :: (Double, Double)
s3 = (Double, Double) -> Double -> (Double, Double)
leapfrogStep (Double, Double)
s2 (Double
yoshidaC3 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
eps)
   in (Double, Double)
s3

-- | N Yoshida-4 macro-steps.
yoshida4 :: Double -> Int -> (Double, Double) -> (Double, Double)
yoshida4 :: Double -> Int -> (Double, Double) -> (Double, Double)
yoshida4 Double
eps Int
n (Double, Double)
state = ((Double, Double) -> (Double, Double))
-> (Double, Double) -> [(Double, Double)]
forall a. (a -> a) -> a -> [a]
iterate ((Double, Double) -> Double -> (Double, Double)
`yoshida4Step` Double
eps) (Double, Double)
state [(Double, Double)] -> Int -> (Double, Double)
forall a. HasCallStack => [a] -> Int -> a
!! Int
n

-- | Run the Yoshida-4 integrator backwards in time.
reverseYoshida4 :: Double -> Int -> (Double, Double) -> (Double, Double)
reverseYoshida4 :: Double -> Int -> (Double, Double) -> (Double, Double)
reverseYoshida4 Double
eps Int
n = (Double, Double) -> (Double, Double)
negateMomentum ((Double, Double) -> (Double, Double))
-> ((Double, Double) -> (Double, Double))
-> (Double, Double)
-> (Double, Double)
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Double -> Int -> (Double, Double) -> (Double, Double)
yoshida4 Double
eps Int
n ((Double, Double) -> (Double, Double))
-> ((Double, Double) -> (Double, Double))
-> (Double, Double)
-> (Double, Double)
forall b c a. (b -> c) -> (a -> b) -> a -> c
. (Double, Double) -> (Double, Double)
negateMomentum

-- | The Yoshida-4 integrator as a cartesian 'System'.
yoshida4System :: Double -> System (->) (Double, Double) (Mono (Double, Double) (Double, Double))
yoshida4System :: Double
-> System
     (->) (Double, Double) (Mono (Double, Double) (Double, Double))
yoshida4System Double
eps = (((Double, Double), Dir (Mono (Double, Double) (Double, Double)))
 -> ((Double, Double),
     Pos (Mono (Double, Double) (Double, Double))))
-> System
     (->) (Double, Double) (Mono (Double, Double) (Double, Double))
forall (arr :: * -> * -> *) s (p :: Poly).
arr (s, Dir p) (s, Pos p) -> System arr s p
system ((((Double, Double), Dir (Mono (Double, Double) (Double, Double)))
  -> ((Double, Double),
      Pos (Mono (Double, Double) (Double, Double))))
 -> System
      (->) (Double, Double) (Mono (Double, Double) (Double, Double)))
-> (((Double, Double),
     Dir (Mono (Double, Double) (Double, Double)))
    -> ((Double, Double),
        Pos (Mono (Double, Double) (Double, Double))))
-> System
     (->) (Double, Double) (Mono (Double, Double) (Double, Double))
forall a b. (a -> b) -> a -> b
$ \case
  ((Double, Double)
_, Left Void
v) -> Void -> ((Double, Double), ((Double, Double), ()))
forall a. Void -> a
absurd Void
v
  ((Double, Double)
s, Right (Double, Double)
_) ->
    let s' :: (Double, Double)
s' = (Double, Double) -> Double -> (Double, Double)
yoshida4Step (Double, Double)
s Double
eps
     in ((Double, Double)
s', ((Double, Double)
s', ()))

-- | The Yoshida-4 integrator as a first-input-seeded 'Process'.
yoshida4Process :: Double -> Process (Double, Double) (Double, Double)
yoshida4Process :: Double -> Process (Double, Double) (Double, Double)
yoshida4Process Double
eps = (Double, Double)
-> ((Double, Double) -> (Double, Double))
-> System
     (->) (Double, Double) (Mono (Double, Double) (Double, Double))
-> Process (Double, Double) (Double, Double)
forall s b a.
s -> (s -> b) -> System (->) s (Mono a b) -> Process a b
systemToProcess (Double
0.0, Double
1.0) (Double, Double) -> (Double, Double)
forall a. a -> a
id (Double
-> System
     (->) (Double, Double) (Mono (Double, Double) (Double, Double))
yoshida4System Double
eps)

-- | Exact oracle: Yoshida-4 integration is reversible up to machine epsilon.
yoshida4Reversible :: Double -> Int -> (Double, Double) -> Bool
yoshida4Reversible :: Double -> Int -> (Double, Double) -> Bool
yoshida4Reversible Double
eps Int
n (Double, Double)
state =
  let forward :: (Double, Double)
forward = Double -> Int -> (Double, Double) -> (Double, Double)
yoshida4 Double
eps Int
n (Double, Double)
state
      recovered :: (Double, Double)
recovered = Double -> Int -> (Double, Double) -> (Double, Double)
reverseYoshida4 Double
eps Int
n (Double, Double)
forward
      tol :: Double
tol = Double
1e-12
      near :: Double -> Double -> Bool
near Double
u Double
v = Double -> Double
forall a. Num a => a -> a
abs (Double
u Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
v) Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
<= Double
tol
   in Double -> Double -> Bool
near ((Double, Double) -> Double
forall a b. (a, b) -> a
fst (Double, Double)
state) ((Double, Double) -> Double
forall a b. (a, b) -> a
fst (Double, Double)
recovered) Bool -> Bool -> Bool
&& Double -> Double -> Bool
near ((Double, Double) -> Double
forall a b. (a, b) -> b
snd (Double, Double)
state) ((Double, Double) -> Double
forall a b. (a, b) -> b
snd (Double, Double)
recovered)

-- | Exact oracle: the Yoshida-4 map preserves phase-space volume.
--
-- The map is linear for the Gaussian target, and a composition of symplectic
-- maps has determinant 1.
yoshida4JacobianDet :: Double -> Int -> Double
yoshida4JacobianDet :: Double -> Int -> Double
yoshida4JacobianDet Double
eps Int
n =
  let f :: (Double, Double) -> (Double, Double)
f = Double -> Int -> (Double, Double) -> (Double, Double)
yoshida4 Double
eps Int
n
      (Double
x1, Double
p1) = (Double, Double) -> (Double, Double)
f (Double
1.0, Double
0.0)
      (Double
x2, Double
p2) = (Double, Double) -> (Double, Double)
f (Double
0.0, Double
1.0)
   in Double
x1 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
p2 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
p1 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
x2

-- | Exact oracle: leapfrog integration is reversible up to machine epsilon.
--
-- Forward integration followed by momentum flip, backward integration, and
-- another momentum flip returns to the initial phase-space point.
leapfrogReversible :: Double -> Int -> (Double, Double) -> Bool
leapfrogReversible :: Double -> Int -> (Double, Double) -> Bool
leapfrogReversible Double
eps Int
n (Double, Double)
state =
  let forward :: (Double, Double)
forward = Double -> Int -> (Double, Double) -> (Double, Double)
leapfrog Double
eps Int
n (Double, Double)
state
      recovered :: (Double, Double)
recovered = Double -> Int -> (Double, Double) -> (Double, Double)
reverseLeapfrog Double
eps Int
n (Double, Double)
forward
      tol :: Double
tol = Double
1e-12
      near :: Double -> Double -> Bool
near Double
u Double
v = Double -> Double
forall a. Num a => a -> a
abs (Double
u Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
v) Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
<= Double
tol
   in Double -> Double -> Bool
near ((Double, Double) -> Double
forall a b. (a, b) -> a
fst (Double, Double)
state) ((Double, Double) -> Double
forall a b. (a, b) -> a
fst (Double, Double)
recovered) Bool -> Bool -> Bool
&& Double -> Double -> Bool
near ((Double, Double) -> Double
forall a b. (a, b) -> b
snd (Double, Double)
state) ((Double, Double) -> Double
forall a b. (a, b) -> b
snd (Double, Double)
recovered)

-- | Exact oracle: the leapfrog map preserves phase-space volume.
--
-- For the Gaussian target the leapfrog step is linear, so the Jacobian is
-- constant.  We estimate it by applying the n-step map to the standard basis
-- vectors and returning the determinant, which is exactly 1.0 for a symplectic
-- integrator.
leapfrogJacobianDet :: Double -> Int -> Double
leapfrogJacobianDet :: Double -> Int -> Double
leapfrogJacobianDet Double
eps Int
n =
  let f :: (Double, Double) -> (Double, Double)
f = Double -> Int -> (Double, Double) -> (Double, Double)
leapfrog Double
eps Int
n
      (Double
x1, Double
p1) = (Double, Double) -> (Double, Double)
f (Double
1.0, Double
0.0)
      (Double
x2, Double
p2) = (Double, Double) -> (Double, Double)
f (Double
0.0, Double
1.0)
   in Double
x1 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
p2 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
p1 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
x2

-- ---------------------------------------------------------------------------
-- Order oracles (harmonic oscillator)
-- ---------------------------------------------------------------------------

-- | Energy of the harmonic oscillator @H(q,p) = (q^2 + p^2) / 2@.
harmonicEnergy :: (Double, Double) -> Double
harmonicEnergy :: (Double, Double) -> Double
harmonicEnergy (Double
q, Double
p) = (Double
q Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
q Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
p Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
p) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
2

-- | Integrate a one-step method for a fixed total time.
integrateForTime ::
  ((Double, Double) -> Double -> (Double, Double)) ->
  Double ->
  Double ->
  (Double, Double) ->
  (Double, Double)
integrateForTime :: ((Double, Double) -> Double -> (Double, Double))
-> Double -> Double -> (Double, Double) -> (Double, Double)
integrateForTime (Double, Double) -> Double -> (Double, Double)
stepper Double
eps Double
totalTime (Double, Double)
state =
  let n :: Int
n = Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
1 (Double -> Int
forall b. Integral b => Double -> b
forall a b. (RealFrac a, Integral b) => a -> b
round (Double
totalTime Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
eps))
   in ((Double, Double) -> (Double, Double))
-> (Double, Double) -> [(Double, Double)]
forall a. (a -> a) -> a -> [a]
iterate ((Double, Double) -> Double -> (Double, Double)
`stepper` Double
eps) (Double, Double)
state [(Double, Double)] -> Int -> (Double, Double)
forall a. HasCallStack => [a] -> Int -> a
!! Int
n

-- | Ordinary-least-squares slope of log(error) vs log(step size).
logLogSlope :: [(Double, Double)] -> Double
logLogSlope :: [(Double, Double)] -> Double
logLogSlope [(Double, Double)]
pairs =
  let n :: Double
n = Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral ([(Double, Double)] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [(Double, Double)]
pairs)
      xs :: [Double]
xs = ((Double, Double) -> Double) -> [(Double, Double)] -> [Double]
forall a b. (a -> b) -> [a] -> [b]
map (Double -> Double
forall a. Floating a => a -> a
log (Double -> Double)
-> ((Double, Double) -> Double) -> (Double, Double) -> Double
forall b c a. (b -> c) -> (a -> b) -> a -> c
. (Double, Double) -> Double
forall a b. (a, b) -> a
fst) [(Double, Double)]
pairs
      ys :: [Double]
ys = ((Double, Double) -> Double) -> [(Double, Double)] -> [Double]
forall a b. (a -> b) -> [a] -> [b]
map (Double -> Double
forall a. Floating a => a -> a
log (Double -> Double)
-> ((Double, Double) -> Double) -> (Double, Double) -> Double
forall b c a. (b -> c) -> (a -> b) -> a -> c
. (Double, Double) -> Double
forall a b. (a, b) -> b
snd) [(Double, Double)]
pairs
      xMean :: Double
xMean = [Double] -> Double
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [Double]
xs Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
n
      yMean :: Double
yMean = [Double] -> Double
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [Double]
ys Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
n
      num :: Double
num = [Double] -> Double
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum ((Double -> Double -> Double) -> [Double] -> [Double] -> [Double]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith (\Double
x Double
y -> (Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
xMean) Double -> Double -> Double
forall a. Num a => a -> a -> a
* (Double
y Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
yMean)) [Double]
xs [Double]
ys)
      den :: Double
den = [Double] -> Double
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum ((Double -> Double) -> [Double] -> [Double]
forall a b. (a -> b) -> [a] -> [b]
map (\Double
x -> (Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
xMean) Double -> Double -> Double
forall a. Num a => a -> a -> a
* (Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
xMean)) [Double]
xs)
   in Double
num Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
den

-- | Measured energy-error slope for leapfrog on the harmonic oscillator.
--
-- Expected value is 2.0.  We use four step sizes and a fixed integration
-- length so that the leading term dominates the observed error.
leapfrogOrderSlope :: Double -> (Double, Double) -> Double
leapfrogOrderSlope :: Double -> (Double, Double) -> Double
leapfrogOrderSlope Double
totalTime (Double, Double)
state =
  let e0 :: Double
e0 = (Double, Double) -> Double
harmonicEnergy (Double, Double)
state
      errorAt :: Double -> Double
errorAt Double
eps =
        Double -> Double
forall a. Num a => a -> a
abs ((Double, Double) -> Double
harmonicEnergy (((Double, Double) -> Double -> (Double, Double))
-> Double -> Double -> (Double, Double) -> (Double, Double)
integrateForTime (Double, Double) -> Double -> (Double, Double)
leapfrogStep Double
eps Double
totalTime (Double, Double)
state) Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
e0)
      pairs :: [(Double, Double)]
pairs = [(Double
eps, Double -> Double
errorAt Double
eps) | Double
eps <- [Double
0.05, Double
0.025, Double
0.0125, Double
0.00625]]
   in [(Double, Double)] -> Double
logLogSlope [(Double, Double)]
pairs

-- | Measured energy-error slope for Yoshida-4 on the harmonic oscillator.
--
-- Expected value is 4.0.
yoshida4OrderSlope :: Double -> (Double, Double) -> Double
yoshida4OrderSlope :: Double -> (Double, Double) -> Double
yoshida4OrderSlope Double
totalTime (Double, Double)
state =
  let e0 :: Double
e0 = (Double, Double) -> Double
harmonicEnergy (Double, Double)
state
      errorAt :: Double -> Double
errorAt Double
eps =
        Double -> Double
forall a. Num a => a -> a
abs ((Double, Double) -> Double
harmonicEnergy (((Double, Double) -> Double -> (Double, Double))
-> Double -> Double -> (Double, Double) -> (Double, Double)
integrateForTime (Double, Double) -> Double -> (Double, Double)
yoshida4Step Double
eps Double
totalTime (Double, Double)
state) Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
e0)
      pairs :: [(Double, Double)]
pairs = [(Double
eps, Double -> Double
errorAt Double
eps) | Double
eps <- [Double
0.05, Double
0.025, Double
0.0125, Double
0.00625]]
   in [(Double, Double)] -> Double
logLogSlope [(Double, Double)]
pairs

-- | One HMC step: sample momentum, run L leapfrog steps, MH accept/reject.
hmcStep :: Double -> Int -> Double -> IO (Double, Bool)
hmcStep :: Double -> Int -> Double -> IO (Double, Bool)
hmcStep Double
eps Int
nLeap Double
x = do
  u1 <- (Double, Double) -> IO Double
forall a (m :: * -> *). (Random a, MonadIO m) => (a, a) -> m a
randomRIO (Double
0, Double
1 :: Double)
  u2 <- randomRIO (0, 1 :: Double)
  let p = Double -> Double
forall a. Floating a => a -> a
sqrt (-(Double
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double -> Double
forall a. Floating a => a -> a
log (Double -> Double -> Double
forall a. Ord a => a -> a -> a
max Double
u1 Double
1e-300))) Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double -> Double
forall a. Floating a => a -> a
cos (Double
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
forall a. Floating a => a
pi Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
u2)
      eCurrent = -Double -> Double
logPTarget Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
+ (Double
p Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
p) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
2
      (xProp, pProp) = leapfrog eps nLeap (x, p)
      eProposed = -Double -> Double
logPTarget Double
xProp Double -> Double -> Double
forall a. Num a => a -> a -> a
+ (Double
pProp Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
pProp) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
2
      logAlpha = Double
eCurrent Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
eProposed
  u <- randomRIO (0, 1 :: Double)
  let accept = Double -> Double
forall a. Floating a => a -> a
log Double
u Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
< Double
logAlpha
      xNext = if Bool
accept then Double
xProp else Double
x
  pure (xNext, accept)

-- | Generate N HMC samples.
hmcSamples :: Int -> Double -> Int -> IO [Double]
hmcSamples :: Int -> Double -> Int -> IO [Double]
hmcSamples Int
nSamples Double
eps Int
nLeap =
  let go :: Int -> Double -> [Double] -> IO [Double]
go Int
0 Double
_ [Double]
xs = [Double] -> IO [Double]
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure ([Double] -> [Double]
forall a. [a] -> [a]
reverse [Double]
xs)
      go Int
k Double
x [Double]
xs = do
        (xNext, _) <- Double -> Int -> Double -> IO (Double, Bool)
hmcStep Double
eps Int
nLeap Double
x
        go (k - 1) xNext (xNext : xs)
   in Int -> Double -> [Double] -> IO [Double]
go Int
nSamples Double
0.0 []

-- | Check empirical moments match N(0,1).
momentTest :: [Double] -> (Bool, Double, Double)
momentTest :: [Double] -> (Bool, Double, Double)
momentTest [Double]
xs =
  let n :: Double
n = Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral ([Double] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [Double]
xs)
      mean :: Double
mean = [Double] -> Double
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [Double]
xs Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
n
      var :: Double
var = [Double] -> Double
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum ((Double -> Double) -> [Double] -> [Double]
forall a b. (a -> b) -> [a] -> [b]
map (\Double
x -> (Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
mean) Double -> Double -> Double
forall a. Floating a => a -> a -> a
** Double
2) [Double]
xs) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
n
      se :: Double
se = Double -> Double
forall a. Floating a => a -> a
sqrt (Double
var Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
n)
      meanOk :: Bool
meanOk = Double -> Double
forall a. Num a => a -> a
abs Double
mean Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
<= Double
3 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
se
      varOk :: Bool
varOk = Double -> Double
forall a. Num a => a -> a
abs (Double
var Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
1) Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
<= Double
0.2
   in (Bool
meanOk Bool -> Bool -> Bool
&& Bool
varOk, Double
mean, Double
var)