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")
|
|||
|
|
}
|
|||
|
|
}
|