103 lines
2.7 KiB
Go
103 lines
2.7 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|||
|
|
// SPDX-License-Identifier: MIT
|
||
|
|
|
||
|
|
// Command hmc samples a correlated two-dimensional Gaussian by
|
||
|
|
// Hamiltonian Monte Carlo on the differentiable log density, and
|
||
|
|
// checks the chain against the distribution's known moments: mean
|
||
|
|
// zero, unit variances, correlation 0.9. The sampler never sees an
|
||
|
|
// analytic gradient, only the autograd's.
|
||
|
|
//
|
||
|
|
// Usage: go run ./examples/hmc
|
||
|
|
package main
|
||
|
|
|
||
|
|
import (
|
||
|
|
"fmt"
|
||
|
|
"log"
|
||
|
|
"math"
|
||
|
|
|
||
|
|
"sourcedock.dev/petrbalvin/tensor"
|
||
|
|
)
|
||
|
|
|
||
|
|
func main() {
|
||
|
|
const rho = 0.9
|
||
|
|
// Precision matrix of the correlated Gaussian (up to scale, which
|
||
|
|
// the unnormalised density does not need).
|
||
|
|
prec, err := tensor.FromFloats([]float64{1, -rho, -rho, 1}, 2, 2)
|
||
|
|
if err != nil {
|
||
|
|
log.Fatal(err)
|
||
|
|
}
|
||
|
|
aT := tensor.FromArray(prec, false)
|
||
|
|
|
||
|
|
logDensity := func(q *tensor.Tensor) (*tensor.Tensor, error) {
|
||
|
|
// log p(q) ∝ −½ qᵀAq: one matrix-vector product on the graph,
|
||
|
|
// then the inner product with q itself.
|
||
|
|
r, err := aT.MatMul(q)
|
||
|
|
if err != nil {
|
||
|
|
return nil, err
|
||
|
|
}
|
||
|
|
quad, err := r.Mul(q)
|
||
|
|
if err != nil {
|
||
|
|
return nil, err
|
||
|
|
}
|
||
|
|
s, err := quad.Sum()
|
||
|
|
if err != nil {
|
||
|
|
return nil, err
|
||
|
|
}
|
||
|
|
return s.Scale(-0.5)
|
||
|
|
}
|
||
|
|
|
||
|
|
q0, err := tensor.FromFloats([]float64{0.5, -0.5}, 2)
|
||
|
|
if err != nil {
|
||
|
|
log.Fatal(err)
|
||
|
|
}
|
||
|
|
samples, err := tensor.SampleHMC(logDensity, q0, tensor.HMCOptions{
|
||
|
|
Step: 0.15,
|
||
|
|
Steps: 12,
|
||
|
|
BurnIn: 500,
|
||
|
|
Thin: 5,
|
||
|
|
Samples: 20000,
|
||
|
|
Seed: 7,
|
||
|
|
})
|
||
|
|
if err != nil {
|
||
|
|
log.Fatal(err)
|
||
|
|
}
|
||
|
|
|
||
|
|
n := samples.Shape()[0]
|
||
|
|
mean0, mean1 := 0.0, 0.0
|
||
|
|
for i := range n {
|
||
|
|
mean0 += samples.FloatAt(i * 2)
|
||
|
|
mean1 += samples.FloatAt(i*2 + 1)
|
||
|
|
}
|
||
|
|
mean0 /= float64(n)
|
||
|
|
mean1 /= float64(n)
|
||
|
|
var0, var1, cov := 0.0, 0.0, 0.0
|
||
|
|
for i := range n {
|
||
|
|
d0 := samples.FloatAt(i*2) - mean0
|
||
|
|
d1 := samples.FloatAt(i*2+1) - mean1
|
||
|
|
var0 += d0 * d0
|
||
|
|
var1 += d1 * d1
|
||
|
|
cov += d0 * d1
|
||
|
|
}
|
||
|
|
var0 /= float64(n - 1)
|
||
|
|
var1 /= float64(n - 1)
|
||
|
|
cov /= float64(n - 1)
|
||
|
|
corr := cov / math.Sqrt(var0*var1)
|
||
|
|
|
||
|
|
fmt.Printf("samples %d\n", n)
|
||
|
|
fmt.Printf("mean (%.3f, %.3f), want (0, 0)\n", mean0, mean1)
|
||
|
|
// The covariance is A⁻¹: unit-over-(1−ρ²) variances around the
|
||
|
|
// correlation rho.
|
||
|
|
wantVar := 1 / (1 - rho*rho)
|
||
|
|
fmt.Printf("variance (%.3f, %.3f), want (%.3f, %.3f)\n", var0, var1, wantVar, wantVar)
|
||
|
|
fmt.Printf("correlation %.3f, want %.3f\n", corr, rho)
|
||
|
|
if math.Abs(mean0) > 0.05 || math.Abs(mean1) > 0.05 {
|
||
|
|
log.Fatal("the sample mean drifted")
|
||
|
|
}
|
||
|
|
if math.Abs(corr-rho) > 0.02 {
|
||
|
|
log.Fatalf("the sample correlation %.3f missed %.3f", corr, rho)
|
||
|
|
}
|
||
|
|
if math.Abs(var0-wantVar) > 0.15*wantVar || math.Abs(var1-wantVar) > 0.15*wantVar {
|
||
|
|
log.Fatalf("the sample variances (%.3f, %.3f) missed %.3f", var0, var1, wantVar)
|
||
|
|
}
|
||
|
|
}
|