Files
tensor/grad/hmc_test.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

218 lines
6.9 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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")
}
}