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