// Copyright (c) 2026 Petr Balvín (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) } }