feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,102 @@
|
||||
// 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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user