Files

103 lines
2.7 KiB
Go
Raw Permalink Normal View History

2026-09-03 10:00:00 +02:00
// 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)
}
}