module Circuit.Inference.HMC
(
hmcSamples,
momentTest,
leapfrogStep,
leapfrog,
negateMomentum,
reverseLeapfrog,
leapfrogSystem,
leapfrogProcess,
yoshida4Step,
yoshida4,
yoshida4System,
yoshida4Process,
leapfrogReversible,
leapfrogJacobianDet,
yoshida4Reversible,
yoshida4JacobianDet,
leapfrogOrderSlope,
yoshida4OrderSlope,
)
where
import Circuit (Mono, Process, System, system)
import Circuit.Process (systemToProcess)
import Data.Void (absurd)
import System.Random (randomRIO)
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)
gradLogP :: Double -> Double
gradLogP :: Double -> Double
gradLogP Double
x = -Double
x
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)
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
negateMomentum :: (Double, Double) -> (Double, Double)
negateMomentum :: (Double, Double) -> (Double, Double)
negateMomentum (Double
x, Double
p) = (Double
x, -Double
p)
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
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', ()))
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)
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
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
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
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
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', ()))
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)
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)
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
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)
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
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
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
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
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
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
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)
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 []
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)