Files
tensor/grad/hmc_test.go
T

218 lines
6.9 KiB
Go
Raw 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
package grad
import (
"math"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// gaussianLogDensity builds the log density of independent standard
// normals: log π(q) = −‖q‖²/2.
func gaussianLogDensity(q *Tensor) (*Tensor, error) {
sq, err := q.Mul(q)
if err != nil {
return nil, err
}
total, err := sq.Sum()
if err != nil {
return nil, err
}
return total.Scale(-0.5)
}
// gammaLogDensity builds log π(x) = log x − x, the unnormalised log
// density of a Gamma(2, 1) distribution.
func gammaLogDensity(q *Tensor) (*Tensor, error) {
logq, err := q.Log()
if err != nil {
return nil, err
}
shifted, err := logq.Sub(q)
if err != nil {
return nil, err
}
return shifted.Sum()
}
// TestSampleHMCNormal runs the chain on a two-dimensional standard
// normal and pins the sample moments: the target's mean is zero, its
// variance one and its components independent. A fixed seed makes the
// draw deterministic, so the bounds are checked facts about this run,
// not hopes about a random one.
func TestSampleHMCNormal(t *testing.T) {
q0, err := core.FromFloats([]float64{2, -2}, 2)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
samples, err := SampleHMC(gaussianLogDensity, q0,
HMCOptions{Step: 0.3, Steps: 20, BurnIn: 500, Samples: 6000, Thin: 1, Seed: 42})
if err != nil {
t.Fatalf("SampleHMC: %v", err)
}
if got := samples.Shape(); got[0] != 6000 || got[1] != 2 {
t.Fatalf("samples shape %v, want [6000 2]", got)
}
rows := samples.Shape()[0]
mean := []float64{0, 0}
variance := []float64{0, 0}
for r := range rows {
for c := range 2 {
v := samples.FloatAt(r*2 + c)
mean[c] += v / float64(rows)
}
}
for r := range rows {
for c := range 2 {
d := samples.FloatAt(r*2+c) - mean[c]
variance[c] += d * d / float64(rows)
}
}
for c := range 2 {
if math.Abs(mean[c]) > 0.15 {
t.Fatalf("mean[%d] = %.4g, want |mean| ≤ 0.15", c, mean[c])
}
if math.Abs(variance[c]-1) > 0.2 {
t.Fatalf("variance[%d] = %.4g, want 1 ± 0.2", c, variance[c])
}
}
// Cross moment of the independent components.
cov := 0.0
for r := range rows {
cov += (samples.FloatAt(r*2) - mean[0]) * (samples.FloatAt(r*2+1) - mean[1]) / float64(rows)
}
if math.Abs(cov) > 0.15 {
t.Fatalf("cross moment = %.4g, want |cov| ≤ 0.15", cov)
}
}
// TestSampleHMCDeterministic checks the seed contract: the same seed
// replays bit-identically, a different seed does not.
func TestSampleHMCDeterministic(t *testing.T) {
q0, _ := core.FromFloats([]float64{1}, 1)
run := func(seed int64) *core.Array {
samples, err := SampleHMC(gaussianLogDensity, q0,
HMCOptions{Step: 0.4, Steps: 16, BurnIn: 100, Samples: 200, Seed: seed})
if err != nil {
t.Fatalf("SampleHMC: %v", err)
}
return samples
}
a, b := run(7), run(7)
c := run(8)
for i := range a.Len() {
if a.FloatAt(i) != b.FloatAt(i) {
t.Fatalf("the same seed produced different samples at %d", i)
}
}
same := true
for i := range a.Len() {
if a.FloatAt(i) != c.FloatAt(i) {
same = false
break
}
}
if same {
t.Fatal("different seeds produced identical samples")
}
}
// TestSampleHMCSupportedDensity samples a Gamma(2, 1) target,
// log π(x) = log x − x on x > 0. Proposals that overshoot into the
// forbidden half-line yield a NaN density and are rejected, so every
// kept sample stays positive and the mean approaches 2.
func TestSampleHMCSupportedDensity(t *testing.T) {
q0, _ := core.FromFloats([]float64{1}, 1)
samples, err := SampleHMC(gammaLogDensity, q0,
HMCOptions{Step: 0.3, Steps: 10, BurnIn: 500, Samples: 6000, Seed: 3})
if err != nil {
t.Fatalf("SampleHMC: %v", err)
}
sum := 0.0
for i := range samples.Len() {
x := samples.FloatAt(i)
if x <= 0 {
t.Fatalf("sample %d = %g left the support", i, x)
}
sum += x
}
mean := sum / float64(samples.Len())
if math.Abs(mean-2) > 0.15 {
t.Fatalf("Gamma(2) mean = %.4g, want 2 ± 0.15", mean)
}
}
// TestSampleHMCRejectedByGuard exercises the explicit rejection path:
// the density errors outside its support instead of returning NaN,
// and the chain still stays inside it.
func TestSampleHMCRejectedByGuard(t *testing.T) {
q0, _ := core.FromFloats([]float64{0.5}, 1)
samples, err := SampleHMC(func(q *Tensor) (*Tensor, error) {
x := q.Data().FloatAt(0)
if x <= 0 {
return nil, errf("outside the support")
}
return gammaLogDensity(q)
}, q0, HMCOptions{Step: 0.5, Steps: 20, BurnIn: 300, Samples: 2000, Seed: 5})
if err != nil {
t.Fatalf("SampleHMC: %v", err)
}
for i := range samples.Len() {
if samples.FloatAt(i) <= 0 {
t.Fatalf("sample %d = %g left the support", i, samples.FloatAt(i))
}
}
}
// TestSampleHMCErrors pins the validation contract, including a
// density that fails at the start, returns a non-scalar, or never
// touches the leaf core.
func TestSampleHMCErrors(t *testing.T) {
q0, _ := core.FromFloats([]float64{1}, 1)
if _, err := SampleHMC(nil, q0, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil {
t.Fatal("expected an error for a nil density")
}
rank2, _ := core.FromFloats([]float64{1, 1}, 1, 2)
if _, err := SampleHMC(gaussianLogDensity, rank2, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil {
t.Fatal("expected an error for a rank-2 state")
}
if _, err := SampleHMC(gaussianLogDensity, nil, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil {
t.Fatal("expected an error for a nil state")
}
if _, err := SampleHMC(gaussianLogDensity, q0, HMCOptions{Steps: 5, Samples: 1}); err == nil {
t.Fatal("expected an error for a non-positive step")
}
if _, err := SampleHMC(gaussianLogDensity, q0, HMCOptions{Step: 0.1, Samples: 1}); err == nil {
t.Fatal("expected an error for a non-positive step count")
}
if _, err := SampleHMC(gaussianLogDensity, q0, HMCOptions{Step: 0.1, Steps: 5}); err == nil {
t.Fatal("expected an error for a non-positive sample count")
}
if _, err := SampleHMC(gaussianLogDensity, q0, HMCOptions{Step: 0.1, Steps: 5, Samples: 1, BurnIn: -3}); err == nil {
t.Fatal("expected an error for a negative BurnIn")
}
nonScalar := func(q *Tensor) (*Tensor, error) {
two, _ := core.FromFloats([]float64{1, 2}, 2)
return FromArray(two, false), nil
}
if _, err := SampleHMC(nonScalar, q0, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil {
t.Fatal("expected an error for a non-scalar density")
}
detached := func(q *Tensor) (*Tensor, error) {
one, _ := core.FromFloats([]float64{1}, 1)
return FromArray(one, false), nil
}
if _, err := SampleHMC(detached, q0, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil {
t.Fatal("expected an error for a density disconnected from the leaf")
}
failsAtStart := func(q *Tensor) (*Tensor, error) {
return nil, errf("no density at the start")
}
if _, err := SampleHMC(failsAtStart, q0, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil {
t.Fatal("expected the start-time density error to be fatal")
}
}