218 lines
6.9 KiB
Go
218 lines
6.9 KiB
Go
// 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")
|
||
}
|
||
}
|